1use crate::models::with_tracing::{linear_b as linear, Linear};
8use candle::{DType, Device, IndexOp, Module, Result, Tensor, D};
9use candle_nn::VarBuilder;
10use serde::de::{self, Deserializer, Visitor};
11use serde::Deserialize;
12use std::fmt;
13
14#[derive(Debug, Clone)]
15pub enum EosTokenId {
16 Single(u32),
17 Multiple(Vec<u32>),
18}
19
20impl<'de> Deserialize<'de> for EosTokenId {
21 fn deserialize<D>(deserializer: D) -> std::result::Result<Self, D::Error>
22 where
23 D: Deserializer<'de>,
24 {
25 struct EosTokenIdVisitor;
26
27 impl<'de> Visitor<'de> for EosTokenIdVisitor {
28 type Value = EosTokenId;
29
30 fn expecting(&self, formatter: &mut fmt::Formatter) -> fmt::Result {
31 formatter.write_str("an integer or a list of integers")
32 }
33
34 fn visit_u64<E>(self, value: u64) -> std::result::Result<Self::Value, E>
35 where
36 E: de::Error,
37 {
38 if value <= u32::MAX as u64 {
39 Ok(EosTokenId::Single(value as u32))
40 } else {
41 Err(de::Error::custom("value too large for u32"))
42 }
43 }
44
45 fn visit_seq<A>(self, mut seq: A) -> std::result::Result<Self::Value, A::Error>
46 where
47 A: serde::de::SeqAccess<'de>,
48 {
49 let mut values = Vec::new();
50 while let Some(value) = seq.next_element::<u32>()? {
51 values.push(value);
52 }
53 Ok(EosTokenId::Multiple(values))
54 }
55 }
56
57 deserializer.deserialize_any(EosTokenIdVisitor)
58 }
59}
60
61fn default_one() -> usize {
62 1
63}
64
65#[derive(Debug, Clone, serde::Deserialize)]
66pub struct Config {
67 pub num_layers: usize,
68 pub padded_vocab_size: usize,
69 pub hidden_size: usize,
70 pub ffn_hidden_size: usize,
71 pub kv_channels: usize,
72 pub num_attention_heads: usize,
73 pub seq_length: usize,
74 pub layernorm_epsilon: f64,
75 pub rmsnorm: bool,
76 pub apply_residual_connection_post_layernorm: bool,
77 pub post_layer_norm: bool,
78 pub add_bias_linear: bool,
79 pub add_qkv_bias: bool,
80 pub bias_dropout_fusion: bool,
81 pub multi_query_attention: bool,
82 pub multi_query_group_num: usize,
83 pub apply_query_key_layer_scaling: bool,
84 pub attention_softmax_in_fp32: bool,
85 pub fp32_residual_connection: bool,
86 #[serde(default = "default_one")]
87 pub rope_ratio: usize,
88 pub eos_token_id: Option<EosTokenId>,
89}
90
91#[derive(Debug, Clone)]
92struct RotaryEmbedding {
93 cache: Tensor,
94}
95
96impl RotaryEmbedding {
97 fn new(cfg: &Config, dtype: DType, dev: &Device) -> Result<Self> {
98 let rotary_dim = cfg.kv_channels;
99 let n_elem = rotary_dim / 2;
100 let base = 10_000f64 * cfg.rope_ratio as f64;
101 let inv_freq: Vec<_> = (0..n_elem)
102 .step_by(2)
103 .map(|i| 1f32 / base.powf(i as f64 / n_elem as f64) as f32)
104 .collect();
105 let inv_freq_len = inv_freq.len();
106 let inv_freq = Tensor::from_vec(inv_freq, (1, inv_freq_len), dev)?.to_dtype(dtype)?;
107 let t = Tensor::arange(0u32, cfg.seq_length as u32, dev)?
108 .to_dtype(dtype)?
109 .reshape((cfg.seq_length, 1))?;
110 let freqs = t.matmul(&inv_freq)?;
111 let cache = Tensor::stack(&[&freqs.cos()?, &freqs.sin()?], D::Minus1)?;
112 Ok(Self { cache })
113 }
114
115 fn apply(&self, xs: &Tensor, seqlen_offset: usize) -> Result<Tensor> {
116 let (seqlen, _b, np, _hn) = xs.dims4()?;
117 let cache = self.cache.narrow(0, seqlen_offset, seqlen)?;
118 let rot_dim = cache.dim(D::Minus2)? * 2;
119 let (xs, xs_pass) = (
120 xs.narrow(D::Minus1, 0, rot_dim)?,
121 xs.narrow(D::Minus1, rot_dim, rot_dim)?,
122 );
123 let xshaped = xs.reshape((seqlen, (), np, rot_dim / 2, 2))?;
124 let cache = cache.reshape((seqlen, (), 1, rot_dim / 2, 2))?;
125 let (xshaped0, xshaped1) = (
126 xshaped.i((.., .., .., .., 0))?,
127 xshaped.i((.., .., .., .., 1))?,
128 );
129 let (cache0, cache1) = (cache.i((.., .., .., .., 0))?, cache.i((.., .., .., .., 1))?);
130 let xs_out = Tensor::stack(
131 &[
132 (xshaped0.broadcast_mul(&cache0)? - xshaped1.broadcast_mul(&cache1)?)?,
133 (xshaped1.broadcast_mul(&cache0)? + xshaped0.broadcast_mul(&cache1)?)?,
134 ],
135 D::Minus1,
136 )?;
137 let xs_out = xs_out.flatten_from(3)?;
138 Tensor::cat(&[xs_out, xs_pass], D::Minus1)
139 }
140}
141
142#[derive(Debug, Clone)]
143struct CoreAttention {
144 coeff: Option<f64>,
145 norm_factor: f64,
146 dtype: DType,
147}
148
149fn masked_fill(on_false: &Tensor, mask: &Tensor, on_true: f32, dtype: DType) -> Result<Tensor> {
150 let shape = mask.shape();
151 let on_true = Tensor::new(on_true, on_false.device())?.broadcast_as(shape.dims())?;
152 let m = mask.where_cond(&on_true.to_dtype(dtype)?, on_false)?;
153 Ok(m)
154}
155
156impl CoreAttention {
157 fn new(layer_number: usize, cfg: &Config, dtype: DType) -> Result<Self> {
158 let norm_factor = (cfg.kv_channels as f64).sqrt();
159 let (norm_factor, coeff) = if cfg.apply_query_key_layer_scaling {
160 let coeff = f64::max(1.0, layer_number as f64);
161 (norm_factor * coeff, Some(coeff))
162 } else {
163 (norm_factor, None)
164 };
165 Ok(Self {
166 coeff,
167 norm_factor,
168 dtype,
169 })
170 }
171
172 fn forward(
173 &self,
174 query_layer: &Tensor,
175 key_layer: &Tensor,
176 value_layer: &Tensor,
177 attention_mask: &Option<Tensor>,
178 ) -> Result<Tensor> {
179 let output_size = (
180 query_layer.dim(1)?, query_layer.dim(2)?, query_layer.dim(0)?, key_layer.dim(0)?, );
185 let query_layer =
186 query_layer.reshape((output_size.2, output_size.0 * output_size.1, ()))?;
187 let key_layer = key_layer.reshape((output_size.3, output_size.0 * output_size.1, ()))?;
188 let matmul_result = Tensor::matmul(
189 &query_layer.transpose(0, 1)?.contiguous()?,
190 &key_layer.transpose(0, 1)?.transpose(1, 2)?.contiguous()?,
191 )?;
192 let matmul_result = (matmul_result / self.norm_factor)?.reshape(output_size)?;
193 let matmul_result = match self.coeff {
194 None => matmul_result,
195 Some(coeff) => (matmul_result * coeff)?,
196 };
197 let attention_scores = match attention_mask {
198 Some(mask) => masked_fill(
199 &matmul_result,
200 &mask.broadcast_left((matmul_result.dim(0)?, matmul_result.dim(1)?))?,
201 f32::NEG_INFINITY,
202 self.dtype,
203 )?,
204 None => matmul_result,
205 };
206 let attention_probs = candle_nn::ops::softmax_last_dim(&attention_scores)?;
207
208 let output_size = (
209 value_layer.dim(1)?,
210 value_layer.dim(2)?,
211 query_layer.dim(0)?,
212 value_layer.dim(3)?,
213 );
214 let value_layer =
215 value_layer.reshape((value_layer.dim(0)?, output_size.0 * output_size.1, ()))?;
216 let attention_probs =
217 attention_probs.reshape((output_size.0 * output_size.1, output_size.2, ()))?;
218 let context_layer = Tensor::matmul(
219 &attention_probs.contiguous()?,
220 &value_layer.transpose(0, 1)?.contiguous()?,
221 )?;
222 let context_layer = context_layer.reshape(output_size)?;
223 let context_layer = context_layer.permute((2, 0, 1, 3))?.contiguous()?;
224 context_layer.flatten_from(D::Minus2)
225 }
226}
227
228#[derive(Debug, Clone)]
229struct SelfAttention {
230 query_key_value: Linear,
231 core_attention: CoreAttention,
232 dense: Linear,
233 multi_query_attention: bool,
234 num_attention_heads_per_partition: usize,
235 num_multi_query_groups_per_partition: usize,
236 hidden_size_per_attention_head: usize,
237 kv_cache: Option<(Tensor, Tensor)>,
238}
239
240impl SelfAttention {
241 fn new(layer_number: usize, cfg: &Config, vb: VarBuilder) -> Result<Self> {
242 let projection_size = cfg.kv_channels * cfg.num_attention_heads;
243 let hidden_size_per_attention_head = projection_size / cfg.num_attention_heads;
244 let qkv_hidden_size = if cfg.multi_query_attention {
245 projection_size + 2 * hidden_size_per_attention_head * cfg.multi_query_group_num
246 } else {
247 3 * projection_size
248 };
249 let query_key_value = linear(
250 cfg.hidden_size,
251 qkv_hidden_size,
252 cfg.add_bias_linear || cfg.add_qkv_bias,
253 vb.pp("query_key_value"),
254 )?;
255 let core_attention = CoreAttention::new(layer_number, cfg, vb.dtype())?;
256 let dense = linear(
257 cfg.hidden_size,
258 cfg.hidden_size,
259 cfg.add_bias_linear,
260 vb.pp("dense"),
261 )?;
262 Ok(Self {
263 query_key_value,
264 core_attention,
265 dense,
266 multi_query_attention: cfg.multi_query_attention,
267 num_attention_heads_per_partition: cfg.num_attention_heads,
268 num_multi_query_groups_per_partition: cfg.multi_query_group_num,
269 hidden_size_per_attention_head: cfg.kv_channels,
270 kv_cache: None,
271 })
272 }
273
274 fn reset_kv_cache(&mut self) {
275 self.kv_cache = None
276 }
277
278 fn forward(
279 &mut self,
280 xs: &Tensor,
281 attention_mask: &Option<Tensor>,
282 rotary_emb: &RotaryEmbedding,
283 ) -> Result<Tensor> {
284 let mixed_x_layer = xs.apply(&self.query_key_value)?;
285 if !self.multi_query_attention {
286 candle::bail!("only multi_query_attention=true is supported")
287 }
288 let hpa = self.hidden_size_per_attention_head;
289 let query_layer =
290 mixed_x_layer.narrow(D::Minus1, 0, self.num_attention_heads_per_partition * hpa)?;
291 let key_layer = mixed_x_layer.narrow(
292 D::Minus1,
293 self.num_attention_heads_per_partition * hpa,
294 self.num_multi_query_groups_per_partition * hpa,
295 )?;
296 let value_layer = mixed_x_layer.narrow(
297 D::Minus1,
298 self.num_attention_heads_per_partition * hpa
299 + self.num_multi_query_groups_per_partition * hpa,
300 self.num_multi_query_groups_per_partition * hpa,
301 )?;
302 let query_layer = query_layer.reshape((
303 query_layer.dim(0)?,
304 query_layer.dim(1)?,
305 self.num_attention_heads_per_partition,
306 hpa,
307 ))?;
308 let key_layer = key_layer.reshape((
309 key_layer.dim(0)?,
310 key_layer.dim(1)?,
311 self.num_multi_query_groups_per_partition,
312 hpa,
313 ))?;
314 let value_layer = value_layer.reshape((
315 value_layer.dim(0)?,
316 value_layer.dim(1)?,
317 self.num_multi_query_groups_per_partition,
318 hpa,
319 ))?;
320
321 let seqlen_offset = match &self.kv_cache {
323 None => 0,
324 Some((prev_k, _)) => prev_k.dim(0)?,
325 };
326 let query_layer = rotary_emb.apply(&query_layer, seqlen_offset)?;
327 let key_layer = rotary_emb.apply(&key_layer, seqlen_offset)?;
328
329 let (key_layer, value_layer) = match &self.kv_cache {
331 None => (key_layer, value_layer),
332 Some((prev_k, prev_v)) => {
333 let k = Tensor::cat(&[prev_k, &key_layer], 0)?;
334 let v = Tensor::cat(&[prev_v, &value_layer], 0)?;
335 (k, v)
336 }
337 };
338 self.kv_cache = Some((key_layer.clone(), value_layer.clone()));
339
340 let ratio =
342 self.num_attention_heads_per_partition / self.num_multi_query_groups_per_partition;
343 let key_layer = {
344 let (d0, d1, d2, d3) = key_layer.dims4()?;
345 key_layer
346 .unsqueeze(D::Minus2)?
347 .expand((d0, d1, d2, ratio, d3))?
348 .reshape((
349 d0,
350 d1,
351 self.num_attention_heads_per_partition,
352 self.hidden_size_per_attention_head,
353 ))?
354 };
355 let value_layer = {
356 let (d0, d1, d2, d3) = value_layer.dims4()?;
357 value_layer
358 .unsqueeze(D::Minus2)?
359 .expand((d0, d1, d2, ratio, d3))?
360 .reshape((
361 d0,
362 d1,
363 self.num_attention_heads_per_partition,
364 self.hidden_size_per_attention_head,
365 ))?
366 };
367
368 let context_layer =
369 self.core_attention
370 .forward(&query_layer, &key_layer, &value_layer, attention_mask)?;
371 let output = context_layer.apply(&self.dense)?;
372 Ok(output)
373 }
374}
375
376#[allow(clippy::upper_case_acronyms)]
377#[derive(Debug, Clone)]
378struct MLP {
379 dense_h_to_4h: Linear,
380 dense_4h_to_h: Linear,
381}
382
383impl MLP {
384 fn new(cfg: &Config, vb: VarBuilder) -> Result<Self> {
385 let dense_h_to_4h = linear(
386 cfg.hidden_size,
387 cfg.ffn_hidden_size * 2,
388 cfg.add_bias_linear,
389 vb.pp("dense_h_to_4h"),
390 )?;
391 let dense_4h_to_h = linear(
392 cfg.ffn_hidden_size,
393 cfg.hidden_size,
394 cfg.add_bias_linear,
395 vb.pp("dense_4h_to_h"),
396 )?;
397 Ok(Self {
398 dense_4h_to_h,
399 dense_h_to_4h,
400 })
401 }
402}
403
404impl Module for MLP {
405 fn forward(&self, xs: &Tensor) -> Result<Tensor> {
406 xs.apply(&self.dense_h_to_4h)?
407 .apply(&candle_nn::Activation::Swiglu)?
408 .apply(&self.dense_4h_to_h)
409 }
410}
411
412#[derive(Debug, Clone)]
413struct Block {
414 input_layernorm: candle_nn::LayerNorm,
415 self_attention: SelfAttention,
416 post_attention_layernorm: candle_nn::LayerNorm,
417 mlp: MLP,
418 apply_residual_connection_post_layernorm: bool,
419}
420
421impl Block {
422 fn new(layer_number: usize, cfg: &Config, vb: VarBuilder) -> Result<Self> {
423 let input_layernorm = if cfg.rmsnorm {
424 candle_nn::rms_norm(
425 cfg.hidden_size,
426 cfg.layernorm_epsilon,
427 vb.pp("input_layernorm"),
428 )?
429 .into_inner()
430 } else {
431 candle_nn::layer_norm(
432 cfg.hidden_size,
433 cfg.layernorm_epsilon,
434 vb.pp("input_layernorm"),
435 )?
436 };
437 let post_attention_layernorm = if cfg.rmsnorm {
438 candle_nn::rms_norm(
439 cfg.hidden_size,
440 cfg.layernorm_epsilon,
441 vb.pp("post_attention_layernorm"),
442 )?
443 .into_inner()
444 } else {
445 candle_nn::layer_norm(
446 cfg.hidden_size,
447 cfg.layernorm_epsilon,
448 vb.pp("post_attention_layernorm"),
449 )?
450 };
451 let self_attention = SelfAttention::new(layer_number, cfg, vb.pp("self_attention"))?;
452 let mlp = MLP::new(cfg, vb.pp("mlp"))?;
453 Ok(Self {
454 input_layernorm,
455 self_attention,
456 post_attention_layernorm,
457 mlp,
458 apply_residual_connection_post_layernorm: cfg.apply_residual_connection_post_layernorm,
459 })
460 }
461
462 fn reset_kv_cache(&mut self) {
463 self.self_attention.reset_kv_cache()
464 }
465
466 fn forward(
467 &mut self,
468 xs: &Tensor,
469 attention_mask: &Option<Tensor>,
470 rotary_emb: &RotaryEmbedding,
471 ) -> Result<Tensor> {
472 let layernorm_output = xs.apply(&self.input_layernorm)?;
473 let attention_output =
474 self.self_attention
475 .forward(&layernorm_output, attention_mask, rotary_emb)?;
476 let residual = if self.apply_residual_connection_post_layernorm {
477 &layernorm_output
478 } else {
479 xs
480 };
481 let layernorm_input = (residual + attention_output)?;
482 let layernorm_output = layernorm_input.apply(&self.post_attention_layernorm)?;
483 let mlp_output = layernorm_output.apply(&self.mlp)?;
484 let residual = if self.apply_residual_connection_post_layernorm {
485 &layernorm_output
486 } else {
487 &layernorm_input
488 };
489 mlp_output + residual
490 }
491}
492
493#[derive(Debug, Clone)]
494struct Transformer {
495 layers: Vec<Block>,
496 final_layernorm: Option<candle_nn::LayerNorm>,
497 rotary_emb: RotaryEmbedding,
498}
499
500impl Transformer {
501 fn new(cfg: &Config, vb: VarBuilder) -> Result<Self> {
502 let vb_l = vb.pp("layers");
503 let mut layers = Vec::with_capacity(cfg.num_layers);
504 for layer_index in 0..cfg.num_layers {
505 let block = Block::new(layer_index + 1, cfg, vb_l.pp(layer_index))?;
506 layers.push(block)
507 }
508 let final_layernorm = if cfg.post_layer_norm {
509 let ln = if cfg.rmsnorm {
510 candle_nn::rms_norm(
511 cfg.hidden_size,
512 cfg.layernorm_epsilon,
513 vb.pp("final_layernorm"),
514 )?
515 .into_inner()
516 } else {
517 candle_nn::layer_norm(
518 cfg.hidden_size,
519 cfg.layernorm_epsilon,
520 vb.pp("final_layernorm"),
521 )?
522 };
523 Some(ln)
524 } else {
525 None
526 };
527 let rotary_emb = RotaryEmbedding::new(cfg, vb.dtype(), vb.device())?;
528 Ok(Self {
529 layers,
530 final_layernorm,
531 rotary_emb,
532 })
533 }
534
535 fn reset_kv_cache(&mut self) {
536 for block in self.layers.iter_mut() {
537 block.reset_kv_cache()
538 }
539 }
540
541 fn forward(&mut self, xs: &Tensor, attention_mask: &Option<Tensor>) -> Result<Tensor> {
542 let mut xs = xs.clone();
543 for block in self.layers.iter_mut() {
544 xs = block.forward(&xs, attention_mask, &self.rotary_emb)?
545 }
546 match self.final_layernorm.as_ref() {
547 None => Ok(xs),
548 Some(ln) => xs.apply(ln),
549 }
550 }
551}
552
553#[derive(Debug, Clone)]
554struct Embedding {
555 word_embeddings: candle_nn::Embedding,
556 fp32_residual_connection: bool,
557}
558
559impl Embedding {
560 fn new(cfg: &Config, vb: VarBuilder) -> Result<Self> {
561 let word_embeddings = candle_nn::embedding(
562 cfg.padded_vocab_size,
563 cfg.hidden_size,
564 vb.pp("word_embeddings"),
565 )?;
566 Ok(Self {
567 word_embeddings,
568 fp32_residual_connection: cfg.fp32_residual_connection,
569 })
570 }
571}
572
573impl Module for Embedding {
574 fn forward(&self, xs: &Tensor) -> Result<Tensor> {
575 let xs = self.word_embeddings.forward(xs)?.transpose(0, 1)?; if self.fp32_residual_connection {
577 xs.to_dtype(candle::DType::F32)
578 } else {
579 xs.contiguous()
580 }
581 }
582}
583
584#[derive(Debug, Clone)]
585pub struct Model {
586 embedding: Embedding,
587 encoder: Transformer,
588 output_layer: Linear,
589}
590
591fn get_mask(size: usize, device: &Device) -> Result<Tensor> {
592 let mask: Vec<_> = (0..size)
593 .flat_map(|i| (0..size).map(move |j| u8::from(j > i)))
594 .collect();
595 Tensor::from_slice(&mask, (size, size), device)
596}
597
598impl Model {
599 pub fn new(cfg: &Config, vb: VarBuilder) -> Result<Self> {
600 let vb = vb.pp("transformer");
601 let embedding = Embedding::new(cfg, vb.pp("embedding"))?;
602 let encoder = Transformer::new(cfg, vb.pp("encoder"))?;
603 let output_layer = linear(
604 cfg.hidden_size,
605 cfg.padded_vocab_size,
606 false,
607 vb.pp("output_layer"),
608 )?;
609
610 Ok(Self {
611 embedding,
612 encoder,
613 output_layer,
614 })
615 }
616
617 pub fn reset_kv_cache(&mut self) {
618 self.encoder.reset_kv_cache()
619 }
620
621 pub fn forward(&mut self, xs: &Tensor) -> Result<Tensor> {
622 let (_b_size, seq_len) = xs.dims2()?;
623 let input_embeds = xs.apply(&self.embedding)?;
624 let attention_mask = if seq_len <= 1 {
625 None
626 } else {
627 Some(get_mask(seq_len, xs.device())?)
628 };
629 let xs = self.encoder.forward(&input_embeds, &attention_mask)?;
630 let lm_logits = xs.i(seq_len - 1)?.apply(&self.output_layer)?;
631 Ok(lm_logits)
632 }
633}