Files
rustytorch/crates/models/rtx-csm/examples/lora_finetune_step.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

243 lines
8.3 KiB
Rust

//! 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::{Generator, Segment, audio_io};
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::<u32>()?;
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(
&[],
&current,
&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::<f32>()?,
);
}
// 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::<f32>()?;
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::<f32>()?;
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(())
}