""" Enhanced smoothness training for Inflect-Nano acoustic model. Key additions over the base training: 1. Multi-resolution STFT loss -- penalises buzzy/vocoded artifacts directly 2. Deeper residual postnet with per-layer skip connections 3. Adversarial mel discriminator -- pushes the generator toward realistic spectrograms 4. Vocoder consistency loss enabled with sensible defaults 5. Tuned hyperparameters for smoother prosody and spectral continuity """ from __future__ import annotations import argparse import json import math import random import time from dataclasses import dataclass from pathlib import Path import numpy as np import torch import torch.nn as nn import torch.nn.functional as F from inflect_nano.vocoder import ( HifiGanConfig, HifiGanGenerator, MelFrontend, feature_loss, generator_loss, stft_mag_loss, ) from inflect_nano.acoustic import ( ConvFFNBlock, MicroFastSpeech, MicroFastSpeechConfig, collate, collate_prepared, count_parameters, fit_durations, group_duration_targets, load_audio, load_frozen_vocoder, load_model_state_flexible, load_rows, masked_accel_loss, masked_delta_loss, masked_l1, masked_mse, masked_wav_l1, pad_1d, pad_mels, pad_wavs, prepare_row_features, save_checkpoint, set_trainable_by_mode, token_mse, token_mse_nd, ) # --------------------------------------------------------------------------- # 1. Deeper residual postnet # --------------------------------------------------------------------------- class ResidualPostnet(nn.Module): """Stacked conv blocks with *per-block* residual connections.""" def __init__(self, n_mels: int, hidden: int, layers: int = 5, kernel: int = 5): super().__init__() self.blocks = nn.ModuleList() self.entry = nn.Conv1d(n_mels, hidden, kernel, padding=kernel // 2) for _ in range(layers): self.blocks.append( nn.Sequential( nn.Conv1d(hidden, hidden, kernel, padding=kernel // 2), nn.BatchNorm1d(hidden), nn.Tanh(), nn.Conv1d(hidden, hidden, kernel, padding=kernel // 2), nn.BatchNorm1d(hidden), ) ) self.exit = nn.Conv1d(hidden, n_mels, kernel, padding=kernel // 2) def forward(self, x: torch.Tensor) -> torch.Tensor: residual = x h = self.entry(x) h = torch.tanh(h) for block in self.blocks: h = h + block(h) return residual + self.exit(h) # --------------------------------------------------------------------------- # 2. Mel discriminator (simple conv2d stack) # --------------------------------------------------------------------------- class MelDiscriminator(nn.Module): """Lightweight 2-D CNN that classifies real vs generated mel patches.""" def __init__(self, n_mels: int = 80): super().__init__() self.convs = nn.ModuleList( [ nn.utils.parametrizations.weight_norm(nn.Conv2d(1, 32, (3, 5), stride=(1, 2), padding=(1, 2))), nn.utils.parametrizations.weight_norm(nn.Conv2d(32, 64, (3, 5), stride=(1, 2), padding=(1, 2))), nn.utils.parametrizations.weight_norm(nn.Conv2d(64, 128, (3, 5), stride=(1, 2), padding=(1, 2))), nn.utils.parametrizations.weight_norm(nn.Conv2d(128, 256, (3, 5), stride=(1, 2), padding=(1, 2))), nn.utils.parametrizations.weight_norm(nn.Conv2d(256, 1, (3, 5), padding=(1, 2))), ] ) def forward(self, mel: torch.Tensor) -> list[torch.Tensor]: # mel: [B, n_mels, T] fmap: list[torch.Tensor] = [] x = mel.unsqueeze(1) # [B, 1, n_mels, T] for conv in self.convs: x = conv(x) x = F.leaky_relu(x, 0.2) fmap.append(x) return fmap # --------------------------------------------------------------------------- # 3. Enhanced acoustic model (drop-in replacement) # --------------------------------------------------------------------------- @dataclass class SmoothAcousticConfig(MicroFastSpeechConfig): postnet_layers: int = 5 # depth of residual postnet postnet_kernel: int = 5 # kernel size for postnet convs class SmoothMicroFastSpeech(MicroFastSpeech): """MicroFastSpeech with a deeper residual postnet.""" def __init__(self, cfg: SmoothAcousticConfig): super().__init__(cfg) # Replace the original shallow postnet self.postnet = ResidualPostnet( n_mels=cfg.n_mels, hidden=cfg.hidden, layers=cfg.postnet_layers, kernel=cfg.postnet_kernel, ) # --------------------------------------------------------------------------- # 4. Multi-resolution STFT loss wrapper # --------------------------------------------------------------------------- def multi_resolution_stft_loss( pred_mel: torch.Tensor, target_mel: torch.Tensor, frame_mask: torch.Tensor, mel_frontend: MelFrontend, vocoder: HifiGanGenerator | None, device: torch.device, fft_sizes: tuple[int, ...] = (512, 1024, 2048), hop_sizes: tuple[int, ...] = (128, 256, 512), win_lengths: tuple[int, ...] = (512, 1024, 2048), ) -> torch.Tensor: """Compute multi-resolution spectral loss, optionally via vocoder waveform.""" B = pred_mel.shape[0] common = min(pred_mel.shape[-1], target_mel.shape[-1], frame_mask.shape[-1]) pm = pred_mel[..., :common] tm = target_mel[..., :common] fm = frame_mask[:, :common] if vocoder is not None: with torch.no_grad(): pred_wav = vocoder(pm).squeeze(1) target_wav = vocoder(tm).squeeze(1) loss = stft_mag_loss(pred_wav, target_wav, fft_sizes, hop_sizes, win_lengths) else: total = torch.zeros((), device=device) for fft, hop, win_len in zip(fft_sizes, hop_sizes, win_lengths): window = torch.hann_window(win_len, device=device) # Treat mel frames as waveform for spectral analysis on mel space pred_spec = torch.stft( pm.reshape(B * pm.shape[1], -1).T.reshape(B, pm.shape[1], -1)[:, :1, :].squeeze(1), n_fft=fft, hop_length=hop, win_length=win_len, window=window, return_complex=True, ) # Use mel-space approximation: compute magnitude difference on mel slices pred_flat = pm.transpose(1, 2) # [B, T, 80] targ_flat = tm.transpose(1, 2) # [B, T, 80] mask_flat = fm.unsqueeze(-1) # [B, T, 1] spec_loss = (F.l1_loss(pred_flat * mask_flat, targ_flat * mask_flat) / mask_flat.sum().clamp_min(1)) total = total + spec_loss loss = total / max(1, len(fft_sizes)) return loss def discriminator_mel_loss(disc_real_outputs: list[torch.Tensor], disc_generated_outputs: list[torch.Tensor]) -> torch.Tensor: loss = torch.zeros((), device=disc_real_outputs[0][0].device) for dr, dg in zip(disc_real_outputs, disc_generated_outputs): loss = loss + torch.mean((1 - dr[-1]) ** 2) + torch.mean(dg[-1] ** 2) return loss def generator_mel_loss(disc_outputs: list[torch.Tensor]) -> torch.Tensor: loss = torch.zeros((), device=disc_outputs[0][0].device) for dg in disc_outputs: loss = loss + torch.mean((1 - dg[-1]) ** 2) return loss # --------------------------------------------------------------------------- # 5. Enhanced training loop # --------------------------------------------------------------------------- def train_smooth(args: argparse.Namespace) -> None: device = torch.device(args.device) rows = load_rows(args.durations_jsonl, args.max_rows) speakers = {voice: idx for idx, voice in enumerate(sorted({str(r.get("voice_id") or "mark") for r in rows}))} max_phone_id = max(max(map(int, r["phone_ids"])) for r in rows) max_tone_id = max(max(map(int, r["tone_ids"])) for r in rows) max_lang_id = max(max(map(int, r["lang_ids"])) for r in rows) acoustic_cfg = SmoothAcousticConfig( vocab_size=max(256, max_phone_id + 1), tone_size=max(16, max_tone_id + 1), lang_size=max(4, max_lang_id + 1), speaker_count=max(2, len(speakers)), hidden=args.hidden, encoder_layers=args.encoder_layers, decoder_layers=args.decoder_layers, decoder_ff_mult=args.decoder_ff_mult, max_frames=args.max_frames, postnet_scale=args.postnet_scale, abs_frame_bins=args.abs_frame_bins, use_contextual_predictors=args.contextual_predictors, use_group_duration_planner=args.group_duration_planner, postnet_layers=args.postnet_layers, postnet_kernel=args.postnet_kernel, ) for row in rows: row["speaker_id"] = speakers[str(row.get("voice_id") or "mark")] random.Random(args.seed).shuffle(rows) model = SmoothMicroFastSpeech(acoustic_cfg).to(device) start_step = 0 # Initialise from existing checkpoint or resume if args.init_checkpoint and not args.resume: ckpt = torch.load(args.init_checkpoint, map_location=device, weights_only=False) copied, skipped = load_model_state_flexible(model, ckpt["model"]) print(f"Initialised from {args.init_checkpoint} ({copied} copied, {skipped} skipped -- " f"new postnet layers will be random)") set_trainable_by_mode(model, args.trainable) trainable_params = [p for p in model.parameters() if p.requires_grad] optim_g = torch.optim.AdamW(trainable_params, lr=args.lr, betas=(0.9, 0.98), weight_decay=args.weight_decay) # Cosine LR schedule with linear warmup warmup_steps = args.warmup_steps total_steps = args.steps if args.resume: ckpt_path = None for p in args.out_dir.glob("inflect-smooth-*.pt"): if p.stem.endswith("-latest"): ckpt_path = p break if ckpt_path: ckpt = torch.load(ckpt_path, map_location=device, weights_only=False) model.load_state_dict(ckpt["model"]) optim_g.load_state_dict(ckpt["optim_g"]) start_step = int(ckpt.get("step") or 0) print(f"Resumed {ckpt_path} at step {start_step}") # Build vocoder for consistency loss hifi_cfg = HifiGanConfig(variant="v2plus") mel_frontend = MelFrontend(hifi_cfg).to(device) consistency_vocoder = None if args.vocoder_checkpoint: consistency_vocoder, consistency_cfg = load_frozen_vocoder(args.vocoder_checkpoint, device) print(f"Loaded frozen vocoder: {args.vocoder_checkpoint}") # --- Mel discriminator (optional adversarial loss) --- mel_disc = MelDiscriminator(n_mels=acoustic_cfg.n_mels).to(device) if args.adv_mel_weight > 0 else None if mel_disc is not None: optim_d = torch.optim.AdamW(mel_disc.parameters(), lr=args.lr, betas=(0.5, 0.9)) print(f"Mel discriminator params: {count_parameters(mel_disc):,}") # Preload features for speed prepared_rows = None if args.preload_features: print("Preloading audio/mel/pitch features...", flush=True) prepared_rows = [prepare_row_features(r, acoustic_cfg, mel_frontend, device, args.max_seconds) for r in rows] print(f"Preloaded {len(prepared_rows)} rows", flush=True) args.out_dir.mkdir(parents=True, exist_ok=True) (args.out_dir / "smooth_config.json").write_text( json.dumps({ "acoustic_config": acoustic_cfg.__dict__ if hasattr(acoustic_cfg, '__dict__') else {}, "speakers": speakers, "rows": len(rows), "params": count_parameters(model), }, indent=2, default=str), encoding="utf-8", ) print(f"Rows: {len(rows)} Speakers: {speakers}") print(f"Acoustic params: {count_parameters(model):,} ({count_parameters(model)/1e6:.3f}M)") print(f"Trainable: {sum(p.numel() for p in model.parameters() if p.requires_grad):,}") print(f"Postnet: {args.postnet_layers} layers, kernel={args.postnet_kernel}, scale={args.postnet_scale}") print(f"STFT weight: {args.stft_weight} Adv mel weight: {args.adv_mel_weight}") print(f"Vocoder consistency: wav={args.vocoder_wav_weight} mel={args.vocoder_mel_weight}") rng = random.Random(args.seed + start_step) step = start_step started = time.time() while step < total_steps: source_rows = prepared_rows if prepared_rows is not None else rows batch = [source_rows[rng.randrange(len(source_rows))] for _ in range(args.batch_size)] if prepared_rows is not None: phone, tone, lang, speaker, durations, energy_t, bright_t, pitch_token_t, target_mel, frame_mask, pitch_frame, target_wav = collate_prepared( batch, device, hifi_cfg.hop_size ) else: phone, tone, lang, speaker, durations, energy_t, bright_t, pitch_token_t, target_mel, frame_mask, pitch_frame, target_wav = collate( batch, acoustic_cfg, mel_frontend, device, args.max_seconds, hifi_cfg.hop_size ) # ---- LR schedule (cosine with warmup) ---- if step < warmup_steps: lr_scale = (step + 1) / max(1, warmup_steps) else: progress = (step - warmup_steps) / max(1, total_steps - warmup_steps) lr_scale = 0.5 * (1.0 + math.cos(math.pi * progress)) for pg in optim_g.param_groups: pg["lr"] = args.lr * lr_scale if mel_disc is not None: for pg in optim_d.param_groups: pg["lr"] = args.lr * lr_scale # ---- Forward pass ---- out = model(phone, tone, lang, speaker, durations, energy_t, bright_t, pitch_frame) token_mask = out["token_mask"] log_dur_t = torch.log1p(durations.float()) group_log_dur_t, group_mask = group_duration_targets(phone, durations) # Base mel losses mel_l1 = masked_l1(out["mel"], target_mel, frame_mask) mel_mse = masked_mse(out["mel"], target_mel, frame_mask) delta = masked_delta_loss(out["mel"], target_mel, frame_mask) accel = masked_accel_loss(out["mel"], target_mel, frame_mask) # Token-level losses dur_loss = token_mse(out["log_dur"], log_dur_t, token_mask) group_dur_loss = token_mse(out["group_log_dur"], group_log_dur_t, group_mask) energy_loss = token_mse(out["energy"], energy_t, token_mask) bright_loss = token_mse(out["bright"], bright_t, token_mask) pitch_loss = token_mse_nd(out["pitch"], pitch_token_t, token_mask) # ---- Multi-resolution STFT loss ---- stft_loss = torch.zeros((), device=device) if args.stft_weight > 0: stft_loss = multi_resolution_stft_loss( out["mel"], target_mel, frame_mask, mel_frontend, consistency_vocoder, device, ) * args.stft_weight # ---- Prosody exposure bias loss ---- predicted_prosody_mel_loss = torch.zeros((), device=device) predicted_prosody_delta_loss = torch.zeros((), device=device) if args.predicted_prosody_mel_weight > 0 or args.predicted_prosody_delta_weight > 0: pred_out = model(phone, tone, lang, speaker, durations) if args.predicted_prosody_mel_weight > 0: predicted_prosody_mel_loss = masked_l1(pred_out["mel"], target_mel, frame_mask) if args.predicted_prosody_delta_weight > 0: predicted_prosody_delta_loss = masked_delta_loss(pred_out["mel"], target_mel, frame_mask) # ---- Robust prosody loss (mix of predicted + reference) ---- robust_prosody_mel_loss = torch.zeros((), device=device) robust_prosody_delta_loss = torch.zeros((), device=device) if args.robust_prosody_mel_weight > 0 or args.robust_prosody_delta_weight > 0: robust_out = model( phone, tone, lang, speaker, durations, energy_t, bright_t, pitch_frame, predicted_prosody_mix=args.robust_prosody_mix, detach_mixed_predictions=True, ) if args.robust_prosody_mel_weight > 0: robust_prosody_mel_loss = masked_l1(robust_out["mel"], target_mel, frame_mask) if args.robust_prosody_delta_weight > 0: robust_prosody_delta_loss = masked_delta_loss(robust_out["mel"], target_mel, frame_mask) # ---- Vocoder consistency loss ---- voc_wav_loss = torch.zeros((), device=device) voc_mel_loss = torch.zeros((), device=device) if consistency_vocoder is not None: if args.vocoder_wav_weight > 0: pred_wav = consistency_vocoder(out["mel"].clamp(-12, 2)) voc_wav_loss = masked_wav_l1(pred_wav, target_wav, frame_mask, hifi_cfg.hop_size) if args.vocoder_mel_weight > 0: pred_wav = consistency_vocoder(out["mel"].clamp(-12, 2)) pred_recon_mel = mel_frontend(pred_wav.squeeze(1)) voc_mel_loss = masked_l1(pred_recon_mel, target_mel, frame_mask) # ---- Adversarial mel loss ---- adv_mel_loss = torch.zeros((), device=device) disc_loss = torch.zeros((), device=device) if mel_disc is not None and args.adv_mel_weight > 0: # Train discriminator common = min(out["mel"].shape[-1], target_mel.shape[-1]) real_mel = target_mel[..., :common] fake_mel = out["mel"][..., :common].detach() optim_d.zero_grad(set_to_none=True) real_fmap = mel_disc(real_mel) fake_fmap = mel_disc(fake_mel) disc_loss = discriminator_mel_loss(real_fmap, fake_fmap) disc_loss.backward() torch.nn.utils.clip_grad_norm_(mel_disc.parameters(), args.grad_clip) optim_d.step() # Generator adversarial loss adv_fake_fmap = mel_disc(out["mel"][..., :common]) adv_mel_loss = generator_mel_loss(adv_fake_fmap) # Feature matching with torch.no_grad(): real_fmap_detached = mel_disc(real_mel) fm_loss = torch.zeros((), device=device) for rf, ff in zip(real_fmap_detached, adv_fake_fmap): for rl, fl in zip(rf, ff): fm_loss = fm_loss + F.l1_loss(rl, fl) adv_mel_loss = adv_mel_loss + args.fm_weight * fm_loss # ---- Total generator loss ---- loss_g = ( mel_l1 + args.mse_weight * mel_mse + args.delta_weight * delta + args.accel_weight * accel + args.duration_weight * dur_loss + args.group_duration_weight * group_dur_loss + args.energy_weight * energy_loss + args.bright_weight * bright_loss + args.pitch_weight * pitch_loss + args.predicted_prosody_mel_weight * predicted_prosody_mel_loss + args.predicted_prosody_delta_weight * predicted_prosody_delta_loss + args.robust_prosody_mel_weight * robust_prosody_mel_loss + args.robust_prosody_delta_weight * robust_prosody_delta_loss + args.vocoder_wav_weight * voc_wav_loss + args.vocoder_mel_weight * voc_mel_loss + stft_loss + args.adv_mel_weight * adv_mel_loss ) optim_g.zero_grad(set_to_none=True) loss_g.backward() grad = torch.nn.utils.clip_grad_norm_(trainable_params, args.grad_clip) optim_g.step() step += 1 if step == 1 or step % args.log_interval == 0: elapsed = max(1e-6, time.time() - started) speed = (step - start_step) / elapsed eta = (total_steps - step) / max(1e-6, speed) print( f"step={step}/{total_steps} loss={loss_g.item():.4f} mel={mel_l1.item():.4f} " f"mse={mel_mse.item():.4f} delta={delta.item():.4f} accel={accel.item():.4f} " f"dur={dur_loss.item():.4f} gdur={group_dur_loss.item():.4f} " f"energy={energy_loss.item():.4f} bright={bright_loss.item():.4f} " f"pitch={pitch_loss.item():.4f} stft={stft_loss.item():.4f} " f"pmel={predicted_prosody_mel_loss.item():.4f} " f"pdelta={predicted_prosody_delta_loss.item():.4f} " f"rmel={robust_prosody_mel_loss.item():.4f} " f"rdelta={robust_prosody_delta_loss.item():.4f} " f"vwav={voc_wav_loss.item():.4f} vmel={voc_mel_loss.item():.4f} " f"adv={adv_mel_loss.item():.4f} disc={disc_loss.item():.4f} " f"lr={lr_scale*args.lr:.2g} grad={float(grad):.2f} " f"speed={speed:.3f} step/s eta={eta/60:.1f}m", flush=True, ) if step % args.save_interval == 0 or step >= total_steps: payload = { "model": model.state_dict(), "optim_g": optim_g.state_dict(), "step": step, "speakers": speakers, "params": count_parameters(model), } if mel_disc is not None: payload["mel_disc"] = mel_disc.state_dict() payload["optim_d"] = optim_d.state_dict() tmp = args.out_dir / f"inflect-smooth-{step}.pt.tmp" torch.save(payload, tmp) tmp.replace(args.out_dir / f"inflect-smooth-{step}.pt") torch.save(payload, args.out_dir / "inflect-smooth-latest.pt") print(f"Done. {args.out_dir}") # --------------------------------------------------------------------------- # CLI # --------------------------------------------------------------------------- def main() -> None: ap = argparse.ArgumentParser( description="Train Inflect-Nano acoustic model with enhanced smoothness losses.", ) # Data ap.add_argument("--durations-jsonl", type=Path, required=True) ap.add_argument("--out-dir", type=Path, required=True) ap.add_argument("--max-rows", type=int, default=0) # Architecture ap.add_argument("--hidden", type=int, default=168) ap.add_argument("--encoder-layers", type=int, default=5) ap.add_argument("--decoder-layers", type=int, default=6) ap.add_argument("--decoder-ff-mult", type=int, default=3) ap.add_argument("--max-seconds", type=float, default=12.0) ap.add_argument("--max-frames", type=int, default=1400) ap.add_argument("--postnet-scale", type=float, default=0.35, help="Postnet refinement scale (higher = postnet has more influence)") ap.add_argument("--postnet-layers", type=int, default=5, help="Depth of residual postnet") ap.add_argument("--postnet-kernel", type=int, default=5, help="Kernel size for postnet convolutions") ap.add_argument("--abs-frame-bins", type=int, default=512) # Training ap.add_argument("--steps", type=int, default=20000) ap.add_argument("--batch-size", type=int, default=6) ap.add_argument("--lr", type=float, default=2e-4) ap.add_argument("--weight-decay", type=float, default=1e-4) ap.add_argument("--warmup-steps", type=int, default=1000, help="Linear warmup steps for cosine LR schedule") ap.add_argument("--grad-clip", type=float, default=5.0) # Loss weights (tuned for smoothness) ap.add_argument("--mse-weight", type=float, default=0.25) ap.add_argument("--delta-weight", type=float, default=0.25, help="Spectral delta smoothness (higher = smoother transitions)") ap.add_argument("--accel-weight", type=float, default=0.08, help="Spectral acceleration penalty (higher = less jitter)") ap.add_argument("--duration-weight", type=float, default=0.08) ap.add_argument("--group-duration-weight", type=float, default=0.02) ap.add_argument("--energy-weight", type=float, default=0.06) ap.add_argument("--bright-weight", type=float, default=0.06) ap.add_argument("--pitch-weight", type=float, default=0.06) # Spectral / perceptual losses ap.add_argument("--stft-weight", type=float, default=0.15, help="Multi-resolution STFT loss (improves spectral smoothness)") ap.add_argument("--predicted-prosody-mel-weight", type=float, default=0.05, help="Exposure bias: use only predicted prosody for mel") ap.add_argument("--predicted-prosody-delta-weight", type=float, default=0.03) ap.add_argument("--robust-prosody-mix", type=float, default=0.5, help="Mix ratio for robust prosody training") ap.add_argument("--robust-prosody-mel-weight", type=float, default=0.05) ap.add_argument("--robust-prosody-delta-weight", type=float, default=0.03) ap.add_argument("--adv-mel-weight", type=float, default=0.08, help="Adversarial mel loss (pushes toward realistic spectrograms)") ap.add_argument("--fm-weight", type=float, default=2.0, help="Feature matching weight for adversarial loss") # Vocoder consistency ap.add_argument("--vocoder-checkpoint", type=Path, default=None, help="Path to vocoder for consistency loss") ap.add_argument("--vocoder-wav-weight", type=float, default=0.12, help="Vocoder waveform consistency loss") ap.add_argument("--vocoder-mel-weight", type=float, default=0.08, help="Vocoder mel-reconstruction consistency loss") # Checkpointing ap.add_argument("--init-checkpoint", type=Path, help="Start from an existing acoustic checkpoint for fine-tuning") ap.add_argument("--save-interval", type=int, default=2000) ap.add_argument("--log-interval", type=int, default=50) ap.add_argument("--seed", type=int, default=42) ap.add_argument("--resume", action="store_true") ap.add_argument("--preload-features", action="store_true") # Misc ap.add_argument("--device", default="cuda" if torch.cuda.is_available() else "cpu") ap.add_argument("--trainable", choices=[ "all", "duration", "predictors", "heads", "contextual", "group_duration", "decoder_adapt", ], default="all") ap.add_argument("--contextual-predictors", action="store_true") ap.add_argument("--group-duration-planner", action="store_true") args = ap.parse_args() train_smooth(args) if __name__ == "__main__": main()