1use 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 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}