Files
soprano-factory/train_decoder.py
T
2026-02-20 07:54:34 -05:00

370 lines
16 KiB
Python

# train_decoder.py - Trains the SopranoDecoder using a frozen LLM backbone with GAN and MR-STFT losses for high-fidelity audio synthesis.
import argparse
import logging
import os
import random
import time
from pathlib import Path
from typing import Tuple, Dict, Any, List
import matplotlib.pyplot as plt
import numpy as np
import torch
import torch.nn.functional as F
import yaml
from huggingface_hub import hf_hub_download
from torch.utils.data import DataLoader
from torch.utils.tensorboard import SummaryWriter
from tqdm import tqdm
from transformers import AutoModelForCausalLM, AutoTokenizer, PreTrainedTokenizer
from dataset import SopranoDecoderDataset, DecoderCollator, SAMPLES_PER_TOKEN
from decoder.decoder import SopranoDecoder
from decoder.discriminator import Discriminator
from decoder.losses import (
MelSpectrogramWrapper, feature_matching_loss, discriminator_loss,
generator_loss, MultiResolutionSTFTLoss
)
def setup_logger(log_file: str):
logger = logging.getLogger("SopranoDecoder")
logger.setLevel(logging.INFO)
if logger.hasHandlers():
logger.handlers.clear()
formatter = logging.Formatter('%(asctime)s - %(message)s')
fh = logging.FileHandler(log_file)
fh.setFormatter(formatter)
logger.addHandler(fh)
ch = logging.StreamHandler()
ch.setFormatter(logging.Formatter('%(message)s'))
logger.addHandler(ch)
return logger
class SopranoDecoderTrainer:
"""
Stage 2: Freezes the LLM, extracts hidden states, and trains a Vocos
decoder with MR-STFT and GAN losses to synthesize high-fidelity audio.
"""
def __init__(self, config: Dict[str, Any], use_disc: bool = True):
self.config = config
self.use_disc = use_disc
self.device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
self.logger = setup_logger(self.config['logging'].get('log_file', 'decoder_training.log'))
self.logger.info("Initializing Soprano Stage-2 Decoder GAN Trainer...")
# Setup Hardware & Seeds
self.seed = self.config['optimizations']['seed']
torch.manual_seed(self.seed)
random.seed(self.seed)
np.random.seed(self.seed)
if torch.cuda.is_available():
torch.cuda.manual_seed(self.seed)
self.allow_tf32 = self.config['optimizations']['allow_tf32']
torch.set_float32_matmul_precision('high' if self.allow_tf32 else 'highest')
self.dtype = getattr(torch, self.config['optimizations']['mixed_precision'])
# Setup TensorBoard
tb_dir = Path(self.config['logging']['tensorboard_dir']) / "decoder_stage"
self.writer = SummaryWriter(log_dir=str(tb_dir))
# Load Components
self.tokenizer = self._setup_tokenizer()
self.llm, self.decoder, self.discriminator = self._setup_models()
self.train_loader, self.val_loader = self._setup_dataloaders()
# Setup Losses
self.mel_fn = MelSpectrogramWrapper().to(self.device)
self.mr_stft = MultiResolutionSTFTLoss().to(self.device)
# Loss Weights
self.lambda_mel = 45.0
self.lambda_fm = 2.0
self.lambda_gen = 1.0
self.lambda_stft = 1.0
# Optimizers
self.max_lr_g = self.config['decoder_training']['learning_rate_g']
self.max_lr_d = self.config['decoder_training']['learning_rate_d']
self.opt_g = torch.optim.AdamW(self.decoder.parameters(), lr=self.max_lr_g, betas=(0.8, 0.99), weight_decay=0.1)
if self.use_disc:
self.opt_d = torch.optim.AdamW(self.discriminator.parameters(), lr=self.max_lr_d, betas=(0.8, 0.99), weight_decay=0.1)
def _setup_tokenizer(self) -> PreTrainedTokenizer:
self.logger.info("Loading Tokenizer...")
tokenizer = AutoTokenizer.from_pretrained(self.config['model']['base_llm'], use_fast=True)
tokenizer.padding_side = 'right'
return tokenizer
def _setup_models(self):
# 1. Load the FINE-TUNED LLM from Stage 1
llm_path = Path(self.config['model']['llm_save_dir']) / "final_llm"
if not llm_path.exists():
# Fallback to the base model if the user skipped Stage 1 for some reason
self.logger.warning(f"Fine-tuned LLM not found at {llm_path}. Falling back to base model.")
llm_path = self.config['model']['base_llm']
self.logger.info(f"Loading LLM from {llm_path} (Frozen)...")
llm = AutoModelForCausalLM.from_pretrained(llm_path, attn_implementation=self.config['optimizations']['attn_implementation'])
llm.to(self.dtype).to(self.device)
llm.eval()
for param in llm.parameters():
param.requires_grad = False
# 2. Load the base Vocos Decoder
self.logger.info("Loading Base Vocos Decoder (Trainable)...")
decoder = SopranoDecoder()
decoder_path = hf_hub_download(repo_id=self.config['model']['base_llm'], filename='decoder.pth')
decoder.load_state_dict(torch.load(decoder_path, map_location='cpu'))
decoder.to(self.device)
decoder.train()
# 3. Initialize GAN Discriminator
discriminator = None
if self.use_disc:
self.logger.info("Initializing GAN Discriminator...")
discriminator = Discriminator()
discriminator.to(self.device)
discriminator.train()
return llm, decoder, discriminator
def _setup_dataloaders(self) -> Tuple[DataLoader, DataLoader]:
input_dir = Path(self.config['dataset']['input_dir'])
batch_size = self.config['decoder_training']['batch_size']
collator = DecoderCollator(tokenizer=self.tokenizer)
train_ds = SopranoDecoderDataset(str(input_dir / 'train.json'), target_sr=self.config['dataset']['sample_rate'])
val_ds = SopranoDecoderDataset(str(input_dir / 'val.json'), target_sr=self.config['dataset']['sample_rate'])
train_loader = DataLoader(train_ds, batch_size=batch_size, shuffle=True, num_workers=4, collate_fn=collator, drop_last=True)
val_loader = DataLoader(val_ds, batch_size=max(1, batch_size // 2), shuffle=False, num_workers=2, collate_fn=collator, drop_last=True)
return train_loader, val_loader
def _crop_for_gan(self, real_audio: torch.Tensor, fake_audio: torch.Tensor, valid_lens: List[int]) -> Tuple[torch.Tensor, torch.Tensor]:
"""Randomly crops audio segments for the Discriminator to grade."""
seg_size = self.config['decoder_training']['segment_size_samples']
bsz = real_audio.size(0)
real_crops, fake_crops = [], []
for b in range(bsz):
v_len = min(valid_lens[b], real_audio.size(1), fake_audio.size(1))
if v_len <= seg_size:
pad_len = seg_size - v_len
r_c = F.pad(real_audio[b, :v_len], (0, pad_len))
f_c = F.pad(fake_audio[b, :v_len], (0, pad_len))
else:
start_idx = random.randint(0, v_len - seg_size)
r_c = real_audio[b, start_idx : start_idx + seg_size]
f_c = fake_audio[b, start_idx : start_idx + seg_size]
real_crops.append(r_c)
fake_crops.append(f_c)
real_tensor = torch.stack(real_crops).unsqueeze(1) # (B, 1, T)
fake_tensor = torch.stack(fake_crops).unsqueeze(1)
return real_tensor, fake_tensor
@torch.no_grad()
def evaluate(self, step: int):
self.decoder.eval()
if self.use_disc: self.discriminator.eval()
val_mel, val_stft = 0.0, 0.0
val_steps = min(5, len(self.val_loader))
val_iter = iter(self.val_loader)
for i in range(val_steps):
x, y, gt_audio, audio_mask = next(val_iter)
x, y, gt_audio, audio_mask = x.to(self.device), y.to(self.device), gt_audio.to(self.device), audio_mask.to(self.device)
with torch.autocast(device_type=self.device.type, dtype=self.dtype):
outputs = self.llm(x, output_hidden_states=True)
hidden_states = outputs.hidden_states[-1].to(torch.float32)
# Gather and pad latents
gathered_list = [hidden_states[b][audio_mask[b]] for b in range(hidden_states.size(0))]
decoder_in = torch.nn.utils.rnn.pad_sequence(gathered_list, batch_first=True).transpose(1, 2)
fake_audio = self.decoder(decoder_in).squeeze(1)
min_len = min(fake_audio.size(1), gt_audio.size(1))
fake_audio, real_audio = fake_audio[:, :min_len], gt_audio[:, :min_len]
# Generate Mel Images on first batch
if i == 0:
gen_mel = self.mel_fn(fake_audio[0:1]).squeeze(0).cpu().numpy()
real_mel = self.mel_fn(real_audio[0:1]).squeeze(0).cpu().numpy()
fig, ax = plt.subplots(2, 1, figsize=(10, 6))
ax[0].imshow(real_mel, aspect='auto', origin='lower')
ax[0].set_title("Ground Truth Mel")
ax[1].imshow(gen_mel, aspect='auto', origin='lower')
ax[1].set_title(f"Generated Mel (Step {step})")
plt.tight_layout()
self.writer.add_figure("Val/Mel_Spectrogram", fig, step)
self.writer.add_audio("Val/Generated_Audio", fake_audio[0], step, sample_rate=self.config['dataset']['sample_rate'])
self.writer.flush()
sc_loss, mag_loss = self.mr_stft(fake_audio, real_audio)
val_stft += (sc_loss + mag_loss).item()
val_stft /= val_steps
self.writer.add_scalar("Val/STFT_Loss", val_stft, step)
self.logger.info(f"[Val Step {step}] MR-STFT Loss: {val_stft:.3f}")
self.decoder.train()
if self.use_disc: self.discriminator.train()
def save_checkpoint(self, step: int, name: str = None):
save_name = name if name else f"checkpoint_{step}"
path = Path(self.config['model']['decoder_save_dir']) / save_name
os.makedirs(path, exist_ok=True)
torch.save(self.decoder.state_dict(), path / "decoder.pth")
if self.use_disc:
torch.save(self.discriminator.state_dict(), path / "discriminator.pth")
self.logger.info(f"Saved Decoder checkpoint to {path}")
def train(self):
max_steps = self.config['decoder_training']['max_steps']
val_freq = self.config['decoder_training']['val_freq']
save_freq = self.config['decoder_training']['save_freq']
self.logger.info(f"Starting Decoder GAN Training for {max_steps} steps...")
train_iter = iter(self.train_loader)
pbar = tqdm(range(max_steps), ncols=150, dynamic_ncols=True)
for step in pbar:
start_time = time.time()
try:
x, y, gt_audio, audio_mask = next(train_iter)
except StopIteration:
train_iter = iter(self.train_loader)
x, y, gt_audio, audio_mask = next(train_iter)
x, y, gt_audio, audio_mask = x.to(self.device), y.to(self.device), gt_audio.to(self.device), audio_mask.to(self.device)
# 1. Forward LLM (Frozen)
with torch.no_grad():
with torch.autocast(device_type=self.device.type, dtype=self.dtype):
hidden_states = self.llm(x, output_hidden_states=True).hidden_states[-1].to(torch.float32)
# 2. Gather active audio latent states and pad for Decoder
gathered_list = []
valid_lens = []
for b_idx in range(hidden_states.size(0)):
states = hidden_states[b_idx][audio_mask[b_idx]]
gathered_list.append(states)
valid_lens.append(states.size(0) * SAMPLES_PER_TOKEN)
decoder_in = torch.nn.utils.rnn.pad_sequence(gathered_list, batch_first=True).transpose(1, 2)
# =======================================================
# TRAIN DISCRIMINATOR
# =======================================================
d_loss_item = 0.0
if self.use_disc:
self.opt_d.zero_grad()
with torch.no_grad(): # Detach generator graph
fake_audio = self.decoder(decoder_in).squeeze(1)
min_len = min(fake_audio.size(1), gt_audio.size(1))
fake_audio, real_audio = fake_audio[:, :min_len], gt_audio[:, :min_len]
real_crops, fake_crops = self._crop_for_gan(real_audio, fake_audio, valid_lens)
y_d_rs, y_d_gs, _, _ = self.discriminator(real_crops, fake_crops)
d_loss, _, _ = discriminator_loss(y_d_rs, y_d_gs)
d_loss.backward()
torch.nn.utils.clip_grad_norm_(self.discriminator.parameters(), 1.0)
self.opt_d.step()
d_loss_item = d_loss.item()
# =======================================================
# TRAIN GENERATOR (DECODER)
# =======================================================
self.opt_g.zero_grad()
fake_audio = self.decoder(decoder_in).squeeze(1)
min_len = min(fake_audio.size(1), gt_audio.size(1))
fake_audio, real_audio = fake_audio[:, :min_len], gt_audio[:, :min_len]
# Mel & STFT Losses (Audio reconstruction)
pred_mel, gt_mel = self.mel_fn(fake_audio), self.mel_fn(real_audio)
# Masking for Mel Loss based on valid sequence length
mel_loss = 0.0
frames_per_token = SAMPLES_PER_TOKEN // 512
for b in range(fake_audio.size(0)):
v_mel_len = min((valid_lens[b] // 512), pred_mel.size(2))
if v_mel_len > 0:
mel_loss += F.l1_loss(pred_mel[b, :, :v_mel_len], gt_mel[b, :, :v_mel_len])
mel_loss /= fake_audio.size(0)
sc_loss, mag_loss = self.mr_stft(fake_audio, real_audio)
# GAN Generator Losses
loss_fm, loss_gen = torch.tensor(0.0, device=self.device), torch.tensor(0.0, device=self.device)
if self.use_disc:
real_crops_g, fake_crops_g = self._crop_for_gan(real_audio, fake_audio, valid_lens)
y_d_rs, y_d_gs, fmap_rs, fmap_gs = self.discriminator(real_crops_g, fake_crops_g)
loss_fm = feature_matching_loss(fmap_rs, fmap_gs)
loss_gen, _ = generator_loss(y_d_gs)
total_loss_g = (self.lambda_mel * mel_loss) + (self.lambda_stft * (sc_loss + mag_loss)) + (self.lambda_gen * loss_gen) + (self.lambda_fm * loss_fm)
total_loss_g.backward()
torch.nn.utils.clip_grad_norm_(self.decoder.parameters(), 1.0)
self.opt_g.step()
# Logging
self.writer.add_scalar("Train/Mel_Loss", mel_loss.item(), step)
self.writer.add_scalar("Train/MR_STFT_Loss", (sc_loss + mag_loss).item(), step)
if self.use_disc:
self.writer.add_scalar("Train/D_Loss", d_loss_item, step)
self.writer.add_scalar("Train/G_Loss", loss_gen.item(), step)
dt_ms = (time.time() - start_time) * 1000
pbar.set_description(
f"STFT: {(sc_loss+mag_loss).item():.3f} | Mel: {mel_loss.item():.3f} | "
f"G: {loss_gen.item():.3f} | D: {d_loss_item:.3f} | {dt_ms:.1f}ms"
)
if step > 0 and step % val_freq == 0:
self.evaluate(step)
if step > 0 and step % save_freq == 0:
self.save_checkpoint(step)
self.save_checkpoint(max_steps, name="final_decoder")
self.writer.close()
self.logger.info("Stage 2 Decoder Training Complete.")
def main():
parser = argparse.ArgumentParser(description="Train Soprano Vocos Decoder (GAN)")
parser.add_argument("--config", type=Path, default=Path("config.yaml"))
parser.add_argument("--no-disc", action="store_true", help="Disable GAN Discriminator")
args = parser.parse_args()
with open(args.config, 'r') as f:
config = yaml.safe_load(f)
trainer = SopranoDecoderTrainer(config, use_disc=not args.no_disc)
trainer.train()
if __name__ == '__main__':
main()