Minor fixes, code parity.

This commit is contained in:
Nighthawk
2026-02-26 14:19:28 -05:00
parent e7ee43de41
commit 504a70f928
4 changed files with 357 additions and 275 deletions
+40
View File
@@ -0,0 +1,40 @@
# Training Strategy & Learnings for LLM-Backbone TTS
---
## 1. The Prompt Format & Token Specialization
Early experiments showed that simply feeding text and discrete audio tokens without clear boundaries leads to confusion. The LLM struggles to differentiate between "understanding text" and "generating audio representations."
**Learning:** We heavily customized the prompt and tokenizer to include explicit, distinct special tokens.
- The format must explicitly anchor the model's intent: `[TEXT]<text prompt>[START]<audio tokens>[STOP]`.
- We added dedicated special tokens (`<|audio_start|>`, `<|audio_end|>`, `<|text|>`, etc.) to the tokenizer vocabulary and resized the embedding matrices. This provides absolute clarity to the causal transformer regarding which modality it is currently processing.
---
## 2. LLM Training Stability Strategies
Teaching a causal language model (Qwen backbone) to map text directly to audio-latent distributions often results in the model outputting garbage audio if not carefully stabilized.
**Learning & Strategy:** The loss calculation must be incredibly intentional.
* **Batching:** Changed the batching from packing to dynamic batching - one sample per row. Felt like this could be more stable. Shall experiment training with packing. I think this should also work.
* **Text Weighting:** Initially trained the llm with text_weight as 0, but turns out that the model wasn't able to learn this way. It was overfitting and wasn't able to generalize. I was guessing even with 0 weights for text tokens, the model should be able to learn the relationship between text and audio tokens and generalize well.
If the text tokens are zero-weighted in the loss calculation aggressively from the start, the model can lose its language understanding capabilities, learning only to spout random audio formats. A critical learning is to experiment with text weighting strategies (e.g., maintaining a small loss weight on the text prediction part for the first N steps) so the LLM retains language comprehension and can generalize on new texts.
* **Masking for Loss:** Ensure that padding tokens (`[PAD]`) or masked sections are strictly ignored in the `CrossEntropyLoss` calculation. Any gradient penalty on padding will rapidly destabilize autoregressive predictions.
* **Gradual Length Training (Curriculum Learning):** The model often crashes or learns poorly if immediately exposed to full 1000+ token sequences. It's recommended starting with shorter text-audio pairs and gradually increasing the sequence context length (`max_length`) as the model stabilizes. Yet to try this.
---
## 3. Two-Stage Decoder (Vocoder) Training
The role of the Decoder (Vocos) is to take the *continuous* hidden states evaluated by the LLM and construct high-fidelity audio waveforms.
**Learning & Strategy:** The Decoder must be trained independently in two distinct stages, with the LLM rigidly completely frozen.
#### Stage 1: Pure Reconstruction (Global Structure)
* **Objective:** Teach the Decoder how to map hidden state dimensions to spectral layout.
* **Method:** Train using purely reconstruction losses: **Mel-Spectrogram L1 Loss** and **Multi-Resolution STFT Loss**.
* **Data:** This phase is executed over the *almost-full* generated sequences to ensure the Decoder learns long-term structural coherence and timing alignment.
#### Stage 2: Adversarial Refinement (Local Fidelity)
* **Objective:** Eliminate robotic artifacts and "muffled" qualities to achieve crystal-clear, natural acoustic texture.
* **Method:** Introduce Discriminator networks (Multi-Period and Multi-Scale).
* **Crucial Learning (Cropping):** Running Discriminators (GANs) on full-length sequences destroys memory and stability. We transitioned to **Segment-Based Training (Random Cropping)**. During this stage, we extract random 1-second chunks (e.g., ~32,000 samples) from both the ground truth and the generated audio, feeding only these small crops into the Discriminator.
* **Combined Loss:** The final phase loss comprises Reconstruction (Full Audio) + Adversarial GAN Loss (Cropped Audio) + Discriminator Feature Matching.
+31 -12
View File
@@ -11,6 +11,7 @@ global:
num_workers: 4
use_wandb: true
wandb_project: "soprano-tts"
tokenizer_name: "ekwek/Soprano-80M"
# ------------------------------------------------------------------------------
# File Paths
@@ -20,23 +21,30 @@ global:
# Examples:
# Linux: "/home/ubuntu/data/lj_speech/LJSpeech-1.1"
# Windows: "C:/Users/Name/Documents/datasets/LJSpeech-1.1"
# Relative: "./datasets/LJSpeech-1.1"
# Relative: "./data/LJSpeech-1.1"
# ------------------------------------------------------------------------------
paths:
# Path to LJSpeech-1.1 directory
dataset_root: "./data/LJSpeech-1.1"
# Base directory to save all checkpoints and logs
# Base directory to save all checkpoints, logs, and generated datasets
save_dir: "./checkpoints"
# Pretrained model paths (set to null if training from scratch)
# Linux ex: "/home/ubuntu/soma/ckpt/suprano/codec/step_42000.pt"
# Pretrained model paths (set to null if training from scratch or not applicable)
pretrained_codec_path: null
pretrained_llm_path: null
pretrained_decoder_path: null
pretrained_discriminator_path: null
# ------------------------------------------------------------------------------
# Codec Training Configuration
# Data Generation Configuration (generate_dataset*.py)
# ------------------------------------------------------------------------------
data_generation:
val_prop: 0.1
val_max: 512
# ------------------------------------------------------------------------------
# Codec Training Configuration (codec_train.py)
# ------------------------------------------------------------------------------
codec:
sample_rate: 32000
@@ -44,9 +52,14 @@ codec:
num_epochs: 100
learning_rate: 1.0e-4
freeze_encoder: false
# Encoder/Decoder architecture params
encoder_dim: 768
encoder_num_layers: 8
bottleneck_channels: 5
decoder_num_layers: 8
# ------------------------------------------------------------------------------
# LLM Training Configuration
# LLM Training Configuration (train_llm.py)
# ------------------------------------------------------------------------------
llm:
from_scratch: false
@@ -65,11 +78,11 @@ llm:
weight_decay: 0.1
# ------------------------------------------------------------------------------
# Decoder (Vocos) Training Configuration
# Decoder (Vocos) Training Configuration (train_decoder.py)
# ------------------------------------------------------------------------------
decoder:
use_discriminator: true
batch_size: 64
batch_size: 8 # Reduced to avoid OOM, adjust based on your GPU
max_steps: 200000
max_lr: 2.0e-4
min_lr_ratio: 0.1
@@ -79,9 +92,12 @@ decoder:
seq_len: 1024
segment_size_samples: 32768 # ~1 sec (16 tokens)
val_freq: 250
val_steps: 10
save_freq: 3000
text_factor: 0.0
betas: [0.8, 0.99]
weight_decay: 0.1
start_step: 0
# Loss Weights
lambda_mel: 45.0
@@ -90,8 +106,11 @@ decoder:
lambda_stft: 1.0
# ------------------------------------------------------------------------------
# Data Generation Configuration (generate_dataset*.py)
# Inference Configuration (simple_inference.py)
# ------------------------------------------------------------------------------
data_generation:
val_prop: 0.1
val_max: 512
inference:
temperature: 0.8
top_k: 50
top_p: 0.95
repetition_penalty: 1.2
max_new_tokens: 512
+223 -215
View File
@@ -2,14 +2,16 @@
Training script for Soprano Decoder (Vocos).
Freezes LLM and trains Decoder with GAN loss.
"""
import os
import random
import time
import os
import io
import wandb
import matplotlib.pyplot as plt
import numpy as np
import torch
import torchaudio
from torch.utils.data import DataLoader
from tqdm import tqdm
from transformers import AutoModelForCausalLM, AutoTokenizer, AutoConfig
@@ -22,35 +24,33 @@ from decoder.losses import MelSpectrogramWrapper, feature_matching_loss, discrim
from config_loader import load_config
# Global Tokenizer for collate function
tokenizer = AutoTokenizer.from_pretrained('ekwek/Soprano-80M')
tokenizer.padding_side = 'right' # Essential for training!
def worker_seed_init(_):
worker_seed = torch.initial_seed() % (2**32-1)
np.random.seed(worker_seed)
random.seed(worker_seed)
def get_lr(it, max_lr, min_lr, warmup_steps, cooldown_steps, max_steps): # WSD schedule
def get_lr(it, max_lr, min_lr, warmup_steps, cooldown_steps, max_steps):
if it < warmup_steps:
return max_lr * (it + 1) / warmup_steps
if it < max_steps - cooldown_steps:
return max_lr
return min_lr + (max_lr - min_lr) * ((max_steps - it) / cooldown_steps)
def collate_pack(batch_in):
# batch_in is list of (text, wav)
def collate_pack(batch_in, tokenizer):
texts = [x[0] for x in batch_in]
wavs = [x[1] for x in batch_in]
aud_token_lens = [x[2] for x in batch_in]
# We need to process each sample to align audio
# Since lengths vary, we process list then pad
tokens_batch = tokenizer(texts, padding=True, return_tensors='pt')
input_ids = tokens_batch['input_ids']
batch_tokens_list = []
batch_audio_list = []
for i in range(len(texts)):
# Get raw tokens without padding for alignment logic
raw_tokens = tokenizer(texts[i], padding=False, truncation=False)['input_ids']
tokens = torch.tensor(raw_tokens, dtype=torch.long)
@@ -73,29 +73,165 @@ def collate_pack(batch_in):
batch_tokens_list.append(tokens)
batch_audio_list.append(aligned_audio)
# Pad Tokens
batch_tokens = torch.nn.utils.rnn.pad_sequence(batch_tokens_list, batch_first=True, padding_value=0)
# Pad Audio
batch_audio = torch.nn.utils.rnn.pad_sequence(batch_audio_list, batch_first=True, padding_value=0.0)
x = batch_tokens[:, :-1]
y = batch_tokens[:, 1:]
# Calculate max seq len of x
max_len_x = x.size(1)
gt_audio = batch_audio[:, :max_len_x * SAMPLES_PER_TOKEN]
# Create Audio Mask (True where token is audio)
audio_mask = (y > 3) & (y <= 8003)
return x, y, gt_audio, audio_mask
if __name__ == '__main__':
# ------------------
# Load Configuration
# ------------------
def evaluate(step, val_dataloader_it, val_dataloader, model, decoder, discriminator,
mel_fn, mr_stft, use_disc, device, device_type, val_steps, segment_size, use_wandb):
decoder.eval()
if use_disc and discriminator is not None:
discriminator.eval()
val_mel_loss_accum = 0.0
val_gen_loss_accum = 0.0
val_fm_loss_accum = 0.0
val_d_loss_accum = 0.0
val_sc_loss_accum = 0.0
val_mag_loss_accum = 0.0
log_dict = {}
with torch.no_grad():
for _ in range(val_steps):
try:
val_batch = next(val_dataloader_it)
except StopIteration:
val_dataloader_it = iter(val_dataloader)
val_batch = next(val_dataloader_it)
vx, vy, vgt_audio, vaudio_mask = val_batch
vx, vy = vx.to(device), vy.to(device)
vgt_audio = vgt_audio.to(device)
vaudio_mask = vaudio_mask.to(device)
with torch.autocast(device_type=device_type, dtype=torch.bfloat16):
voutputs = model(vx, output_hidden_states=True)
v_hidden = voutputs.hidden_states[-1].to(torch.float32)
v_gathered_states_list = []
for b_idx in range(v_hidden.size(0)):
mask = vaudio_mask[b_idx]
v_valid_states = v_hidden[b_idx][mask]
v_gathered_states_list.append(v_valid_states)
v_in_padded = torch.nn.utils.rnn.pad_sequence(v_gathered_states_list, batch_first=True, padding_value=0.0)
v_bsz = v_in_padded.size(0)
v_max_aud_len = v_in_padded.size(1)
v_audio_loss_mask = torch.zeros((v_bsz, v_max_aud_len), dtype=torch.bool, device=device)
for b_idx in range(v_bsz):
length = v_gathered_states_list[b_idx].size(0)
v_audio_loss_mask[b_idx, :length] = True
v_in = v_in_padded.transpose(1, 2)
v_fake_audio = decoder(v_in)
if v_fake_audio.size(1) == 1: v_fake_audio = v_fake_audio.squeeze(1)
min_len_v = min(v_fake_audio.size(1), vgt_audio.size(1))
v_fake_audio = v_fake_audio[:, :min_len_v]
v_real_audio = vgt_audio[:, :min_len_v]
frames_per_token_v = SAMPLES_PER_TOKEN // 512
v_mel_mask = v_audio_loss_mask.repeat_interleave(frames_per_token_v, dim=1)
v_pred_mel = mel_fn(v_fake_audio)
v_gt_mel = mel_fn(v_real_audio)
min_mel_len_v = min(v_pred_mel.size(2), v_gt_mel.size(2), v_mel_mask.size(1))
v_pred_mel = v_pred_mel[:, :, :min_mel_len_v]
v_gt_mel = v_gt_mel[:, :, :min_mel_len_v]
v_mel_mask = v_mel_mask[:, :min_mel_len_v]
v_mel_loss_raw = torch.nn.functional.l1_loss(v_pred_mel, v_gt_mel, reduction='none')
v_mel_loss = (v_mel_loss_raw * v_mel_mask.unsqueeze(1)).sum() / (v_mel_mask.sum() * v_pred_mel.size(1) + 1e-6)
val_mel_loss_accum += v_mel_loss.item()
v_sample_mask = v_audio_loss_mask.repeat_interleave(SAMPLES_PER_TOKEN, dim=1)[:, :min_len_v]
v_sc_loss, v_mag_loss = mr_stft(v_fake_audio * v_sample_mask, v_real_audio * v_sample_mask)
val_sc_loss_accum += v_sc_loss.item()
val_mag_loss_accum += v_mag_loss.item()
if use_disc and discriminator is not None:
v_real_crop_list = []
v_fake_crop_list = []
v_min_len = min(v_fake_audio.size(1), v_real_audio.size(1))
for b_idx in range(v_bsz):
v_valid_len = v_gathered_states_list[b_idx].size(0) * SAMPLES_PER_TOKEN
v_valid_len = min(v_valid_len, v_min_len)
if v_valid_len <= segment_size:
v_pad_len = segment_size - v_valid_len
vr_c = torch.nn.functional.pad(v_real_audio[b_idx, :v_valid_len], (0, v_pad_len))
vf_c = torch.nn.functional.pad(v_fake_audio[b_idx, :v_valid_len], (0, v_pad_len))
else:
v_start_idx = random.randint(0, v_valid_len - segment_size)
vr_c = v_real_audio[b_idx, v_start_idx : v_start_idx + segment_size]
vf_c = v_fake_audio[b_idx, v_start_idx : v_start_idx + segment_size]
v_real_crop_list.append(vr_c)
v_fake_crop_list.append(vf_c)
v_real_crops = torch.stack(v_real_crop_list).unsqueeze(1)
v_fake_crops = torch.stack(v_fake_crop_list).unsqueeze(1)
vy_d_rs, vy_d_gs, vfmap_rs, vfmap_gs = discriminator(v_real_crops, v_fake_crops)
v_fm_loss = feature_matching_loss(vfmap_rs, vfmap_gs)
v_gen_loss, _ = generator_loss(vy_d_gs)
v_d_loss, _, _ = discriminator_loss(vy_d_rs, vy_d_gs)
val_gen_loss_accum += v_gen_loss.item()
val_fm_loss_accum += v_fm_loss.item()
val_d_loss_accum += v_d_loss.item()
log_dict.update({
"val/loss_mel": val_mel_loss_accum / val_steps,
"val/loss_sc": val_sc_loss_accum / val_steps,
"val/loss_mag": val_mag_loss_accum / val_steps
})
if use_disc and discriminator is not None:
log_dict.update({
"val/loss_gen": val_gen_loss_accum / val_steps,
"val/loss_fm": val_fm_loss_accum / val_steps,
"val/loss_d": val_d_loss_accum / val_steps,
})
if use_wandb:
# Generate Mel Images (from last val batch)
gen_mel = mel_fn(v_fake_audio[0:1]).squeeze(0).cpu().numpy()
gt_mel = mel_fn(v_real_audio[0:1]).squeeze(0).cpu().numpy()
fig, ax = plt.subplots(2, 1, figsize=(10, 6))
ax[0].imshow(gt_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("Generated Mel (Val)")
plt.tight_layout()
log_dict["val/mel_spectrograms"] = wandb.Image(fig)
plt.close(fig)
decoder.train()
if use_disc and discriminator is not None:
discriminator.train()
return log_dict, val_dataloader_it
def main():
config = load_config("config.yaml")
cfg_global = config["global"]
cfg_paths = config["paths"]
@@ -104,24 +240,28 @@ if __name__ == '__main__':
device = cfg_global["device"]
seed = cfg_global["seed"]
device_type = "cuda" if device.startswith("cuda") else "cpu"
tokenizer_name = cfg_global.get("tokenizer_name", "ekwek/Soprano-80M")
use_wandb = cfg_global["use_wandb"]
torch.manual_seed(seed)
if torch.cuda.is_available():
torch.cuda.manual_seed(seed)
torch.set_float32_matmul_precision('high')
# Setup directories
train_dataset_path = os.path.join(cfg_paths["dataset_root"], "train.json")
val_dataset_path = os.path.join(cfg_paths["dataset_root"], "val.json")
save_path = os.path.join(cfg_paths["save_dir"], "decoder")
os.makedirs(save_path, exist_ok=True)
print(f"Save Path: {save_path}")
if cfg_global["use_wandb"]:
if use_wandb:
wandb.init(project=cfg_global["wandb_project"], config=config)
# Initialize Mel Spectrogram Wrapper dynamically
tokenizer = AutoTokenizer.from_pretrained(tokenizer_name)
tokenizer.padding_side = 'right'
mel_fn = MelSpectrogramWrapper().to(device)
mr_stft = MultiResolutionSTFTLoss().to(device)
# ------------------
# Hyperparameters
@@ -135,12 +275,13 @@ if __name__ == '__main__':
batch_size = cfg_decoder["batch_size"]
segment_size_samples = cfg_decoder["segment_size_samples"]
val_freq = cfg_decoder["val_freq"]
val_steps = cfg_decoder.get("val_steps", 10)
save_freq = cfg_decoder["save_freq"]
betas = tuple(cfg_decoder["betas"])
weight_decay = cfg_decoder["weight_decay"]
start_step = cfg_decoder.get("start_step", 0)
use_disc = cfg_decoder["use_discriminator"]
# Loss Weights
lambda_mel = cfg_decoder["lambda_mel"]
lambda_fm = cfg_decoder["lambda_fm"]
lambda_gen = cfg_decoder["lambda_gen"]
@@ -150,19 +291,18 @@ if __name__ == '__main__':
# 1. Load LLM and Freeze
# ------------------
print("Loading LLM...")
llm_config = AutoConfig.from_pretrained('ekwek/Soprano-80M')
llm_config = AutoConfig.from_pretrained(tokenizer_name)
model = AutoModelForCausalLM.from_config(llm_config)
pretrained_llm_path = cfg_paths["pretrained_llm_path"]
if pretrained_llm_path and os.path.exists(pretrained_llm_path):
print(f"Loading custom LLM checkpoint from {pretrained_llm_path}")
if pretrained_llm_path.endswith('.safetensors'):
state_dict = load_file(pretrained_llm_path)
model.load_state_dict(state_dict)
else:
model = AutoModelForCausalLM.from_pretrained(pretrained_llm_path)
else:
print("Warning: No pretrained LLM provided for Decoder training. Using random init (not recommended).")
print("Warning: Training Decoder without a pre-trained LLM. Make sure this is intended.")
model.to(torch.bfloat16).to(device)
model.eval()
@@ -173,7 +313,6 @@ if __name__ == '__main__':
# ------------------
# 2. Load Decoder
# ------------------
print("Loading Decoder...")
decoder = SopranoDecoder()
pretrained_decoder_path = cfg_paths["pretrained_decoder_path"]
@@ -185,16 +324,12 @@ if __name__ == '__main__':
decoder.to(device)
decoder.train()
print("Decoder loaded.")
# Initialize MR-STFT Loss
mr_stft = MultiResolutionSTFTLoss().to(device)
# ------------------
# 3. Load Discriminator
# ------------------
discriminator = None
if cfg_decoder["use_discriminator"]:
if use_disc:
print("Initializing Discriminator...")
discriminator = Discriminator()
pretrained_disc_path = cfg_paths["pretrained_discriminator_path"]
@@ -202,6 +337,8 @@ if __name__ == '__main__':
if pretrained_disc_path and os.path.exists(pretrained_disc_path):
print(f"Loading custom Discriminator checkpoint from {pretrained_disc_path}")
discriminator.load_state_dict(torch.load(pretrained_disc_path, map_location='cpu'))
else:
print("Training Discriminator from scratch.")
discriminator.to(device)
discriminator.train()
@@ -220,40 +357,36 @@ if __name__ == '__main__':
pin_memory=True,
persistent_workers=True,
worker_init_fn=worker_seed_init,
collate_fn=collate_pack,
collate_fn=lambda batch_in: collate_pack(batch_in, tokenizer),
)
dataloader_it = iter(dataloader)
val_dataset = AudioDataset(val_dataset_path)
val_dataloader = DataLoader(
val_dataset,
batch_size=max(1, batch_size // 4), # Reduce val batch size to prevent OOM
batch_size=max(1, batch_size // 4),
shuffle=False,
num_workers=max(1, cfg_global["num_workers"] // 2),
pin_memory=True,
persistent_workers=True,
worker_init_fn=worker_seed_init,
collate_fn=collate_pack,
collate_fn=lambda batch_in: collate_pack(batch_in, tokenizer),
)
val_dataloader_it = iter(val_dataloader)
# ------------------
# 5. Optimizers
# ------------------
opt_g = torch.optim.AdamW(decoder.parameters(), max_lr, betas=betas, weight_decay=weight_decay)
opt_d = None
if cfg_decoder["use_discriminator"]:
if use_disc:
opt_d = torch.optim.AdamW(discriminator.parameters(), max_lr, betas=betas, weight_decay=weight_decay)
# ------------------
# Training Loop
# ------------------
pbar = tqdm(range(start_step + 1, max_steps + 1), ncols=200, dynamic_ncols=True)
pbar = tqdm(range(start_step, max_steps), ncols=200, dynamic_ncols=True)
for step in pbar:
start = time.time()
# Get Data
try:
batch_data = next(dataloader_it)
if batch_data[0] is None:
@@ -266,17 +399,15 @@ if __name__ == '__main__':
x, y, gt_audio, audio_mask = batch_data
x, y = x.to(device), y.to(device)
gt_audio = gt_audio.to(device) # (B, T_audio_samples)
gt_audio = gt_audio.to(device)
audio_mask = audio_mask.to(device)
# Forward LLM (No Grad)
with torch.no_grad():
with torch.autocast(device_type=device_type, dtype=torch.bfloat16):
outputs = model(x, output_hidden_states=True)
hidden_states = outputs.hidden_states[-1] # (B, T_total, D)
hidden_states = outputs.hidden_states[-1]
hidden_states = hidden_states.to(torch.float32)
# GATHER AUDIO LATENTS Logic
gathered_states_list = []
for b_idx in range(hidden_states.size(0)):
mask = audio_mask[b_idx]
@@ -292,25 +423,19 @@ if __name__ == '__main__':
length = gathered_states_list[b_idx].size(0)
audio_loss_mask[b_idx, :length] = True
# ---------------------
# Generator Forward
# ---------------------
decoder_in = decoder_in_padded.transpose(1, 2) # (B, C, T)
fake_audio = decoder(decoder_in) # (B, 1, T_audio_gen)
if fake_audio.size(1) == 1: fake_audio = fake_audio.squeeze(1)
min_len = min(fake_audio.size(1), gt_audio.size(1))
fake_audio = fake_audio[:, :min_len]
real_audio = gt_audio[:, :min_len]
# ---------------------
# Train Discriminator
# ---------------------
d_loss_item = 0.0
if cfg_decoder["use_discriminator"]:
if use_disc:
opt_d.zero_grad()
# --- Random Cropping Logic ---
decoder_in = decoder_in_padded.transpose(1, 2)
fake_audio = decoder(decoder_in)
if fake_audio.size(1) == 1: fake_audio = fake_audio.squeeze(1)
min_len = min(fake_audio.size(1), gt_audio.size(1))
fake_audio = fake_audio[:, :min_len]
real_audio = gt_audio[:, :min_len]
real_crop_list = []
fake_crop_list = []
@@ -341,15 +466,19 @@ if __name__ == '__main__':
opt_d.step()
d_loss_item = d_loss.item()
# ---------------------
# Train Generator
# ---------------------
opt_g.zero_grad()
# We need "fake_crops_g" (with grad) for generator loss
decoder_in = decoder_in_padded.transpose(1, 2)
fake_audio = decoder(decoder_in)
if fake_audio.size(1) == 1: fake_audio = fake_audio.squeeze(1)
min_len = min(fake_audio.size(1), gt_audio.size(1))
fake_audio = fake_audio[:, :min_len]
real_audio = gt_audio[:, :min_len]
real_crop_list_g = []
fake_crop_list_g = []
if cfg_decoder["use_discriminator"]:
if use_disc:
for b_idx in range(bsz):
valid_len = gathered_states_list[b_idx].size(0) * SAMPLES_PER_TOKEN
valid_len = min(valid_len, min_len)
@@ -369,7 +498,6 @@ if __name__ == '__main__':
real_crops_g = torch.stack(real_crop_list_g).unsqueeze(1)
fake_crops_g = torch.stack(fake_crop_list_g).unsqueeze(1)
# Mel Loss
frames_per_token = SAMPLES_PER_TOKEN // 512
mel_mask = audio_loss_mask.repeat_interleave(frames_per_token, dim=1)
@@ -384,16 +512,13 @@ if __name__ == '__main__':
loss_mel_raw = torch.nn.functional.l1_loss(pred_mel, gt_mel, reduction='none')
loss_mel = (loss_mel_raw * mel_mask.unsqueeze(1)).sum() / (mel_mask.sum() * pred_mel.size(1) + 1e-6)
# Multi-Resolution STFT Loss
sample_mask = audio_loss_mask.repeat_interleave(SAMPLES_PER_TOKEN, dim=1)
sample_mask = sample_mask[:, :min_len]
sample_mask = audio_loss_mask.repeat_interleave(SAMPLES_PER_TOKEN, dim=1)[:, :min_len]
sc_loss, mag_loss = mr_stft(fake_audio * sample_mask, real_audio * sample_mask)
loss_fm = torch.tensor(0.0, device=device)
loss_gen = torch.tensor(0.0, device=device)
if cfg_decoder["use_discriminator"]:
if use_disc:
y_d_rs, y_d_gs, fmap_rs, fmap_gs = discriminator(real_crops_g, fake_crops_g)
loss_fm = feature_matching_loss(fmap_rs, fmap_gs)
loss_gen, _ = generator_loss(y_d_gs)
@@ -403,10 +528,9 @@ if __name__ == '__main__':
total_loss_g.backward()
norm_g = torch.nn.utils.clip_grad_norm_(decoder.parameters(), 1.0)
# LR Update
lr = get_lr(step, max_lr, min_lr, warmup_steps, cooldown_steps, max_steps)
for param_group in opt_g.param_groups: param_group['lr'] = lr
if cfg_decoder["use_discriminator"]:
if use_disc:
for param_group in opt_d.param_groups: param_group['lr'] = lr / 2
opt_g.step()
@@ -417,7 +541,6 @@ if __name__ == '__main__':
tqdm_log = f'mel: {loss_mel.item():.3f} | gen: {loss_gen.item():.3f} | sc: {sc_loss.item():.3f} | mag: {mag_loss.item():.3f} | fm: {loss_fm.item():.3f} | d: {d_loss_item:.3f} | lr: {lr:.2e} | time: {dt:.2f} ms'
pbar.set_description(tqdm_log)
# WandB Logging
log_dict = {
"train/loss_mel": loss_mel.item(),
"train/loss_gen": loss_gen.item(),
@@ -429,144 +552,25 @@ if __name__ == '__main__':
"train/loss_mag": mag_loss.item()
}
# ---------------------
# Validation Loop
# ---------------------
if step % val_freq == 0:
decoder.eval()
if discriminator: discriminator.eval()
if step > 0 and step % val_freq == 0:
val_log_dict, val_dataloader_it = evaluate(
step=step,
val_dataloader_it=val_dataloader_it,
val_dataloader=val_dataloader,
model=model,
decoder=decoder,
discriminator=discriminator,
mel_fn=mel_fn,
mr_stft=mr_stft,
use_disc=use_disc,
device=device,
device_type=device_type,
val_steps=val_steps,
segment_size=segment_size_samples,
use_wandb=use_wandb
)
log_dict.update(val_log_dict)
val_mel_loss_accum = 0.0
val_gen_loss_accum = 0.0
val_fm_loss_accum = 0.0
val_d_loss_accum = 0.0
val_sc_loss_accum = 0.0
val_mag_loss_accum = 0.0
val_steps = 10
with torch.no_grad():
for _ in range(val_steps):
try:
val_batch = next(val_dataloader_it)
except StopIteration:
val_dataloader_it = iter(val_dataloader)
val_batch = next(val_dataloader_it)
vx, vy, vgt_audio, vaudio_mask = val_batch
vx, vy = vx.to(device), vy.to(device)
vgt_audio = vgt_audio.to(device)
vaudio_mask = vaudio_mask.to(device)
with torch.autocast(device_type=device_type, dtype=torch.bfloat16):
voutputs = model(vx, output_hidden_states=True)
v_hidden = voutputs.hidden_states[-1].to(torch.float32)
v_gathered_states_list = []
for b_idx in range(v_hidden.size(0)):
mask = vaudio_mask[b_idx]
v_valid_states = v_hidden[b_idx][mask]
v_gathered_states_list.append(v_valid_states)
v_in_padded = torch.nn.utils.rnn.pad_sequence(v_gathered_states_list, batch_first=True, padding_value=0.0)
v_bsz = v_in_padded.size(0)
v_max_aud_len = v_in_padded.size(1)
v_audio_loss_mask = torch.zeros((v_bsz, v_max_aud_len), dtype=torch.bool, device=device)
for b_idx in range(v_bsz):
length = v_gathered_states_list[b_idx].size(0)
v_audio_loss_mask[b_idx, :length] = True
v_in = v_in_padded.transpose(1, 2)
v_fake_audio = decoder(v_in)
if v_fake_audio.size(1) == 1: v_fake_audio = v_fake_audio.squeeze(1)
min_len_v = min(v_fake_audio.size(1), vgt_audio.size(1))
v_fake_audio = v_fake_audio[:, :min_len_v]
v_real_audio = vgt_audio[:, :min_len_v]
frames_per_token_v = SAMPLES_PER_TOKEN // 512
v_mel_mask = v_audio_loss_mask.repeat_interleave(frames_per_token_v, dim=1)
v_pred_mel = mel_fn(v_fake_audio)
v_gt_mel = mel_fn(v_real_audio)
min_mel_len_v = min(v_pred_mel.size(2), v_gt_mel.size(2), v_mel_mask.size(1))
v_pred_mel = v_pred_mel[:, :, :min_mel_len_v]
v_gt_mel = v_gt_mel[:, :, :min_mel_len_v]
v_mel_mask = v_mel_mask[:, :min_mel_len_v]
v_mel_loss_raw = torch.nn.functional.l1_loss(v_pred_mel, v_gt_mel, reduction='none')
v_mel_loss = (v_mel_loss_raw * v_mel_mask.unsqueeze(1)).sum() / (v_mel_mask.sum() * v_pred_mel.size(1) + 1e-6)
val_mel_loss_accum += v_mel_loss.item()
v_sample_mask = v_audio_loss_mask.repeat_interleave(SAMPLES_PER_TOKEN, dim=1)[:, :min_len_v]
v_sc_loss, v_mag_loss = mr_stft(v_fake_audio * v_sample_mask, v_real_audio * v_sample_mask)
val_sc_loss_accum += v_sc_loss.item()
val_mag_loss_accum += v_mag_loss.item()
if cfg_decoder["use_discriminator"]:
v_real_crop_list = []
v_fake_crop_list = []
v_min_len = min(v_fake_audio.size(1), v_real_audio.size(1))
for b_idx in range(v_bsz):
v_valid_len = v_gathered_states_list[b_idx].size(0) * SAMPLES_PER_TOKEN
v_valid_len = min(v_valid_len, v_min_len)
if v_valid_len <= segment_size_samples:
v_pad_len = segment_size_samples - v_valid_len
vr_c = torch.nn.functional.pad(v_real_audio[b_idx, :v_valid_len], (0, v_pad_len))
vf_c = torch.nn.functional.pad(v_fake_audio[b_idx, :v_valid_len], (0, v_pad_len))
else:
v_start_idx = random.randint(0, v_valid_len - segment_size_samples)
vr_c = v_real_audio[b_idx, v_start_idx : v_start_idx + segment_size_samples]
vf_c = v_fake_audio[b_idx, v_start_idx : v_start_idx + segment_size_samples]
v_real_crop_list.append(vr_c)
v_fake_crop_list.append(vf_c)
v_real_crops = torch.stack(v_real_crop_list).unsqueeze(1)
v_fake_crops = torch.stack(v_fake_crop_list).unsqueeze(1)
vy_d_rs, vy_d_gs, vfmap_rs, vfmap_gs = discriminator(v_real_crops, v_fake_crops)
v_fm_loss = feature_matching_loss(vfmap_rs, vfmap_gs)
v_gen_loss, _ = generator_loss(vy_d_gs)
v_d_loss, _, _ = discriminator_loss(vy_d_rs, vy_d_gs)
val_gen_loss_accum += v_gen_loss.item()
val_fm_loss_accum += v_fm_loss.item()
val_d_loss_accum += v_d_loss.item()
# Average metrics
val_log = {
"val/loss_mel": val_mel_loss_accum / val_steps,
"val/loss_gen": val_gen_loss_accum / val_steps,
"val/loss_fm": val_fm_loss_accum / val_steps,
"val/loss_d": val_d_loss_accum / val_steps,
"val/loss_sc": val_sc_loss_accum / val_steps,
"val/loss_mag": val_mag_loss_accum / val_steps
}
log_dict.update(val_log)
# Generate Mel Images (from last val batch)
if cfg_global["use_wandb"]:
gen_mel = mel_fn(v_fake_audio[0:1]).squeeze(0).cpu().numpy()
gt_mel = mel_fn(v_real_audio[0:1]).squeeze(0).cpu().numpy()
fig, ax = plt.subplots(2, 1, figsize=(10, 6))
ax[0].imshow(gt_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("Generated Mel (Val)")
plt.tight_layout()
log_dict["val/mel_spectrograms"] = wandb.Image(fig)
plt.close(fig)
decoder.train()
if discriminator: discriminator.train()
# Save Checkpoint
if step > 0 and step % save_freq == 0:
print(f"\nSaving checkpoint at step {step} to {save_path}...")
ckpt_name_dec = f"decoder_step_{step}.pth"
@@ -575,13 +579,17 @@ if __name__ == '__main__':
if discriminator:
torch.save(discriminator.state_dict(), os.path.join(save_path, ckpt_name_disc))
if cfg_global["use_wandb"]:
if use_wandb:
wandb.log(log_dict, step=step)
print(f"Training complete. Saving model at {save_path}")
print(f"\nTraining complete. Saving model at {save_path}")
torch.save(decoder.state_dict(), os.path.join(save_path, "decoder_trained.pth"))
if discriminator:
torch.save(discriminator.state_dict(), os.path.join(save_path, "discriminator_trained.pth"))
if cfg_global["use_wandb"]:
if use_wandb:
wandb.finish()
if __name__ == '__main__':
main()
+59 -44
View File
@@ -3,8 +3,6 @@ Training script for Soprano LLM backbone.
Usage:
python train_llm.py
Adapted from https://github.com/karpathy/nanoGPT
"""
import os
import random
@@ -21,36 +19,72 @@ from safetensors.torch import load_file
from dataset import AudioDataset
from config_loader import load_config
# Initialize tokenizer globally so it can be used in the collate functions
tokenizer = AutoTokenizer.from_pretrained('ekwek/Soprano-80M')
tokenizer.padding_side = 'right' # Essential for training!
def worker_seed_init(_):
worker_seed = torch.initial_seed() % (2**32-1)
np.random.seed(worker_seed)
random.seed(worker_seed)
def get_lr(it, max_lr, min_lr, warmup_steps, cooldown_steps, max_steps): # WSD schedule
def get_lr(it, max_lr, min_lr, warmup_steps, cooldown_steps, max_steps):
if it < warmup_steps:
return max_lr * (it + 1) / warmup_steps
if it < max_steps - cooldown_steps:
return max_lr
return min_lr + (max_lr - min_lr) * ((max_steps - it) / cooldown_steps)
def collate_dynamic(texts):
# Dynamic Batching: Pad to the longest in this batch (max 2048 safety)
def collate_pack(texts, tokenizer, seq_len, batch_size):
tokens_batch = tokenizer(texts, padding=False, truncation=False)
batch = []
cur_sample, cur_size = [], 0
for i in range(len(texts)):
tokens = torch.tensor(tokens_batch['input_ids'][i][:-1], dtype=torch.long)
cur_size += tokens.size(0)
cur_sample.append(tokens)
if cur_size >= seq_len + 1:
batch.append(torch.cat(cur_sample)[: seq_len + 1])
cur_sample, cur_size = [], 0
if len(batch) == batch_size:
break
if cur_sample and not batch:
batch_item = torch.cat(cur_sample + [torch.zeros(seq_len, dtype=torch.long)])[: seq_len + 1]
batch.append(batch_item)
if len(batch) < batch_size:
pad = batch[-1]
while len(batch) < batch_size:
batch.append(pad)
batch = torch.stack(batch)
x = batch[:, :-1]
y = batch[:, 1:]
return x, y
def collate_dynamic(texts, tokenizer):
tokenized = tokenizer(texts, padding=True, truncation=True, max_length=2048, return_tensors='pt', add_special_tokens=False)
batch = tokenized['input_ids']
attn_mask = tokenized['attention_mask']
x = batch[:, :-1]
y = batch[:, 1:]
# Attention mask needs to align with x. Since we shift x by removing the last token,
# we should also remove the last token from the mask.
attn_mask = attn_mask[:, :-1]
return x, y, attn_mask
def collate_pack_val(texts, tokenizer, seq_len):
out = tokenizer(texts, padding=True, truncation=True, max_length=seq_len+1, return_tensors='pt', add_special_tokens=False)
batch = out['input_ids']
if batch.size(1) < seq_len + 1:
pad_len = seq_len + 1 - batch.size(1)
batch = torch.nn.functional.pad(batch, (0, pad_len), value=tokenizer.pad_token_id)
x = batch[:, :-1]
y = batch[:, 1:]
return x, y
def compute_loss(x, logits, y, num_steps, mask=None):
pred = logits.view(-1, logits.size(-1))
labels = y.reshape(-1)
@@ -60,28 +94,20 @@ def compute_loss(x, logits, y, num_steps, mask=None):
mask = mask.reshape(-1)
loss = loss * mask
# Audio tokens: >=3 and <=8003.
# NOTE: If [STOP] is 3, it counts as audio.
# We apply the mask to filter out padding.
audio_mask_cond = torch.logical_and(labels >= 3, labels <= 8003)
if mask is not None:
audio_mask = audio_mask_cond & (mask > 0)
else:
audio_mask = audio_mask_cond
# Text tokens: The rest, BUT excluding masked (padding) tokens
if mask is not None:
text_mask = (~audio_mask_cond) & (mask > 0)
else:
text_mask = ~audio_mask_cond
# Avoid division by zero
audio_mean = loss[audio_mask].mean() if audio_mask.sum() > 0 else torch.tensor(0.0, device=loss.device)
text_mean = loss[text_mask].mean() if text_mask.sum() > 0 else torch.tensor(0.0, device=loss.device)
# Acc: only on non-masked tokens.
# Current logic: (logits.argmax(dim=-1) == y).view(-1)[audio_mask]
# This correctly calculates accuracy only on valid audio tokens.
acc = (logits.argmax(dim=-1).view(-1) == labels).view(-1)[audio_mask].to(torch.float32).mean()
if torch.isnan(acc): acc = torch.tensor(0.0, device=loss.device)
@@ -90,6 +116,7 @@ def compute_loss(x, logits, y, num_steps, mask=None):
acc = acc / num_steps
return audio_loss, text_loss, acc
def evaluate(model, val_dataloader, step, device, use_wandb):
model.eval()
val_dataloader_it = iter(val_dataloader)
@@ -122,10 +149,7 @@ def evaluate(model, val_dataloader, step, device, use_wandb):
model.train()
if __name__ == '__main__':
# ------------------
# Load Configuration
# ------------------
def main():
config = load_config("config.yaml")
cfg_global = config["global"]
cfg_paths = config["paths"]
@@ -134,25 +158,26 @@ if __name__ == '__main__':
device = cfg_global["device"]
seed = cfg_global["seed"]
device_type = "cuda" if device.startswith("cuda") else "cpu"
tokenizer_name = cfg_global.get("tokenizer_name", "ekwek/Soprano-80M")
torch.manual_seed(seed)
if torch.cuda.is_available():
torch.cuda.manual_seed(seed)
torch.set_float32_matmul_precision('high')
# Setup directories
train_dataset_path = os.path.join(cfg_paths["dataset_root"], "train.json")
val_dataset_path = os.path.join(cfg_paths["dataset_root"], "val.json")
save_path = os.path.join(cfg_paths["save_dir"], "llm")
os.makedirs(save_path, exist_ok=True)
print(f"Save Path: {save_path}")
if cfg_global["use_wandb"]:
wandb.init(project=cfg_global["wandb_project"], config=config)
# ------------------
# Hyperparameters
# ------------------
tokenizer = AutoTokenizer.from_pretrained(tokenizer_name)
tokenizer.padding_side = 'right'
max_steps = cfg_llm["max_steps"]
max_lr = float(cfg_llm["max_lr"])
min_lr = cfg_llm["min_lr_ratio"] * max_lr
@@ -168,31 +193,25 @@ if __name__ == '__main__':
betas = tuple(cfg_llm["betas"])
weight_decay = cfg_llm["weight_decay"]
# ------------------
# Model Setup
# ------------------
if cfg_llm["from_scratch"]:
print("Initializing model from scratch (random weights)...")
m_config = AutoConfig.from_pretrained('ekwek/Soprano-80M')
m_config = AutoConfig.from_pretrained(tokenizer_name)
model = AutoModelForCausalLM.from_config(m_config)
else:
pretrained_path = cfg_paths["pretrained_llm_path"]
if pretrained_path and os.path.exists(pretrained_path):
print(f"Loading pretrained model weights from {pretrained_path}...")
m_config = AutoConfig.from_pretrained('ekwek/Soprano-80M')
m_config = AutoConfig.from_pretrained(tokenizer_name)
model = AutoModelForCausalLM.from_config(m_config)
state_dict = load_file(pretrained_path)
model.load_state_dict(state_dict)
else:
print("Loading default pretrained model weights from HF...")
model = AutoModelForCausalLM.from_pretrained('ekwek/Soprano-80M')
model = AutoModelForCausalLM.from_pretrained(tokenizer_name)
model.to(device)
model.train()
# ------------------
# Dataset Setup
# ------------------
dataset = AudioDataset(train_dataset_path)
dataloader = DataLoader(
dataset,
@@ -202,7 +221,7 @@ if __name__ == '__main__':
pin_memory=True,
persistent_workers=True,
worker_init_fn=worker_seed_init,
collate_fn=collate_dynamic,
collate_fn=lambda texts: collate_dynamic(texts, tokenizer),
)
dataloader_it = iter(dataloader)
@@ -215,18 +234,11 @@ if __name__ == '__main__':
pin_memory=True,
persistent_workers=True,
worker_init_fn=worker_seed_init,
collate_fn=collate_dynamic,
collate_fn=lambda texts: collate_dynamic(texts, tokenizer),
)
# ------------------
# Optimizer
# ------------------
opt = torch.optim.AdamW(model.parameters(), max_lr, betas=betas, weight_decay=weight_decay, fused=True)
# ------------------
# Training Loop
# ------------------
# Determine start step based on loaded checkpoint if needed, defaulting to 1 for new runs
start_step = 1
pbar = tqdm(range(start_step, max_steps + 1), ncols=200, dynamic_ncols=True)
@@ -301,3 +313,6 @@ if __name__ == '__main__':
if cfg_global["use_wandb"]:
wandb.finish()
if __name__ == '__main__':
main()