Feat: make style vectors w.r.t. subdirs structure
This commit is contained in:
@@ -109,7 +109,7 @@ def run():
|
||||
envs = config.train_ms_config.env
|
||||
for env_name, env_value in envs.items():
|
||||
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)
|
||||
logger.info(
|
||||
"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(
|
||||
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()
|
||||
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:
|
||||
@@ -193,13 +193,9 @@ def run():
|
||||
os.makedirs(config.out_dir, exist_ok=True)
|
||||
|
||||
if not args.skip_default_style:
|
||||
# Save default style to out_dir
|
||||
default_style.set_style_config(
|
||||
args.config, os.path.join(config.out_dir, "config.json")
|
||||
)
|
||||
default_style.save_neutral_vector(
|
||||
default_style.save_styles_by_dirs(
|
||||
os.path.join(args.model, "wavs"),
|
||||
os.path.join(config.out_dir, "style_vectors.npy"),
|
||||
config.out_dir,
|
||||
)
|
||||
|
||||
torch.manual_seed(hps.train.seed)
|
||||
@@ -215,24 +211,25 @@ def run():
|
||||
writer = SummaryWriter(log_dir=model_dir)
|
||||
writer_eval = SummaryWriter(log_dir=os.path.join(model_dir, "eval"))
|
||||
train_dataset = TextAudioSpeakerLoader(hps.data.training_files, hps.data)
|
||||
train_sampler = DistributedBucketSampler(
|
||||
train_dataset,
|
||||
hps.train.batch_size,
|
||||
[32, 300, 400, 500, 600, 700, 800, 900, 1000],
|
||||
num_replicas=n_gpus,
|
||||
rank=rank,
|
||||
shuffle=True,
|
||||
)
|
||||
# train_sampler = DistributedBucketSampler(
|
||||
# train_dataset,
|
||||
# hps.train.batch_size,
|
||||
# [32, 300, 400, 500, 600, 700, 800, 900, 1000],
|
||||
# num_replicas=n_gpus,
|
||||
# rank=rank,
|
||||
# shuffle=True,
|
||||
# )
|
||||
collate_fn = TextAudioSpeakerCollate(use_jp_extra=True)
|
||||
train_loader = DataLoader(
|
||||
train_dataset,
|
||||
# メモリ消費量を減らそうとnum_workersを1にしてみる
|
||||
# num_workers=min(config.train_ms_config.num_workers, os.cpu_count() // 2),
|
||||
num_workers=1,
|
||||
shuffle=False,
|
||||
shuffle=True,
|
||||
pin_memory=True,
|
||||
collate_fn=collate_fn,
|
||||
batch_sampler=train_sampler,
|
||||
# batch_sampler=train_sampler,
|
||||
batch_size=hps.train.batch_size,
|
||||
persistent_workers=True,
|
||||
# これもメモリ消費量を減らそうとしてコメントアウト
|
||||
# prefetch_factor=6,
|
||||
@@ -579,7 +576,7 @@ def run():
|
||||
optim_g,
|
||||
hps.train.learning_rate,
|
||||
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
|
||||
utils.checkpoints.save_checkpoint(
|
||||
@@ -587,7 +584,7 @@ def run():
|
||||
optim_d,
|
||||
hps.train.learning_rate,
|
||||
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:
|
||||
assert optim_dur_disc is not None
|
||||
@@ -596,7 +593,7 @@ def run():
|
||||
optim_dur_disc,
|
||||
hps.train.learning_rate,
|
||||
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:
|
||||
assert optim_wd is not None
|
||||
@@ -605,7 +602,7 @@ def run():
|
||||
optim_wd,
|
||||
hps.train.learning_rate,
|
||||
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(
|
||||
net_g,
|
||||
@@ -663,7 +660,7 @@ def train_and_evaluate(
|
||||
if writers is not None:
|
||||
writer, writer_eval = writers
|
||||
|
||||
train_loader.batch_sampler.set_epoch(epoch)
|
||||
# train_loader.batch_sampler.set_epoch(epoch)
|
||||
global global_step
|
||||
|
||||
net_g.train()
|
||||
@@ -869,14 +866,12 @@ def train_and_evaluate(
|
||||
"loss/g/kl": loss_kl,
|
||||
}
|
||||
)
|
||||
scalar_dict.update({f"loss/g/{i}": v for i, v in enumerate(losses_gen)})
|
||||
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(
|
||||
{"loss/d_r/{}".format(i): v for i, v in enumerate(losses_disc_r)}
|
||||
)
|
||||
scalar_dict.update(
|
||||
{"loss/d_g/{}".format(i): v for i, v in enumerate(losses_disc_g)}
|
||||
{f"loss/d_g/{i}": v for i, v in enumerate(losses_disc_g)}
|
||||
)
|
||||
|
||||
if net_dur_disc is not None:
|
||||
@@ -884,23 +879,20 @@ def train_and_evaluate(
|
||||
|
||||
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)
|
||||
}
|
||||
)
|
||||
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)
|
||||
}
|
||||
)
|
||||
|
||||
scalar_dict.update({"loss/g/dur_gen": loss_dur_gen})
|
||||
scalar_dict.update(
|
||||
{
|
||||
"loss/g/dur_gen_{}".format(i): v
|
||||
for i, v in enumerate(losses_dur_gen)
|
||||
}
|
||||
{f"loss/g/dur_gen_{i}": v for i, v in enumerate(losses_dur_gen)}
|
||||
)
|
||||
|
||||
if net_wd is not None:
|
||||
@@ -945,14 +937,14 @@ def train_and_evaluate(
|
||||
optim_g,
|
||||
hps.train.learning_rate,
|
||||
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(
|
||||
net_d,
|
||||
optim_d,
|
||||
hps.train.learning_rate,
|
||||
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:
|
||||
utils.checkpoints.save_checkpoint(
|
||||
@@ -960,7 +952,7 @@ def train_and_evaluate(
|
||||
optim_dur_disc,
|
||||
hps.train.learning_rate,
|
||||
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:
|
||||
utils.checkpoints.save_checkpoint(
|
||||
@@ -968,7 +960,7 @@ def train_and_evaluate(
|
||||
optim_wd,
|
||||
hps.train.learning_rate,
|
||||
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
|
||||
if keep_ckpts > 0:
|
||||
@@ -1006,9 +998,7 @@ def train_and_evaluate(
|
||||
global_step += 1
|
||||
if pbar is not None:
|
||||
pbar.set_description(
|
||||
"Epoch {}({:.0f}%)/{}".format(
|
||||
epoch, 100.0 * batch_idx / len(train_loader), hps.train.epochs
|
||||
)
|
||||
f"Epoch {epoch}({100.0 * batch_idx / len(train_loader):.0f}%)/{hps.train.epochs}"
|
||||
)
|
||||
pbar.update()
|
||||
|
||||
@@ -1022,6 +1012,7 @@ def evaluate(hps, generator, eval_loader, writer_eval):
|
||||
generator.eval()
|
||||
image_dict = {}
|
||||
audio_dict = {}
|
||||
print()
|
||||
logger.info("Evaluating ...")
|
||||
with torch.no_grad():
|
||||
for batch_idx, (
|
||||
|
||||
Reference in New Issue
Block a user