SFlowTTS / train_codec_speaker.py
FashionFlora's picture
Upload full repo excluding dump_40, dump_100, precomputed_tokens, precomputed_data
fb0011a verified
Raw History Blame Contribute Delete
30.4 kB
# train_codec_speaker.py
# ==============================================================================
# Training script for Hybrid TTS Codec with LEARNABLE SPEAKER EMBEDDINGS
#
# KEY DIFFERENCE from train_codec_hybrid_temporal.py:
# - NO mel-based style encoder
# - Uses learnable speaker embeddings: nn.Embedding(11, 128)
# - Speaker IDs 0-10 for 11 speakers
# - Speaker conditioning via AdaIN1d throughout decoder
#
# Architecture:
# SpeakerEmbedding: speaker_id [B] -> embedding [B, 128]
# ProsodyEncoder: (pitch, energy) -> latent (NO speaker to force codebook usage)
# FSQ: latent -> tokens (prosody ONLY quantized)
# TextEncoder: text -> text_emb (continuous)
# Decoder: (quantized_prosody, text_emb, speaker_emb) -> waveform
#
# Usage:
# python train_codec_speaker.py -p Configs/config_codec_speaker.yml
# ==============================================================================
import os
import os.path as osp
import sys
import yaml
import shutil
import numpy as np
import torch
import click
import warnings
warnings.simplefilter('ignore')
import random
from munch import Munch
from torch import nn
import torch.nn.functional as F
import torchaudio
from torch.nn.utils.rnn import pad_sequence
from models_speaker import build_model_speaker, load_checkpoint, gaussian_upsample
from meldataset import build_dataloader
from utils import get_data_path_list, log_print, length_to_mask, recursive_munch
from losses import MultiResolutionSTFTLoss, MagPhaseLoss, GeneratorLoss, DiscriminatorLoss, WavLMLoss
from optimizers import build_optimizer
import time
from accelerate import Accelerator
from accelerate.utils import LoggerType
from accelerate import DistributedDataParallelKwargs
from torch.utils.tensorboard import SummaryWriter
import logging
from accelerate.logging import get_logger
logger = get_logger(__name__, log_level="DEBUG")
# Fix "Too many open files" error with DataLoader workers
import torch.multiprocessing
torch.multiprocessing.set_sharing_strategy('file_system')
# Increase file descriptor limit
import resource
soft, hard = resource.getrlimit(resource.RLIMIT_NOFILE)
resource.setrlimit(resource.RLIMIT_NOFILE, (min(65536, hard), hard))
import gc
@click.command()
@click.option('-p', '--config_path', default='Configs/config_codec_speaker.yml', type=str)
def main(config_path):
config = yaml.safe_load(open(config_path))
log_dir = config['log_dir']
if not osp.exists(log_dir):
os.makedirs(log_dir, exist_ok=True)
shutil.copy(config_path, osp.join(log_dir, osp.basename(config_path)))
ddp_kwargs = DistributedDataParallelKwargs(find_unused_parameters=True)
accelerator = Accelerator(
project_dir=log_dir,
split_batches=True,
kwargs_handlers=[ddp_kwargs],
mixed_precision='bf16'
)
if accelerator.is_main_process:
writer = SummaryWriter(log_dir + "/tensorboard")
file_handler = logging.FileHandler(osp.join(log_dir, 'train_codec_speaker.log'))
file_handler.setLevel(logging.DEBUG)
file_handler.setFormatter(logging.Formatter('%(levelname)s:%(asctime)s: %(message)s'))
logger.logger.addHandler(file_handler)
# Config
batch_size = config.get('batch_size', 10)
device = accelerator.device
epochs = config.get('epochs_codec', 100)
save_freq = config.get('save_freq', 2)
log_interval = config.get('log_interval', 10)
# n_mels - default to 40
n_mels = config.get('n_mels', 40)
data_params = config.get('data_params', None)
sr = config['preprocess_params'].get('sr', 44100)
train_path = data_params['train_data']
val_path = data_params['val_data']
root_path = data_params['root_path']
min_length = data_params['min_length']
max_len = config.get('max_len', 7500)
train_list, val_list = get_data_path_list(train_path, val_path)
# Determine dump directory based on n_mels
dump_dir = f"./dump_{n_mels}"
# Dataloaders
train_dataloader = build_dataloader(
train_list,
root_path,
alignment_path="../data_preparation_v2/alignments_train.safetensors",
feature_dir=f"{dump_dir}/train",
stats_path=f"{dump_dir}/speaker_stats.json",
min_length=min_length,
batch_size=batch_size,
num_workers=16,
dataset_config={},
device=device
)
val_dataloader = build_dataloader(
val_list,
root_path,
alignment_path="../data_preparation_v2/alignments_val.safetensors",
feature_dir=f"{dump_dir}/val",
stats_path=f"{dump_dir}/speaker_stats.json",
min_length=min_length,
batch_size=20,
validation=True,
num_workers=0,
device=device,
dataset_config={}
)
# Model params
model_params = recursive_munch(config['model_params'])
# Get hop_length from config
hop_length = config['preprocess_params']['spect_params'].get('hop_length', 441)
# Scheduler params
scheduler_params = {
"max_lr": float(config['optimizer_params'].get('lr', 1e-5)),
"min_lr": float(config['optimizer_params'].get('min_lr', 1e-5)),
"pct_start": float(config['optimizer_params'].get('pct_start', 0.33)),
"epochs": epochs,
"steps_per_epoch": len(train_dataloader),
}
decay_epochs = int(epochs * scheduler_params["pct_start"])
print(f"[LR Schedule] Pre-TMA (epochs 0-{decay_epochs}): {scheduler_params['max_lr']:.2e} -> {scheduler_params['min_lr']:.2e}")
print(f"[LR Schedule] TMA (epochs {decay_epochs}-{epochs}): constant {scheduler_params['min_lr']:.2e}")
# Build model with speaker embeddings
model = build_model_speaker(model_params)
# Log speaker embedding info
num_speakers = model_params.get('num_speakers', 11)
speaker_dim = model_params.get('speaker_dim', 128)
print(f"[Speaker Codec] Number of speakers: {num_speakers} (IDs 0-{num_speakers-1})")
print(f"[Speaker Codec] Speaker embedding dim: {speaker_dim}")
print(f"[Speaker Codec] Speaker embedding shape: {model.codec.speaker_embedding.weight.shape}")
scheduler_params_dict = {key: scheduler_params.copy() for key in model}
# Prepare models
for k in model:
model[k] = accelerator.prepare(model[k])
model[k] = model[k].to(device)
train_dataloader, val_dataloader = accelerator.prepare(train_dataloader, val_dataloader)
optimizer = build_optimizer(
{key: model[key].parameters() for key in model},
scheduler_params_dict={key: scheduler_params.copy() for key in model},
lr=float(config['optimizer_params'].get('lr', 1e-5))
)
for k, v in optimizer.optimizers.items():
optimizer.optimizers[k] = accelerator.prepare(optimizer.optimizers[k])
optimizer.schedulers[k] = accelerator.prepare(optimizer.schedulers[k])
# Load pretrained
start_epoch = 0
iters = 0
with accelerator.main_process_first():
if config.get('pretrained_model', '') != '':
model, optimizer, start_epoch, iters = load_checkpoint(
model, optimizer, config['pretrained_model'],
load_only_params=config.get('load_only_params', False)
)
print(f"Loaded pretrained model from {config['pretrained_model']}")
desired_lr = float(config['optimizer_params'].get('lr', 1e-5))
for key in optimizer.optimizers:
for param_group in optimizer.optimizers[key].param_groups:
param_group['lr'] = desired_lr
print(f"Reset learning rate to {desired_lr:.2e}")
# Losses
loss_params = Munch(config.get('loss_params', {}))
TMA_epoch = loss_params.get('TMA_epoch', 5)
stft_loss = MultiResolutionSTFTLoss().to(device)
mag_phase_loss = MagPhaseLoss(
n_fft=model_params.decoder.gen_istft_n_fft,
hop_length=model_params.decoder.gen_istft_hop_size
).to(device)
gl = GeneratorLoss(model.mpd, model.msd).to(device)
dl = DiscriminatorLoss(model.mpd, model.msd).to(device)
wl = WavLMLoss(model_sr=sr, slm_sr=16000).to(device)
# Loss weights
lambda_mel = loss_params.get('lambda_mel', 45.0)
lambda_gen = loss_params.get('lambda_gen', 1.0)
lambda_slm = loss_params.get('lambda_slm', 1.0)
lambda_mag = loss_params.get('lambda_mag', 0.1)
lambda_entropy = loss_params.get('lambda_entropy', 10.0)
lambda_f0 = loss_params.get('lambda_f0', 1.0)
best_loss = float('inf')
# Compression level
codec_strides = model_params.codec_strides
upsample_rates = model_params.decoder.upsample_rates
codec_compression = 1
for s in codec_strides:
codec_compression *= s
istft_hop = model_params.decoder.gen_istft_hop_size
print(f"[Speaker Codec] Compression: {codec_compression}x")
print(f"[Speaker Codec] Upsample Rates: {upsample_rates} (Total: {np.prod(upsample_rates)})")
print(f"[Speaker Codec] Token Rate: {44100 / (hop_length * codec_compression):.2f} Hz")
print(f"[Speaker Codec] n_mels: {model_params.n_mels}")
print(f"[Speaker Codec] max_len: {max_len}")
# FSQ levels
fsq_levels = model_params.get('fsq_levels', [4] * 6)
codebook_size = np.prod(fsq_levels)
print(f"[Speaker Codec] FSQ Levels: {fsq_levels} ({codebook_size} codes)")
for epoch in range(start_epoch, epochs):
running_loss = 0
start_time = time.time()
_ = [model[key].train() for key in model]
for i, batch in enumerate(train_dataloader):
try:
# ===============================================================
# START OF TRAINING STEP
# ===============================================================
waves = batch[1]
tensors = [b.to(device) for b in batch[2:13] if isinstance(b, torch.Tensor)]
(
speakers_ids, # [B] - speaker IDs (0-10)
languages_ids,
mels,
mel_input_length,
_target_ph,
_target_ph_l,
context_ph,
context_lens,
durations_tg,
pitches,
energies,
) = tensors
with torch.no_grad():
text_pad_mask = length_to_mask(_target_ph_l)
# Extract language embedding
lang_emb = accelerator.unwrap_model(model.text_encoder).language_emb(languages_ids).detach()
# Text encoding
t_en = model.text_encoder(
_target_ph,
_target_ph_l,
language_emb=lang_emb
)
asr_features, _ = gaussian_upsample(
t_en, durations_tg, mel_input_length,
sigma_sq=10, token_mask=~text_pad_mask
)
# Gather lengths
mel_input_length_all = accelerator.gather(mel_input_length)
mel_len = min([int(mel_input_length_all.min().item() / 2 - 1), max_len // 2])
mel_len_st = int(mel_input_length.min().item() / 2 - 1)
# Prepare batches
en, gt, wav, f0_list, n0_list, spk_list = [], [], [], [], [], []
for bib in range(len(mel_input_length)):
mel_length = int(mel_input_length[bib].item() / 2)
random_start = np.random.randint(0, mel_length - mel_len)
en.append(asr_features[bib, :, (random_start * 2):((random_start+mel_len) * 2)])
gt.append(mels[bib, :, (random_start * 2):((random_start+mel_len) * 2)])
y = waves[bib][(random_start * 2) * 441:((random_start+mel_len) * 2) * 441]
wav.append(y.to(device))
f0_list.append(pitches[bib, (random_start * 2):((random_start+mel_len) * 2)])
n0_list.append(energies[bib, (random_start * 2):((random_start+mel_len) * 2)])
spk_list.append(speakers_ids[bib]) # Append speaker ID for this sample
en = torch.stack(en) # [B, 512, T]
gt = torch.stack(gt).detach() # [B, n_mels, T]
F0 = torch.stack(f0_list).float().detach() # [B, T]
N0 = torch.stack(n0_list).float().detach() # [B, T]
wav = torch.stack(wav).float().detach() # [B, T_audio]
speaker_ids_batch = torch.stack(spk_list) # [B] - speaker IDs
if gt.shape[-1] < 80:
continue
# Scheduled sampling for F0 prediction
current_prob_pitch = min(0.8, 0.3 + (epoch / 20))
use_predicted_f0 = random.random() < current_prob_pitch
# ===============================================================
# CODEC FORWARD WITH SPEAKER IDS
# ===============================================================
codec_out = model.codec(
pitch=F0,
energy=N0,
text_emb=en,
speaker_ids=speaker_ids_batch, # Pass speaker IDs instead of mel
use_predicted_f0=use_predicted_f0,
language_emb=lang_emb
)
y_rec = codec_out['wav']
mag = codec_out['mag']
phase = codec_out['phase']
commitment_loss = codec_out['commitment_loss']
tokens = codec_out['tokens']
f0_pred = codec_out['f0_pred']
speaker_emb = codec_out['speaker_emb'] # [B, 128] learned speaker embedding
# Match lengths
min_len = min(y_rec.shape[-1], wav.shape[-1])
y_rec = y_rec[..., :min_len]
wav = wav[..., :min_len]
# Match F0 lengths
min_len_f0 = min(f0_pred.shape[-1], F0.shape[-1])
f0_pred = f0_pred[..., :min_len_f0]
F0_target = F0[..., :min_len_f0]
# F0 Loss
loss_f0 = F.mse_loss(f0_pred, F0_target)
# === Discriminator Step ===
if epoch >= TMA_epoch:
optimizer.zero_grad()
d_loss = dl(wav.detach().unsqueeze(1).float(), y_rec.detach()).mean()
accelerator.backward(d_loss)
optimizer.step('msd')
optimizer.step('mpd')
else:
d_loss = torch.tensor(0.0, device=device)
# === Generator Step ===
optimizer.zero_grad()
# STFT loss
loss_mel = stft_loss(y_rec.squeeze(), wav.detach())
# Mag/Phase loss
loss_mag_phase = mag_phase_loss(mag, phase, wav.detach())
# Codebook utilization
with torch.no_grad():
B, n_q, T = tokens.shape
flat_tokens = tokens.reshape(-1)
unique_codes = len(torch.unique(flat_tokens))
codebook_size = np.prod(model.codec.fsq_levels)
perplexity = unique_codes / codebook_size * 100
if epoch >= TMA_epoch:
loss_gen = gl(wav.detach().unsqueeze(1).float(), y_rec).mean()
loss_slm = wl(wav.detach(), y_rec).mean()
g_loss = (
lambda_mel * loss_mel +
lambda_gen * loss_gen +
lambda_slm * loss_slm +
lambda_mag * loss_mag_phase +
lambda_entropy * commitment_loss +
lambda_f0 * loss_f0
)
else:
loss_gen = torch.tensor(0.0, device=device)
loss_slm = torch.tensor(0.0, device=device)
g_loss = (
lambda_mel * loss_mel +
lambda_mag * loss_mag_phase +
lambda_entropy * commitment_loss +
lambda_f0 * loss_f0
)
running_loss += accelerator.gather(loss_mel).mean().item()
accelerator.backward(g_loss)
optimizer.step('text_encoder')
optimizer.step('codec')
iters += 1
entropy_loss = commitment_loss.item() if torch.is_tensor(commitment_loss) else commitment_loss
if (i + 1) % log_interval == 0 and accelerator.is_main_process:
current_lr = optimizer.optimizers['codec'].param_groups[0]['lr']
# Log unique speakers in batch
unique_speakers = torch.unique(speaker_ids_batch).tolist()
log_print(
f'Epoch [{epoch+1}/{epochs}], Step [{i+1}/{len(train_dataloader)}], '
f'LR: {current_lr:.2e}, '
f'Mel: {running_loss / log_interval:.5f}, '
f'Gen: {loss_gen.item():.5f}, '
f'Disc: {d_loss.item():.5f}, '
f'F0: {loss_f0.item():.5f}, '
f'SLM: {loss_slm.item():.5f}, '
f'Entropy: {entropy_loss:.5f}, '
f'FSQ Util: {perplexity:.1f}% ({unique_codes}/{codebook_size}), '
f'Speakers: {unique_speakers}',
logger
)
writer.add_scalar('train/learning_rate', current_lr, iters)
writer.add_scalar('train/mel_loss', running_loss / log_interval, iters)
writer.add_scalar('train/gen_loss', loss_gen.item(), iters)
writer.add_scalar('train/d_loss', d_loss.item(), iters)
writer.add_scalar('train/f0_loss', loss_f0.item(), iters)
writer.add_scalar('train/slm_loss', loss_slm.item(), iters)
writer.add_scalar('train/entropy_loss', entropy_loss, iters)
writer.add_scalar('train/fsq_utilization', perplexity, iters)
writer.add_scalar('train/unique_codes', unique_codes, iters)
writer.add_scalar('train/mag_phase_loss', loss_mag_phase.item(), iters)
running_loss = 0
except torch.cuda.OutOfMemoryError:
if accelerator.is_main_process:
print(f"[Warning] OOM at epoch {epoch}, step {i}. Flushing cache.")
for key in model:
model[key].zero_grad()
torch.cuda.empty_cache()
gc.collect()
continue
# === Validation ===
loss_test = 0
_ = [model[key].eval() for key in model]
eval_samples = [] # Store (mel_len, mel, asr, f0, n, lang, wav, speaker_id)
with torch.no_grad():
iters_test = 0
for batch_idx, batch in enumerate(val_dataloader):
try:
waves = batch[1]
tensors = [b.to(device) for b in batch[2:13] if isinstance(b, torch.Tensor)]
(
speakers_ids,
languages_ids,
mels,
mel_input_length,
_target_ph,
_target_ph_l,
context_ph,
context_lens,
durations_tg,
pitches,
energies,
) = tensors
text_pad_mask = length_to_mask(_target_ph_l)
lang_emb = model.text_encoder.language_emb(languages_ids)
t_en = model.text_encoder(
_target_ph,
_target_ph_l,
language_emb=lang_emb
)
asr_features, _ = gaussian_upsample(
t_en, durations_tg, mel_input_length,
sigma_sq=10, token_mask=~text_pad_mask
)
# Collect samples for evaluation
if len(eval_samples) < 20:
for bib in range(len(mel_input_length)):
if len(eval_samples) >= 20: break
eval_samples.append((
mel_input_length[bib].item(),
mels[bib],
asr_features[bib],
pitches[bib],
energies[bib],
lang_emb[bib] if lang_emb is not None else None,
waves[bib],
speakers_ids[bib], # Include speaker ID
))
mel_len = min([int(mel_input_length.min().item() / 2 - 1), max_len // 2])
en, gt, wav_list, f0_list, n_list, spk_list = [], [], [], [], [], []
for bib in range(len(mel_input_length)):
mel_length = int(mel_input_length[bib].item() / 2)
random_start = np.random.randint(0, mel_length - mel_len)
en.append(asr_features[bib, :, (random_start * 2):((random_start + mel_len) * 2)])
gt.append(mels[bib, :, (random_start * 2):((random_start + mel_len) * 2)])
y = waves[bib][(random_start * 2) * 441:((random_start + mel_len) * 2) * 441]
wav_list.append(y.to(device))
f0_list.append(pitches[bib, (random_start * 2):((random_start + mel_len) * 2)])
n_list.append(energies[bib, (random_start * 2):((random_start + mel_len) * 2)])
spk_list.append(speakers_ids[bib])
wav = torch.stack(wav_list).float().detach()
F0_curve = torch.stack(f0_list).detach()
N_curve = torch.stack(n_list).detach()
en = torch.stack(en)
gt = torch.stack(gt).detach()
speaker_ids_batch = torch.stack(spk_list)
# Forward with speaker IDs
codec_out = model.codec(
pitch=F0_curve,
energy=N_curve,
text_emb=en,
speaker_ids=speaker_ids_batch,
language_emb=lang_emb
)
y_rec = codec_out['wav']
min_len = min(y_rec.shape[-1], wav.shape[-1])
y_rec = y_rec[..., :min_len]
wav = wav[..., :min_len]
loss_mel = stft_loss(y_rec.squeeze(), wav.detach())
loss_test += accelerator.gather(loss_mel).mean().item()
iters_test += 1
except Exception as e:
print(f"Error in validation loop: {e}")
continue
if accelerator.is_main_process:
print(f'Epoch: {epoch + 1}')
current_val_loss = loss_test / iters_test if iters_test > 0 else 0
log_print(f'Validation loss: {current_val_loss:.3f}\n', logger)
writer.add_scalar('eval/mel_loss', current_val_loss, epoch + 1)
# Generate audio samples
with torch.no_grad():
for bib, sample in enumerate(eval_samples):
mel_len_val, gt_mel, en_sample, f0_sample, n_sample, lang_emb_sample, gt_wav, speaker_id = sample
mel_length = int(mel_len_val)
gt_mel = gt_mel[:, :mel_length].unsqueeze(0).to(device)
en_sample = en_sample[:, :mel_length].unsqueeze(0).to(device)
f0_sample = f0_sample[:mel_length].unsqueeze(0).float().to(device)
n_sample = n_sample[:mel_length].unsqueeze(0).float().to(device)
lang_emb_sample = lang_emb_sample.unsqueeze(0).to(device) if lang_emb_sample is not None else None
speaker_id_sample = speaker_id.unsqueeze(0).to(device)
# Forward with speaker ID
codec_out = model.codec(
f0_sample, n_sample, en_sample, speaker_id_sample,
language_emb=lang_emb_sample
)
y_rec = codec_out['wav']
speaker_emb = codec_out['speaker_emb']
writer.add_audio(
f'eval/speaker_codec_{bib}_spk{speaker_id.item()}',
y_rec.cpu().numpy().squeeze(), epoch, sample_rate=sr
)
print(f" Sample {bib}: speaker_id={speaker_id.item()}, speaker_emb shape={speaker_emb.shape}")
# Test tokenize + decode path
tokens, text_down, spk_emb = model.codec.tokenize(
f0_sample, n_sample, en_sample, speaker_id_sample
)
y_from_tokens = model.codec.decode_tokens(
tokens, text_down, speaker_id_sample, language_emb=lang_emb_sample
)
writer.add_audio(
f'eval/from_tokens_{bib}_spk{speaker_id.item()}',
y_from_tokens.cpu().numpy().squeeze(), epoch, sample_rate=sr
)
# Save GT audio
gt_wav_len = int(mel_length * hop_length)
if len(gt_wav) > gt_wav_len:
gt_wav = gt_wav[:gt_wav_len]
writer.add_audio(f'eval/gt_{bib}', gt_wav.cpu().numpy(), epoch, sample_rate=sr)
# Test speaker interpolation
if len(eval_samples) >= 2:
sample1 = eval_samples[0]
sample2 = eval_samples[1]
spk_id_1 = sample1[7].item()
spk_id_2 = sample2[7].item()
if spk_id_1 != spk_id_2:
# Use first sample's content with interpolated speaker
mel_len_val, _, en_sample, f0_sample, n_sample, lang_emb_sample, _, _ = sample1
mel_length = int(mel_len_val)
en_sample = en_sample[:, :mel_length].unsqueeze(0).to(device)
f0_sample = f0_sample[:mel_length].unsqueeze(0).float().to(device)
n_sample = n_sample[:mel_length].unsqueeze(0).float().to(device)
lang_emb_sample = lang_emb_sample.unsqueeze(0).to(device) if lang_emb_sample is not None else None
# Get interpolated speaker embedding
interp_emb = model.codec.interpolate_speakers(spk_id_1, spk_id_2, alpha=0.5)
# Tokenize and decode with interpolated speaker
tokens, text_down, _ = model.codec.tokenize(
f0_sample, n_sample, en_sample,
torch.tensor([spk_id_1], device=device)
)
y_interp = model.codec.decode_tokens_with_speaker_emb(
tokens, text_down, interp_emb, language_emb=lang_emb_sample
)
writer.add_audio(
f'eval/interpolated_spk{spk_id_1}_to_spk{spk_id_2}',
y_interp.cpu().numpy().squeeze(), epoch, sample_rate=sr
)
print(f" Generated interpolated audio: speaker {spk_id_1} -> {spk_id_2}")
# Save checkpoint
if epoch % save_freq == 0:
if accelerator.is_main_process:
print('Saving checkpoint...')
state = {
'net': {key: model[key].state_dict() for key in model},
'optimizer': optimizer.state_dict(),
'iters': iters,
'val_loss': loss_test / iters_test if iters_test > 0 else 0,
'epoch': epoch,
}
save_path = osp.join(log_dir, f'epoch_speaker_codec_{epoch:05d}.pth')
torch.save(state, save_path)
gc.collect()
torch.cuda.empty_cache()
# Final save
if accelerator.is_main_process:
print('Saving final checkpoint...')
net_state = {
key: accelerator.unwrap_model(model[key]).state_dict()
for key in model
}
state = {
'net': net_state,
'optimizer': optimizer.state_dict(),
'iters': iters,
'val_loss': loss_test / iters_test if iters_test > 0 else 0,
'epoch': epoch,
}
torch.save(state, osp.join(log_dir, 'speaker_codec_final.pth'))
if __name__ == "__main__":
main()