# 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()