Skip to main content

candle_transformers/models/
csm.rs

1//! Implementation of the Conversational Speech Model (CSM) from Sesame
2//!
3//! See: [CSM](Conversational Speech Model)
4//!
5/// CSM (Conversational Speech Model) is a speech generation model from Sesame that generates RVQ
6/// audio codes from text and audio inputs. The model architecture employs a Llama backbone and a
7/// smaller audio decoder that produces Mimi audio codes.
8///
9use crate::generation::LogitsProcessor;
10use candle::{DType, Device, IndexOp, Module, Result, Tensor, D};
11use candle_nn::{embedding, linear_b, Embedding, Linear, RmsNorm, VarBuilder};
12use std::sync::Arc;
13
14#[derive(serde::Deserialize, Debug, Clone, Copy, PartialEq, Eq)]
15pub enum Flavor {
16    #[serde(rename = "llama-1B")]
17    Llama1B,
18    #[serde(rename = "llama-100M")]
19    Llama100M,
20}
21
22#[derive(serde::Deserialize, Debug, Clone)]
23pub struct Config {
24    pub audio_num_codebooks: usize,
25    pub audio_vocab_size: usize,
26    pub backbone_flavor: Flavor,
27    pub decoder_flavor: Flavor,
28    pub text_vocab_size: usize,
29}
30
31#[allow(unused)]
32#[derive(Debug, Clone)]
33pub struct LlamaConfig {
34    vocab_size: usize,
35    num_layers: usize,
36    num_heads: usize,
37    num_kv_heads: usize,
38    embed_dim: usize,
39    max_seq_len: usize,
40    intermediate_dim: usize,
41    norm_eps: f64,
42    rope_base: f32,
43    scale_factor: usize,
44}
45
46impl LlamaConfig {
47    pub fn from_flavor(flavor: Flavor) -> Self {
48        match flavor {
49            Flavor::Llama1B => Self {
50                vocab_size: 128256,
51                num_layers: 16,
52                num_heads: 32,
53                num_kv_heads: 8,
54                embed_dim: 2048,
55                max_seq_len: 2048,
56                intermediate_dim: 8192,
57                norm_eps: 1e-5,
58                rope_base: 500_000.,
59                scale_factor: 32,
60            },
61            Flavor::Llama100M => Self {
62                vocab_size: 128256,
63                num_layers: 4,
64                num_heads: 8,
65                num_kv_heads: 2,
66                embed_dim: 1024,
67                max_seq_len: 2048,
68                intermediate_dim: 8192,
69                norm_eps: 1e-5,
70                rope_base: 500_000.,
71                scale_factor: 32,
72            },
73        }
74    }
75}
76
77#[derive(Debug, Clone)]
78struct RotaryEmbedding {
79    sin: Tensor,
80    cos: Tensor,
81}
82
83fn calculate_default_inv_freq(cfg: &LlamaConfig) -> Vec<f32> {
84    let head_dim = cfg.embed_dim / cfg.num_heads;
85    (0..head_dim)
86        .step_by(2)
87        .map(|i| 1f32 / cfg.rope_base.powf(i as f32 / head_dim as f32))
88        .collect()
89}
90
91impl RotaryEmbedding {
92    fn new(dtype: DType, cfg: &LlamaConfig, dev: &Device) -> Result<Self> {
93        let low_freq_factor = 1.0;
94        let high_freq_factor = 4.0;
95        let original_max_position_embeddings = 8192;
96        let scale_factor = cfg.scale_factor as f32;
97        let theta = {
98            let low_freq_wavelen = original_max_position_embeddings as f32 / low_freq_factor;
99            let high_freq_wavelen = original_max_position_embeddings as f32 / high_freq_factor;
100
101            calculate_default_inv_freq(cfg)
102                .into_iter()
103                .map(|freq| {
104                    let wavelen = 2. * std::f32::consts::PI / freq;
105                    if wavelen < high_freq_wavelen {
106                        freq
107                    } else if wavelen > low_freq_wavelen {
108                        freq / scale_factor
109                    } else {
110                        let smooth = (original_max_position_embeddings as f32 / wavelen
111                            - low_freq_factor)
112                            / (high_freq_factor - low_freq_factor);
113                        (1. - smooth) * freq / scale_factor + smooth * freq
114                    }
115                })
116                .collect::<Vec<_>>()
117        };
118
119        let theta = Tensor::new(theta, dev)?;
120        let idx_theta = Tensor::arange(0, cfg.max_seq_len as u32, dev)?
121            .to_dtype(DType::F32)?
122            .reshape((cfg.max_seq_len, 1))?
123            .matmul(&theta.reshape((1, theta.elem_count()))?)?;
124        // This is different from the paper, see:
125        // https://github.com/huggingface/transformers/blob/6112b1c6442aaf7affd2b0676a1cd4eee30c45cf/src/transformers/models/llama/modeling_llama.py#L112
126        let cos = idx_theta.cos()?.to_dtype(dtype)?;
127        let sin = idx_theta.sin()?.to_dtype(dtype)?;
128        Ok(Self { cos, sin })
129    }
130
131    fn apply_rotary_emb_qkv(
132        &self,
133        q: &Tensor,
134        k: &Tensor,
135        seqlen_offset: usize,
136    ) -> Result<(Tensor, Tensor)> {
137        let (_b_sz, _h, seq_len, _n_embd) = q.dims4()?;
138        let cos = self.cos.narrow(0, seqlen_offset, seq_len)?;
139        let sin = self.sin.narrow(0, seqlen_offset, seq_len)?;
140        let q_embed = candle_nn::rotary_emb::rope_i(q, &cos, &sin)?;
141        let k_embed = candle_nn::rotary_emb::rope_i(k, &cos, &sin)?;
142        Ok((q_embed, k_embed))
143    }
144}
145fn rms_norm(hidden_size: usize, eps: f64, vb: VarBuilder) -> Result<RmsNorm> {
146    let weight = vb.get((hidden_size,), "scale")?;
147    Ok(RmsNorm::new(weight, eps))
148}
149
150#[derive(Debug, Clone)]
151struct Attention {
152    q_proj: Linear,
153    k_proj: Linear,
154    v_proj: Linear,
155    o_proj: Linear,
156    rotary_emb: Arc<RotaryEmbedding>,
157    kv_cache: Option<(Tensor, Tensor)>,
158    num_heads: usize,
159    head_dim: usize,
160    num_kv_heads: usize,
161    num_kv_groups: usize,
162}
163
164impl Attention {
165    fn new(cfg: &LlamaConfig, rotary_emb: Arc<RotaryEmbedding>, vb: VarBuilder) -> Result<Self> {
166        let head_dim = cfg.embed_dim / cfg.num_heads;
167        let kv_dim = cfg.num_kv_heads * head_dim;
168
169        let q_proj = linear_b(cfg.embed_dim, cfg.embed_dim, false, vb.pp("q_proj"))?;
170        let k_proj = linear_b(cfg.embed_dim, kv_dim, false, vb.pp("k_proj"))?;
171        let v_proj = linear_b(cfg.embed_dim, kv_dim, false, vb.pp("v_proj"))?;
172        let o_proj = linear_b(cfg.embed_dim, cfg.embed_dim, false, vb.pp("output_proj"))?;
173        Ok(Self {
174            q_proj,
175            k_proj,
176            v_proj,
177            o_proj,
178            rotary_emb,
179            kv_cache: None,
180            num_heads: cfg.num_heads,
181            num_kv_heads: cfg.num_kv_heads,
182            num_kv_groups: cfg.num_heads / cfg.num_kv_heads,
183            head_dim,
184        })
185    }
186
187    fn forward(
188        &mut self,
189        xs: &Tensor,
190        attention_mask: Option<&Tensor>,
191        seqlen_offset: usize,
192    ) -> Result<Tensor> {
193        let (b_sz, q_len, _) = xs.dims3()?;
194
195        let query_states = self.q_proj.forward(xs)?;
196        let key_states = self.k_proj.forward(xs)?;
197        let value_states = self.v_proj.forward(xs)?;
198
199        let query_states = query_states
200            .reshape((b_sz, q_len, self.num_heads, self.head_dim))?
201            .transpose(1, 2)?
202            .contiguous()?;
203        let key_states = key_states
204            .reshape((b_sz, q_len, self.num_kv_heads, self.head_dim))?
205            .transpose(1, 2)?
206            .contiguous()?;
207        let value_states = value_states
208            .reshape((b_sz, q_len, self.num_kv_heads, self.head_dim))?
209            .transpose(1, 2)?
210            .contiguous()?;
211
212        let (query_states, key_states) =
213            self.rotary_emb
214                .apply_rotary_emb_qkv(&query_states, &key_states, seqlen_offset)?;
215
216        let (key_states, value_states) = match &self.kv_cache {
217            None => (key_states, value_states),
218            Some((prev_k, prev_v)) => {
219                let key_states = Tensor::cat(&[prev_k, &key_states], 2)?;
220                let value_states = Tensor::cat(&[prev_v, &value_states], 2)?;
221                (key_states, value_states)
222            }
223        };
224        self.kv_cache = Some((key_states.clone(), value_states.clone()));
225
226        let key_states = crate::utils::repeat_kv(key_states, self.num_kv_groups)?;
227        let value_states = crate::utils::repeat_kv(value_states, self.num_kv_groups)?;
228
229        let attn_output = {
230            let scale = 1f64 / f64::sqrt(self.head_dim as f64);
231            let attn_weights = (query_states.matmul(&key_states.transpose(2, 3)?)? * scale)?;
232
233            let attn_weights = match attention_mask {
234                None => attn_weights,
235                Some(mask) => attn_weights.broadcast_add(mask)?,
236            };
237            let attn_weights = candle_nn::ops::softmax_last_dim(&attn_weights)?;
238            attn_weights.matmul(&value_states)?
239        };
240        attn_output
241            .transpose(1, 2)?
242            .reshape((b_sz, q_len, self.num_heads * self.head_dim))?
243            .apply(&self.o_proj)
244    }
245
246    fn clear_kv_cache(&mut self) {
247        self.kv_cache = None
248    }
249}
250
251#[derive(Debug, Clone)]
252struct Mlp {
253    w1: Linear,
254    w2: Linear,
255    w3: Linear,
256}
257
258impl Mlp {
259    fn new(cfg: &LlamaConfig, vb: VarBuilder) -> Result<Self> {
260        let w1 = linear_b(cfg.embed_dim, cfg.intermediate_dim, false, vb.pp("w1"))?;
261        let w2 = linear_b(cfg.intermediate_dim, cfg.embed_dim, false, vb.pp("w2"))?;
262        let w3 = linear_b(cfg.embed_dim, cfg.intermediate_dim, false, vb.pp("w3"))?;
263        Ok(Self { w1, w2, w3 })
264    }
265}
266
267impl Module for Mlp {
268    fn forward(&self, xs: &Tensor) -> Result<Tensor> {
269        let lhs = xs.apply(&self.w1)?.silu()?;
270        let rhs = xs.apply(&self.w3)?;
271        (lhs * rhs)?.apply(&self.w2)
272    }
273}
274
275#[derive(Debug, Clone)]
276struct Layer {
277    mlp_norm: RmsNorm,
278    sa_norm: RmsNorm,
279    attn: Attention,
280    mlp: Mlp,
281}
282
283impl Layer {
284    fn new(cfg: &LlamaConfig, rotary_emb: Arc<RotaryEmbedding>, vb: VarBuilder) -> Result<Self> {
285        let mlp_norm = rms_norm(cfg.embed_dim, cfg.norm_eps, vb.pp("mlp_norm"))?;
286        let sa_norm = rms_norm(cfg.embed_dim, cfg.norm_eps, vb.pp("sa_norm"))?;
287        let attn = Attention::new(cfg, rotary_emb, vb.pp("attn"))?;
288        let mlp = Mlp::new(cfg, vb.pp("mlp"))?;
289        Ok(Self {
290            mlp_norm,
291            sa_norm,
292            attn,
293            mlp,
294        })
295    }
296
297    fn forward(
298        &mut self,
299        xs: &Tensor,
300        attention_mask: Option<&Tensor>,
301        seqlen_offset: usize,
302    ) -> Result<Tensor> {
303        let residual = xs;
304        let xs = self.sa_norm.forward(xs)?;
305        let xs = self.attn.forward(&xs, attention_mask, seqlen_offset)?;
306        let xs = (xs + residual)?;
307        let residual = &xs;
308        let xs = xs.apply(&self.mlp_norm)?.apply(&self.mlp)?;
309        residual + xs
310    }
311
312    fn clear_kv_cache(&mut self) {
313        self.attn.clear_kv_cache()
314    }
315}
316
317#[derive(Debug, Clone)]
318pub struct LlamaModel {
319    layers: Vec<Layer>,
320    norm: RmsNorm,
321    device: Device,
322    dtype: DType,
323}
324
325impl LlamaModel {
326    pub fn new(cfg: &LlamaConfig, vb: VarBuilder) -> Result<Self> {
327        let rotary_emb = Arc::new(RotaryEmbedding::new(vb.dtype(), cfg, vb.device())?);
328        let mut layers = Vec::with_capacity(cfg.num_layers);
329        let vb_l = vb.pp("layers");
330        for layer_idx in 0..cfg.num_layers {
331            let layer = Layer::new(cfg, rotary_emb.clone(), vb_l.pp(layer_idx))?;
332            layers.push(layer);
333        }
334        let norm = rms_norm(cfg.embed_dim, cfg.norm_eps, vb.pp("norm"))?;
335        Ok(Self {
336            layers,
337            norm,
338            device: vb.device().clone(),
339            dtype: vb.dtype(),
340        })
341    }
342
343    pub fn clear_kv_cache(&mut self) {
344        for layer in self.layers.iter_mut() {
345            layer.clear_kv_cache()
346        }
347    }
348
349    fn prepare_decoder_attention_mask(
350        &self,
351        tgt_len: usize,
352        seqlen_offset: usize,
353    ) -> Result<Tensor> {
354        let mask: Vec<_> = (0..tgt_len)
355            .flat_map(|i| (0..tgt_len).map(move |j| if i < j { f32::NEG_INFINITY } else { 0. }))
356            .collect();
357        let mask = Tensor::from_slice(&mask, (tgt_len, tgt_len), &self.device)?;
358        let mask = if seqlen_offset > 0 {
359            let mask0 = Tensor::zeros((tgt_len, seqlen_offset), DType::F32, &self.device)?;
360            Tensor::cat(&[&mask0, &mask], D::Minus1)?
361        } else {
362            mask
363        };
364        mask.expand((1, 1, tgt_len, tgt_len + seqlen_offset))?
365            .to_dtype(self.dtype)
366    }
367
368    pub fn forward(&mut self, xs: &Tensor, seqlen_offset: usize) -> Result<Tensor> {
369        let (_b_size, seq_len, _embed_dim) = xs.dims3()?;
370        let attention_mask = if seq_len <= 1 {
371            None
372        } else {
373            let mask = self.prepare_decoder_attention_mask(seq_len, seqlen_offset)?;
374            Some(mask)
375        };
376        let mut xs = xs.clone();
377        for layer in self.layers.iter_mut() {
378            xs = layer.forward(&xs, attention_mask.as_ref(), seqlen_offset)?;
379        }
380        let ys = xs.narrow(1, seq_len - 1, 1)?.apply(&self.norm)?;
381        Ok(ys)
382    }
383}
384
385#[derive(Debug, Clone)]
386pub struct Model {
387    backbone: LlamaModel,
388    decoder: LlamaModel,
389    codebook0_head: Linear,
390    audio_embeddings: Embedding,
391    text_embeddings: Embedding,
392    projection: Linear,
393    audio_head: Tensor,
394    config: Config,
395}
396
397impl Model {
398    pub fn new(cfg: &Config, vb: VarBuilder) -> Result<Self> {
399        let backbone_cfg = LlamaConfig::from_flavor(cfg.backbone_flavor);
400        let backbone = LlamaModel::new(&backbone_cfg, vb.pp("backbone"))?;
401        let decoder_cfg = LlamaConfig::from_flavor(cfg.decoder_flavor);
402        let decoder = LlamaModel::new(&decoder_cfg, vb.pp("decoder"))?;
403        let backbone_dim = backbone_cfg.embed_dim;
404        let decoder_dim = decoder_cfg.embed_dim;
405        let audio_embeddings = embedding(
406            cfg.audio_vocab_size * cfg.audio_num_codebooks,
407            backbone_dim,
408            vb.pp("audio_embeddings"),
409        )?;
410        let text_embeddings =
411            embedding(cfg.text_vocab_size, backbone_dim, vb.pp("text_embeddings"))?;
412        let projection = linear_b(backbone_dim, decoder_dim, false, vb.pp("projection"))?;
413        let codebook0_head = linear_b(
414            backbone_dim,
415            cfg.audio_vocab_size,
416            false,
417            vb.pp("codebook0_head"),
418        )?;
419        let audio_head = vb.get(
420            (
421                cfg.audio_num_codebooks - 1,
422                decoder_dim,
423                cfg.audio_vocab_size,
424            ),
425            "audio_head",
426        )?;
427        Ok(Self {
428            backbone,
429            decoder,
430            codebook0_head,
431            audio_embeddings,
432            text_embeddings,
433            projection,
434            audio_head,
435            config: cfg.clone(),
436        })
437    }
438
439    pub fn clear_kv_cache(&mut self) {
440        self.backbone.clear_kv_cache();
441        self.decoder.clear_kv_cache();
442    }
443
444    pub fn generate_frame(
445        &mut self,
446        tokens: &Tensor,
447        tokens_mask: &Tensor,
448        input_pos: usize,
449        lp: &mut LogitsProcessor,
450    ) -> Result<Vec<u32>> {
451        let (b_sz, seq_len, _cb_plus_one) = tokens.dims3()?;
452        let audio_tokens = tokens.narrow(2, 0, self.config.audio_num_codebooks)?;
453        let text_tokens = tokens.narrow(2, self.config.audio_num_codebooks, 1)?;
454        let text_embeds = self.text_embeddings.forward(&text_tokens)?;
455        let arange = (Tensor::arange(
456            0u32,
457            self.config.audio_num_codebooks as u32,
458            &self.decoder.device,
459        )? * self.config.audio_vocab_size as f64)?;
460        let audio_tokens = audio_tokens.broadcast_add(&arange.reshape((1, 1, ()))?)?;
461        let audio_embeds = self.audio_embeddings.forward(&audio_tokens)?.reshape((
462            b_sz,
463            seq_len,
464            self.config.audio_num_codebooks,
465            (),
466        ))?;
467        let embeds = Tensor::cat(&[&audio_embeds, &text_embeds], D::Minus2)?;
468        let embeds = embeds.broadcast_mul(
469            &tokens_mask
470                .to_dtype(self.backbone.dtype)?
471                .unsqueeze(D::Minus1)?,
472        )?;
473        let embeds = embeds.sum(2)?;
474        let h = self.backbone.forward(&embeds, input_pos)?;
475        let c0_logits = h.apply(&self.codebook0_head)?;
476        let c0_sample = lp.sample(&c0_logits.i((0, 0))?)?;
477        let mut all_samples = vec![c0_sample];
478        let c0_sample = Tensor::from_slice(&[c0_sample], (1, 1), &self.decoder.device)?;
479        let c0_embed = self.audio_embeddings.forward(&c0_sample)?;
480        let mut curr_h = Tensor::cat(&[h, c0_embed], 1)?;
481
482        self.decoder.clear_kv_cache();
483        let mut decoder_pos = 0;
484        for i in 1..self.config.audio_num_codebooks {
485            let proj_h = curr_h.apply(&self.projection)?;
486            let decoder_h = self.decoder.forward(&proj_h, decoder_pos)?;
487            decoder_pos += curr_h.dim(1)?;
488            let ci_logits = decoder_h.broadcast_matmul(&self.audio_head.get(i - 1)?)?;
489            let ci_sample = lp.sample(&ci_logits.i((0, 0))?)?;
490            all_samples.push(ci_sample);
491            let ci_sample = Tensor::from_slice(
492                &[ci_sample + (i * self.config.audio_vocab_size) as u32],
493                (1, 1),
494                &self.decoder.device,
495            )?;
496            let ci_embed = self.audio_embeddings.forward(&ci_sample)?;
497            curr_h = ci_embed
498        }
499        Ok(all_samples)
500    }
501
502    pub fn audio_tokens_and_mask(&self, mut frame: Vec<u32>) -> Result<(Tensor, Tensor)> {
503        let cb = self.config.audio_num_codebooks;
504        let device = &self.backbone.device;
505        let mut mask = vec![1u8; cb];
506        mask.push(0);
507        let mask = Tensor::from_vec(mask, (1, 1, cb + 1), device)?;
508
509        frame.push(0);
510        let tokens = Tensor::from_vec(frame, (1, 1, cb + 1), device)?;
511        Ok((tokens, mask))
512    }
513
514    pub fn text_tokens_and_mask(&self, ids: &[u32]) -> Result<(Tensor, Tensor)> {
515        let cb = self.config.audio_num_codebooks;
516        let device = &self.backbone.device;
517        let mut tokens = vec![];
518        let mut mask = vec![];
519        for &v in ids.iter() {
520            let mut token = vec![0; cb];
521            token.push(v);
522            let token = Tensor::from_vec(token, (1, 1, cb + 1), device)?;
523            tokens.push(token);
524            let mut m = vec![0u8; cb];
525            m.push(1);
526            let m = Tensor::from_vec(m, (1, 1, cb + 1), device)?;
527            mask.push(m);
528        }
529        let tokens = Tensor::cat(&tokens, 1)?;
530        let mask = Tensor::cat(&mask, 1)?;
531        Ok((tokens, mask))
532    }
533}