init (not checked bat script yet)

This commit is contained in:
litagin02
2023-12-27 05:37:46 +09:00
parent 11a1e7e80d
commit 58fab45b84
235 changed files with 2831 additions and 773509 deletions

142
utils.py
View File

@@ -1,39 +1,22 @@
import os
import glob
import argparse
import logging
import glob
import json
import shutil
import subprocess
import numpy as np
from huggingface_hub import hf_hub_download
from scipy.io.wavfile import read
import torch
import logging
import os
import re
from collections import OrderedDict
import subprocess
import numpy as np
import torch
from huggingface_hub import hf_hub_download
from safetensors import safe_open
from safetensors.torch import save_file
from scipy.io.wavfile import read
from tools.log import logger
MATPLOTLIB_FLAG = False
logger = logging.getLogger(__name__)
def download_emo_models(mirror, repo_id, model_name):
if mirror == "openi":
import openi
openi.model.download_model(
"Stardust_minus/Bert-VITS2",
repo_id.split("/")[-1],
"./emotional",
)
else:
hf_hub_download(
repo_id,
"pytorch_model.bin",
local_dir=model_name,
local_dir_use_symlinks=False,
)
def download_checkpoint(
dir_path, repo_config, token=None, regex="G_*.pth", mirror="openi"
@@ -43,31 +26,20 @@ def download_checkpoint(
if f_list:
print("Use existed model, skip downloading.")
return
if mirror.lower() == "openi":
import openi
kwargs = {"token": token} if token else {}
openi.login(**kwargs)
model_image = repo_config["model_image"]
openi.model.download_model(repo_id, model_image, dir_path)
fs = glob.glob(os.path.join(dir_path, model_image, "*.pth"))
for file in fs:
shutil.move(file, dir_path)
shutil.rmtree(os.path.join(dir_path, model_image))
else:
for file in ["DUR_0.pth", "D_0.pth", "G_0.pth"]:
hf_hub_download(
repo_id, file, local_dir=dir_path, local_dir_use_symlinks=False
)
for file in ["DUR_0.pth", "D_0.pth", "G_0.pth"]:
hf_hub_download(repo_id, file, local_dir=dir_path, local_dir_use_symlinks=False)
def load_checkpoint(checkpoint_path, model, optimizer=None, skip_optimizer=False):
def load_checkpoint(
checkpoint_path, model, optimizer=None, skip_optimizer=False, for_infer=False
):
assert os.path.isfile(checkpoint_path)
checkpoint_dict = torch.load(checkpoint_path, map_location="cpu")
iteration = checkpoint_dict["iteration"]
learning_rate = checkpoint_dict["learning_rate"]
logger.info(
f"Loading model and optimizer at iteration {iteration} from {checkpoint_path}"
)
if (
optimizer is not None
and not skip_optimizer
@@ -104,8 +76,10 @@ def load_checkpoint(checkpoint_path, model, optimizer=None, skip_optimizer=False
logger.warn(
f"Seems you are using the old version of the model, the {k} is automatically set to zero for backward compatibility"
)
elif "enc_q" in k and for_infer:
continue
else:
logger.error(f"{k} is not in the checkpoint")
logger.error(f"{k} is not in the checkpoint {checkpoint_path}")
new_state_dict[k] = v
@@ -142,33 +116,59 @@ def save_checkpoint(model, optimizer, learning_rate, iteration, checkpoint_path)
)
def save_compressed_models_checkpoint(model, iteration, checkpoint_path, ishalf=False):
logger.info(
f"Saving compressed model at iteration {iteration} to {checkpoint_path}"
)
def save_safetensors(model, iteration, checkpoint_path, is_half=False, for_infer=False):
"""
Save model with safetensors.
"""
if hasattr(model, "module"):
state_dict = model.module.state_dict()
else:
state_dict = model.state_dict()
keys = []
for k in state_dict:
if "enc_q" in k:
if "enc_q" in k and for_infer:
continue # noqa: E701
keys.append(k)
new_dict_g = (
new_dict = (
{k: state_dict[k].half() for k in keys}
if ishalf
if is_half
else {k: state_dict[k] for k in keys}
)
torch.save(
{
"model": new_dict_g,
"iteration": iteration,
"learning_rate": 0,
},
checkpoint_path,
)
new_dict["iteration"] = torch.LongTensor([iteration])
logger.info(f"Saved safetensors to {checkpoint_path}")
save_file(new_dict, checkpoint_path)
def load_safetensors(checkpoint_path, model, for_infer=False):
"""
Load safetensors model.
"""
tensors = {}
iteration = None
with safe_open(checkpoint_path, framework="pt", device="cpu") as f:
for key in f.keys():
if key == "iteration":
iteration = f.get_tensor(key).item()
tensors[key] = f.get_tensor(key)
if hasattr(model, "module"):
result = model.module.load_state_dict(tensors, strict=False)
else:
result = model.load_state_dict(tensors, strict=False)
for key in result.missing_keys:
if key.startswith("enc_q") and for_infer:
continue
logger.warning(f"Missing key: {key}")
for key in result.unexpected_keys:
if key == "iteration":
continue
logger.warning(f"Unexpected key: {key}")
if iteration is None:
logger.info(f"Loaded safetensors '{checkpoint_path}'")
else:
logger.info(f"Loaded safetensors '{checkpoint_path}' (iteration {iteration})")
return model, iteration
def summarize(
@@ -190,10 +190,20 @@ def summarize(
writer.add_audio(k, v, global_step, audio_sampling_rate)
def is_resuming(dir_path):
g_list = glob.glob(os.path.join(dir_path, "G_*.pth"))
d_list = glob.glob(os.path.join(dir_path, "D_*.pth"))
dur_list = glob.glob(os.path.join(dir_path, "DUR_*.pth"))
return len(g_list) > 0 and len(d_list) > 0 and len(dur_list) > 0
def latest_checkpoint_path(dir_path, regex="G_*.pth"):
f_list = glob.glob(os.path.join(dir_path, regex))
f_list.sort(key=lambda f: int("".join(filter(str.isdigit, f))))
x = f_list[-1]
try:
x = f_list[-1]
except IndexError:
raise ValueError(f"No checkpoint found in {dir_path} with regex {regex}")
return x