1use crate::models::with_tracing::{linear_no_bias as linear, Linear, RmsNorm};
23use candle::{DType, Device, IndexOp, Module, Result, Tensor, D};
24use candle_nn::VarBuilder;
25use std::sync::Arc;
26
27#[derive(Debug, Clone, serde::Deserialize)]
28pub enum RopeScalingType {
29 #[serde(rename = "longrope")]
30 LongRope,
31}
32
33#[derive(Debug, Clone, serde::Deserialize)]
34pub struct RopeScaling {
35 pub short_factor: Vec<f32>,
36 pub long_factor: Vec<f32>,
37 #[serde(rename = "type")]
38 pub type_: RopeScalingType,
39}
40
41#[derive(Debug, Clone, serde::Deserialize)]
43pub struct Config {
44 pub vocab_size: usize,
45 pub hidden_act: candle_nn::Activation,
46 pub hidden_size: usize,
47 pub intermediate_size: usize,
48 pub num_hidden_layers: usize,
49 pub num_attention_heads: usize,
50 pub num_key_value_heads: usize,
51 pub rms_norm_eps: f64,
52 pub rope_theta: f64,
53 pub bos_token_id: Option<u32>,
54 pub eos_token_id: Option<u32>,
55 pub rope_scaling: Option<RopeScaling>,
56 pub max_position_embeddings: usize,
57 pub original_max_position_embeddings: Option<usize>,
58 pub partial_rotary_factor: Option<f64>,
59 #[serde(default)]
60 pub tie_word_embeddings: bool,
61}
62
63impl Config {
64 pub fn head_dim(&self) -> usize {
65 self.hidden_size / self.num_attention_heads
66 }
67}
68
69#[derive(Debug, Clone)]
70pub struct RotaryEmbedding {
71 partial_dim: Option<usize>,
72 sin: Tensor,
73 cos: Tensor,
74}
75
76impl RotaryEmbedding {
77 pub fn new(dtype: DType, cfg: &Config, dev: &Device) -> Result<Self> {
78 let partial_dim = cfg
79 .partial_rotary_factor
80 .as_ref()
81 .map(|v| (v * cfg.head_dim() as f64) as usize);
82 let dim = partial_dim.unwrap_or(cfg.head_dim());
83 let freqs = match cfg.rope_scaling.as_ref() {
84 None => {
85 let max_seq_len = cfg.max_position_embeddings;
86 let inv_freq: Vec<_> = (0..dim)
87 .step_by(2)
88 .map(|i| 1f32 / cfg.rope_theta.powf(i as f64 / dim as f64) as f32)
89 .collect();
90 let inv_freq = Tensor::from_vec(inv_freq, (1, ()), dev)?.to_dtype(dtype)?;
91 let t = Tensor::arange(0u32, max_seq_len as u32, dev)?
92 .to_dtype(dtype)?
93 .reshape((max_seq_len, 1))?;
94 t.matmul(&inv_freq)?
95 }
96 Some(rope_scaling) => {
97 let inv_freq_s: Vec<_> = (0..dim)
98 .step_by(2)
99 .zip(rope_scaling.short_factor.iter())
100 .map(|(i, &f)| f / cfg.rope_theta.powf(i as f64 / dim as f64) as f32)
101 .collect();
102 let inv_freq_s = Tensor::from_vec(inv_freq_s, (1, ()), dev)?.to_dtype(dtype)?;
103 let max_seq_len = cfg.max_position_embeddings;
104 match cfg.original_max_position_embeddings {
105 None => {
106 let t = Tensor::arange(0u32, max_seq_len as u32, dev)?
107 .to_dtype(dtype)?
108 .reshape((max_seq_len, 1))?;
109 t.matmul(&inv_freq_s)?
110 }
111 Some(original_max_seq_len) => {
112 let t_s = Tensor::arange(0u32, original_max_seq_len as u32, dev)?
113 .to_dtype(dtype)?
114 .reshape((original_max_seq_len, 1))?;
115 let freq_s = t_s.matmul(&inv_freq_s)?;
116 let inv_freq_l: Vec<_> = (0..dim)
117 .step_by(2)
118 .zip(rope_scaling.long_factor.iter())
119 .map(|(i, &f)| f / cfg.rope_theta.powf(i as f64 / dim as f64) as f32)
120 .collect();
121 let inv_freq_l =
122 Tensor::from_vec(inv_freq_l, (1, ()), dev)?.to_dtype(dtype)?;
123 let t_l =
124 Tensor::arange(original_max_seq_len as u32, max_seq_len as u32, dev)?
125 .to_dtype(dtype)?
126 .reshape(((), 1))?;
127 let freq_l = t_l.matmul(&inv_freq_l)?;
128 Tensor::cat(&[&freq_s, &freq_l], 0)?
129 }
130 }
131 }
132 };
133 Ok(Self {
134 partial_dim,
135 sin: freqs.sin()?,
136 cos: freqs.cos()?,
137 })
138 }
139
140 fn rope(&self, xs: &Tensor, cos: &Tensor, sin: &Tensor) -> Result<Tensor> {
141 let x = match self.partial_dim {
142 None => candle_nn::rotary_emb::rope(&xs.contiguous()?, cos, sin)?,
143 Some(dim) => {
144 let xs_rot = xs.i((.., .., .., ..dim))?.contiguous()?;
145 let xs_pass = xs.i((.., .., .., dim..))?;
146 let xs_rot = candle_nn::rotary_emb::rope(&xs_rot, cos, sin)?;
147 Tensor::cat(&[&xs_rot, &xs_pass], D::Minus1)?.contiguous()?
148 }
149 };
150 Ok(x)
151 }
152
153 pub fn apply_rotary_emb_qkv(
154 &self,
155 q: &Tensor,
156 k: &Tensor,
157 seqlen_offset: usize,
158 ) -> Result<(Tensor, Tensor)> {
159 let (_b_sz, _h, seq_len, _n_embd) = q.dims4()?;
160 let cos = self.cos.narrow(0, seqlen_offset, seq_len)?;
161 let sin = self.sin.narrow(0, seqlen_offset, seq_len)?;
162 let q_embed = self.rope(&q.contiguous()?, &cos, &sin)?;
163 let k_embed = self.rope(&k.contiguous()?, &cos, &sin)?;
164 Ok((q_embed, k_embed))
165 }
166}
167
168#[derive(Debug, Clone)]
169struct Attention {
170 qkv_proj: Linear,
171 o_proj: Linear,
172 num_heads: usize,
173 num_kv_heads: usize,
174 num_kv_groups: usize,
175 head_dim: usize,
176 rotary_emb: Arc<RotaryEmbedding>,
177 kv_cache: Option<(Tensor, Tensor)>,
178}
179
180impl Attention {
181 fn new(rotary_emb: Arc<RotaryEmbedding>, cfg: &Config, vb: VarBuilder) -> Result<Self> {
182 let num_heads = cfg.num_attention_heads;
183 let num_kv_heads = cfg.num_key_value_heads;
184 let head_dim = cfg.head_dim();
185 let op_size = num_heads * head_dim + 2 * num_kv_heads * head_dim;
186 let qkv_proj = linear(cfg.hidden_size, op_size, vb.pp("qkv_proj"))?;
187 let o_proj = linear(num_heads * head_dim, cfg.hidden_size, vb.pp("o_proj"))?;
188 Ok(Self {
189 qkv_proj,
190 o_proj,
191 rotary_emb,
192 kv_cache: None,
193 num_heads,
194 num_kv_heads,
195 num_kv_groups: num_heads / num_kv_heads,
196 head_dim,
197 })
198 }
199
200 fn forward(
201 &mut self,
202 xs: &Tensor,
203 attention_mask: Option<&Tensor>,
204 seqlen_offset: usize,
205 ) -> Result<Tensor> {
206 let (b_sz, q_len, _) = xs.dims3()?;
207
208 let qkv = self.qkv_proj.forward(xs)?;
209 let query_pos = self.num_heads * self.head_dim;
210 let query_states = qkv.narrow(D::Minus1, 0, query_pos)?;
211 let key_states = qkv.narrow(D::Minus1, query_pos, self.num_kv_heads * self.head_dim)?;
212 let value_states = qkv.narrow(
213 D::Minus1,
214 query_pos + self.num_kv_heads * self.head_dim,
215 self.num_kv_heads * self.head_dim,
216 )?;
217
218 let query_states = query_states
219 .reshape((b_sz, q_len, self.num_heads, self.head_dim))?
220 .transpose(1, 2)?;
221 let key_states = key_states
222 .reshape((b_sz, q_len, self.num_kv_heads, self.head_dim))?
223 .transpose(1, 2)?;
224 let value_states = value_states
225 .reshape((b_sz, q_len, self.num_kv_heads, self.head_dim))?
226 .transpose(1, 2)?;
227
228 let (query_states, key_states) =
229 self.rotary_emb
230 .apply_rotary_emb_qkv(&query_states, &key_states, seqlen_offset)?;
231
232 let (key_states, value_states) = match &self.kv_cache {
233 None => (key_states, value_states),
234 Some((prev_k, prev_v)) => {
235 let key_states = Tensor::cat(&[prev_k, &key_states], 2)?;
236 let value_states = Tensor::cat(&[prev_v, &value_states], 2)?;
237 (key_states, value_states)
238 }
239 };
240 self.kv_cache = Some((key_states.clone(), value_states.clone()));
241
242 let key_states = crate::utils::repeat_kv(key_states, self.num_kv_groups)?.contiguous()?;
243 let value_states =
244 crate::utils::repeat_kv(value_states, self.num_kv_groups)?.contiguous()?;
245
246 let attn_output = {
247 let scale = 1f64 / f64::sqrt(self.head_dim as f64);
248 let attn_weights = (query_states.matmul(&key_states.transpose(2, 3)?)? * scale)?;
249
250 let attn_weights = match attention_mask {
251 None => attn_weights,
252 Some(mask) => attn_weights.broadcast_add(mask)?,
253 };
254 let attn_weights = candle_nn::ops::softmax_last_dim(&attn_weights)?;
255 attn_weights.matmul(&value_states)?
256 };
257 attn_output
258 .transpose(1, 2)?
259 .reshape((b_sz, q_len, ()))?
260 .apply(&self.o_proj)
261 }
262
263 fn clear_kv_cache(&mut self) {
264 self.kv_cache = None
265 }
266}
267
268#[derive(Debug, Clone)]
269struct Mlp {
270 gate_up_proj: Linear,
271 down_proj: Linear,
272 act_fn: candle_nn::Activation,
273 i_size: usize,
274}
275
276impl Mlp {
277 fn new(cfg: &Config, vb: VarBuilder) -> Result<Self> {
278 let hidden_size = cfg.hidden_size;
279 let i_size = cfg.intermediate_size;
280 let gate_up_proj = linear(hidden_size, 2 * i_size, vb.pp("gate_up_proj"))?;
281 let down_proj = linear(i_size, hidden_size, vb.pp("down_proj"))?;
282 Ok(Self {
283 gate_up_proj,
284 down_proj,
285 act_fn: cfg.hidden_act,
286 i_size,
287 })
288 }
289}
290
291impl Module for Mlp {
292 fn forward(&self, xs: &Tensor) -> Result<Tensor> {
293 let up_states = xs.apply(&self.gate_up_proj)?;
294 let gate = up_states.narrow(D::Minus1, 0, self.i_size)?;
295 let up_states = up_states.narrow(D::Minus1, self.i_size, self.i_size)?;
296 let up_states = (up_states * gate.apply(&self.act_fn))?;
297 up_states.apply(&self.down_proj)
298 }
299}
300
301#[derive(Debug, Clone)]
302struct DecoderLayer {
303 self_attn: Attention,
304 mlp: Mlp,
305 input_layernorm: RmsNorm,
306 post_attention_layernorm: RmsNorm,
307}
308
309impl DecoderLayer {
310 fn new(rotary_emb: Arc<RotaryEmbedding>, cfg: &Config, vb: VarBuilder) -> Result<Self> {
311 let self_attn = Attention::new(rotary_emb, cfg, vb.pp("self_attn"))?;
312 let mlp = Mlp::new(cfg, vb.pp("mlp"))?;
313 let input_layernorm =
314 RmsNorm::new(cfg.hidden_size, cfg.rms_norm_eps, vb.pp("input_layernorm"))?;
315 let post_attention_layernorm = RmsNorm::new(
316 cfg.hidden_size,
317 cfg.rms_norm_eps,
318 vb.pp("post_attention_layernorm"),
319 )?;
320 Ok(Self {
321 self_attn,
322 mlp,
323 input_layernorm,
324 post_attention_layernorm,
325 })
326 }
327
328 fn forward(
329 &mut self,
330 xs: &Tensor,
331 attention_mask: Option<&Tensor>,
332 seqlen_offset: usize,
333 ) -> Result<Tensor> {
334 let residual = xs;
335 let xs = self.input_layernorm.forward(xs)?;
336 let xs = self.self_attn.forward(&xs, attention_mask, seqlen_offset)?;
337 let xs = (xs + residual)?;
338 let residual = &xs;
339 let xs = xs.apply(&self.post_attention_layernorm)?.apply(&self.mlp)?;
340 residual + xs
341 }
342
343 fn clear_kv_cache(&mut self) {
344 self.self_attn.clear_kv_cache()
345 }
346}
347
348#[derive(Debug, Clone)]
349pub struct Model {
350 embed_tokens: candle_nn::Embedding,
351 layers: Vec<DecoderLayer>,
352 norm: RmsNorm,
353 lm_head: Linear,
354 device: Device,
355 dtype: DType,
356}
357
358impl Model {
359 pub fn new(cfg: &Config, vb: VarBuilder) -> Result<Self> {
360 let vb_m = vb.pp("model");
361 let embed_tokens =
362 candle_nn::embedding(cfg.vocab_size, cfg.hidden_size, vb_m.pp("embed_tokens"))?;
363 let rotary_emb = Arc::new(RotaryEmbedding::new(vb.dtype(), cfg, vb_m.device())?);
364 let mut layers = Vec::with_capacity(cfg.num_hidden_layers);
365 let vb_l = vb_m.pp("layers");
366 for layer_idx in 0..cfg.num_hidden_layers {
367 let layer = DecoderLayer::new(rotary_emb.clone(), cfg, vb_l.pp(layer_idx))?;
368 layers.push(layer)
369 }
370 let norm = RmsNorm::new(cfg.hidden_size, cfg.rms_norm_eps, vb_m.pp("norm"))?;
371 let lm_head = if cfg.tie_word_embeddings {
372 Linear::from_weights(embed_tokens.embeddings().clone(), None)
373 } else {
374 linear(cfg.hidden_size, cfg.vocab_size, vb.pp("lm_head"))?
375 };
376 Ok(Self {
377 embed_tokens,
378 layers,
379 norm,
380 lm_head,
381 device: vb.device().clone(),
382 dtype: vb.dtype(),
383 })
384 }
385
386 fn prepare_decoder_attention_mask(
387 &self,
388 b_size: usize,
389 tgt_len: usize,
390 seqlen_offset: usize,
391 ) -> Result<Tensor> {
392 let mask: Vec<_> = (0..tgt_len)
393 .flat_map(|i| (0..tgt_len).map(move |j| if i < j { f32::NEG_INFINITY } else { 0. }))
394 .collect();
395 let mask = Tensor::from_slice(&mask, (tgt_len, tgt_len), &self.device)?;
396 let mask = if seqlen_offset > 0 {
397 let mask0 = Tensor::zeros((tgt_len, seqlen_offset), DType::F32, &self.device)?;
398 Tensor::cat(&[&mask0, &mask], D::Minus1)?
399 } else {
400 mask
401 };
402 mask.expand((b_size, 1, tgt_len, tgt_len + seqlen_offset))?
403 .to_dtype(self.dtype)
404 }
405
406 pub fn forward(&mut self, input_ids: &Tensor, seqlen_offset: usize) -> Result<Tensor> {
407 let (b_size, seq_len) = input_ids.dims2()?;
408 let attention_mask = if seq_len <= 1 {
409 None
410 } else {
411 let mask = self.prepare_decoder_attention_mask(b_size, seq_len, seqlen_offset)?;
412 Some(mask)
413 };
414 let mut xs = self.embed_tokens.forward(input_ids)?;
415 for layer in self.layers.iter_mut() {
416 xs = layer.forward(&xs, attention_mask.as_ref(), seqlen_offset)?
417 }
418 xs.narrow(1, seq_len - 1, 1)?
419 .apply(&self.norm)?
420 .apply(&self.lm_head)
421 }
422
423 pub fn clear_kv_cache(&mut self) {
424 for layer in self.layers.iter_mut() {
425 layer.clear_kv_cache()
426 }
427 }
428}