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

938 lines
34 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.
//! Multi-epoch LoRA training infrastructure.
//!
//! Provides:
//! - [`TrainingExample`] — one `(transcript, audio)` pair, with the audio
//! pre-encoded to Mimi codes so each training step doesn't re-encode.
//! - [`TrainingDataset`] — scan a directory of `(*.wav, *.txt)` pairs and
//! prepare a list of [`TrainingExample`]s.
//! - [`Trainer`] — wraps a [`Generator`] (with LoRA already injected) +
//! AdamW + LR schedule + gradient clipping. `train_one_epoch` and
//! `train_n_epochs` drive the optimizer.
//! - [`save_lora_adapter`] / [`load_lora_adapter`] — round-trip the LoRA
//! A,B matrices through safetensors.
use crate::audio_io;
use crate::error::{CsmError, Result};
use crate::generator::Generator;
use crate::prompt::{Segment, build_prompt};
use candle_core::Tensor;
use candle_nn::{AdamW, Optimizer, ParamsAdamW, VarMap};
use std::path::{Path, PathBuf};
/// One training example: a transcript paired with the Mimi-encoded audio
/// codes for that utterance. Pre-encoding avoids running Mimi on every
/// training step.
#[derive(Debug, Clone)]
pub struct TrainingExample {
pub speaker: u32,
pub text: String,
/// Per-frame target codes, shape `(num_frames, num_codebooks=32)`.
pub frame_codes: Vec<Vec<u32>>,
/// Path the audio came from, for logging.
pub source: Option<PathBuf>,
/// Optional control-token tag prepended to `text` at training time, so
/// the LoRA learns `<tag> <text>` → matching prosody. Same format as
/// `GenerateOptions::emotion_hint` at inference. None = unconditioned.
pub emotion_tag: Option<String>,
/// Optional curriculum stage label (e.g. "audiobook", "podcast", "va").
/// Consumed by [`CurriculumTrainer`] to bucket examples; ignored by the
/// flat [`Trainer`].
pub stage: Option<String>,
}
impl TrainingExample {
/// Build from raw 24 kHz mono samples by Mimi-encoding the audio.
///
/// Mimi's encoder carries streaming state across calls — when
/// processing a manifest of N clips back-to-back its internal
/// frame counter grows monotonically and eventually overflows the
/// transformer's max-seq buffer (`narrow invalid args ... [8192, ...]`
/// at clip ~30 in a 200-clip corpus). Reset before each encode so
/// every clip starts with a fresh state.
pub fn from_audio(
speaker: u32,
text: impl Into<String>,
audio_24k: &[f32],
generator: &mut Generator,
) -> Result<Self> {
generator.mimi.reset_state();
let codes = generator.mimi.encode(audio_24k)?;
let (_b, num_codebooks, num_frames) = codes.dims3()?;
let mut frame_codes = Vec::with_capacity(num_frames);
for frame in 0..num_frames {
let row = codes
.narrow(2, frame, 1)?
.squeeze(2)?
.flatten_all()?
.to_vec1::<u32>()?;
assert_eq!(row.len(), num_codebooks);
frame_codes.push(row);
}
Ok(Self {
speaker,
text: text.into(),
frame_codes,
source: None,
emotion_tag: None,
stage: None,
})
}
}
/// One row of a JSONL training manifest. Field shapes:
/// ```jsonl
/// {"wav": "data/0001.wav", "transcript": "...", "emotion_tag": "[whisper]", "stage": "audiobook", "speaker": 0}
/// ```
/// `emotion_tag` and `stage` are optional; `speaker` defaults to 0 when absent.
#[derive(Debug, Clone, serde::Deserialize)]
pub struct ManifestRow {
pub wav: PathBuf,
pub transcript: String,
#[serde(default)]
pub emotion_tag: Option<String>,
#[serde(default)]
pub stage: Option<String>,
#[serde(default)]
pub speaker: Option<u32>,
}
/// A dataset of training examples loaded from disk.
pub struct TrainingDataset {
pub examples: Vec<TrainingExample>,
}
impl TrainingDataset {
/// Scan a directory for `*.wav` files. For each `foo.wav` look up the
/// matching transcript at `foo.txt`. Skip files that don't have a pair.
/// Mimi-encode each audio file once at load time.
pub fn load_from_dir<P: AsRef<Path>>(
dir: P,
speaker: u32,
generator: &mut Generator,
) -> Result<Self> {
let dir = dir.as_ref();
let mut examples = Vec::new();
for entry in std::fs::read_dir(dir)? {
let entry = entry?;
let path = entry.path();
let ext = path.extension().and_then(|s| s.to_str()).unwrap_or("");
if ext != "wav" {
continue;
}
let txt_path = path.with_extension("txt");
if !txt_path.exists() {
tracing::warn!("skipping {}: no matching .txt transcript", path.display());
continue;
}
let text = std::fs::read_to_string(&txt_path)?.trim().to_string();
if text.is_empty() {
tracing::warn!("skipping {}: empty transcript", path.display());
continue;
}
let audio = audio_io::load_mono_24k(&path)?;
let mut ex = TrainingExample::from_audio(speaker, text, &audio, generator)?;
ex.source = Some(path);
tracing::info!(
"loaded example: {} chars, {} frames",
ex.text.len(),
ex.frame_codes.len()
);
examples.push(ex);
}
if examples.is_empty() {
return Err(CsmError::Config(format!(
"no (.wav, .txt) pairs found in {}",
dir.display()
)));
}
Ok(Self { examples })
}
/// Load a JSONL manifest of `ManifestRow` rows. Each row's WAV is
/// resolved relative to `manifest_dir` (the manifest's parent directory)
/// when not absolute, so manifests stay portable. Skips rows with empty
/// transcripts or missing audio. Mimi-encodes each audio file once.
pub fn load_from_manifest<P: AsRef<Path>>(
manifest: P,
default_speaker: u32,
generator: &mut Generator,
) -> Result<Self> {
let manifest = manifest.as_ref();
let manifest_dir = manifest.parent().unwrap_or(Path::new("."));
let body = std::fs::read_to_string(manifest)?;
let mut examples = Vec::new();
for (lineno, line) in body.lines().enumerate() {
let line = line.trim();
if line.is_empty() || line.starts_with('#') {
continue;
}
let row: ManifestRow = serde_json::from_str(line)
.map_err(|e| CsmError::Config(format!("manifest line {}: {e}", lineno + 1)))?;
let text = row.transcript.trim().to_string();
if text.is_empty() {
tracing::warn!("manifest line {}: empty transcript, skipping", lineno + 1);
continue;
}
let wav_path = if row.wav.is_absolute() {
row.wav.clone()
} else {
manifest_dir.join(&row.wav)
};
if !wav_path.exists() {
tracing::warn!(
"manifest line {}: wav not found at {}, skipping",
lineno + 1,
wav_path.display()
);
continue;
}
let speaker = row.speaker.unwrap_or(default_speaker);
let audio = audio_io::load_mono_24k(&wav_path)?;
// Mimi's transformer carries a position counter across encode()
// calls that `reset_state()` does NOT fully clear. Reload every
// 10 clips — empirical floor; 25 still tipped over occasionally
// (single-window cumulative reaches ~8102 frames before the
// 25-clip mark). Cost: ~200ms × N/10 reload overhead at load.
if !examples.is_empty() && examples.len() % 10 == 0 {
generator.mimi.reload().map_err(|e| {
CsmError::Config(format!("mimi reload at clip {}: {e}", examples.len()))
})?;
}
let mut ex = TrainingExample::from_audio(speaker, text, &audio, generator)?;
ex.source = Some(wav_path);
ex.emotion_tag = row.emotion_tag;
ex.stage = row.stage;
tracing::debug!(
"manifest example: tag={:?} stage={:?} frames={}",
ex.emotion_tag,
ex.stage,
ex.frame_codes.len()
);
examples.push(ex);
}
if examples.is_empty() {
return Err(CsmError::Config(format!(
"no usable rows in manifest {}",
manifest.display()
)));
}
tracing::info!(
"loaded {} examples from manifest {}",
examples.len(),
manifest.display()
);
Ok(Self { examples })
}
pub fn len(&self) -> usize {
self.examples.len()
}
pub fn is_empty(&self) -> bool {
self.examples.is_empty()
}
}
/// Cosine LR schedule with linear warmup. `step` is 0-indexed.
pub fn cosine_lr(step: usize, total_steps: usize, warmup: usize, peak: f64, end: f64) -> f64 {
if step < warmup {
peak * (step as f64 + 1.0) / (warmup.max(1) as f64)
} else {
let progress = ((step - warmup) as f64) / ((total_steps - warmup).max(1) as f64);
let cos = (1.0 + (progress * std::f64::consts::PI).cos()) * 0.5;
end + (peak - end) * cos
}
}
#[derive(Debug, Clone)]
pub struct TrainingConfig {
pub epochs: usize,
pub peak_lr: f64,
pub end_lr: f64,
pub warmup_steps: usize,
pub grad_clip: Option<f64>,
/// Number of frames sampled per example per step (random subset of frames).
/// Smaller = faster step but noisier; larger = slower but more stable.
pub frames_per_step: usize,
pub seed: u64,
}
impl Default for TrainingConfig {
fn default() -> Self {
Self {
epochs: 5,
peak_lr: 5e-4,
end_lr: 1e-5,
warmup_steps: 16,
grad_clip: Some(1.0),
frames_per_step: 4,
seed: 42,
}
}
}
pub struct Trainer<'a> {
pub generator: &'a mut Generator,
pub vm: &'a VarMap,
pub dataset: &'a TrainingDataset,
pub config: TrainingConfig,
}
impl<'a> Trainer<'a> {
pub fn new(
generator: &'a mut Generator,
vm: &'a VarMap,
dataset: &'a TrainingDataset,
config: TrainingConfig,
) -> Self {
Self {
generator,
vm,
dataset,
config,
}
}
/// Run a full training schedule: `epochs * dataset_size * frames_per_step`
/// gradient steps with the cosine LR schedule. Returns the per-step loss
/// trace so callers can plot / log.
pub fn train(&mut self) -> Result<Vec<f32>> {
use rand::rngs::StdRng;
use rand::seq::SliceRandom;
use rand::{Rng, SeedableRng};
let total_steps =
self.config.epochs * self.dataset.len() * self.config.frames_per_step.max(1);
if total_steps == 0 {
return Ok(Vec::new());
}
tracing::info!(
"training: epochs={} dataset={} fps={} → total_steps={} peak_lr={}",
self.config.epochs,
self.dataset.len(),
self.config.frames_per_step,
total_steps,
self.config.peak_lr
);
let mut rng = StdRng::seed_from_u64(self.config.seed);
let mut optim = AdamW::new(
self.vm.all_vars(),
ParamsAdamW {
lr: self.config.peak_lr,
..ParamsAdamW::default()
},
)?;
let mut losses = Vec::with_capacity(total_steps);
let mut step = 0usize;
for epoch in 0..self.config.epochs {
let mut order: Vec<usize> = (0..self.dataset.len()).collect();
order.shuffle(&mut rng);
for &idx in order.iter() {
let ex = &self.dataset.examples[idx];
if ex.frame_codes.is_empty() {
continue;
}
// Build the prompt once for this example. If the example
// carries an emotion_tag, prepend it with the same format the
// inference path uses (`apply_emotion_hint`) so the LoRA
// learns the same `<tag> <text>` pattern at training time.
let prompt_text = crate::generator::apply_emotion_hint(
ex.text.clone(),
ex.emotion_tag.as_deref(),
);
let current = Segment::new_text(ex.speaker, prompt_text);
let prompt = build_prompt(
&[],
&current,
&self.generator.model,
&mut self.generator.mimi,
&self.generator.tokenizer,
)?;
let prompt_len = prompt.tokens.dim(1)?;
for _ in 0..self.config.frames_per_step.max(1) {
let lr = cosine_lr(
step,
total_steps,
self.config.warmup_steps,
self.config.peak_lr,
self.config.end_lr,
);
optim.set_learning_rate(lr);
// Pick a random frame from this example as the target.
let frame_idx = rng.gen_range(0..ex.frame_codes.len());
let target = &ex.frame_codes[frame_idx];
self.generator.model.clear_kv_cache();
let loss = self.generator.model.inner.forward_loss(
&prompt.tokens,
&prompt.mask,
0,
target,
)?;
let loss_val = loss.to_scalar::<f32>()?;
losses.push(loss_val);
let mut grads = loss.backward()?;
if let Some(clip) = self.config.grad_clip {
clip_grads(self.vm, &mut grads, clip as f32)?;
}
optim.step(&grads)?;
self.generator.model.inner.refresh_lora(self.vm)?;
step += 1;
if step.is_multiple_of(10) || step == total_steps - 1 {
tracing::info!(
" step {step:>4}/{total_steps} epoch {epoch} ex {idx} frame {frame_idx} \
prompt_len={prompt_len} lr={lr:.2e} loss={loss_val:.4}"
);
}
}
}
}
Ok(losses)
}
}
/// One stage of a curriculum: a label that selects which examples this stage
/// trains on, the training schedule for the stage, and an optional checkpoint
/// path for the adapter snapshot taken AFTER the stage finishes.
///
/// The label matching rule:
/// - `"*"` matches every example regardless of `ex.stage`
/// - any other string matches `ex.stage == Some(label)` exactly
///
/// Per `personal_voice_training_guide.md` §4 the canonical 3-stage recipe is:
///
/// - audiobook: 3 epochs, peak_lr 1e-4
/// - podcast: 1 epoch, peak_lr 3e-5
/// - va: 1 epoch, peak_lr 1e-5
///
/// each over a growing pool. The `CurriculumTrainer` enforces the pool growth
/// implicitly by virtue of how you label your manifest (a podcast example
/// labeled `stage: "podcast"` is only matched by the podcast stage; if you
/// want it included in the va stage too, label it `"va"` and the va stage
/// will pick it up).
#[derive(Debug, Clone)]
pub struct CurriculumStage {
pub name: String,
pub config: TrainingConfig,
pub checkpoint: Option<PathBuf>,
}
/// Multi-stage LoRA trainer driven by per-example stage labels in the
/// dataset. After each stage the adapter is optionally written to a
/// safetensors checkpoint so callers can compare intermediate snapshots
/// (per the literature recipe: clean → mixed → noisy at decreasing LR).
pub struct CurriculumTrainer<'a> {
pub generator: &'a mut Generator,
pub vm: &'a VarMap,
pub dataset: &'a TrainingDataset,
pub stages: Vec<CurriculumStage>,
}
impl<'a> CurriculumTrainer<'a> {
pub fn new(
generator: &'a mut Generator,
vm: &'a VarMap,
dataset: &'a TrainingDataset,
stages: Vec<CurriculumStage>,
) -> Self {
Self {
generator,
vm,
dataset,
stages,
}
}
/// Run all stages sequentially. Returns one loss-trace `Vec<f32>` per
/// stage in `self.stages` order. The adapter VarMap is shared across
/// stages so each stage's training picks up where the previous one
/// left off — the curriculum is cumulative, not reset.
pub fn run(&mut self) -> Result<Vec<Vec<f32>>> {
let mut all_traces = Vec::with_capacity(self.stages.len());
for stage in self.stages.clone().into_iter() {
// Bucket examples for this stage. "*" is the global catch-all.
let mut stage_examples = Vec::new();
for ex in self.dataset.examples.iter() {
let pick = if stage.name == "*" {
true
} else {
matches!(ex.stage.as_deref(), Some(s) if s == stage.name)
};
if pick {
stage_examples.push(ex.clone());
}
}
if stage_examples.is_empty() {
tracing::warn!(
"curriculum stage '{}' matched 0 examples — skipping",
stage.name
);
all_traces.push(Vec::new());
continue;
}
tracing::info!(
"curriculum stage '{}': {} examples, epochs={} peak_lr={}",
stage.name,
stage_examples.len(),
stage.config.epochs,
stage.config.peak_lr,
);
let stage_dataset = TrainingDataset {
examples: stage_examples,
};
let mut trainer = Trainer::new(
self.generator,
self.vm,
&stage_dataset,
stage.config.clone(),
);
let losses = trainer.train()?;
if let Some(out) = stage.checkpoint.as_ref() {
save_lora_adapter(self.vm, out)?;
tracing::info!(
"curriculum stage '{}': checkpoint saved → {}",
stage.name,
out.display()
);
}
all_traces.push(losses);
}
Ok(all_traces)
}
}
/// Global L2 norm gradient clipping. Walks all Vars in the VarMap, computes
/// `||g||_2` across all gradients, and scales every gradient by
/// `min(1, max_norm / ||g||_2)`.
fn clip_grads(
vm: &VarMap,
grads: &mut candle_core::backprop::GradStore,
max_norm: f32,
) -> Result<()> {
let mut total_sq: f32 = 0.0;
for v in vm.all_vars() {
if let Some(g) = grads.get(&v) {
let s = g
.flatten_all()?
.to_dtype(candle_core::DType::F32)?
.sqr()?
.sum_all()?
.to_scalar::<f32>()?;
total_sq += s;
}
}
let norm = total_sq.sqrt();
if norm > max_norm {
let scale = max_norm / norm;
for v in vm.all_vars() {
if let Some(g) = grads.get(&v) {
let scaled = (g * scale as f64)?;
grads.insert(v.as_tensor(), scaled);
}
}
}
Ok(())
}
/// One-example metric from the held-out evaluator.
#[derive(Debug, Clone, serde::Serialize)]
pub struct EvalRow {
/// Position in the input dataset.
pub idx: usize,
/// Source wav (when known) — round-trips through the manifest loader.
pub source: Option<PathBuf>,
/// Number of frames in the example's Mimi-encoded audio.
pub frames: usize,
/// Mean teacher-forced cross-entropy across all 32 codebooks, averaged
/// over `frames_sampled` frames drawn at random from the example.
pub mean_loss: f32,
/// Number of frames the loss was averaged over (capped at the example's
/// available frames).
pub frames_sampled: usize,
}
/// Aggregate summary of an `evaluate_held_out` run. Useful for
/// machine-readable A/B comparison between base and LoRA-applied runs.
#[derive(Debug, Clone, serde::Serialize)]
pub struct EvalSummary {
pub n_examples: usize,
pub n_frames_total: usize,
pub mean_loss: f32,
pub median_loss: f32,
pub p90_loss: f32,
pub min_loss: f32,
pub max_loss: f32,
}
impl EvalSummary {
pub fn from_rows(rows: &[EvalRow]) -> Self {
if rows.is_empty() {
return Self {
n_examples: 0,
n_frames_total: 0,
mean_loss: 0.0,
median_loss: 0.0,
p90_loss: 0.0,
min_loss: 0.0,
max_loss: 0.0,
};
}
let mut losses: Vec<f32> = rows.iter().map(|r| r.mean_loss).collect();
losses.sort_by(|a, b| a.partial_cmp(b).unwrap_or(std::cmp::Ordering::Equal));
let n_frames_total: usize = rows.iter().map(|r| r.frames_sampled).sum();
let mean = losses.iter().sum::<f32>() / losses.len() as f32;
let median = losses[losses.len() / 2];
let p90_idx = ((losses.len() as f64) * 0.9).ceil() as usize;
let p90 = losses[p90_idx.saturating_sub(1).min(losses.len() - 1)];
Self {
n_examples: rows.len(),
n_frames_total,
mean_loss: mean,
median_loss: median,
p90_loss: p90,
min_loss: *losses.first().unwrap(),
max_loss: *losses.last().unwrap(),
}
}
}
/// Run the teacher-forced `forward_loss` over a held-out dataset to get a
/// scalar quality signal. Reuses the same prompt-build path the trainer uses
/// (with `apply_emotion_hint` so eval and inference apply tags identically).
///
/// `frames_per_example` caps how many frames are scored per example;
/// uniformly sampled with the given `seed` so two runs (base vs LoRA) over
/// the same dataset score the same frames in the same order. Set to 0 to
/// score every frame (slow on large examples).
pub fn evaluate_held_out(
generator: &mut Generator,
dataset: &TrainingDataset,
frames_per_example: usize,
seed: u64,
) -> Result<Vec<EvalRow>> {
use rand::rngs::StdRng;
use rand::{Rng, SeedableRng};
let mut rng = StdRng::seed_from_u64(seed);
let mut rows = Vec::with_capacity(dataset.len());
for (idx, ex) in dataset.examples.iter().enumerate() {
if ex.frame_codes.is_empty() {
continue;
}
let prompt_text =
crate::generator::apply_emotion_hint(ex.text.clone(), ex.emotion_tag.as_deref());
let current = Segment::new_text(ex.speaker, prompt_text);
let prompt = build_prompt(
&[],
&current,
&generator.model,
&mut generator.mimi,
&generator.tokenizer,
)?;
let total_frames = ex.frame_codes.len();
let n_sample = if frames_per_example == 0 {
total_frames
} else {
frames_per_example.min(total_frames)
};
let mut chosen = if frames_per_example == 0 {
(0..total_frames).collect::<Vec<_>>()
} else {
(0..n_sample)
.map(|_| rng.gen_range(0..total_frames))
.collect()
};
chosen.sort();
let mut sum = 0.0f32;
for frame_idx in chosen.iter().copied() {
generator.model.clear_kv_cache();
let target = &ex.frame_codes[frame_idx];
let loss =
generator
.model
.inner
.forward_loss(&prompt.tokens, &prompt.mask, 0, target)?;
sum += loss.to_scalar::<f32>()?;
}
let mean = if n_sample > 0 {
sum / n_sample as f32
} else {
0.0
};
rows.push(EvalRow {
idx,
source: ex.source.clone(),
frames: total_frames,
mean_loss: mean,
frames_sampled: n_sample,
});
if (idx + 1) % 10 == 0 {
tracing::info!("eval progress: {}/{}", idx + 1, dataset.len());
}
}
Ok(rows)
}
/// Self-describing metadata embedded in safetensors at save time so the
/// adapter file knows its own LoRA hyperparameters. Lets `apply_lora_adapter`
/// auto-configure at inference without the user having to remember matching
/// `--lora-rank` / `--lora-alpha` / `--extended-lora` flags.
///
/// Stored under the safetensors `__metadata__` map as a single JSON-encoded
/// string under key `rtx_csm_lora`. Backward compat: adapters without this
/// key fall back to caller-supplied defaults (rank=8, alpha=16, extended=false).
#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)]
pub struct LoraAdapterMetadata {
pub rank: usize,
pub alpha: f32,
/// Snapshot of `LoraConfig.target_modules` at training time. Reconstructs
/// `extended()` vs `default()` reliably when the user adds new presets.
pub target_modules: Vec<String>,
/// rtx-csm crate version that produced the adapter (just for diagnostics
/// — no compatibility logic gates on this today).
#[serde(default)]
pub crate_version: Option<String>,
}
const LORA_METADATA_KEY: &str = "rtx_csm_lora";
impl LoraAdapterMetadata {
pub fn from_lora_config(cfg: &crate::lora::LoraConfig) -> Self {
Self {
rank: cfg.rank,
alpha: cfg.alpha,
target_modules: cfg.target_modules.clone(),
crate_version: Some(env!("CARGO_PKG_VERSION").to_string()),
}
}
/// Heuristic: was the adapter trained with Phase 12.1 extended coverage?
/// True if any MLP module is in `target_modules`. The exact preset name
/// isn't stored — only the resolved patterns — so we infer from content.
pub fn is_extended(&self) -> bool {
self.target_modules
.iter()
.any(|m| m.contains("mlp.") || m.contains("output_proj") || m == "k_proj")
}
}
/// Read just the LoRA self-description from a safetensors header, without
/// loading any tensor data. Returns `None` when the file has no
/// `rtx_csm_lora` metadata key (older adapters trained before Phase 12.5).
pub fn read_lora_adapter_metadata<P: AsRef<Path>>(path: P) -> Result<Option<LoraAdapterMetadata>> {
let data = std::fs::read(path.as_ref())?;
let (_n, meta) = safetensors::SafeTensors::read_metadata(&data)
.map_err(|e| CsmError::Other(anyhow::anyhow!("read safetensors header: {e}")))?;
let map = match meta.metadata() {
Some(m) => m,
None => return Ok(None),
};
let json = match map.get(LORA_METADATA_KEY) {
Some(s) => s,
None => return Ok(None),
};
let parsed: LoraAdapterMetadata = serde_json::from_str(json)
.map_err(|e| CsmError::Other(anyhow::anyhow!("parse {LORA_METADATA_KEY}: {e}")))?;
Ok(Some(parsed))
}
/// One-shot inference-time LoRA loader: inject adapters with the matching
/// shape, populate from a safetensors file, and refresh the model's tensor
/// handles so the next forward pass picks up the trained weights.
///
/// Each parameter accepts an explicit override; when `None`, the value is
/// auto-detected from the adapter file's embedded `LoraAdapterMetadata`
/// (Phase 12.5). For files without metadata (older adapters), pass explicit
/// values: defaults are rank=8, alpha=16, extended=false.
///
/// Adapters trained with the narrow q+v recipe still load cleanly when
/// `extended = true` — the unmatched k/o/MLP slots are initialized B=0 and
/// stay no-ops, at the cost of a small extra VarMap.
///
/// Returns the `VarMap` so callers can keep it alive (the model holds plain
/// tensor handles refreshed from this VarMap; dropping it doesn't break a
/// loaded model but `refresh_lora` after another `optim.step` would). For
/// inference-only usage callers can drop the VarMap immediately after this
/// returns — `refresh_lora` snapshots the values into the model's owned
/// tensors.
pub fn apply_lora_adapter<P: AsRef<Path>>(
generator: &mut Generator,
path: P,
rank: Option<usize>,
alpha: Option<f32>,
extended: Option<bool>,
device: &candle_core::Device,
) -> Result<VarMap> {
let path = path.as_ref();
let meta = read_lora_adapter_metadata(path)?;
if meta.is_some() {
tracing::info!(
"adapter has embedded metadata; CLI rank/alpha/extended flags act as overrides"
);
}
let resolved_rank = rank.or_else(|| meta.as_ref().map(|m| m.rank)).unwrap_or(8);
let resolved_alpha = alpha
.or_else(|| meta.as_ref().map(|m| m.alpha))
.unwrap_or(16.0);
let resolved_extended = extended
.or_else(|| meta.as_ref().map(|m| m.is_extended()))
.unwrap_or(false);
let base = if resolved_extended {
crate::lora::LoraConfig::extended()
} else {
crate::lora::LoraConfig::default()
};
let cfg = crate::lora::LoraConfig {
rank: resolved_rank,
alpha: resolved_alpha,
..base
};
let vm = VarMap::new();
generator
.model
.inner
.add_lora_to_backbone(&cfg, &vm)
.map_err(|e| CsmError::Other(anyhow::anyhow!("add_lora_to_backbone: {e}")))?;
load_lora_adapter(&vm, path, device)?;
generator
.model
.inner
.refresh_lora(&vm)
.map_err(|e| CsmError::Other(anyhow::anyhow!("refresh_lora: {e}")))?;
tracing::info!(
"applied LoRA adapter from {} (rank={} alpha={} extended={})",
path.display(),
resolved_rank,
resolved_alpha,
resolved_extended,
);
Ok(vm)
}
/// Save the LoRA adapter parameters (A,B for every layer) as a safetensors
/// file. Reload via `load_lora_adapter`. No embedded metadata — callers that
/// know the LoRA hyperparameters should prefer
/// [`save_lora_adapter_with_metadata`] so the trained adapter is
/// self-describing at inference time.
pub fn save_lora_adapter<P: AsRef<Path>>(vm: &VarMap, out: P) -> Result<()> {
save_lora_adapter_inner(vm, out, None)
}
/// Save with embedded `LoraAdapterMetadata` (Phase 12.5). The metadata is
/// written into the safetensors header under `__metadata__["rtx_csm_lora"]`
/// as a JSON string; `apply_lora_adapter` will auto-pick it up at inference.
pub fn save_lora_adapter_with_metadata<P: AsRef<Path>>(
vm: &VarMap,
out: P,
metadata: &LoraAdapterMetadata,
) -> Result<()> {
save_lora_adapter_inner(vm, out, Some(metadata))
}
fn save_lora_adapter_inner<P: AsRef<Path>>(
vm: &VarMap,
out: P,
metadata: Option<&LoraAdapterMetadata>,
) -> Result<()> {
use std::collections::HashMap;
let vars = vm.data().lock().unwrap();
let mut tensors: HashMap<String, Tensor> = HashMap::new();
for (name, var) in vars.iter() {
tensors.insert(name.clone(), var.as_tensor().clone());
}
drop(vars);
let n_tensors = tensors.len();
let header_meta = match metadata {
Some(m) => {
let json = serde_json::to_string(m).map_err(|e| {
CsmError::Other(anyhow::anyhow!("serialize {LORA_METADATA_KEY}: {e}"))
})?;
let mut map = HashMap::new();
map.insert(LORA_METADATA_KEY.to_string(), json);
Some(map)
}
None => None,
};
safetensors::serialize_to_file(&tensors, header_meta.clone(), out.as_ref())
.map_err(|e| CsmError::Other(anyhow::anyhow!("safetensors save: {e}")))?;
tracing::info!(
"saved {n_tensors} LoRA tensors → {}{}",
out.as_ref().display(),
if header_meta.is_some() {
" (with metadata)"
} else {
""
}
);
Ok(())
}
/// Load LoRA adapter weights from a safetensors file, copying into the
/// matching Vars in the VarMap. Each Var must already exist (i.e. the model
/// must have been constructed with `add_lora_to_backbone` before calling).
pub fn load_lora_adapter<P: AsRef<Path>>(
vm: &VarMap,
path: P,
device: &candle_core::Device,
) -> Result<()> {
let tensors = candle_core::safetensors::load(path.as_ref(), device)?;
let vars = vm.data().lock().unwrap();
let mut count = 0usize;
for (name, t) in tensors.iter() {
if let Some(var) = vars.get(name) {
var.set(t)?;
count += 1;
} else {
tracing::warn!("safetensors has tensor `{name}` but VarMap has no matching Var");
}
}
tracing::info!(
"loaded {count} LoRA tensors from {}",
path.as_ref().display()
);
Ok(())
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn metadata_is_extended_picks_up_extended_target_modules() {
let cfg = crate::lora::LoraConfig::extended();
let m = LoraAdapterMetadata::from_lora_config(&cfg);
assert!(m.is_extended());
}
#[test]
fn metadata_is_extended_false_for_classic_qv() {
let cfg = crate::lora::LoraConfig::default();
let m = LoraAdapterMetadata::from_lora_config(&cfg);
assert!(!m.is_extended());
}
#[test]
fn metadata_round_trips_through_json() {
let m = LoraAdapterMetadata {
rank: 16,
alpha: 32.0,
target_modules: vec!["q_proj".into(), "v_proj".into()],
crate_version: Some("9.9.9".into()),
};
let s = serde_json::to_string(&m).unwrap();
let back: LoraAdapterMetadata = serde_json::from_str(&s).unwrap();
assert_eq!(back.rank, 16);
assert_eq!(back.alpha, 32.0);
assert_eq!(back.target_modules.len(), 2);
assert!(!back.is_extended());
}
}