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]>
243 lines
8.3 KiB
Rust
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(
|
|
&[],
|
|
¤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::<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(())
|
|
}
|