//! Vendored fork of `candle_transformers::models::csm` (candle 0.9.2). //! //! Why fork: the upstream `Model::generate_frame` runs sampling internally on //! a single conditional pass — we need finer-grained access to: //! - the c0 logits BEFORE sampling (for Classifier-Free Guidance, where we //! combine logits from a conditioned and an unconditioned forward pass) //! - the swap of `Linear` for `QMatMul` on the backbone's quantizable //! projections (Item 10 of the optimization roadmap) //! //! The base behavior is preserved bit-for-bit; new features are gated behind //! optional parameters so existing callers see no change. //! //! Original source: candle-transformers 0.9.2, //! `huggingface/candle/candle-transformers/src/models/csm.rs`. //! Apache-2.0 license preserved per upstream. use candle_core::{D, DType, Device, IndexOp, Module, Result, Tensor}; use candle_nn::{Embedding, Linear, RmsNorm, VarBuilder, embedding, linear_b}; use candle_transformers::generation::LogitsProcessor; use std::sync::Arc; #[derive(serde::Deserialize, Debug, Clone, Copy, PartialEq, Eq)] pub enum Flavor { #[serde(rename = "llama-1B")] Llama1B, #[serde(rename = "llama-100M")] Llama100M, } #[derive(serde::Deserialize, Debug, Clone)] pub struct Config { pub audio_num_codebooks: usize, pub audio_vocab_size: usize, pub backbone_flavor: Flavor, pub decoder_flavor: Flavor, pub text_vocab_size: usize, } #[allow(unused)] #[derive(Debug, Clone)] pub struct LlamaConfig { vocab_size: usize, num_layers: usize, num_heads: usize, num_kv_heads: usize, embed_dim: usize, max_seq_len: usize, intermediate_dim: usize, norm_eps: f64, rope_base: f32, scale_factor: usize, } impl LlamaConfig { pub fn from_flavor(flavor: Flavor) -> Self { match flavor { Flavor::Llama1B => Self { vocab_size: 128256, num_layers: 16, num_heads: 32, num_kv_heads: 8, embed_dim: 2048, max_seq_len: 2048, intermediate_dim: 8192, norm_eps: 1e-5, rope_base: 500_000., scale_factor: 32, }, Flavor::Llama100M => Self { vocab_size: 128256, num_layers: 4, num_heads: 8, num_kv_heads: 2, embed_dim: 1024, max_seq_len: 2048, intermediate_dim: 8192, norm_eps: 1e-5, rope_base: 500_000., scale_factor: 32, }, } } } #[derive(Debug, Clone)] struct RotaryEmbedding { sin: Tensor, cos: Tensor, } fn calculate_default_inv_freq(cfg: &LlamaConfig) -> Vec { let head_dim = cfg.embed_dim / cfg.num_heads; (0..head_dim) .step_by(2) .map(|i| 1f32 / cfg.rope_base.powf(i as f32 / head_dim as f32)) .collect() } impl RotaryEmbedding { fn new(dtype: DType, cfg: &LlamaConfig, dev: &Device) -> Result { let low_freq_factor = 1.0; let high_freq_factor = 4.0; let original_max_position_embeddings = 8192; let scale_factor = cfg.scale_factor as f32; let theta = { let low_freq_wavelen = original_max_position_embeddings as f32 / low_freq_factor; let high_freq_wavelen = original_max_position_embeddings as f32 / high_freq_factor; calculate_default_inv_freq(cfg) .into_iter() .map(|freq| { let wavelen = 2. * std::f32::consts::PI / freq; if wavelen < high_freq_wavelen { freq } else if wavelen > low_freq_wavelen { freq / scale_factor } else { let smooth = (original_max_position_embeddings as f32 / wavelen - low_freq_factor) / (high_freq_factor - low_freq_factor); (1. - smooth) * freq / scale_factor + smooth * freq } }) .collect::>() }; let theta = Tensor::new(theta, dev)?; let idx_theta = Tensor::arange(0, cfg.max_seq_len as u32, dev)? .to_dtype(DType::F32)? .reshape((cfg.max_seq_len, 1))? .matmul(&theta.reshape((1, theta.elem_count()))?)?; // This is different from the paper, see: // https://github.com/huggingface/transformers/blob/6112b1c6442aaf7affd2b0676a1cd4eee30c45cf/src/transformers/models/llama/modeling_llama.py#L112 let cos = idx_theta.cos()?.to_dtype(dtype)?; let sin = idx_theta.sin()?.to_dtype(dtype)?; Ok(Self { cos, sin }) } fn apply_rotary_emb_qkv( &self, q: &Tensor, k: &Tensor, seqlen_offset: usize, ) -> Result<(Tensor, Tensor)> { let (_b_sz, _h, seq_len, _n_embd) = q.dims4()?; let cos = self.cos.narrow(0, seqlen_offset, seq_len)?; let sin = self.sin.narrow(0, seqlen_offset, seq_len)?; let q_embed = candle_nn::rotary_emb::rope_i(q, &cos, &sin)?; let k_embed = candle_nn::rotary_emb::rope_i(k, &cos, &sin)?; Ok((q_embed, k_embed)) } } fn rms_norm(hidden_size: usize, eps: f64, vb: VarBuilder) -> Result { let weight = vb.get((hidden_size,), "scale")?; Ok(RmsNorm::new(weight, eps)) } #[derive(Clone, Debug)] pub(crate) struct Attention { q_proj: Linear, k_proj: Linear, v_proj: Linear, o_proj: Linear, /// Optional LoRA adapters on the four attention projections. None means /// inference-only; Some adds an additive delta path. Inactive (B=0) at /// init so behavior is identical to base. `o_lora` corresponds to the /// `output_proj` weight in safetensors (the field is named `o_proj` here /// for brevity but the upstream key is `output_proj`). pub(crate) q_lora: Option, pub(crate) k_lora: Option, pub(crate) v_lora: Option, pub(crate) o_lora: Option, rotary_emb: Arc, kv_cache: Option<(Tensor, Tensor)>, num_heads: usize, head_dim: usize, num_kv_heads: usize, num_kv_groups: usize, } impl Attention { fn new(cfg: &LlamaConfig, rotary_emb: Arc, vb: VarBuilder) -> Result { let head_dim = cfg.embed_dim / cfg.num_heads; let kv_dim = cfg.num_kv_heads * head_dim; let q_proj = linear_b(cfg.embed_dim, cfg.embed_dim, false, vb.pp("q_proj"))?; let k_proj = linear_b(cfg.embed_dim, kv_dim, false, vb.pp("k_proj"))?; let v_proj = linear_b(cfg.embed_dim, kv_dim, false, vb.pp("v_proj"))?; let o_proj = linear_b(cfg.embed_dim, cfg.embed_dim, false, vb.pp("output_proj"))?; Ok(Self { q_proj, k_proj, v_proj, o_proj, q_lora: None, k_lora: None, v_lora: None, o_lora: None, rotary_emb, kv_cache: None, num_heads: cfg.num_heads, num_kv_heads: cfg.num_kv_heads, num_kv_groups: cfg.num_heads / cfg.num_kv_heads, head_dim, }) } fn forward( &mut self, xs: &Tensor, attention_mask: Option<&Tensor>, seqlen_offset: usize, ) -> Result { let (b_sz, q_len, _) = xs.dims3()?; let query_states = self.q_proj.forward(xs)?; let query_states = match &self.q_lora { Some(l) => { let l_out = l.forward(xs)?; if std::env::var("CSM_LORA_DEBUG").is_ok() { eprintln!( "LoRA: a.is_var={} a.track_op={} l_out.track_op={} q_proj_out.track_op={}", l.a.is_variable(), l.a.track_op(), l_out.track_op(), query_states.track_op(), ); } let combined = (query_states + l_out)?; if std::env::var("CSM_LORA_DEBUG").is_ok() { eprintln!("LoRA: combined.track_op={}", combined.track_op()); } combined } None => query_states, }; let key_states = self.k_proj.forward(xs)?; let key_states = match &self.k_lora { Some(l) => (key_states + l.forward(xs)?)?, None => key_states, }; let value_states = self.v_proj.forward(xs)?; let value_states = match &self.v_lora { Some(l) => (value_states + l.forward(xs)?)?, None => value_states, }; let query_states = query_states .reshape((b_sz, q_len, self.num_heads, self.head_dim))? .transpose(1, 2)? .contiguous()?; let key_states = key_states .reshape((b_sz, q_len, self.num_kv_heads, self.head_dim))? .transpose(1, 2)? .contiguous()?; let value_states = value_states .reshape((b_sz, q_len, self.num_kv_heads, self.head_dim))? .transpose(1, 2)? .contiguous()?; let (query_states, key_states) = self.rotary_emb .apply_rotary_emb_qkv(&query_states, &key_states, seqlen_offset)?; let (key_states, value_states) = match &self.kv_cache { None => (key_states, value_states), Some((prev_k, prev_v)) => { let key_states = Tensor::cat(&[prev_k, &key_states], 2)?; let value_states = Tensor::cat(&[prev_v, &value_states], 2)?; (key_states, value_states) } }; self.kv_cache = Some((key_states.clone(), value_states.clone())); let key_states = candle_transformers::utils::repeat_kv(key_states, self.num_kv_groups)?; let value_states = candle_transformers::utils::repeat_kv(value_states, self.num_kv_groups)?; let attn_output = { let scale = 1f64 / f64::sqrt(self.head_dim as f64); let attn_weights = (query_states.matmul(&key_states.transpose(2, 3)?)? * scale)?; let attn_weights = match attention_mask { None => attn_weights, Some(mask) => attn_weights.broadcast_add(mask)?, }; let attn_weights = candle_nn::ops::softmax_last_dim(&attn_weights)?; attn_weights.matmul(&value_states)? }; let pre_o = attn_output .transpose(1, 2)? .reshape((b_sz, q_len, self.num_heads * self.head_dim))?; let out = self.o_proj.forward(&pre_o)?; let out = match &self.o_lora { Some(l) => (out + l.forward(&pre_o)?)?, None => out, }; if std::env::var("CSM_LORA_DEBUG").is_ok() { eprintln!("Attn::forward returns: track_op={}", out.track_op()); } Ok(out) } fn clear_kv_cache(&mut self) { self.kv_cache = None } } #[derive(Debug, Clone)] pub(crate) struct Mlp { w1: Linear, w2: Linear, w3: Linear, /// Optional LoRA adapters on the SwiGLU MLP projections. CSM uses Llama's /// `w1` / `w2` / `w3` naming where `w1` is the gate, `w3` is the up, and /// `w2` is the down projection. Each is independently optional; B=0 init /// makes them no-ops until trained. pub(crate) gate_lora: Option, pub(crate) up_lora: Option, pub(crate) down_lora: Option, } impl Mlp { fn new(cfg: &LlamaConfig, vb: VarBuilder) -> Result { let w1 = linear_b(cfg.embed_dim, cfg.intermediate_dim, false, vb.pp("w1"))?; let w2 = linear_b(cfg.intermediate_dim, cfg.embed_dim, false, vb.pp("w2"))?; let w3 = linear_b(cfg.embed_dim, cfg.intermediate_dim, false, vb.pp("w3"))?; Ok(Self { w1, w2, w3, gate_lora: None, up_lora: None, down_lora: None, }) } } impl Module for Mlp { fn forward(&self, xs: &Tensor) -> Result { let gate = self.w1.forward(xs)?; let gate = match &self.gate_lora { Some(l) => (gate + l.forward(xs)?)?, None => gate, }; let up = self.w3.forward(xs)?; let up = match &self.up_lora { Some(l) => (up + l.forward(xs)?)?, None => up, }; let mid = (gate.silu()? * up)?; let down = self.w2.forward(&mid)?; match &self.down_lora { Some(l) => down + l.forward(&mid)?, None => Ok(down), } } } #[derive(Debug, Clone)] pub(crate) struct Layer { mlp_norm: RmsNorm, sa_norm: RmsNorm, pub(crate) attn: Attention, pub(crate) mlp: Mlp, } impl Layer { fn new(cfg: &LlamaConfig, rotary_emb: Arc, vb: VarBuilder) -> Result { let mlp_norm = rms_norm(cfg.embed_dim, cfg.norm_eps, vb.pp("mlp_norm"))?; let sa_norm = rms_norm(cfg.embed_dim, cfg.norm_eps, vb.pp("sa_norm"))?; let attn = Attention::new(cfg, rotary_emb, vb.pp("attn"))?; let mlp = Mlp::new(cfg, vb.pp("mlp"))?; Ok(Self { mlp_norm, sa_norm, attn, mlp, }) } fn forward( &mut self, xs: &Tensor, attention_mask: Option<&Tensor>, seqlen_offset: usize, ) -> Result { let residual = xs; let xs = self.sa_norm.forward(xs)?; let xs = self.attn.forward(&xs, attention_mask, seqlen_offset)?; if std::env::var("CSM_LORA_DEBUG").is_ok() { eprintln!("Layer: post-attn track_op={}", xs.track_op()); } let xs = (xs + residual)?; if std::env::var("CSM_LORA_DEBUG").is_ok() { eprintln!("Layer: post-residual track_op={}", xs.track_op()); } let residual = &xs; let xs = xs.apply(&self.mlp_norm)?.apply(&self.mlp)?; if std::env::var("CSM_LORA_DEBUG").is_ok() { eprintln!("Layer: post-mlp track_op={}", xs.track_op()); } let out = (residual + xs)?; if std::env::var("CSM_LORA_DEBUG").is_ok() { eprintln!("Layer::forward returns: track_op={}", out.track_op()); } Ok(out) } fn clear_kv_cache(&mut self) { self.attn.clear_kv_cache() } } #[derive(Debug, Clone)] pub struct LlamaModel { pub(crate) layers: Vec, norm: RmsNorm, pub(crate) device: Device, pub(crate) dtype: DType, /// Optional per-layer activation-steering vectors. When set, the /// `apply()` hook runs after each Layer's forward to add a steering /// vector to the residual stream. See `crate::steering`. steering: Option, /// When `Some`, every forward pass pushes one mean-pooled-over-seq /// activation vector per layer into this buffer. Used by /// `examples/steering_extract` to derive ActAdd-style steering /// vectors from emotion-labeled audio. Caller calls `start_capture` /// before forward and `take_capture` after. capture: Option>, } impl LlamaModel { pub fn new(cfg: &LlamaConfig, vb: VarBuilder) -> Result { let rotary_emb = Arc::new(RotaryEmbedding::new(vb.dtype(), cfg, vb.device())?); let mut layers = Vec::with_capacity(cfg.num_layers); let vb_l = vb.pp("layers"); for layer_idx in 0..cfg.num_layers { let layer = Layer::new(cfg, rotary_emb.clone(), vb_l.pp(layer_idx))?; layers.push(layer); } let norm = rms_norm(cfg.embed_dim, cfg.norm_eps, vb.pp("norm"))?; Ok(Self { layers, norm, device: vb.device().clone(), dtype: vb.dtype(), steering: None, capture: None, }) } /// Install activation-steering vectors. Called once before generation /// starts; subsequent `forward` calls apply the steering at every /// step. Pass `None` to remove steering. pub fn set_steering(&mut self, steering: Option) { self.steering = steering; } pub fn steering(&self) -> Option<&crate::steering::LayerSteering> { self.steering.as_ref() } /// Begin capturing per-layer mean-pooled activations on the next /// `forward()` call. After `forward` returns, retrieve them via /// `take_capture`. No-op if `forward` is not called between these. pub fn start_capture(&mut self) { self.capture = Some(Vec::with_capacity(self.layers.len())); } /// Drain captured activations and disable capture. Returns one /// `(embed_dim,)` tensor per layer (in layer order) when capture was /// active and a forward pass ran, else `None`. pub fn take_capture(&mut self) -> Option> { self.capture.take() } pub fn clear_kv_cache(&mut self) { for layer in self.layers.iter_mut() { layer.clear_kv_cache() } } fn prepare_decoder_attention_mask( &self, tgt_len: usize, seqlen_offset: usize, ) -> Result { let mask: Vec<_> = (0..tgt_len) .flat_map(|i| (0..tgt_len).map(move |j| if i < j { f32::NEG_INFINITY } else { 0. })) .collect(); let mask = Tensor::from_slice(&mask, (tgt_len, tgt_len), &self.device)?; let mask = if seqlen_offset > 0 { let mask0 = Tensor::zeros((tgt_len, seqlen_offset), DType::F32, &self.device)?; Tensor::cat(&[&mask0, &mask], D::Minus1)? } else { mask }; mask.expand((1, 1, tgt_len, tgt_len + seqlen_offset))? .to_dtype(self.dtype) } pub fn forward(&mut self, xs: &Tensor, seqlen_offset: usize) -> Result { let (_b_size, seq_len, _embed_dim) = xs.dims3()?; let attention_mask = if seq_len <= 1 { None } else { let mask = self.prepare_decoder_attention_mask(seq_len, seqlen_offset)?; Some(mask) }; let mut xs = xs.clone(); for (layer_idx, layer) in self.layers.iter_mut().enumerate() { xs = layer.forward(&xs, attention_mask.as_ref(), seqlen_offset)?; if let Some(steering) = self.steering.as_ref() { xs = steering.apply(layer_idx, &xs)?; } if let Some(capture) = self.capture.as_mut() { // Mean over seq dim → (batch, embed_dim) → squeeze batch=1 → (embed_dim,). let pooled = xs.mean(1)?; let pooled = if pooled.dims().first().copied() == Some(1) { pooled.squeeze(0)? } else { // batch>1: pool over batch too so we always store (embed_dim,). pooled.mean(0)? }; capture.push(pooled); } } if std::env::var("CSM_LORA_DEBUG").is_ok() { eprintln!("LlamaModel: post-loop xs.track_op={}", xs.track_op()); } let narrowed = xs.narrow(1, seq_len - 1, 1)?; // candle's RmsNorm `Module::forward` uses a fused custom op that DROPS // the autograd chain (returns a Tensor with `op = None`). For training // (LoRA fine-tune via `forward_loss`) we need the chain preserved, so // call `forward_diff` which routes through the unfused LayerNorm-style // implementation. The numerical result is identical; only the autograd // graph differs. let ys = self.norm.forward_diff(&narrowed)?; if std::env::var("CSM_LORA_DEBUG").is_ok() { eprintln!("LlamaModel::forward returns: track_op={}", ys.track_op()); } Ok(ys) } } #[derive(Debug, Clone)] pub struct Model { backbone: LlamaModel, /// Optional second backbone instance with the same weights but an /// independent KV cache, used as the unconditional branch for /// Classifier-Free Guidance (Koel-TTS style). `None` until `enable_cfg` /// is called. Memory cost is just the additional KV cache (a few MB); /// weight tensors are Arc-shared with the conditional backbone. cfg_backbone: Option, decoder: LlamaModel, codebook0_head: Linear, audio_embeddings: Embedding, text_embeddings: Embedding, projection: Linear, audio_head: Tensor, config: Config, } impl Model { pub fn new(cfg: &Config, vb: VarBuilder) -> Result { let backbone_cfg = LlamaConfig::from_flavor(cfg.backbone_flavor); let backbone = LlamaModel::new(&backbone_cfg, vb.pp("backbone"))?; let decoder_cfg = LlamaConfig::from_flavor(cfg.decoder_flavor); let decoder = LlamaModel::new(&decoder_cfg, vb.pp("decoder"))?; let backbone_dim = backbone_cfg.embed_dim; let decoder_dim = decoder_cfg.embed_dim; let audio_embeddings = embedding( cfg.audio_vocab_size * cfg.audio_num_codebooks, backbone_dim, vb.pp("audio_embeddings"), )?; let text_embeddings = embedding(cfg.text_vocab_size, backbone_dim, vb.pp("text_embeddings"))?; let projection = linear_b(backbone_dim, decoder_dim, false, vb.pp("projection"))?; let codebook0_head = linear_b( backbone_dim, cfg.audio_vocab_size, false, vb.pp("codebook0_head"), )?; let audio_head = vb.get( ( cfg.audio_num_codebooks - 1, decoder_dim, cfg.audio_vocab_size, ), "audio_head", )?; Ok(Self { backbone, cfg_backbone: None, decoder, codebook0_head, audio_embeddings, text_embeddings, projection, audio_head, config: cfg.clone(), }) } /// Initialize the unconditional backbone for CFG. Call once after `new`, /// passing the same VarBuilder so weights resolve to the same tensors. pub fn enable_cfg(&mut self, vb: VarBuilder) -> Result<()> { let backbone_cfg = LlamaConfig::from_flavor(self.config.backbone_flavor); self.cfg_backbone = Some(LlamaModel::new(&backbone_cfg, vb.pp("backbone"))?); Ok(()) } pub fn cfg_enabled(&self) -> bool { self.cfg_backbone.is_some() } /// Inject trainable LoRA adapters on the backbone, with per-module /// targeting controlled by `cfg.target_modules`. The classic recipe is /// q+v only (rank 8, alpha 16); the extended recipe (Phase 12.1) covers /// q, k, v, output_proj plus the SwiGLU MLP (w1/w2/w3) for higher capacity /// when ~hours of training audio are available — this gives the adapter /// real prosody control rather than just a thin attention nudge. /// /// Decoder + heads stay frozen by construction (we only walk the backbone /// layers; `cfg.exclude_patterns` is the additional safety net). Adapters /// are registered into `vm` for AdamW pickup; B is zero-initialized so /// initial behavior is identical to the un-adapted base. pub fn add_lora_to_backbone( &mut self, cfg: &crate::lora::LoraConfig, vm: &candle_nn::VarMap, ) -> candle_core::Result<()> { let backbone_cfg = LlamaConfig::from_flavor(self.config.backbone_flavor); let head_dim = backbone_cfg.embed_dim / backbone_cfg.num_heads; let kv_dim = backbone_cfg.num_kv_heads * head_dim; let embed_dim = backbone_cfg.embed_dim; let inter_dim = backbone_cfg.intermediate_dim; let device = self.backbone.device.clone(); let dtype = self.backbone.dtype; // Closure to keep each per-target block readable and uniform. let make = |name: &str, in_dim: usize, out_dim: usize| -> candle_core::Result { crate::lora::LoraDelta::new( cfg.rank, cfg.alpha as f64, in_dim, out_dim, name, vm, &device, dtype, ) .map_err(|e| candle_core::Error::Msg(e.to_string())) }; let mut counts = [0usize; 7]; // q, k, v, o, w1(gate), w3(up), w2(down) for (i, layer) in self.backbone.layers.iter_mut().enumerate() { // attn.q_proj: embed_dim → embed_dim let key = format!("backbone.layers.{i}.attn.q_proj"); if cfg.matches(&format!("{key}.weight")) { layer.attn.q_lora = Some(make(&key, embed_dim, embed_dim)?); counts[0] += 1; } // attn.k_proj: embed_dim → kv_dim let key = format!("backbone.layers.{i}.attn.k_proj"); if cfg.matches(&format!("{key}.weight")) { layer.attn.k_lora = Some(make(&key, embed_dim, kv_dim)?); counts[1] += 1; } // attn.v_proj: embed_dim → kv_dim let key = format!("backbone.layers.{i}.attn.v_proj"); if cfg.matches(&format!("{key}.weight")) { layer.attn.v_lora = Some(make(&key, embed_dim, kv_dim)?); counts[2] += 1; } // attn.output_proj: embed_dim → embed_dim. Note: safetensors key is // `output_proj` (per upstream csm/torchtune); the struct field is // named `o_proj` but the LoRA prefix follows the on-disk name so // saved adapter files are inspectable. let key = format!("backbone.layers.{i}.attn.output_proj"); if cfg.matches(&format!("{key}.weight")) { layer.attn.o_lora = Some(make(&key, embed_dim, embed_dim)?); counts[3] += 1; } // mlp.w1 (SwiGLU gate): embed_dim → intermediate_dim let key = format!("backbone.layers.{i}.mlp.w1"); if cfg.matches(&format!("{key}.weight")) { layer.mlp.gate_lora = Some(make(&key, embed_dim, inter_dim)?); counts[4] += 1; } // mlp.w3 (SwiGLU up): embed_dim → intermediate_dim let key = format!("backbone.layers.{i}.mlp.w3"); if cfg.matches(&format!("{key}.weight")) { layer.mlp.up_lora = Some(make(&key, embed_dim, inter_dim)?); counts[5] += 1; } // mlp.w2 (SwiGLU down): intermediate_dim → embed_dim let key = format!("backbone.layers.{i}.mlp.w2"); if cfg.matches(&format!("{key}.weight")) { layer.mlp.down_lora = Some(make(&key, inter_dim, embed_dim)?); counts[6] += 1; } } tracing::info!( "LoRA injected: q={} k={} v={} o={} mlp_gate={} mlp_up={} mlp_down={}", counts[0], counts[1], counts[2], counts[3], counts[4], counts[5], counts[6], ); Ok(()) } /// After an optimizer step, refresh the LoRA Tensor handles inside each /// Attention / Mlp from the VarMap so the next forward pass sees the /// updated values. Required because LoraDelta holds plain Tensors (not /// Vars), and candle's optimizer mutates the underlying Var storage /// in place. pub fn refresh_lora(&mut self, vm: &candle_nn::VarMap) -> candle_core::Result<()> { for (i, layer) in self.backbone.layers.iter_mut().enumerate() { let do_refresh = |slot: Option<&mut crate::lora::LoraDelta>, key: String| -> candle_core::Result<()> { if let Some(d) = slot { d.refresh_from(vm, &key) .map_err(|e| candle_core::Error::Msg(e.to_string()))?; } Ok(()) }; do_refresh( layer.attn.q_lora.as_mut(), format!("backbone.layers.{i}.attn.q_proj"), )?; do_refresh( layer.attn.k_lora.as_mut(), format!("backbone.layers.{i}.attn.k_proj"), )?; do_refresh( layer.attn.v_lora.as_mut(), format!("backbone.layers.{i}.attn.v_proj"), )?; do_refresh( layer.attn.o_lora.as_mut(), format!("backbone.layers.{i}.attn.output_proj"), )?; do_refresh( layer.mlp.gate_lora.as_mut(), format!("backbone.layers.{i}.mlp.w1"), )?; do_refresh( layer.mlp.up_lora.as_mut(), format!("backbone.layers.{i}.mlp.w3"), )?; do_refresh( layer.mlp.down_lora.as_mut(), format!("backbone.layers.{i}.mlp.w2"), )?; } Ok(()) } pub fn clear_kv_cache(&mut self) { self.backbone.clear_kv_cache(); self.decoder.clear_kv_cache(); if let Some(b) = self.cfg_backbone.as_mut() { b.clear_kv_cache(); } } /// Install activation-steering vectors on the conditional backbone. /// The unconditional CFG backbone (when present) is intentionally left /// unmodified — steering should bias the conditional pathway, not the /// unconditional baseline that CFG subtracts. Pass `None` to remove. pub fn set_backbone_steering(&mut self, steering: Option) { self.backbone.set_steering(steering); } pub fn backbone_steering(&self) -> Option<&crate::steering::LayerSteering> { self.backbone.steering() } pub fn backbone_num_layers(&self) -> usize { self.backbone.layers.len() } /// Install activation-steering vectors on the depth decoder. The /// decoder generates the acoustic codebooks (c1..N-1) given a /// sampled c0 + the backbone hidden state — so steering here should /// affect prosody/timbre without changing word-level content. The /// decoder for CSM-1B is Llama100M (4 layers × 1024 embed_dim); /// vectors must match that shape. pub fn set_decoder_steering(&mut self, steering: Option) { self.decoder.set_steering(steering); } pub fn decoder_num_layers(&self) -> usize { self.decoder.layers.len() } pub fn decoder_embed_dim(&self) -> usize { LlamaConfig::from_flavor(self.config.decoder_flavor).embed_dim } pub fn backbone_embed_dim(&self) -> usize { LlamaConfig::from_flavor(self.config.backbone_flavor).embed_dim } /// Teacher-forced forward through the backbone over a built-prompt /// `(tokens, mask)` and return one mean-pooled-over-seq activation /// vector per layer. Used by `examples/steering_extract` to derive /// per-emotion difference-of-means steering vectors. /// /// Steering is disabled for the duration of this call so the captured /// activations are baseline (not already-steered). pub fn capture_backbone_activations( &mut self, tokens: &Tensor, tokens_mask: &Tensor, ) -> Result> { let saved_steering = self.backbone.steering.take(); self.backbone.clear_kv_cache(); self.backbone.start_capture(); let embeds = self.build_embeds(tokens, tokens_mask)?; let _h = self.backbone.forward(&embeds, 0)?; let captured = self .backbone .take_capture() .ok_or_else(|| crate::error::CsmError::Config("capture buffer empty".into()))?; self.backbone.steering = saved_steering; Ok(captured) } /// Teacher-forced single-frame capture at the depth decoder. Builds /// the same prompt as `capture_backbone_activations`, runs the /// backbone (no capture), then teacher-forces the decoder with the /// ground-truth c0 and captures one mean-pooled-over-seq activation /// per decoder layer (4 vectors @ 1024 dim for CSM-1B). /// /// `target_c0` should be the Mimi-encoded codebook-0 token for the /// audio frame the model would predict next given the prompt — i.e. /// the FIRST audio token AFTER the prompt's audio prefix. The caller /// is responsible for slicing the manifest's frame_codes accordingly. pub fn capture_decoder_activations( &mut self, tokens: &Tensor, tokens_mask: &Tensor, target_c0: u32, ) -> Result> { let saved_b_steering = self.backbone.steering.take(); let saved_d_steering = self.decoder.steering.take(); self.backbone.clear_kv_cache(); self.decoder.clear_kv_cache(); let embeds = self.build_embeds(tokens, tokens_mask)?; let h = self.backbone.forward(&embeds, 0)?; let c0_t = Tensor::from_slice(&[target_c0], (1, 1), &self.decoder.device)?; let c0_embed = self.audio_embeddings.forward(&c0_t)?; let curr_h = Tensor::cat(&[h, c0_embed], 1)?; let proj_h = curr_h.apply(&self.projection)?; self.decoder.start_capture(); let _decoder_h = self.decoder.forward(&proj_h, 0)?; let captured = self .decoder .take_capture() .ok_or_else(|| crate::error::CsmError::Config("decoder capture empty".into()))?; self.backbone.steering = saved_b_steering; self.decoder.steering = saved_d_steering; Ok(captured) } /// Build the per-frame embedding tensor `(B, S, D)` from packed token slots. /// Shared by `generate_frame` and `generate_frame_cfg`. fn build_embeds(&self, tokens: &Tensor, tokens_mask: &Tensor) -> Result { let (b_sz, seq_len, _cb_plus_one) = tokens.dims3()?; let audio_tokens = tokens.narrow(2, 0, self.config.audio_num_codebooks)?; let text_tokens = tokens.narrow(2, self.config.audio_num_codebooks, 1)?; let text_embeds = self.text_embeddings.forward(&text_tokens)?; let arange = (Tensor::arange( 0u32, self.config.audio_num_codebooks as u32, &self.decoder.device, )? * self.config.audio_vocab_size as f64)?; let audio_tokens = audio_tokens.broadcast_add(&arange.reshape((1, 1, ()))?)?; let audio_embeds = self.audio_embeddings.forward(&audio_tokens)?.reshape(( b_sz, seq_len, self.config.audio_num_codebooks, (), ))?; let embeds = Tensor::cat(&[&audio_embeds, &text_embeds], D::Minus2)?; let embeds = embeds.broadcast_mul( &tokens_mask .to_dtype(self.backbone.dtype)? .unsqueeze(D::Minus1)?, )?; embeds.sum(2) } /// Run the c1..cN-1 decoder loop using `h` (the backbone hidden state) and /// the sampled `c0`. Shared by `generate_frame` and `generate_frame_cfg`. fn run_decoder( &mut self, h: Tensor, c0_sample: u32, lp: &mut LogitsProcessor, ) -> Result> { let mut all_samples = vec![c0_sample]; let c0_sample_t = Tensor::from_slice(&[c0_sample], (1, 1), &self.decoder.device)?; let c0_embed = self.audio_embeddings.forward(&c0_sample_t)?; let mut curr_h = Tensor::cat(&[h, c0_embed], 1)?; self.decoder.clear_kv_cache(); let mut decoder_pos = 0; #[allow(clippy::needless_range_loop)] for i in 1..self.config.audio_num_codebooks { let proj_h = curr_h.apply(&self.projection)?; let decoder_h = self.decoder.forward(&proj_h, decoder_pos)?; decoder_pos += curr_h.dim(1)?; let ci_logits = decoder_h.broadcast_matmul(&self.audio_head.get(i - 1)?)?; let ci_sample = lp.sample(&ci_logits.i((0, 0))?)?; all_samples.push(ci_sample); let ci_sample_t = Tensor::from_slice( &[ci_sample + (i * self.config.audio_vocab_size) as u32], (1, 1), &self.decoder.device, )?; curr_h = self.audio_embeddings.forward(&ci_sample_t)?; } Ok(all_samples) } pub fn generate_frame( &mut self, tokens: &Tensor, tokens_mask: &Tensor, input_pos: usize, lp: &mut LogitsProcessor, ) -> Result> { let embeds = self.build_embeds(tokens, tokens_mask)?; let h = self.backbone.forward(&embeds, input_pos)?; let c0_logits = h.apply(&self.codebook0_head)?; let c0_sample = lp.sample(&c0_logits.i((0, 0))?)?; self.run_decoder(h, c0_sample, lp) } /// Teacher-forced training loss for one frame. /// /// Given input `tokens` / `tokens_mask` of shape `(1, S, cb+1)` and the /// ground-truth audio codes `target_codes` of shape `(num_codebooks,)` /// (the 32 Mimi tokens for the frame the model should predict next), /// returns the scalar mean cross-entropy across all codebooks. /// /// Decoder uses teacher forcing on the previous codebook tokens (i.e. the /// targets, not sampled predictions) so the loss for codebook i is /// independent of the model's current behavior on codebooks 0..i-1. /// This is the standard setup for AR-codec model fine-tuning (Koel-TTS, /// VoiceCraft, et al.). /// /// The full backward pass through this loss ALL parameters in the model /// will accumulate gradients — so for LoRA fine-tuning you need to wrap /// the trainable layers (e.g. backbone q/v projections) with `LoraLinear` /// before calling this. The pretrained Linear weights you don't want to /// update should be loaded as plain (non-Var) Tensors so candle's /// autograd treats them as constants. pub fn forward_loss( &mut self, tokens: &Tensor, tokens_mask: &Tensor, input_pos: usize, target_codes: &[u32], ) -> Result { if target_codes.len() != self.config.audio_num_codebooks { candle_core::bail!( "target_codes length {} != audio_num_codebooks {}", target_codes.len(), self.config.audio_num_codebooks ); } let embeds = self.build_embeds(tokens, tokens_mask)?; let h = self.backbone.forward(&embeds, input_pos)?; if std::env::var("CSM_GRAD_DEBUG").is_ok() { eprintln!( "FL: embeds.track_op={}, h.track_op={}", embeds.track_op(), h.track_op() ); } // c0 loss. cross_entropy expects F32 logits; cast if model runs in F16/BF16. let c0_logits = h.apply(&self.codebook0_head)?; // (1, 1, vocab) let c0_logits_2d = c0_logits.i((0, 0))?.unsqueeze(0)?.to_dtype(DType::F32)?; // (1, vocab) if std::env::var("CSM_GRAD_DEBUG").is_ok() { eprintln!( "FL: c0_logits.track_op={}, c0_logits_2d.track_op={}", c0_logits.track_op(), c0_logits_2d.track_op() ); } let c0_target = Tensor::from_slice(&[target_codes[0]], (1,), &self.decoder.device)?; let mut total_loss = candle_nn::loss::cross_entropy(&c0_logits_2d, &c0_target)?; if std::env::var("CSM_GRAD_DEBUG").is_ok() { eprintln!("FL: c0_loss.track_op={}", total_loss.track_op()); } // Teacher-forced decoder: feed ground-truth previous tokens. let c0_target_t = Tensor::from_slice(&[target_codes[0]], (1, 1), &self.decoder.device)?; let c0_embed = self.audio_embeddings.forward(&c0_target_t)?; let mut curr_h = Tensor::cat(&[h, c0_embed], 1)?; self.decoder.clear_kv_cache(); let mut decoder_pos = 0usize; #[allow(clippy::needless_range_loop)] for i in 1..self.config.audio_num_codebooks { let proj_h = curr_h.apply(&self.projection)?; let decoder_h = self.decoder.forward(&proj_h, decoder_pos)?; decoder_pos += curr_h.dim(1)?; let ci_logits = decoder_h.broadcast_matmul(&self.audio_head.get(i - 1)?)?; let ci_logits_2d = ci_logits.i((0, 0))?.unsqueeze(0)?.to_dtype(DType::F32)?; let ci_target = Tensor::from_slice(&[target_codes[i]], (1,), &self.decoder.device)?; let ci_loss = candle_nn::loss::cross_entropy(&ci_logits_2d, &ci_target)?; total_loss = (total_loss + ci_loss)?; // Teacher-force the next decoder input with the GT codebook id. let next_id = target_codes[i] + (i * self.config.audio_vocab_size) as u32; let next_t = Tensor::from_slice(&[next_id], (1, 1), &self.decoder.device)?; curr_h = self.audio_embeddings.forward(&next_t)?; } // Mean across codebooks. let n = self.config.audio_num_codebooks as f64; total_loss / n } /// Classifier-Free Guidance variant of `generate_frame`. /// /// Runs the backbone twice — once on the conditional prompt, once on the /// unconditional prompt (typically empty / no-context) — and combines the /// codebook-0 logits as `uncond + scale * (cond - uncond)` before sampling. /// The decoder loop for c1..cN-1 uses the conditional hidden state only. /// /// `enable_cfg(vb)` MUST be called before this — it allocates the second /// backbone instance with a separate KV cache. /// /// `cond_input_pos` and `uncond_input_pos` track each backbone's KV state /// independently. Caller is responsible for incrementing them after the /// call (they advance by `tokens.dim(1)` and `uncond_tokens.dim(1)` /// respectively). /// /// Reference: Koel-TTS (NVIDIA, arXiv 2502.05236), `cfg_scale ∈ [1.5, 3.0]`. #[allow(clippy::too_many_arguments)] pub fn generate_frame_cfg( &mut self, cond_tokens: &Tensor, cond_mask: &Tensor, cond_input_pos: usize, uncond_tokens: &Tensor, uncond_mask: &Tensor, uncond_input_pos: usize, cfg_scale: f64, lp: &mut LogitsProcessor, ) -> Result> { if self.cfg_backbone.is_none() { return Err(candle_core::Error::Msg( "generate_frame_cfg: enable_cfg(vb) must be called first".into(), )); } // Build both embeds before any mutable borrow on self. let cond_embeds = self.build_embeds(cond_tokens, cond_mask)?; let uncond_embeds = self.build_embeds(uncond_tokens, uncond_mask)?; let cond_h = self.backbone.forward(&cond_embeds, cond_input_pos)?; let uncond_h = self .cfg_backbone .as_mut() .unwrap() .forward(&uncond_embeds, uncond_input_pos)?; let cond_c0 = cond_h.apply(&self.codebook0_head)?; let uncond_c0 = uncond_h.apply(&self.codebook0_head)?; // out = uncond + scale * (cond - uncond) let diff = (cond_c0 - &uncond_c0)?; let combined = (uncond_c0 + (diff * cfg_scale)?)?; let c0_sample = lp.sample(&combined.i((0, 0))?)?; // Decoder operates on the conditional hidden state — c1..cN-1 don't // need CFG (intuitively: c0 picks the syllable, c1..N-1 colour it). self.run_decoder(cond_h, c0_sample, lp) } pub fn audio_tokens_and_mask(&self, mut frame: Vec) -> Result<(Tensor, Tensor)> { let cb = self.config.audio_num_codebooks; let device = &self.backbone.device; let mut mask = vec![1u8; cb]; mask.push(0); let mask = Tensor::from_vec(mask, (1, 1, cb + 1), device)?; frame.push(0); let tokens = Tensor::from_vec(frame, (1, 1, cb + 1), device)?; Ok((tokens, mask)) } pub fn text_tokens_and_mask(&self, ids: &[u32]) -> Result<(Tensor, Tensor)> { let cb = self.config.audio_num_codebooks; let device = &self.backbone.device; let mut tokens = vec![]; let mut mask = vec![]; for &v in ids.iter() { let mut token = vec![0; cb]; token.push(v); let token = Tensor::from_vec(token, (1, 1, cb + 1), device)?; tokens.push(token); let mut m = vec![0u8; cb]; m.push(1); let m = Tensor::from_vec(m, (1, 1, cb + 1), device)?; mask.push(m); } let tokens = Tensor::cat(&tokens, 1)?; let mask = Tensor::cat(&mask, 1)?; Ok((tokens, mask)) } }