8-gen bench (4 emotions × 2 corpora) at seed=42 against firdhokk Whisper-LV3: target RAVDESS CREMA-D happy happy (0.999) ✓ happy (0.999) ✓ angry neutral (0.92) sad (0.99) fearful happy (0.998) fearful (0.984) ✓ sad angry (0.99) fearful (0.99) CREMA-D 2/4 vs RAVDESS 1/4. Larger / more naturalistic corpus produces more class-pure fearful direction. Neither corpus solves angry or sad — recipe shifts into 'vague expressivity' rather than class-specific corners. Practical: prefer CREMA-D when available; A/B both per emotion if class precision matters. Co-Authored-By: Claude Opus 4.7 (1M context) <[email protected]>
158 lines
4.9 KiB
Rust
158 lines
4.9 KiB
Rust
//! End-to-end LoRA fine-tuning CLI.
|
||
//!
|
||
//! Scans a directory of `(audio.wav, audio.txt)` pairs, fine-tunes a rank-r
|
||
//! LoRA adapter on the backbone q/v projections, saves the adapter as a
|
||
//! safetensors file. Then `examples/generate --lora <path>` (Phase 12.4)
|
||
//! loads and applies it at inference.
|
||
//!
|
||
//! Usage:
|
||
//! mkdir /tmp/voice && place foo.wav + foo.txt pairs in there
|
||
//! cargo run -p rtx-csm --release --features metal --example lora_train -- \
|
||
//! --data-dir /tmp/voice --output /tmp/voice.safetensors \
|
||
//! --epochs 5 --rank 8 --alpha 16
|
||
|
||
use anyhow::Result;
|
||
use candle_nn::VarMap;
|
||
use clap::Parser;
|
||
use rtx_csm::Generator;
|
||
use rtx_csm::lora::LoraConfig;
|
||
use rtx_csm::training::{
|
||
LoraAdapterMetadata, Trainer, TrainingConfig, TrainingDataset, save_lora_adapter_with_metadata,
|
||
};
|
||
use std::path::PathBuf;
|
||
|
||
#[derive(Debug, Parser)]
|
||
struct Cli {
|
||
/// Directory containing (.wav, .txt) pairs.
|
||
#[arg(long)]
|
||
data_dir: PathBuf,
|
||
|
||
/// Output safetensors path for the trained LoRA adapter.
|
||
#[arg(long)]
|
||
output: PathBuf,
|
||
|
||
#[arg(long, default_value_t = 0)]
|
||
speaker: u32,
|
||
|
||
#[arg(long, default_value_t = 5)]
|
||
epochs: usize,
|
||
|
||
#[arg(long, default_value_t = 8)]
|
||
rank: usize,
|
||
|
||
#[arg(long, default_value_t = 16.0)]
|
||
alpha: f32,
|
||
|
||
#[arg(long, default_value_t = 5e-4)]
|
||
peak_lr: f64,
|
||
|
||
#[arg(long, default_value_t = 1e-5)]
|
||
end_lr: f64,
|
||
|
||
#[arg(long, default_value_t = 16)]
|
||
warmup_steps: usize,
|
||
|
||
#[arg(long, default_value_t = 1.0)]
|
||
grad_clip: f64,
|
||
|
||
/// Frames sampled per training step (1 = one frame per step).
|
||
#[arg(long, default_value_t = 4)]
|
||
frames_per_step: usize,
|
||
|
||
#[arg(long, default_value_t = 42)]
|
||
seed: u64,
|
||
|
||
#[arg(long)]
|
||
cpu: bool,
|
||
|
||
/// Use extended LoRA coverage (q+k+v+output_proj + MLP w1/w2/w3) instead
|
||
/// of the default q+v only. Recommended when fine-tuning on hours of
|
||
/// paired audio for stronger prosody adaptation; the default is safer
|
||
/// when data is scarce (~30 min). Adapter file grows ~6× (still tiny
|
||
/// vs the base model).
|
||
#[arg(long, default_value_t = false)]
|
||
extended_lora: bool,
|
||
}
|
||
|
||
fn main() -> Result<()> {
|
||
tracing_subscriber::fmt().init();
|
||
let cli = Cli::parse();
|
||
|
||
let device = if cli.cpu {
|
||
candle_core::Device::Cpu
|
||
} else {
|
||
Generator::default_device()?
|
||
};
|
||
println!("device: {device:?}");
|
||
|
||
let mut generator = Generator::load_csm_1b(&device)?;
|
||
println!("model loaded");
|
||
|
||
println!("scanning dataset at {}", cli.data_dir.display());
|
||
let dataset = TrainingDataset::load_from_dir(&cli.data_dir, cli.speaker, &mut generator)?;
|
||
println!("dataset: {} examples", dataset.len());
|
||
|
||
// Inject LoRA into the backbone. `--extended-lora` opts into the Phase 12.1
|
||
// recipe (q+k+v+output_proj + MLP w1/w2/w3); default keeps q+v only.
|
||
let base = if cli.extended_lora {
|
||
LoraConfig::extended()
|
||
} else {
|
||
LoraConfig::default()
|
||
};
|
||
let lora_cfg = LoraConfig {
|
||
rank: cli.rank,
|
||
alpha: cli.alpha,
|
||
..base
|
||
};
|
||
println!("LoRA config: target_modules={:?}", lora_cfg.target_modules);
|
||
let vm = VarMap::new();
|
||
generator.model.inner.add_lora_to_backbone(&lora_cfg, &vm)?;
|
||
let n_params: usize = vm.all_vars().iter().map(|v| v.shape().elem_count()).sum();
|
||
println!(
|
||
"LoRA injected: {} trainable params ({} adapter Vars, rank={} alpha={})",
|
||
n_params,
|
||
vm.all_vars().len(),
|
||
cli.rank,
|
||
cli.alpha,
|
||
);
|
||
|
||
let train_cfg = TrainingConfig {
|
||
epochs: cli.epochs,
|
||
peak_lr: cli.peak_lr,
|
||
end_lr: cli.end_lr,
|
||
warmup_steps: cli.warmup_steps,
|
||
grad_clip: Some(cli.grad_clip),
|
||
frames_per_step: cli.frames_per_step,
|
||
seed: cli.seed,
|
||
};
|
||
let mut trainer = Trainer::new(&mut generator, &vm, &dataset, train_cfg);
|
||
let losses = trainer.train()?;
|
||
|
||
// Print loss curve summary.
|
||
if !losses.is_empty() {
|
||
let n = losses.len();
|
||
let first10: f32 = losses.iter().take(10).sum::<f32>() / 10.0_f32.min(n as f32);
|
||
let last10: f32 = losses.iter().rev().take(10).sum::<f32>() / 10.0_f32.min(n as f32);
|
||
println!(
|
||
"\nloss curve: first 10 avg = {first10:.4}, last 10 avg = {last10:.4}, change = {:+.2}%",
|
||
100.0 * (last10 - first10) / first10
|
||
);
|
||
println!(
|
||
"min={:.4} max={:.4}",
|
||
losses.iter().cloned().fold(f32::INFINITY, f32::min),
|
||
losses.iter().cloned().fold(f32::NEG_INFINITY, f32::max)
|
||
);
|
||
}
|
||
|
||
let metadata = LoraAdapterMetadata::from_lora_config(&lora_cfg);
|
||
save_lora_adapter_with_metadata(&vm, &cli.output, &metadata)?;
|
||
println!("\n✓ trained LoRA adapter saved to {}", cli.output.display());
|
||
println!(
|
||
"Use it at inference: examples/generate --lora {} \\\n\
|
||
\t--text \"...\" --out /tmp/out.wav\n\
|
||
(rank/alpha/extended auto-detected from embedded metadata)",
|
||
cli.output.display(),
|
||
);
|
||
Ok(())
|
||
}
|