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