mirror of
https://github.com/Nighthawk42/soprano-factory.git
synced 2026-08-30 04:30:21 +00:00
370 lines
16 KiB
Python
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() |