Files
rustytorch/crates/models/rtx-csm/src/csm_fork.rs
T
2026-04-30 05:24:58 +00:00

1132 lines
44 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.
//! 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<f32> {
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<Self> {
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::<Vec<_>>()
};
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<RmsNorm> {
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<crate::lora::LoraDelta>,
pub(crate) k_lora: Option<crate::lora::LoraDelta>,
pub(crate) v_lora: Option<crate::lora::LoraDelta>,
pub(crate) o_lora: Option<crate::lora::LoraDelta>,
rotary_emb: Arc<RotaryEmbedding>,
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<RotaryEmbedding>, vb: VarBuilder) -> Result<Self> {
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<Tensor> {
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<crate::lora::LoraDelta>,
pub(crate) up_lora: Option<crate::lora::LoraDelta>,
pub(crate) down_lora: Option<crate::lora::LoraDelta>,
}
impl Mlp {
fn new(cfg: &LlamaConfig, vb: VarBuilder) -> Result<Self> {
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<Tensor> {
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<RotaryEmbedding>, vb: VarBuilder) -> Result<Self> {
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<Tensor> {
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<Layer>,
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<crate::steering::LayerSteering>,
/// 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<Vec<Tensor>>,
}
impl LlamaModel {
pub fn new(cfg: &LlamaConfig, vb: VarBuilder) -> Result<Self> {
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<crate::steering::LayerSteering>) {
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<Vec<Tensor>> {
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<Tensor> {
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<Tensor> {
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<LlamaModel>,
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<Self> {
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> {
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<crate::steering::LayerSteering>) {
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<crate::steering::LayerSteering>) {
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<Vec<Tensor>> {
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<Vec<Tensor>> {
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<Tensor> {
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<Vec<u32>> {
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<Vec<u32>> {
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<Tensor> {
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<Vec<u32>> {
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<u32>) -> 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))
}
}