Files
rustytorch/crates/models/rtx-csm/examples/lora_train.rs
T
osobhandClaude Opus 4.7 a5cedfb46a rtx-csm: emotional_speech_guide — CREMA-D vs RAVDESS firdhokk verdict
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]>
2026-04-30 00:01:02 -07:00

158 lines
4.9 KiB
Rust
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
//! 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(())
}