//! 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 ` (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::() / 10.0_f32.min(n as f32); let last10: f32 = losses.iter().rev().take(10).sum::() / 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(()) }