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
This commit is contained in:
Nighthawk
2026-01-18 01:13:38 -05:00
parent 8230002e0b
commit 9b7dd275c7
10 changed files with 202 additions and 130 deletions
+1
View File
@@ -4,3 +4,4 @@ test.py
*.pth
dist/
*.egg-info/
.vs
+14
View File
@@ -33,6 +33,20 @@
## Installation
This project uses **[uv](https://github.com/astral-sh/uv)** for high-performance dependency management.
```bash
git clone [https://github.com/ekwek1/soprano-factory.git](https://github.com/ekwek1/soprano-factory.git)
cd soprano-factory
# Install dependencies (CPU)
uv sync
# Install dependencies (CUDA 12.6)
uv sync --extra gpu
## Manual Installation
```bash
git clone https://github.com/ekwek1/soprano-factory.git
cd soprano-factory
+20
View File
@@ -0,0 +1,20 @@
# Hardware & Reproducibility
device: "cuda:0"
seed: 1337
# Learning Rate Schedule
max_lr: 5.0e-4 # Use 5.0e-4 (with dot) to ensure it loads as a float
warmup_ratio: 0.1
cooldown_ratio: 0.1
# Training Dynamics
batch_size: 4
grad_accum_steps: 1
seq_len: 1024
max_steps: 10000
val_freq: 250
# Optimizer & Model Config
betas: [0.9, 0.95]
weight_decay: 0.1
text_factor: 0.0 # Increase to train on text inputs
+9 -2
View File
@@ -1,10 +1,13 @@
import json
import pathlib
from torch.utils.data import Dataset
class AudioDataset(Dataset):
def __init__(self, path):
with open(path, encoding='utf-8') as f:
# Convert string path to Path object if necessary for consistency
self.path = pathlib.Path(path)
with open(self.path, encoding='utf-8') as f:
self.dataset = json.load(f)
def __len__(self):
@@ -12,6 +15,10 @@ class AudioDataset(Dataset):
def __getitem__(self, idx):
text, audio = self.dataset[idx]
# Format: [STOP][TEXT]<text prompt>[START]<audio tokens>[STOP]
res = f"[STOP][TEXT]{text}[START]{''.join(list(map(lambda x: f'[{x}]', audio)))}[STOP]"
# Optimization: Use a generator expression for joining tokens
audio_tokens = ''.join(f'[{x}]' for x in audio)
res = f"[STOP][TEXT]{text}[START]{audio_tokens}[STOP]"
return res
+9 -37
View File
@@ -14,10 +14,7 @@ def safe_log(x: torch.Tensor, clip_val: float = 5e-3) -> torch.Tensor:
class SimpleMLP(nn.Module):
def __init__(self,
dim,
intermediate_dim,
):
def __init__(self, dim, intermediate_dim):
super().__init__()
self.pwconv1 = nn.Linear(dim, intermediate_dim)
self.act = nn.GELU()
@@ -31,14 +28,7 @@ class SimpleMLP(nn.Module):
class ConvNeXtBlock(nn.Module):
"""ConvNeXt Block adapted from https://github.com/facebookresearch/ConvNeXt to 1D audio signal.
Args:
dim (int): Number of input channels.
intermediate_dim (int): Dimensionality of the intermediate layer.
layer_scale_init_value (float, optional): Initial value for the layer scale. None means no scaling.
Defaults to None.
"""
"""ConvNeXt Block adapted from https://github.com/facebookresearch/ConvNeXt to 1D audio signal."""
def __init__(
self,
@@ -48,7 +38,7 @@ class ConvNeXtBlock(nn.Module):
dw_kernel_size: int = 7,
):
super().__init__()
self.dwconv = nn.Conv1d(dim, dim, kernel_size=dw_kernel_size, padding=dw_kernel_size//2, groups=dim) # depthwise conv
self.dwconv = nn.Conv1d(dim, dim, kernel_size=dw_kernel_size, padding=dw_kernel_size//2, groups=dim)
self.norm = nn.LayerNorm(dim, eps=1e-6)
self.mlp = SimpleMLP(dim, intermediate_dim)
self.gamma = (
@@ -72,17 +62,6 @@ class ConvNeXtBlock(nn.Module):
class VocosBackbone(nn.Module):
"""
Vocos backbone module built with ConvNeXt blocks.
Args:
input_channels (int): Number of input features channels.
dim (int): Hidden dimension of the model.
intermediate_dim (int): Intermediate dimension used in ConvNeXtBlock.
num_layers (int): Number of ConvNeXtBlock layers.
layer_scale_init_value (float, optional): Initial value for layer scaling.
"""
def __init__(
self,
input_channels: int,
@@ -124,16 +103,7 @@ class VocosBackbone(nn.Module):
nn.init.constant_(m.bias, 0)
def forward(self, x: torch.Tensor) -> torch.Tensor:
"""
Args:
x (Tensor): Input tensor of shape (B, C, L), where B is the batch size,
C denotes output features, and L is the sequence length.
Returns:
Tensor: Output of shape (B, L, H), where B is the batch size, L is the sequence length,
and H denotes the model dimension.
"""
x = self.embed(x) # (B, C, L)
x = self.embed(x)
x = self.norm(x.transpose(1, 2))
x = x.transpose(1, 2)
for conv_block in self.convnext:
@@ -185,17 +155,19 @@ class Encoder(nn.Module):
def encode(self, x):
x = self.encoder(x)
# Sequence slicing for downsampling
x = x[:, :, ::self.downsample_scale]
x = x.transpose(1,2)
x = x.transpose(1, 2)
x = self.downsampler(x)
x = self.quant(x)
return x
def preprocess(self, audio):
if audio.dim() == 2: # raw audio
# Ensure mel_spec is on the same device as the input audio
if audio.dim() == 2: # raw audio (B, T)
x = self.mel_spec(audio)
x = safe_log(x)
elif audio.dim() == 3: # mel spectrogram
elif audio.dim() == 3: # mel spectrogram (B, C, T)
x = audio
return x
+18 -14
View File
@@ -1,7 +1,6 @@
"""
Adapted from https://github.com/duchenzhuang/FSQ-pytorch/blob/main/quantizers/fsq.py#L41
"""
import torch
from torch import nn
from einops import rearrange
@@ -12,16 +11,24 @@ class FSQSTE(nn.Module):
super().__init__()
if levels:
self.dim = len(levels)
self.levels = torch.tensor(levels, dtype=torch.int32).view(1, 1, self.dim)
self.half_levels = (self.levels - 1) * (1 - 1e-3) / 2
self.offset = 0.5 - 0.5 * (self.levels % 2)
self.shift = torch.tan(self.offset / self.half_levels)
else:
self.levels = levels
self._basis = torch.cumprod(torch.tensor([1] + levels[:-1]),
dim=0,
dtype=torch.int32)
# Use register_buffer so these move to GPU automatically
levels_tensor = torch.tensor(levels, dtype=torch.int32).view(1, 1, self.dim)
self.register_buffer("levels", levels_tensor)
half_levels = (self.levels - 1) * (1 - 1e-3) / 2
self.register_buffer("half_levels", half_levels)
offset = 0.5 - 0.5 * (self.levels % 2)
self.register_buffer("offset", offset)
shift = torch.tan(self.offset / self.half_levels)
self.register_buffer("shift", shift)
basis = torch.cumprod(torch.tensor([1] + levels[:-1]), dim=0, dtype=torch.int32)
self.register_buffer("_basis", basis)
else:
self.levels = None
def _scale_and_shift(self, zhat_normalized):
half_width = self.levels // 2
@@ -32,20 +39,17 @@ class FSQSTE(nn.Module):
return (zhat - half_width) / half_width
def indices_to_level_indices(self, indices):
""" Converts indices to indices at each level, perhaps needed for a transformer with factorized embeddings """
indices = rearrange(indices, '... -> ... 1')
codes_non_centered = (indices // self._basis) % self.levels
return codes_non_centered
def to_codebook_index(self, zhat):
""" Converts a `code` to an index in the codebook. """
assert zhat.shape[-1] == self.dim
zhat = self._scale_and_shift(zhat)
indices = (zhat * self._basis).sum(dim = -1).round().to(torch.int32)
indices = (zhat * self._basis).sum(dim=-1).round().to(torch.int32)
return indices
def from_codebook_index(self, indices):
""" Inverse of `codes_to_indices`. """
level_indices = self.indices_to_level_indices(indices)
codes = self._scale_and_shift_inverse(level_indices)
return codes
+31 -17
View File
@@ -21,7 +21,7 @@ from huggingface_hub import hf_hub_download
from encoder.codec import Encoder
# Constants
SAMPLE_RATE = 32000
SEED = 42
VAL_PROP = 0.1
@@ -39,35 +39,50 @@ def get_args():
def main():
args = get_args()
input_dir = args.input_dir
device = "cuda" if torch.cuda.is_available() else "cpu"
print("Loading model.")
encoder = Encoder()
print(f"Loading model onto {device}.")
encoder = Encoder().to(device)
encoder_path = hf_hub_download(repo_id='ekwek/Soprano-Encoder', filename='encoder.pth')
encoder.load_state_dict(torch.load(encoder_path))
encoder.load_state_dict(torch.load(encoder_path, map_location=device))
encoder.eval()
print("Model loaded.")
print("Reading metadata.")
files = []
with open(f'{input_dir}/metadata.txt', encoding='utf-8') as f:
data = f.read().split('\n')
metadata_path = input_dir / 'metadata.txt'
with open(metadata_path, encoding='utf-8') as f:
data = f.read().strip().split('\n')
for line in data:
if '|' not in line:
continue
filename, transcript = line.split('|', maxsplit=1)
files.append((filename, transcript))
print(f'{len(files)} samples located in directory.')
print("Encoding audio.")
print(f"Encoding audio on {device}...")
dataset = []
for sample in tqdm(files):
filename, transcript = sample
sr, audio = wavfile.read(f'{input_dir}/wavs/{filename}.wav')
audio = torch.from_numpy(audio)
for filename, transcript in tqdm(files):
wav_path = input_dir / 'wavs' / f'{filename}.wav'
try:
sr, audio = wavfile.read(wav_path)
except FileNotFoundError:
continue
# Convert to float tensor for torchaudio/encoder compatibility
audio = torch.from_numpy(audio).float()
if sr != SAMPLE_RATE:
audio = torchaudio.functional.resample(audio, sr, SAMPLE_RATE)
audio = audio.unsqueeze(0)
# Prepare for encoder
audio = audio.unsqueeze(0).to(device)
with torch.no_grad():
audio_tokens = encoder(audio)
dataset.append([transcript, audio_tokens.squeeze(0).tolist()])
# Store results (move back to CPU for JSON serialization)
dataset.append([transcript, audio_tokens.squeeze(0).cpu().tolist()])
print("Generating train/test splits.")
random.seed(SEED)
@@ -79,12 +94,11 @@ def main():
print(f'# val samples: {len(val_dataset)}')
print("Saving datasets.")
with open(f'{input_dir}/train.json', 'w', encoding='utf-8') as f:
with open(input_dir / 'train.json', 'w', encoding='utf-8') as f:
json.dump(train_dataset, f, indent=2)
with open(f'{input_dir}/val.json', 'w', encoding='utf-8') as f:
with open(input_dir / 'val.json', 'w', encoding='utf-8') as f:
json.dump(val_dataset, f, indent=2)
print("Datasets saved.")
if __name__ == '__main__':
main()
+35
View File
@@ -0,0 +1,35 @@
[project]
name = "soprano-factory"
version = "0.1.0"
dependencies = [
"transformers",
"numpy",
"scipy",
"tqdm",
"einops",
"pyyaml",
"huggingface-hub",
"torch",
"torchaudio",
]
[project.optional-dependencies]
# Use uv sync --extra gpu to install CUDA support
gpu = [
"torch",
"torchaudio",
]
[[tool.uv.index]]
name = "pytorch-cu126"
url = "https://download.pytorch.org/whl/cu126"
explicit = true
[tool.uv.sources]
# Only use the CUDA index when the 'gpu' extra is requested
torch = [
{ index = "pytorch-cu126", extra = "gpu" },
]
torchaudio = [
{ index = "pytorch-cu126", extra = "gpu" },
]
+1
View File
@@ -6,3 +6,4 @@ torch
torchaudio
tqdm
transformers
pyyaml
+58 -54
View File
@@ -14,7 +14,8 @@ import argparse
import pathlib
import random
import time
import shutil
import yaml
import numpy as np
import torch
from torch.utils.data import DataLoader
@@ -23,52 +24,41 @@ from transformers import AutoModelForCausalLM, AutoTokenizer
from dataset import AudioDataset
def get_args():
parser = argparse.ArgumentParser()
parser.add_argument("--input-dir",
required=False,
default="./example_dataset",
type=pathlib.Path
)
parser.add_argument("--save-dir",
required=True,
type=pathlib.Path
)
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()
# training hyperparameters
device = 'cuda:0'
seed = 1337
max_lr = 5e-4
warmup_ratio = 0.1
cooldown_ratio = 0.1
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
batch_size = 4
grad_accum_steps = 1
seq_len = 1024
val_freq = 250
text_factor = 0.0 # currently does not train on text inputs, you can increase to change this
max_steps = 10000
betas = (0.9, 0.95)
weight_decay = 0.1
train_dataset_path = f'{args.input_dir}/train.json'
val_dataset_path = f'{args.input_dir}/val.json'
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:
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)
return min_lr + (max_lr - min_lr) * ((max_steps - it) / cooldown_steps)
def collate_pack(texts):
tokens_batch = tokenizer(texts, padding=False, truncation=False)
@@ -83,10 +73,9 @@ def collate_pack(texts):
cur_sample, cur_size = [], 0
if len(batch) == batch_size:
break
if cur_sample and not batch: # add partial sample if there isn't enough data
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 up to batch_size for consistency
pad = batch[-1]
while len(batch) < batch_size:
batch.append(pad)
@@ -99,7 +88,7 @@ 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_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()
@@ -108,7 +97,7 @@ def compute_loss(logits, y, num_steps):
acc = acc / num_steps
return audio_loss, text_loss, acc
def evaluate(val_dataloader):
def evaluate(val_dataloader, model, device, device_type):
model.eval()
val_dataloader_it = iter(val_dataloader)
with torch.no_grad():
@@ -128,29 +117,35 @@ def evaluate(val_dataloader):
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()
tokenizer = AutoTokenizer.from_pretrained('ekwek/Soprano-80M')
# --- 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
# LR schedule steps
warmup_steps = int(max_steps * warmup_ratio)
cooldown_steps = int(max_steps * cooldown_ratio)
# model
# Model
model = AutoModelForCausalLM.from_pretrained('ekwek/Soprano-80M')
model.to(torch.bfloat16).to(device)
model.train()
# dataset
dataset = AudioDataset(train_dataset_path)
# we need batch_size * 16 to have enough tokens after packing
dataloader = DataLoader(dataset,
# Datasets
dataloader = DataLoader(
AudioDataset(train_dataset_path),
batch_size=batch_size * 16,
shuffle=True,
num_workers=4,
@@ -160,8 +155,9 @@ if __name__ == '__main__':
collate_fn=collate_pack,
)
dataloader_it = iter(dataloader)
val_dataset = AudioDataset(val_dataset_path)
val_dataloader = DataLoader(val_dataset,
val_dataloader = DataLoader(
AudioDataset(val_dataset_path),
batch_size=batch_size * 16,
shuffle=False,
num_workers=1,
@@ -171,50 +167,58 @@ if __name__ == '__main__':
collate_fn=collate_pack,
)
# optimizer
# 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)
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:
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 = 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()
total_tokens = step * batch_size*seq_len*grad_accum_steps
end = time.time()
dt = (end-start)*1000
tokens_per_second = (batch_size*seq_len*grad_accum_steps) / (end-start)
tqdm_log = f'text loss: {text_loss_accum.item():.3f} | audio loss: {audio_loss_accum.item():.3f} | acc: {acc_accum.item():.4f} | lr: {lr:.2e} | norm: {norm:.3f} | time: {dt:.2f} ms | {tokens_per_second:.2f} t/s'
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.")