Feat: make style vectors w.r.t. subdirs structure
This commit is contained in:
@@ -9,7 +9,7 @@ from style_bert_vits2.logging import logger
|
|||||||
|
|
||||||
|
|
||||||
def set_style_config(json_path: Path, output_path: Path):
|
def set_style_config(json_path: Path, output_path: Path):
|
||||||
with open(json_path, "r", encoding="utf-8") as f:
|
with open(json_path, encoding="utf-8") as f:
|
||||||
json_dict = json.load(f)
|
json_dict = json.load(f)
|
||||||
json_dict["data"]["num_styles"] = 1
|
json_dict["data"]["num_styles"] = 1
|
||||||
json_dict["data"]["style2id"] = {DEFAULT_STYLE: 0}
|
json_dict["data"]["style2id"] = {DEFAULT_STYLE: 0}
|
||||||
@@ -21,6 +21,7 @@ def set_style_config(json_path: Path, output_path: Path):
|
|||||||
def save_neutral_vector(wav_dir: Union[Path, str], output_path: Union[Path, str]):
|
def save_neutral_vector(wav_dir: Union[Path, str], output_path: Union[Path, str]):
|
||||||
wav_dir = Path(wav_dir)
|
wav_dir = Path(wav_dir)
|
||||||
output_path = Path(output_path)
|
output_path = Path(output_path)
|
||||||
|
json_path = output_path / "config.json"
|
||||||
embs = []
|
embs = []
|
||||||
for file in wav_dir.rglob("*.npy"):
|
for file in wav_dir.rglob("*.npy"):
|
||||||
xvec = np.load(file)
|
xvec = np.load(file)
|
||||||
@@ -31,3 +32,63 @@ def save_neutral_vector(wav_dir: Union[Path, str], output_path: Union[Path, str]
|
|||||||
only_mean = np.stack([mean]) # (1, 256)
|
only_mean = np.stack([mean]) # (1, 256)
|
||||||
np.save(output_path, only_mean)
|
np.save(output_path, only_mean)
|
||||||
logger.info(f"Saved mean style vector to {output_path}")
|
logger.info(f"Saved mean style vector to {output_path}")
|
||||||
|
|
||||||
|
with open(json_path, "r", encoding="utf-8") as f:
|
||||||
|
json_dict = json.load(f)
|
||||||
|
json_dict["data"]["num_styles"] = 1
|
||||||
|
json_dict["data"]["style2id"] = {DEFAULT_STYLE: 0}
|
||||||
|
with open(json_path, "w", encoding="utf-8") as f:
|
||||||
|
json.dump(json_dict, f, indent=2, ensure_ascii=False)
|
||||||
|
logger.info(f"Saved style config to {json_path}")
|
||||||
|
|
||||||
|
|
||||||
|
def save_styles_by_dirs(wav_dir: Union[Path, str], output_dir: Union[Path, str]):
|
||||||
|
wav_dir = Path(wav_dir)
|
||||||
|
output_dir = Path(output_dir)
|
||||||
|
output_dir.mkdir(parents=True, exist_ok=True)
|
||||||
|
json_path = output_dir / "config.json"
|
||||||
|
|
||||||
|
subdirs = [d for d in wav_dir.iterdir() if d.is_dir()]
|
||||||
|
subdirs.sort()
|
||||||
|
if len(subdirs) in (0, 1):
|
||||||
|
logger.warning("No style directories found. Saving only neutral style.")
|
||||||
|
save_neutral_vector(wav_dir, output_dir)
|
||||||
|
|
||||||
|
# First get mean of all for Neutral
|
||||||
|
embs = []
|
||||||
|
for file in wav_dir.rglob("*.npy"):
|
||||||
|
xvec = np.load(file)
|
||||||
|
embs.append(np.expand_dims(xvec, axis=0))
|
||||||
|
x = np.concatenate(embs, axis=0) # (N, 256)
|
||||||
|
mean = np.mean(x, axis=0) # (256,)
|
||||||
|
style_vectors = [mean]
|
||||||
|
|
||||||
|
names = [DEFAULT_STYLE]
|
||||||
|
for style_dir in subdirs:
|
||||||
|
npy_files = list(style_dir.rglob("*.npy"))
|
||||||
|
if not npy_files:
|
||||||
|
continue
|
||||||
|
embs = []
|
||||||
|
for file in npy_files:
|
||||||
|
xvec = np.load(file)
|
||||||
|
embs.append(np.expand_dims(xvec, axis=0))
|
||||||
|
|
||||||
|
x = np.concatenate(embs, axis=0) # (N, 256)
|
||||||
|
mean = np.mean(x, axis=0) # (256,)
|
||||||
|
style_vectors.append(mean)
|
||||||
|
names.append(style_dir.name)
|
||||||
|
|
||||||
|
# Stack them to make (num_styles, 256)
|
||||||
|
style_vectors_npy = np.stack(style_vectors, axis=0)
|
||||||
|
np.save(output_dir / "style_vectors.npy", style_vectors_npy)
|
||||||
|
logger.info(f"Saved style vectors to {output_dir / 'style_vectors.npy'}")
|
||||||
|
|
||||||
|
# Save style2id config to json
|
||||||
|
style2id = {name: i for i, name in enumerate(names)}
|
||||||
|
with open(json_path, "r", encoding="utf-8") as f:
|
||||||
|
json_dict = json.load(f)
|
||||||
|
json_dict["data"]["num_styles"] = len(names)
|
||||||
|
json_dict["data"]["style2id"] = style2id
|
||||||
|
with open(json_path, "w", encoding="utf-8") as f:
|
||||||
|
json.dump(json_dict, f, indent=2, ensure_ascii=False)
|
||||||
|
logger.info(f"Saved style config to {json_path}")
|
||||||
|
|||||||
37
train_ms.py
37
train_ms.py
@@ -108,7 +108,7 @@ def run():
|
|||||||
envs = config.train_ms_config.env
|
envs = config.train_ms_config.env
|
||||||
for env_name, env_value in envs.items():
|
for env_name, env_value in envs.items():
|
||||||
if env_name not in os.environ.keys():
|
if env_name not in os.environ.keys():
|
||||||
logger.info("Loading configuration from config {}".format(str(env_value)))
|
logger.info(f"Loading configuration from config {env_value!s}")
|
||||||
os.environ[env_name] = str(env_value)
|
os.environ[env_name] = str(env_value)
|
||||||
logger.info(
|
logger.info(
|
||||||
"Loading environment variables \nMASTER_ADDR: {},\nMASTER_PORT: {},\nWORLD_SIZE: {},\nRANK: {},\nLOCAL_RANK: {}".format(
|
"Loading environment variables \nMASTER_ADDR: {},\nMASTER_PORT: {},\nWORLD_SIZE: {},\nRANK: {},\nLOCAL_RANK: {}".format(
|
||||||
@@ -142,7 +142,7 @@ def run():
|
|||||||
if os.path.realpath(args.config) != os.path.realpath(
|
if os.path.realpath(args.config) != os.path.realpath(
|
||||||
config.train_ms_config.config_path
|
config.train_ms_config.config_path
|
||||||
):
|
):
|
||||||
with open(args.config, "r", encoding="utf-8") as f:
|
with open(args.config, encoding="utf-8") as f:
|
||||||
data = f.read()
|
data = f.read()
|
||||||
os.makedirs(os.path.dirname(config.train_ms_config.config_path), exist_ok=True)
|
os.makedirs(os.path.dirname(config.train_ms_config.config_path), exist_ok=True)
|
||||||
with open(config.train_ms_config.config_path, "w", encoding="utf-8") as f:
|
with open(config.train_ms_config.config_path, "w", encoding="utf-8") as f:
|
||||||
@@ -192,13 +192,9 @@ def run():
|
|||||||
os.makedirs(config.out_dir, exist_ok=True)
|
os.makedirs(config.out_dir, exist_ok=True)
|
||||||
|
|
||||||
if not args.skip_default_style:
|
if not args.skip_default_style:
|
||||||
# Save default style to out_dir
|
default_style.save_styles_by_dirs(
|
||||||
default_style.set_style_config(
|
|
||||||
args.config, os.path.join(config.out_dir, "config.json")
|
|
||||||
)
|
|
||||||
default_style.save_neutral_vector(
|
|
||||||
os.path.join(args.model, "wavs"),
|
os.path.join(args.model, "wavs"),
|
||||||
os.path.join(config.out_dir, "style_vectors.npy"),
|
config.out_dir,
|
||||||
)
|
)
|
||||||
|
|
||||||
torch.manual_seed(hps.train.seed)
|
torch.manual_seed(hps.train.seed)
|
||||||
@@ -505,7 +501,7 @@ def run():
|
|||||||
optim_g,
|
optim_g,
|
||||||
hps.train.learning_rate,
|
hps.train.learning_rate,
|
||||||
epoch,
|
epoch,
|
||||||
os.path.join(model_dir, "G_{}.pth".format(global_step)),
|
os.path.join(model_dir, f"G_{global_step}.pth"),
|
||||||
)
|
)
|
||||||
assert optim_d is not None
|
assert optim_d is not None
|
||||||
utils.checkpoints.save_checkpoint(
|
utils.checkpoints.save_checkpoint(
|
||||||
@@ -513,7 +509,7 @@ def run():
|
|||||||
optim_d,
|
optim_d,
|
||||||
hps.train.learning_rate,
|
hps.train.learning_rate,
|
||||||
epoch,
|
epoch,
|
||||||
os.path.join(model_dir, "D_{}.pth".format(global_step)),
|
os.path.join(model_dir, f"D_{global_step}.pth"),
|
||||||
)
|
)
|
||||||
if net_dur_disc is not None:
|
if net_dur_disc is not None:
|
||||||
assert optim_dur_disc is not None
|
assert optim_dur_disc is not None
|
||||||
@@ -522,7 +518,7 @@ def run():
|
|||||||
optim_dur_disc,
|
optim_dur_disc,
|
||||||
hps.train.learning_rate,
|
hps.train.learning_rate,
|
||||||
epoch,
|
epoch,
|
||||||
os.path.join(model_dir, "DUR_{}.pth".format(global_step)),
|
os.path.join(model_dir, f"DUR_{global_step}.pth"),
|
||||||
)
|
)
|
||||||
utils.safetensors.save_safetensors(
|
utils.safetensors.save_safetensors(
|
||||||
net_g,
|
net_g,
|
||||||
@@ -757,14 +753,12 @@ def train_and_evaluate(
|
|||||||
"loss/g/kl": loss_kl,
|
"loss/g/kl": loss_kl,
|
||||||
}
|
}
|
||||||
)
|
)
|
||||||
|
scalar_dict.update({f"loss/g/{i}": v for i, v in enumerate(losses_gen)})
|
||||||
scalar_dict.update(
|
scalar_dict.update(
|
||||||
{"loss/g/{}".format(i): v for i, v in enumerate(losses_gen)}
|
{f"loss/d_r/{i}": v for i, v in enumerate(losses_disc_r)}
|
||||||
)
|
)
|
||||||
scalar_dict.update(
|
scalar_dict.update(
|
||||||
{"loss/d_r/{}".format(i): v for i, v in enumerate(losses_disc_r)}
|
{f"loss/d_g/{i}": v for i, v in enumerate(losses_disc_g)}
|
||||||
)
|
|
||||||
scalar_dict.update(
|
|
||||||
{"loss/d_g/{}".format(i): v for i, v in enumerate(losses_disc_g)}
|
|
||||||
)
|
)
|
||||||
|
|
||||||
image_dict = {
|
image_dict = {
|
||||||
@@ -801,14 +795,14 @@ def train_and_evaluate(
|
|||||||
optim_g,
|
optim_g,
|
||||||
hps.train.learning_rate,
|
hps.train.learning_rate,
|
||||||
epoch,
|
epoch,
|
||||||
os.path.join(hps.model_dir, "G_{}.pth".format(global_step)),
|
os.path.join(hps.model_dir, f"G_{global_step}.pth"),
|
||||||
)
|
)
|
||||||
utils.checkpoints.save_checkpoint(
|
utils.checkpoints.save_checkpoint(
|
||||||
net_d,
|
net_d,
|
||||||
optim_d,
|
optim_d,
|
||||||
hps.train.learning_rate,
|
hps.train.learning_rate,
|
||||||
epoch,
|
epoch,
|
||||||
os.path.join(hps.model_dir, "D_{}.pth".format(global_step)),
|
os.path.join(hps.model_dir, f"D_{global_step}.pth"),
|
||||||
)
|
)
|
||||||
if net_dur_disc is not None:
|
if net_dur_disc is not None:
|
||||||
utils.checkpoints.save_checkpoint(
|
utils.checkpoints.save_checkpoint(
|
||||||
@@ -816,7 +810,7 @@ def train_and_evaluate(
|
|||||||
optim_dur_disc,
|
optim_dur_disc,
|
||||||
hps.train.learning_rate,
|
hps.train.learning_rate,
|
||||||
epoch,
|
epoch,
|
||||||
os.path.join(hps.model_dir, "DUR_{}.pth".format(global_step)),
|
os.path.join(hps.model_dir, f"DUR_{global_step}.pth"),
|
||||||
)
|
)
|
||||||
keep_ckpts = config.train_ms_config.keep_ckpts
|
keep_ckpts = config.train_ms_config.keep_ckpts
|
||||||
if keep_ckpts > 0:
|
if keep_ckpts > 0:
|
||||||
@@ -853,9 +847,7 @@ def train_and_evaluate(
|
|||||||
global_step += 1
|
global_step += 1
|
||||||
if pbar is not None:
|
if pbar is not None:
|
||||||
pbar.set_description(
|
pbar.set_description(
|
||||||
"Epoch {}({:.0f}%)/{}".format(
|
f"Epoch {epoch}({100.0 * batch_idx / len(train_loader):.0f}%)/{hps.train.epochs}"
|
||||||
epoch, 100.0 * batch_idx / len(train_loader), hps.train.epochs
|
|
||||||
)
|
|
||||||
)
|
)
|
||||||
pbar.update()
|
pbar.update()
|
||||||
# 本家ではこれをスピードアップのために消すと書かれていたので、一応消してみる
|
# 本家ではこれをスピードアップのために消すと書かれていたので、一応消してみる
|
||||||
@@ -870,6 +862,7 @@ def evaluate(hps, generator, eval_loader, writer_eval):
|
|||||||
generator.eval()
|
generator.eval()
|
||||||
image_dict = {}
|
image_dict = {}
|
||||||
audio_dict = {}
|
audio_dict = {}
|
||||||
|
print()
|
||||||
logger.info("Evaluating ...")
|
logger.info("Evaluating ...")
|
||||||
with torch.no_grad():
|
with torch.no_grad():
|
||||||
for batch_idx, (
|
for batch_idx, (
|
||||||
|
|||||||
@@ -109,7 +109,7 @@ def run():
|
|||||||
envs = config.train_ms_config.env
|
envs = config.train_ms_config.env
|
||||||
for env_name, env_value in envs.items():
|
for env_name, env_value in envs.items():
|
||||||
if env_name not in os.environ.keys():
|
if env_name not in os.environ.keys():
|
||||||
logger.info("Loading configuration from config {}".format(str(env_value)))
|
logger.info(f"Loading configuration from config {env_value!s}")
|
||||||
os.environ[env_name] = str(env_value)
|
os.environ[env_name] = str(env_value)
|
||||||
logger.info(
|
logger.info(
|
||||||
"Loading environment variables \nMASTER_ADDR: {},\nMASTER_PORT: {},\nWORLD_SIZE: {},\nRANK: {},\nLOCAL_RANK: {}".format(
|
"Loading environment variables \nMASTER_ADDR: {},\nMASTER_PORT: {},\nWORLD_SIZE: {},\nRANK: {},\nLOCAL_RANK: {}".format(
|
||||||
@@ -143,7 +143,7 @@ def run():
|
|||||||
if os.path.realpath(args.config) != os.path.realpath(
|
if os.path.realpath(args.config) != os.path.realpath(
|
||||||
config.train_ms_config.config_path
|
config.train_ms_config.config_path
|
||||||
):
|
):
|
||||||
with open(args.config, "r", encoding="utf-8") as f:
|
with open(args.config, encoding="utf-8") as f:
|
||||||
data = f.read()
|
data = f.read()
|
||||||
os.makedirs(os.path.dirname(config.train_ms_config.config_path), exist_ok=True)
|
os.makedirs(os.path.dirname(config.train_ms_config.config_path), exist_ok=True)
|
||||||
with open(config.train_ms_config.config_path, "w", encoding="utf-8") as f:
|
with open(config.train_ms_config.config_path, "w", encoding="utf-8") as f:
|
||||||
@@ -193,13 +193,9 @@ def run():
|
|||||||
os.makedirs(config.out_dir, exist_ok=True)
|
os.makedirs(config.out_dir, exist_ok=True)
|
||||||
|
|
||||||
if not args.skip_default_style:
|
if not args.skip_default_style:
|
||||||
# Save default style to out_dir
|
default_style.save_styles_by_dirs(
|
||||||
default_style.set_style_config(
|
|
||||||
args.config, os.path.join(config.out_dir, "config.json")
|
|
||||||
)
|
|
||||||
default_style.save_neutral_vector(
|
|
||||||
os.path.join(args.model, "wavs"),
|
os.path.join(args.model, "wavs"),
|
||||||
os.path.join(config.out_dir, "style_vectors.npy"),
|
config.out_dir,
|
||||||
)
|
)
|
||||||
|
|
||||||
torch.manual_seed(hps.train.seed)
|
torch.manual_seed(hps.train.seed)
|
||||||
@@ -215,24 +211,25 @@ def run():
|
|||||||
writer = SummaryWriter(log_dir=model_dir)
|
writer = SummaryWriter(log_dir=model_dir)
|
||||||
writer_eval = SummaryWriter(log_dir=os.path.join(model_dir, "eval"))
|
writer_eval = SummaryWriter(log_dir=os.path.join(model_dir, "eval"))
|
||||||
train_dataset = TextAudioSpeakerLoader(hps.data.training_files, hps.data)
|
train_dataset = TextAudioSpeakerLoader(hps.data.training_files, hps.data)
|
||||||
train_sampler = DistributedBucketSampler(
|
# train_sampler = DistributedBucketSampler(
|
||||||
train_dataset,
|
# train_dataset,
|
||||||
hps.train.batch_size,
|
# hps.train.batch_size,
|
||||||
[32, 300, 400, 500, 600, 700, 800, 900, 1000],
|
# [32, 300, 400, 500, 600, 700, 800, 900, 1000],
|
||||||
num_replicas=n_gpus,
|
# num_replicas=n_gpus,
|
||||||
rank=rank,
|
# rank=rank,
|
||||||
shuffle=True,
|
# shuffle=True,
|
||||||
)
|
# )
|
||||||
collate_fn = TextAudioSpeakerCollate(use_jp_extra=True)
|
collate_fn = TextAudioSpeakerCollate(use_jp_extra=True)
|
||||||
train_loader = DataLoader(
|
train_loader = DataLoader(
|
||||||
train_dataset,
|
train_dataset,
|
||||||
# メモリ消費量を減らそうとnum_workersを1にしてみる
|
# メモリ消費量を減らそうとnum_workersを1にしてみる
|
||||||
# num_workers=min(config.train_ms_config.num_workers, os.cpu_count() // 2),
|
# num_workers=min(config.train_ms_config.num_workers, os.cpu_count() // 2),
|
||||||
num_workers=1,
|
num_workers=1,
|
||||||
shuffle=False,
|
shuffle=True,
|
||||||
pin_memory=True,
|
pin_memory=True,
|
||||||
collate_fn=collate_fn,
|
collate_fn=collate_fn,
|
||||||
batch_sampler=train_sampler,
|
# batch_sampler=train_sampler,
|
||||||
|
batch_size=hps.train.batch_size,
|
||||||
persistent_workers=True,
|
persistent_workers=True,
|
||||||
# これもメモリ消費量を減らそうとしてコメントアウト
|
# これもメモリ消費量を減らそうとしてコメントアウト
|
||||||
# prefetch_factor=6,
|
# prefetch_factor=6,
|
||||||
@@ -579,7 +576,7 @@ def run():
|
|||||||
optim_g,
|
optim_g,
|
||||||
hps.train.learning_rate,
|
hps.train.learning_rate,
|
||||||
epoch,
|
epoch,
|
||||||
os.path.join(model_dir, "G_{}.pth".format(global_step)),
|
os.path.join(model_dir, f"G_{global_step}.pth"),
|
||||||
)
|
)
|
||||||
assert optim_d is not None
|
assert optim_d is not None
|
||||||
utils.checkpoints.save_checkpoint(
|
utils.checkpoints.save_checkpoint(
|
||||||
@@ -587,7 +584,7 @@ def run():
|
|||||||
optim_d,
|
optim_d,
|
||||||
hps.train.learning_rate,
|
hps.train.learning_rate,
|
||||||
epoch,
|
epoch,
|
||||||
os.path.join(model_dir, "D_{}.pth".format(global_step)),
|
os.path.join(model_dir, f"D_{global_step}.pth"),
|
||||||
)
|
)
|
||||||
if net_dur_disc is not None:
|
if net_dur_disc is not None:
|
||||||
assert optim_dur_disc is not None
|
assert optim_dur_disc is not None
|
||||||
@@ -596,7 +593,7 @@ def run():
|
|||||||
optim_dur_disc,
|
optim_dur_disc,
|
||||||
hps.train.learning_rate,
|
hps.train.learning_rate,
|
||||||
epoch,
|
epoch,
|
||||||
os.path.join(model_dir, "DUR_{}.pth".format(global_step)),
|
os.path.join(model_dir, f"DUR_{global_step}.pth"),
|
||||||
)
|
)
|
||||||
if net_wd is not None:
|
if net_wd is not None:
|
||||||
assert optim_wd is not None
|
assert optim_wd is not None
|
||||||
@@ -605,7 +602,7 @@ def run():
|
|||||||
optim_wd,
|
optim_wd,
|
||||||
hps.train.learning_rate,
|
hps.train.learning_rate,
|
||||||
epoch,
|
epoch,
|
||||||
os.path.join(model_dir, "WD_{}.pth".format(global_step)),
|
os.path.join(model_dir, f"WD_{global_step}.pth"),
|
||||||
)
|
)
|
||||||
utils.safetensors.save_safetensors(
|
utils.safetensors.save_safetensors(
|
||||||
net_g,
|
net_g,
|
||||||
@@ -663,7 +660,7 @@ def train_and_evaluate(
|
|||||||
if writers is not None:
|
if writers is not None:
|
||||||
writer, writer_eval = writers
|
writer, writer_eval = writers
|
||||||
|
|
||||||
train_loader.batch_sampler.set_epoch(epoch)
|
# train_loader.batch_sampler.set_epoch(epoch)
|
||||||
global global_step
|
global global_step
|
||||||
|
|
||||||
net_g.train()
|
net_g.train()
|
||||||
@@ -869,14 +866,12 @@ def train_and_evaluate(
|
|||||||
"loss/g/kl": loss_kl,
|
"loss/g/kl": loss_kl,
|
||||||
}
|
}
|
||||||
)
|
)
|
||||||
|
scalar_dict.update({f"loss/g/{i}": v for i, v in enumerate(losses_gen)})
|
||||||
scalar_dict.update(
|
scalar_dict.update(
|
||||||
{"loss/g/{}".format(i): v for i, v in enumerate(losses_gen)}
|
{f"loss/d_r/{i}": v for i, v in enumerate(losses_disc_r)}
|
||||||
)
|
)
|
||||||
scalar_dict.update(
|
scalar_dict.update(
|
||||||
{"loss/d_r/{}".format(i): v for i, v in enumerate(losses_disc_r)}
|
{f"loss/d_g/{i}": v for i, v in enumerate(losses_disc_g)}
|
||||||
)
|
|
||||||
scalar_dict.update(
|
|
||||||
{"loss/d_g/{}".format(i): v for i, v in enumerate(losses_disc_g)}
|
|
||||||
)
|
)
|
||||||
|
|
||||||
if net_dur_disc is not None:
|
if net_dur_disc is not None:
|
||||||
@@ -884,23 +879,20 @@ def train_and_evaluate(
|
|||||||
|
|
||||||
scalar_dict.update(
|
scalar_dict.update(
|
||||||
{
|
{
|
||||||
"loss/dur_disc_g/{}".format(i): v
|
f"loss/dur_disc_g/{i}": v
|
||||||
for i, v in enumerate(losses_dur_disc_g)
|
for i, v in enumerate(losses_dur_disc_g)
|
||||||
}
|
}
|
||||||
)
|
)
|
||||||
scalar_dict.update(
|
scalar_dict.update(
|
||||||
{
|
{
|
||||||
"loss/dur_disc_r/{}".format(i): v
|
f"loss/dur_disc_r/{i}": v
|
||||||
for i, v in enumerate(losses_dur_disc_r)
|
for i, v in enumerate(losses_dur_disc_r)
|
||||||
}
|
}
|
||||||
)
|
)
|
||||||
|
|
||||||
scalar_dict.update({"loss/g/dur_gen": loss_dur_gen})
|
scalar_dict.update({"loss/g/dur_gen": loss_dur_gen})
|
||||||
scalar_dict.update(
|
scalar_dict.update(
|
||||||
{
|
{f"loss/g/dur_gen_{i}": v for i, v in enumerate(losses_dur_gen)}
|
||||||
"loss/g/dur_gen_{}".format(i): v
|
|
||||||
for i, v in enumerate(losses_dur_gen)
|
|
||||||
}
|
|
||||||
)
|
)
|
||||||
|
|
||||||
if net_wd is not None:
|
if net_wd is not None:
|
||||||
@@ -945,14 +937,14 @@ def train_and_evaluate(
|
|||||||
optim_g,
|
optim_g,
|
||||||
hps.train.learning_rate,
|
hps.train.learning_rate,
|
||||||
epoch,
|
epoch,
|
||||||
os.path.join(hps.model_dir, "G_{}.pth".format(global_step)),
|
os.path.join(hps.model_dir, f"G_{global_step}.pth"),
|
||||||
)
|
)
|
||||||
utils.checkpoints.save_checkpoint(
|
utils.checkpoints.save_checkpoint(
|
||||||
net_d,
|
net_d,
|
||||||
optim_d,
|
optim_d,
|
||||||
hps.train.learning_rate,
|
hps.train.learning_rate,
|
||||||
epoch,
|
epoch,
|
||||||
os.path.join(hps.model_dir, "D_{}.pth".format(global_step)),
|
os.path.join(hps.model_dir, f"D_{global_step}.pth"),
|
||||||
)
|
)
|
||||||
if net_dur_disc is not None:
|
if net_dur_disc is not None:
|
||||||
utils.checkpoints.save_checkpoint(
|
utils.checkpoints.save_checkpoint(
|
||||||
@@ -960,7 +952,7 @@ def train_and_evaluate(
|
|||||||
optim_dur_disc,
|
optim_dur_disc,
|
||||||
hps.train.learning_rate,
|
hps.train.learning_rate,
|
||||||
epoch,
|
epoch,
|
||||||
os.path.join(hps.model_dir, "DUR_{}.pth".format(global_step)),
|
os.path.join(hps.model_dir, f"DUR_{global_step}.pth"),
|
||||||
)
|
)
|
||||||
if net_wd is not None:
|
if net_wd is not None:
|
||||||
utils.checkpoints.save_checkpoint(
|
utils.checkpoints.save_checkpoint(
|
||||||
@@ -968,7 +960,7 @@ def train_and_evaluate(
|
|||||||
optim_wd,
|
optim_wd,
|
||||||
hps.train.learning_rate,
|
hps.train.learning_rate,
|
||||||
epoch,
|
epoch,
|
||||||
os.path.join(hps.model_dir, "WD_{}.pth".format(global_step)),
|
os.path.join(hps.model_dir, f"WD_{global_step}.pth"),
|
||||||
)
|
)
|
||||||
keep_ckpts = config.train_ms_config.keep_ckpts
|
keep_ckpts = config.train_ms_config.keep_ckpts
|
||||||
if keep_ckpts > 0:
|
if keep_ckpts > 0:
|
||||||
@@ -1006,9 +998,7 @@ def train_and_evaluate(
|
|||||||
global_step += 1
|
global_step += 1
|
||||||
if pbar is not None:
|
if pbar is not None:
|
||||||
pbar.set_description(
|
pbar.set_description(
|
||||||
"Epoch {}({:.0f}%)/{}".format(
|
f"Epoch {epoch}({100.0 * batch_idx / len(train_loader):.0f}%)/{hps.train.epochs}"
|
||||||
epoch, 100.0 * batch_idx / len(train_loader), hps.train.epochs
|
|
||||||
)
|
|
||||||
)
|
)
|
||||||
pbar.update()
|
pbar.update()
|
||||||
|
|
||||||
@@ -1022,6 +1012,7 @@ def evaluate(hps, generator, eval_loader, writer_eval):
|
|||||||
generator.eval()
|
generator.eval()
|
||||||
image_dict = {}
|
image_dict = {}
|
||||||
audio_dict = {}
|
audio_dict = {}
|
||||||
|
print()
|
||||||
logger.info("Evaluating ...")
|
logger.info("Evaluating ...")
|
||||||
with torch.no_grad():
|
with torch.no_grad():
|
||||||
for batch_idx, (
|
for batch_idx, (
|
||||||
|
|||||||
Reference in New Issue
Block a user