//! End-to-end LoRA fine-tuning step demo on real CSM-1B. //! //! Pipeline: //! 1. Load CSM-1B (FP) //! 2. Inject LoRA adapters on backbone q_proj/v_proj (StyleSpeech recipe) //! 3. Mimi-encode reference audio → target codebook tokens //! 4. Tokenize transcript //! 5. Repeated forward_loss → backward → AdamW step → refresh_lora //! 6. Verify loss decreases //! //! This proves the LoRA fine-tuning path works end-to-end. The next step is //! to scale up: paired-data loader (multiple `(transcript, wav)` pairs), //! multi-epoch driver, checkpoint export, and a generation script that loads //! the trained adapter. //! //! Usage: //! cargo run -p rtx-csm --release --example lora_finetune_step -- \ //! --wav /tmp/csm/hello.wav --text "Hello from Rust." --steps 30 use anyhow::Result; use candle_nn::{AdamW, Optimizer, ParamsAdamW, VarMap}; use clap::Parser; use rtx_csm::lora::LoraConfig; use rtx_csm::{audio_io, Generator, Segment}; use std::path::PathBuf; #[derive(Debug, Parser)] struct Cli { #[arg(long)] wav: PathBuf, #[arg(long)] text: String, #[arg(long, default_value_t = 0)] speaker: u32, #[arg(long, default_value_t = 30)] steps: usize, #[arg(long, default_value_t = 5e-4)] lr: f64, #[arg(long, default_value_t = 8)] rank: usize, #[arg(long, default_value_t = 16.0)] alpha: f32, #[arg(long)] cpu: bool, /// Use Phase 12.1 extended coverage (q+k+v+output_proj + MLP) instead of /// q+v only. Useful for sanity-checking the wider adapter set on a single /// utterance before launching a real training run. #[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)?; // Audio side: encode reference WAV through Mimi. let audio = audio_io::load_mono_24k(&cli.wav)?; let codes = generator.mimi.encode(&audio)?; let (_b, _cb, t_frames) = codes.dims3()?; if t_frames < 1 { anyhow::bail!("audio too short; need at least one Mimi frame"); } let target_codes = codes .narrow(2, 0, 1)? .squeeze(2)? .flatten_all()? .to_vec1::()?; println!("target codes (frame 0): {} tokens", target_codes.len()); // Build prompt. let current = Segment::new_text(cli.speaker, &cli.text); let prompt = rtx_csm::prompt::build_prompt( &[], ¤t, &generator.model, &mut generator.mimi, &generator.tokenizer, )?; // Inject LoRA adapters into the backbone. `--extended-lora` opts into // Phase 12.1 coverage (q+k+v+output_proj + MLP). let base = if cli.extended_lora { LoraConfig::extended() } else { LoraConfig::default() }; let lora_cfg = LoraConfig { rank: cli.rank, alpha: cli.alpha, ..base }; 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: {} trainable params across {} adapter Vars (rank={} alpha={})", n_params, vm.all_vars().len(), cli.rank, cli.alpha, ); // Direct sanity: take an actual LoRA Var that was injected into the model, // compute a trivial loss involving it, see if backward populates its gradient. { let some_var = &vm.all_vars()[0]; let direct_loss = some_var.as_tensor().sum_all()?; let g = direct_loss.backward()?; println!( "DIRECT-on-injected-Var: {} grad entries; var has grad: {}", g.get_ids().count(), g.get(some_var).is_some(), ); } // Sanity: isolated forward through a single LoRA delta, verify backward gives gradients. { use candle_core::{DType, Tensor}; use candle_nn::VarMap; let test_vm = VarMap::new(); // Note: LoraDelta::new init B=0, so initially the chain produces 0. // We need to force non-zero B for A to have nonzero gradient OR rely on // B getting the only grad. Either way, at least one Var should appear // in the GradStore after backward. let test_lora = rtx_csm::lora::LoraDelta::new(8, 16.0, 64, 64, "test", &test_vm, &device, DType::F32)?; let xs = Tensor::randn(0.0f32, 1.0, (4, 64), &device)?; let target = Tensor::randn(0.0f32, 1.0, (4, 64), &device)?; let pred = test_lora.forward(&xs)?; let diff = (pred - &target)?; let loss = diff.sqr()?.mean_all()?; let g = loss.backward()?; let n_with_grads = test_vm .all_vars() .iter() .filter(|v| g.get(v).is_some()) .count(); println!( "ISOLATED LoraDelta: {}/{} Vars have grads (loss={:.4})", n_with_grads, test_vm.all_vars().len(), loss.to_scalar::()?, ); } // Initial loss (sanity check — should match the un-adapted forward_loss // since B is zero-initialized → adapter is a no-op at step 0). generator.model.clear_kv_cache(); let init_loss_tensor = generator .model .inner .forward_loss(&prompt.tokens, &prompt.mask, 0, &target_codes)?; let init_loss = init_loss_tensor.to_scalar::()?; println!("\ninitial loss (LoRA-init zero): {init_loss:.4}"); println!( "loss tensor: track_op={} is_variable={}", init_loss_tensor.track_op(), init_loss_tensor.is_variable() ); // Optimizer. let mut optim = AdamW::new( vm.all_vars(), ParamsAdamW { lr: cli.lr, ..ParamsAdamW::default() }, )?; let mut last_loss = init_loss; for step in 0..cli.steps { generator.model.clear_kv_cache(); let loss = generator.model.inner.forward_loss( &prompt.tokens, &prompt.mask, 0, &target_codes, )?; let grads = loss.backward()?; if step == 0 { let n_grads_present = vm .all_vars() .iter() .filter(|v| grads.get(v).is_some()) .count(); let n_total = vm.all_vars().len(); // How many gradient entries does the GradStore have TOTAL? let total_grad_entries = grads.get_ids().count(); println!( "DEBUG step 0: {n_grads_present}/{n_total} LoRA Vars have gradients (total grad entries: {total_grad_entries})" ); // is_variable check on LoRA Vars let n_var_flagged = vm .all_vars() .iter() .filter(|v| v.as_tensor().is_variable()) .count(); println!(" LoRA Vars with is_variable=true: {n_var_flagged}/{n_total}"); // Print first few var tensor ids vs first few grad ids print!(" LoRA Var ids (first 4): "); for v in vm.all_vars().iter().take(4) { print!("{:?} ", v.as_tensor().id()); } println!(); print!(" GradStore ids (first 4): "); for id in grads.get_ids().take(4) { print!("{:?} ", id); } println!(); } optim.step(&grads)?; // Pull updated A/B values into the LoraDelta tensors held by Attention. generator.model.inner.refresh_lora(&vm)?; let l = loss.to_scalar::()?; last_loss = l; if step % 5 == 0 || step == cli.steps - 1 { println!("step {step:>3}: loss = {l:.4}"); } } println!( "\nfinal loss: {last_loss:.4} (start: {init_loss:.4}, change: {:+.2}%)", 100.0 * (last_loss - init_loss) / init_loss, ); if last_loss < init_loss * 0.95 { println!("✓ LoRA fine-tune is learning — loss dropped >5%"); } else if last_loss < init_loss { println!("▲ loss dropped slightly — try more steps or higher LR"); } else { println!("⚠ loss did not decrease — check gradient flow / LR"); } Ok(()) }