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