VoxModel.training_step(batch) performs:
- Build conditioning via
ConditionAggregator+WhisperAwareConditioning. - Sample
t ~ Uniform[1, num_steps]per batch item. x_t, _, v_target = schedule.add_noise(mel, t)— the schedule is either Flow Matching (x_t = (1-t)·mel + t·noise,v = noise - mel, default) or cosine-diffusion v-prediction, selected byVoxModelConfig.schedule_type.v_pred = DiffusionDecoder(x_t, t, cond).- Masked L2 between
v_predandv_targetover valid frames.
The mask divisor is mask.sum() * n_mels so loss magnitudes are comparable
to an unmasked .mean() reduction (validated by tests/training/test_losses.py).
| Loss | When to use |
|---|---|
diffusion_loss(v_pred, v_target, mask) |
primary objective |
mel_l1_loss(mel_pred, mel_target, mask) |
optional auxiliary (design spec weight 0.5) |
f0_consistency_loss(mel_pred, f0_target, f0_extractor) |
dependency-injected — pass None to disable |
vox.training.optim.OptimConfig (design defaults):
OptimConfig(
lr=2.0e-4,
betas=(0.9, 0.98),
weight_decay=0.01,
warmup_steps=2000,
max_steps=100_000,
min_lr_ratio=0.01,
)build_scheduler(opt, cfg) returns a LambdaLR that does linear warmup →
cosine decay → min_lr_ratio floor.
Minimal main loop:
trainer = VoxTrainer(
cfg=TrainConfig(max_steps=100_000, log_interval=100, val_interval=1000,
ckpt_interval=5000, grad_clip=1.0, ckpt_dir="ckpts",
optim=OptimConfig()),
model=VoxModel(cfg),
train_loader=DataLoader(...),
val_loader=DataLoader(...),
logger=StdoutLogger(),
)
trainer.train()Per step:
- Moves batch to device.
- Runs
model.training_step. optimizer.zero_grad(set_to_none=True),loss.backward(),clip_grad_norm_(grad_clip),optimizer.step(),scheduler.step().- Logs every
log_interval(loss + current LR). - Validates every
val_interval(averagesmodel.training_steplosses acrossval_batchesitems). - Saves every
ckpt_intervaltockpt_dir/step_########.pt.
{
"step": int,
"model": state_dict,
"optimizer": state_dict,
"scheduler": state_dict,
}
Round-trip is tested (test_trainer_checkpoint_roundtrip). Reload with
VoxTrainer.load_checkpoint(path).
The Logger Protocol is implemented by three classes; pick via
build_logger(kind):
| kind | Class | Extra deps |
|---|---|---|
"stdout" |
StdoutLogger |
— (always available) |
"tensorboard" |
TensorBoardLogger |
tensorboard |
"wandb" |
WandbLogger |
wandb |
Imports are lazy so missing optional deps don't block CI.
A future scripts/train.py will resolve configs/train/{debug,smoke,production}.yaml
into the dataclasses above and call trainer.train(). The trainer itself is
fully usable from a notebook today.