Feat: make style vectors w.r.t. subdirs structure

This commit is contained in:
litagin02
2024-05-25 17:14:28 +09:00
parent 0c87c6ffd4
commit 175aa62a26
3 changed files with 109 additions and 64 deletions

View File

@@ -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}")

View File

@@ -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, (

View File

@@ -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, (