Co-authored-by: Omar Sobh <[email protected]> Co-committed-by: Omar Sobh <[email protected]>
1132 lines
44 KiB
Rust
1132 lines
44 KiB
Rust
//! 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))
|
||
}
|
||
}
|