Files
Nighthawk 9b7dd275c7 Refactored training pipeline to use YAML config and uv
- Move hyperparameters from hardcoded script values to `config.yaml`
- Replace pip requirements with `pyproject.toml` and `uv` support (CUDA 12.6)
- Refactor all scripts to use `pathlib` for robust path handling
- Optimize `generate_dataset.py` with GPU acceleration
- Register quantizer constants as buffers for proper device mapping
- Update README with new installation and usage instructions
2026-01-18 01:13:38 -05:00

224 lines
7.7 KiB
Python

"""
Training script for Soprano.
Usage:
python train.py --input-dir path/to/files --save-dir path/to/weights
Args:
--input-dir: Path to directory of LJSpeech-style dataset. If none is provided this defaults to the provided example dataset.
--save-dir: Path to directory to save weights
Adapted from https://github.com/karpathy/nanoGPT
"""
import argparse
import pathlib
import random
import time
import shutil
import yaml
import numpy as np
import torch
from torch.utils.data import DataLoader
from tqdm import tqdm
from transformers import AutoModelForCausalLM, AutoTokenizer
from dataset import AudioDataset
def get_args():
parser = argparse.ArgumentParser()
parser.add_argument("--config", type=str, default="config.yaml", help="Path to config file")
parser.add_argument("--input-dir", required=False, default="./example_dataset", type=pathlib.Path)
parser.add_argument("--save-dir", required=True, type=pathlib.Path)
return parser.parse_args()
# --- Load Configuration ---
args = get_args()
with open(args.config, 'r') as f:
config = yaml.safe_load(f)
# Inject YAML keys into global
globals().update(config)
# Handle derived values and type casting
betas = tuple(betas) if isinstance(betas, list) else betas
min_lr = 0.1 * max_lr
train_dataset_path = args.input_dir / 'train.json'
val_dataset_path = args.input_dir / 'val.json'
save_path = args.save_dir
# --- Setup Utilities ---
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): # WSD schedule
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(texts):
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.append(torch.cat(cur_sample + [torch.zeros(seq_len, dtype=torch.long)])[: seq_len + 1])
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 compute_loss(logits, y, num_steps):
pred = logits.view(-1, logits.size(-1))
labels = y.reshape(-1)
loss = torch.nn.functional.cross_entropy(pred, labels, reduction='none')
audio_mask = torch.logical_and(y >= 3, y <= 8003).view(-1)
audio_loss = loss[audio_mask].mean()
text_loss = loss[~audio_mask].mean()
acc = (logits.argmax(dim=-1) == y).view(-1)[audio_mask].to(torch.float32).mean()
audio_loss = audio_loss / num_steps
text_loss = text_loss / num_steps
acc = acc / num_steps
return audio_loss, text_loss, acc
def evaluate(val_dataloader, model, device, device_type):
model.eval()
val_dataloader_it = iter(val_dataloader)
with torch.no_grad():
val_audio_loss_accum = torch.tensor(0.0).to(device)
val_text_loss_accum = torch.tensor(0.0).to(device)
val_acc_accum = torch.tensor(0.0).to(device)
val_loss_steps = 1
for _ in range(val_loss_steps):
x, y = next(val_dataloader_it)
x, y = x.to(device), y.to(device)
with torch.autocast(device_type=device_type, dtype=torch.bfloat16):
logits = model(x).logits
audio_loss, text_loss, acc = compute_loss(logits, y, val_loss_steps)
val_audio_loss_accum += audio_loss.detach()
val_text_loss_accum += text_loss.detach()
val_acc_accum += acc.detach()
print(f"validation text loss: {val_text_loss_accum.item():.4f}\tvalidation audio loss: {val_audio_loss_accum.item():.4f}\tvalidation acc: {val_acc_accum.item():.4f}")
model.train()
# --- Main Training Loop ---
if __name__ == '__main__':
# Initialize tokenizer
tokenizer = AutoTokenizer.from_pretrained('ekwek/Soprano-80M')
# Environment Setup
device_type = "cuda" if device.startswith("cuda") else "cpu"
torch.manual_seed(seed)
if torch.cuda.is_available():
torch.cuda.manual_seed(seed)
torch.set_float32_matmul_precision('high')
# Create Save Directory & archive the config used
save_path.mkdir(parents=True, exist_ok=True)
shutil.copy(args.config, save_path / "config_used.yaml")
print(f"Save Path: {save_path}")
# LR schedule steps
warmup_steps = int(max_steps * warmup_ratio)
cooldown_steps = int(max_steps * cooldown_ratio)
# Model
model = AutoModelForCausalLM.from_pretrained('ekwek/Soprano-80M')
model.to(torch.bfloat16).to(device)
model.train()
# Datasets
dataloader = DataLoader(
AudioDataset(train_dataset_path),
batch_size=batch_size * 16,
shuffle=True,
num_workers=4,
pin_memory=True,
persistent_workers=True,
worker_init_fn=worker_seed_init,
collate_fn=collate_pack,
)
dataloader_it = iter(dataloader)
val_dataloader = DataLoader(
AudioDataset(val_dataset_path),
batch_size=batch_size * 16,
shuffle=False,
num_workers=1,
pin_memory=True,
persistent_workers=True,
worker_init_fn=worker_seed_init,
collate_fn=collate_pack,
)
# Optimizer
opt = torch.optim.AdamW(model.parameters(), max_lr, betas=betas, weight_decay=weight_decay, fused=True)
pbar = tqdm(range(0, max_steps), ncols=200, dynamic_ncols=True)
for step in pbar:
start = time.time()
if val_freq > 0 and (step % val_freq == 0 or step == max_steps - 1):
evaluate(val_dataloader, model, device, device_type)
opt.zero_grad()
audio_loss_accum = 0.0
text_loss_accum = 0.0
acc_accum = 0.0
for micro_step in range(grad_accum_steps):
try:
x, y = next(dataloader_it)
except StopIteration:
dataloader_it = iter(dataloader)
x, y = next(dataloader_it)
x, y = x.to(device), y.to(device)
with torch.autocast(device_type=device_type, dtype=torch.bfloat16):
logits = model(x).logits
audio_loss, text_loss, acc = compute_loss(logits, y, grad_accum_steps)
audio_loss_accum += audio_loss.detach()
text_loss_accum += text_loss.detach()
acc_accum += acc.detach()
total_loss = audio_loss + text_factor * text_loss
total_loss.backward()
norm = torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)
lr = get_lr(step)
for param_group in opt.param_groups:
param_group['lr'] = lr
opt.step()
if torch.cuda.is_available():
torch.cuda.synchronize()
end = time.time()
dt = (end - start) * 1000
tokens_per_second = (batch_size * seq_len * grad_accum_steps) / (end - start)
tqdm_log = f'text: {text_loss_accum.item():.3f} | audio: {audio_loss_accum.item():.3f} | acc: {acc_accum.item():.4f} | lr: {lr:.2e} | ms: {dt:.1f} | t/s: {tokens_per_second:.1f}'
pbar.set_description(tqdm_log)
print(f"Training complete. Saving model at {save_path}")
model.save_pretrained(save_path)
tokenizer.save_pretrained(save_path)
print("Saving done.")