1use super::with_tracing::{linear_no_bias as linear, Linear, RmsNorm};
8use candle::{DType, Device, IndexOp, Result, Tensor, D};
9use candle_nn::{embedding, Embedding, Module, VarBuilder};
10use std::{collections::HashMap, f32::consts::PI};
11
12pub const DEFAULT_MAX_SEQ_LEN: usize = 4096;
13
14#[derive(Debug, Clone, serde::Deserialize, Default)]
15pub enum Llama3RopeType {
16 #[serde(rename = "llama3")]
17 Llama3,
18 #[default]
19 #[serde(rename = "default")]
20 Default,
21}
22
23#[derive(Debug, Clone, serde::Deserialize, Default)]
24pub struct Llama3RopeConfig {
25 pub factor: f32,
26 pub low_freq_factor: f32,
27 pub high_freq_factor: f32,
28 pub original_max_position_embeddings: usize,
29 pub rope_type: Llama3RopeType,
30}
31#[derive(Debug, Clone, serde::Deserialize)]
32#[serde(untagged)]
33pub enum LlamaEosToks {
34 Single(u32),
35 Multiple(Vec<u32>),
36}
37
38#[derive(Debug, Clone, serde::Deserialize)]
39pub struct LlamaConfig {
40 pub hidden_size: usize,
41 pub intermediate_size: usize,
42 pub vocab_size: usize,
43 pub num_hidden_layers: usize,
44 pub num_attention_heads: usize,
45 pub num_key_value_heads: Option<usize>,
46 pub rms_norm_eps: f64,
47 #[serde(default = "default_rope")]
48 pub rope_theta: f32,
49 pub bos_token_id: Option<u32>,
50 pub eos_token_id: Option<LlamaEosToks>,
51 pub rope_scaling: Option<Llama3RopeConfig>,
52 pub max_position_embeddings: usize,
53 pub tie_word_embeddings: Option<bool>,
54}
55
56impl LlamaConfig {
57 pub fn num_key_value_heads(&self) -> usize {
58 self.num_key_value_heads.unwrap_or(self.num_attention_heads)
59 }
60}
61
62fn default_rope() -> f32 {
63 10_000.0
64}
65
66impl LlamaConfig {
67 pub fn into_config(self, use_flash_attn: bool) -> Config {
68 Config {
69 hidden_size: self.hidden_size,
70 intermediate_size: self.intermediate_size,
71 vocab_size: self.vocab_size,
72 num_hidden_layers: self.num_hidden_layers,
73 num_attention_heads: self.num_attention_heads,
74 num_key_value_heads: self.num_key_value_heads(),
75 rms_norm_eps: self.rms_norm_eps,
76 rope_theta: self.rope_theta,
77 use_flash_attn,
78 bos_token_id: self.bos_token_id,
79 eos_token_id: self.eos_token_id,
80 rope_scaling: self.rope_scaling,
81 max_position_embeddings: self.max_position_embeddings,
82 tie_word_embeddings: self.tie_word_embeddings.unwrap_or(false),
83 }
84 }
85}
86
87#[derive(Debug, Clone)]
88pub struct Config {
89 pub hidden_size: usize,
90 pub intermediate_size: usize,
91 pub vocab_size: usize,
92 pub num_hidden_layers: usize,
93 pub num_attention_heads: usize,
94 pub num_key_value_heads: usize,
95 pub use_flash_attn: bool,
96 pub rms_norm_eps: f64,
97 pub rope_theta: f32,
98 pub bos_token_id: Option<u32>,
99 pub eos_token_id: Option<LlamaEosToks>,
100 pub rope_scaling: Option<Llama3RopeConfig>,
101 pub max_position_embeddings: usize,
102 pub tie_word_embeddings: bool,
103}
104
105impl Config {
106 pub fn config_7b_v1(use_flash_attn: bool) -> Self {
107 Self {
108 hidden_size: 4096,
109 intermediate_size: 11008,
110 vocab_size: 32000,
111 num_hidden_layers: 32,
112 num_attention_heads: 32,
113 num_key_value_heads: 32,
114 use_flash_attn,
115 rms_norm_eps: 1e-6,
116 rope_theta: 10_000.0,
117 bos_token_id: None,
118 eos_token_id: None,
119 rope_scaling: None,
120 max_position_embeddings: DEFAULT_MAX_SEQ_LEN,
121 tie_word_embeddings: false,
122 }
123 }
124
125 pub fn config_7b_v2(use_flash_attn: bool) -> Self {
126 Self {
127 hidden_size: 4096,
128 intermediate_size: 11008,
129 vocab_size: 32000,
130 num_hidden_layers: 32,
131 num_attention_heads: 32,
132 num_key_value_heads: 32,
133 use_flash_attn,
134 rms_norm_eps: 1e-5,
135 rope_theta: 10_000.0,
136 bos_token_id: None,
137 eos_token_id: None,
138 rope_scaling: None,
139 max_position_embeddings: DEFAULT_MAX_SEQ_LEN,
140 tie_word_embeddings: false,
141 }
142 }
143}
144
145#[derive(Debug, Clone)]
146pub struct Cache {
147 masks: HashMap<(usize, usize), Tensor>,
148 pub use_kv_cache: bool,
149 kvs: Vec<Option<(Tensor, Tensor)>>,
150 cos: Tensor,
151 sin: Tensor,
152 device: Device,
153}
154
155fn calculate_default_inv_freq(cfg: &Config) -> Vec<f32> {
156 let head_dim = cfg.hidden_size / cfg.num_attention_heads;
157 (0..head_dim)
158 .step_by(2)
159 .map(|i| 1f32 / cfg.rope_theta.powf(i as f32 / head_dim as f32))
160 .collect()
161}
162
163impl Cache {
164 pub fn new(use_kv_cache: bool, dtype: DType, config: &Config, device: &Device) -> Result<Self> {
165 let theta = match &config.rope_scaling {
167 None
168 | Some(Llama3RopeConfig {
169 rope_type: Llama3RopeType::Default,
170 ..
171 }) => calculate_default_inv_freq(config),
172 Some(rope_scaling) => {
173 let low_freq_wavelen = rope_scaling.original_max_position_embeddings as f32
174 / rope_scaling.low_freq_factor;
175 let high_freq_wavelen = rope_scaling.original_max_position_embeddings as f32
176 / rope_scaling.high_freq_factor;
177
178 calculate_default_inv_freq(config)
179 .into_iter()
180 .map(|freq| {
181 let wavelen = 2. * PI / freq;
182 if wavelen < high_freq_wavelen {
183 freq
184 } else if wavelen > low_freq_wavelen {
185 freq / rope_scaling.factor
186 } else {
187 let smooth = (rope_scaling.original_max_position_embeddings as f32
188 / wavelen
189 - rope_scaling.low_freq_factor)
190 / (rope_scaling.high_freq_factor - rope_scaling.low_freq_factor);
191 (1. - smooth) * freq / rope_scaling.factor + smooth * freq
192 }
193 })
194 .collect::<Vec<_>>()
195 }
196 };
197
198 let theta = Tensor::new(theta, device)?;
199
200 let idx_theta = Tensor::arange(0, config.max_position_embeddings as u32, device)?
201 .to_dtype(DType::F32)?
202 .reshape((config.max_position_embeddings, 1))?
203 .matmul(&theta.reshape((1, theta.elem_count()))?)?;
204 let cos = idx_theta.cos()?.to_dtype(dtype)?;
207 let sin = idx_theta.sin()?.to_dtype(dtype)?;
208 Ok(Self {
209 masks: HashMap::new(),
210 use_kv_cache,
211 kvs: vec![None; config.num_hidden_layers],
212 device: device.clone(),
213 cos,
214 sin,
215 })
216 }
217
218 fn mask(&mut self, seq_len: usize, index_pos: usize) -> Result<Tensor> {
219 let kv_len = index_pos + seq_len;
220 if let Some(mask) = self.masks.get(&(seq_len, kv_len)) {
221 Ok(mask.clone())
222 } else {
223 let mask = crate::utils::build_causal_mask(seq_len, index_pos, &self.device)?;
224 self.masks.insert((seq_len, kv_len), mask.clone());
225 Ok(mask)
226 }
227 }
228}
229
230#[derive(Debug, Clone)]
231struct CausalSelfAttention {
232 q_proj: Linear,
233 k_proj: Linear,
234 v_proj: Linear,
235 o_proj: Linear,
236 num_attention_heads: usize,
237 num_key_value_heads: usize,
238 head_dim: usize,
239 use_flash_attn: bool,
240 span: tracing::Span,
241 span_rot: tracing::Span,
242 max_position_embeddings: usize,
243}
244
245#[cfg(feature = "flash-attn")]
246fn flash_attn(
247 q: &Tensor,
248 k: &Tensor,
249 v: &Tensor,
250 softmax_scale: f32,
251 causal: bool,
252) -> Result<Tensor> {
253 candle_flash_attn::flash_attn(q, k, v, softmax_scale, causal)
254}
255
256#[cfg(not(feature = "flash-attn"))]
257fn flash_attn(_: &Tensor, _: &Tensor, _: &Tensor, _: f32, _: bool) -> Result<Tensor> {
258 unimplemented!("compile with '--features flash-attn'")
259}
260
261impl CausalSelfAttention {
262 fn apply_rotary_emb(&self, x: &Tensor, index_pos: usize, cache: &Cache) -> Result<Tensor> {
263 let _enter = self.span_rot.enter();
264 let (_b_sz, _, seq_len, _hidden_size) = x.dims4()?;
265 let cos = cache.cos.narrow(0, index_pos, seq_len)?;
266 let sin = cache.sin.narrow(0, index_pos, seq_len)?;
267 candle_nn::rotary_emb::rope(x, &cos, &sin)
268 }
269
270 fn forward(
271 &self,
272 x: &Tensor,
273 index_pos: usize,
274 block_idx: usize,
275 cache: &mut Cache,
276 ) -> Result<Tensor> {
277 let _enter = self.span.enter();
278 let (b_sz, seq_len, hidden_size) = x.dims3()?;
279 let q = self.q_proj.forward(x)?;
280 let k = self.k_proj.forward(x)?;
281 let v = self.v_proj.forward(x)?;
282
283 let q = q
284 .reshape((b_sz, seq_len, self.num_attention_heads, self.head_dim))?
285 .transpose(1, 2)?
286 .contiguous()?;
287 let k = k
288 .reshape((b_sz, seq_len, self.num_key_value_heads, self.head_dim))?
289 .transpose(1, 2)?
290 .contiguous()?;
291 let mut v = v
292 .reshape((b_sz, seq_len, self.num_key_value_heads, self.head_dim))?
293 .transpose(1, 2)?;
294
295 let q = self.apply_rotary_emb(&q, index_pos, cache)?;
296 let mut k = self.apply_rotary_emb(&k, index_pos, cache)?;
297
298 if cache.use_kv_cache {
299 if let Some((cache_k, cache_v)) = &cache.kvs[block_idx] {
300 k = Tensor::cat(&[cache_k, &k], 2)?.contiguous()?;
301 v = Tensor::cat(&[cache_v, &v], 2)?.contiguous()?;
302 let k_seq_len = k.dims()[1];
303 if k_seq_len > self.max_position_embeddings {
304 k = k
305 .narrow(
306 D::Minus1,
307 k_seq_len - self.max_position_embeddings,
308 self.max_position_embeddings,
309 )?
310 .contiguous()?
311 }
312 let v_seq_len = v.dims()[1];
313 if v_seq_len > 2 * self.max_position_embeddings {
314 v = v
315 .narrow(
316 D::Minus1,
317 v_seq_len - self.max_position_embeddings,
318 self.max_position_embeddings,
319 )?
320 .contiguous()?
321 }
322 }
323 cache.kvs[block_idx] = Some((k.clone(), v.clone()))
324 }
325
326 let k = self.repeat_kv(k)?;
327 let v = self.repeat_kv(v)?;
328
329 let y = if self.use_flash_attn {
330 let q = q.transpose(1, 2)?;
332 let k = k.transpose(1, 2)?;
333 let v = v.transpose(1, 2)?;
334 let softmax_scale = 1f32 / (self.head_dim as f32).sqrt();
335 flash_attn(&q, &k, &v, softmax_scale, seq_len > 1)?.transpose(1, 2)?
336 } else {
337 let in_dtype = q.dtype();
338 let q = q.to_dtype(DType::F32)?;
339 let k = k.to_dtype(DType::F32)?;
340 let v = v.to_dtype(DType::F32)?;
341 let att = (q.matmul(&k.t()?)? / (self.head_dim as f64).sqrt())?;
342 let att = if seq_len == 1 {
343 att
344 } else {
345 let mask = cache.mask(seq_len, index_pos)?.broadcast_as(att.shape())?;
346 masked_fill(&att, &mask, f32::NEG_INFINITY)?
347 };
348
349 let att = candle_nn::ops::softmax_last_dim(&att)?;
350 att.matmul(&v.contiguous()?)?.to_dtype(in_dtype)?
352 };
353 let y = y.transpose(1, 2)?.reshape(&[b_sz, seq_len, hidden_size])?;
354 let y = self.o_proj.forward(&y)?;
355 Ok(y)
356 }
357
358 fn repeat_kv(&self, x: Tensor) -> Result<Tensor> {
359 crate::utils::repeat_kv(x, self.num_attention_heads / self.num_key_value_heads)
360 }
361
362 fn load(vb: VarBuilder, cfg: &Config) -> Result<Self> {
363 let span = tracing::span!(tracing::Level::TRACE, "attn");
364 let span_rot = tracing::span!(tracing::Level::TRACE, "attn-rot");
365 let size_in = cfg.hidden_size;
366 let size_q = (cfg.hidden_size / cfg.num_attention_heads) * cfg.num_attention_heads;
367 let size_kv = (cfg.hidden_size / cfg.num_attention_heads) * cfg.num_key_value_heads;
368 let q_proj = linear(size_in, size_q, vb.pp("q_proj"))?;
369 let k_proj = linear(size_in, size_kv, vb.pp("k_proj"))?;
370 let v_proj = linear(size_in, size_kv, vb.pp("v_proj"))?;
371 let o_proj = linear(size_q, size_in, vb.pp("o_proj"))?;
372 Ok(Self {
373 q_proj,
374 k_proj,
375 v_proj,
376 o_proj,
377 num_attention_heads: cfg.num_attention_heads,
378 num_key_value_heads: cfg.num_key_value_heads,
379 head_dim: cfg.hidden_size / cfg.num_attention_heads,
380 use_flash_attn: cfg.use_flash_attn,
381 span,
382 span_rot,
383 max_position_embeddings: cfg.max_position_embeddings,
384 })
385 }
386}
387
388fn masked_fill(on_false: &Tensor, mask: &Tensor, on_true: f32) -> Result<Tensor> {
389 let shape = mask.shape();
390 let on_true = Tensor::new(on_true, on_false.device())?.broadcast_as(shape.dims())?;
391 let m = mask.where_cond(&on_true, on_false)?;
392 Ok(m)
393}
394
395#[derive(Debug, Clone)]
396struct Mlp {
397 c_fc1: Linear,
398 c_fc2: Linear,
399 c_proj: Linear,
400 span: tracing::Span,
401}
402
403impl Mlp {
404 fn forward(&self, x: &Tensor) -> Result<Tensor> {
405 let _enter = self.span.enter();
406 let x = (candle_nn::ops::silu(&self.c_fc1.forward(x)?)? * self.c_fc2.forward(x)?)?;
407 self.c_proj.forward(&x)
408 }
409
410 fn load(vb: VarBuilder, cfg: &Config) -> Result<Self> {
411 let span = tracing::span!(tracing::Level::TRACE, "mlp");
412 let h_size = cfg.hidden_size;
413 let i_size = cfg.intermediate_size;
414 let c_fc1 = linear(h_size, i_size, vb.pp("gate_proj"))?;
415 let c_fc2 = linear(h_size, i_size, vb.pp("up_proj"))?;
416 let c_proj = linear(i_size, h_size, vb.pp("down_proj"))?;
417 Ok(Self {
418 c_fc1,
419 c_fc2,
420 c_proj,
421 span,
422 })
423 }
424}
425
426#[derive(Debug, Clone)]
427struct Block {
428 rms_1: RmsNorm,
429 attn: CausalSelfAttention,
430 rms_2: RmsNorm,
431 mlp: Mlp,
432 span: tracing::Span,
433}
434
435impl Block {
436 fn forward(
437 &self,
438 x: &Tensor,
439 index_pos: usize,
440 block_idx: usize,
441 cache: &mut Cache,
442 ) -> Result<Tensor> {
443 let _enter = self.span.enter();
444 let residual = x;
445 let x = self.rms_1.forward(x)?;
446 let x = (self.attn.forward(&x, index_pos, block_idx, cache)? + residual)?;
447 let residual = &x;
448 let x = (self.mlp.forward(&self.rms_2.forward(&x)?)? + residual)?;
449 Ok(x)
450 }
451
452 fn load(vb: VarBuilder, cfg: &Config) -> Result<Self> {
453 let span = tracing::span!(tracing::Level::TRACE, "block");
454 let attn = CausalSelfAttention::load(vb.pp("self_attn"), cfg)?;
455 let mlp = Mlp::load(vb.pp("mlp"), cfg)?;
456 let rms_1 = RmsNorm::new(cfg.hidden_size, cfg.rms_norm_eps, vb.pp("input_layernorm"))?;
457 let rms_2 = RmsNorm::new(
458 cfg.hidden_size,
459 cfg.rms_norm_eps,
460 vb.pp("post_attention_layernorm"),
461 )?;
462 Ok(Self {
463 rms_1,
464 attn,
465 rms_2,
466 mlp,
467 span,
468 })
469 }
470}
471
472#[derive(Debug, Clone)]
473pub struct Llama {
474 wte: Embedding,
475 blocks: Vec<Block>,
476 ln_f: RmsNorm,
477 lm_head: Linear,
478}
479
480impl Llama {
481 pub fn embed(&self, x: &Tensor) -> Result<Tensor> {
483 self.wte.forward(x)
484 }
485 pub fn forward_input_embed(
487 &self,
488 input_embed: &Tensor,
489 index_pos: usize,
490 cache: &mut Cache,
491 ) -> Result<Tensor> {
492 let (_, seq_len, _) = input_embed.dims3()?;
493 let mut x = input_embed.clone();
494 for (block_idx, block) in self.blocks.iter().enumerate() {
495 x = block.forward(&x, index_pos, block_idx, cache)?;
496 }
497 let x = self.ln_f.forward(&x)?;
498 let x = x.i((.., seq_len - 1, ..))?.contiguous()?;
499 let logits = self.lm_head.forward(&x)?;
500 logits.to_dtype(DType::F32)
501 }
502
503 pub fn forward(&self, x: &Tensor, index_pos: usize, cache: &mut Cache) -> Result<Tensor> {
504 let (_b_sz, seq_len) = x.dims2()?;
505 let mut x = self.wte.forward(x)?;
506 for (block_idx, block) in self.blocks.iter().enumerate() {
507 x = block.forward(&x, index_pos, block_idx, cache)?;
508 }
509 let x = self.ln_f.forward(&x)?;
510 let x = x.i((.., seq_len - 1, ..))?.contiguous()?;
511 let logits = self.lm_head.forward(&x)?;
512 logits.to_dtype(DType::F32)
513 }
514
515 pub fn load(vb: VarBuilder, cfg: &Config) -> Result<Self> {
516 let wte = embedding(cfg.vocab_size, cfg.hidden_size, vb.pp("model.embed_tokens"))?;
517 let lm_head = if cfg.tie_word_embeddings {
518 Linear::from_weights(wte.embeddings().clone(), None)
519 } else {
520 linear(cfg.hidden_size, cfg.vocab_size, vb.pp("lm_head"))?
521 };
522 let ln_f = RmsNorm::new(cfg.hidden_size, cfg.rms_norm_eps, vb.pp("model.norm"))?;
523 let blocks: Vec<_> = (0..cfg.num_hidden_layers)
524 .map(|i| Block::load(vb.pp(format!("model.layers.{i}")), cfg).unwrap())
525 .collect();
526
527 Ok(Self {
528 wte,
529 blocks,
530 ln_f,
531 lm_head,
532 })
533 }
534}