1use candle::{DType, Device, Module, Result, Tensor, D};
9use candle_nn::{linear_b, linear_no_bias, rms_norm, Activation, Linear, RmsNorm, VarBuilder};
10use std::sync::Arc;
11
12#[derive(Debug, Clone, serde::Deserialize)]
13pub struct Config {
14 pub vocab_size: usize,
15 pub hidden_size: usize,
16 pub intermediate_size: usize,
17 pub attention_bias: bool,
18 pub num_hidden_layers: usize,
19 pub num_attention_heads: usize,
20 pub num_key_value_heads: usize,
21 pub rms_norm_eps: f64,
22 pub hidden_act: candle_nn::Activation,
23 pub max_position_embeddings: usize,
24 pub rope_theta: f64,
25 pub tie_word_embeddings: bool,
26 pub clip_qkv: Option<f64>,
27}
28
29#[derive(Debug, Clone)]
30struct RotaryEmbedding {
31 sin: Tensor,
32 cos: Tensor,
33}
34
35impl RotaryEmbedding {
36 fn new(dtype: DType, cfg: &Config, dev: &Device) -> Result<Self> {
37 let dim = cfg.hidden_size / cfg.num_attention_heads;
38 let max_seq_len = cfg.max_position_embeddings;
39 let inv_freq: Vec<_> = (0..dim)
40 .step_by(2)
41 .map(|i| 1f32 / cfg.rope_theta.powf(i as f64 / dim as f64) as f32)
42 .collect();
43 let inv_freq_len = inv_freq.len();
44 let inv_freq = Tensor::from_vec(inv_freq, (1, inv_freq_len), dev)?.to_dtype(dtype)?;
45 let t = Tensor::arange(0u32, max_seq_len as u32, dev)?
46 .to_dtype(dtype)?
47 .reshape((max_seq_len, 1))?;
48 let freqs = t.matmul(&inv_freq)?;
49 Ok(Self {
50 sin: freqs.sin()?,
51 cos: freqs.cos()?,
52 })
53 }
54
55 fn apply_rotary_emb_qkv(
56 &self,
57 q: &Tensor,
58 k: &Tensor,
59 seqlen_offset: usize,
60 ) -> Result<(Tensor, Tensor)> {
61 let (_b_sz, _h, seq_len, _n_embd) = q.dims4()?;
62 let cos = self.cos.narrow(0, seqlen_offset, seq_len)?;
63 let sin = self.sin.narrow(0, seqlen_offset, seq_len)?;
64 let q_embed = candle_nn::rotary_emb::rope(&q.contiguous()?, &cos, &sin)?;
65 let k_embed = candle_nn::rotary_emb::rope(&k.contiguous()?, &cos, &sin)?;
66 Ok((q_embed, k_embed))
67 }
68}
69
70#[derive(Debug, Clone)]
71#[allow(clippy::upper_case_acronyms)]
72struct MLP {
73 gate_proj: Linear,
74 up_proj: Linear,
75 down_proj: Linear,
76 act_fn: Activation,
77}
78
79impl MLP {
80 fn new(cfg: &Config, vb: VarBuilder) -> Result<Self> {
81 let hidden_sz = cfg.hidden_size;
82 let intermediate_sz = cfg.intermediate_size;
83 let gate_proj = linear_no_bias(hidden_sz, intermediate_sz, vb.pp("gate_proj"))?;
84 let up_proj = linear_no_bias(hidden_sz, intermediate_sz, vb.pp("up_proj"))?;
85 let down_proj = linear_no_bias(intermediate_sz, hidden_sz, vb.pp("down_proj"))?;
86 Ok(Self {
87 gate_proj,
88 up_proj,
89 down_proj,
90 act_fn: cfg.hidden_act,
91 })
92 }
93}
94
95impl Module for MLP {
96 fn forward(&self, xs: &Tensor) -> Result<Tensor> {
97 let lhs = xs.apply(&self.gate_proj)?.apply(&self.act_fn)?;
98 let rhs = xs.apply(&self.up_proj)?;
99 (lhs * rhs)?.apply(&self.down_proj)
100 }
101}
102
103#[derive(Debug, Clone)]
104struct Attention {
105 q_proj: Linear,
106 k_proj: Linear,
107 v_proj: Linear,
108 o_proj: Linear,
109 q_norm: RmsNorm,
110 k_norm: RmsNorm,
111 num_heads: usize,
112 num_kv_heads: usize,
113 num_kv_groups: usize,
114 head_dim: usize,
115 hidden_size: usize,
116 rotary_emb: Arc<RotaryEmbedding>,
117 kv_cache: Option<(Tensor, Tensor)>,
118}
119
120impl Attention {
121 fn new(rotary_emb: Arc<RotaryEmbedding>, cfg: &Config, vb: VarBuilder) -> Result<Self> {
122 let hidden_sz = cfg.hidden_size;
123 let num_heads = cfg.num_attention_heads;
124 let num_kv_heads = cfg.num_key_value_heads;
125 let num_kv_groups = num_heads / num_kv_heads;
126 let head_dim = hidden_sz / num_heads;
127 let b = cfg.attention_bias;
128 let q_proj = linear_b(hidden_sz, num_heads * head_dim, b, vb.pp("q_proj"))?;
129 let k_proj = linear_b(hidden_sz, num_kv_heads * head_dim, b, vb.pp("k_proj"))?;
130 let v_proj = linear_b(hidden_sz, num_kv_heads * head_dim, b, vb.pp("v_proj"))?;
131 let o_proj = linear_b(num_heads * head_dim, hidden_sz, b, vb.pp("o_proj"))?;
132 let q_norm = rms_norm(hidden_sz, cfg.rms_norm_eps, vb.pp("q_norm"))?;
133 let k_norm = rms_norm(num_kv_heads * head_dim, cfg.rms_norm_eps, vb.pp("k_norm"))?;
134 Ok(Self {
135 q_proj,
136 k_proj,
137 v_proj,
138 o_proj,
139 q_norm,
140 k_norm,
141 num_heads,
142 num_kv_heads,
143 num_kv_groups,
144 head_dim,
145 hidden_size: hidden_sz,
146 rotary_emb,
147 kv_cache: None,
148 })
149 }
150
151 fn forward(
152 &mut self,
153 xs: &Tensor,
154 attention_mask: Option<&Tensor>,
155 seqlen_offset: usize,
156 ) -> Result<Tensor> {
157 let (b_sz, q_len, _) = xs.dims3()?;
158
159 let query_states = self.q_proj.forward(xs)?;
160 let key_states = self.k_proj.forward(xs)?;
161 let value_states = self.v_proj.forward(xs)?;
162
163 let query_states = self.q_norm.forward(&query_states)?;
164 let key_states = self.k_norm.forward(&key_states)?;
165
166 let query_states = query_states
167 .reshape((b_sz, q_len, self.num_heads, self.head_dim))?
168 .transpose(1, 2)?;
169 let key_states = key_states
170 .reshape((b_sz, q_len, self.num_kv_heads, self.head_dim))?
171 .transpose(1, 2)?;
172 let value_states = value_states
173 .reshape((b_sz, q_len, self.num_kv_heads, self.head_dim))?
174 .transpose(1, 2)?;
175
176 let (query_states, key_states) =
177 self.rotary_emb
178 .apply_rotary_emb_qkv(&query_states, &key_states, seqlen_offset)?;
179
180 let (key_states, value_states) = match &self.kv_cache {
181 None => (key_states, value_states),
182 Some((prev_k, prev_v)) => {
183 let key_states = Tensor::cat(&[prev_k, &key_states], 2)?;
184 let value_states = Tensor::cat(&[prev_v, &value_states], 2)?;
185 (key_states, value_states)
186 }
187 };
188 self.kv_cache = Some((key_states.clone(), value_states.clone()));
189
190 let key_states = crate::utils::repeat_kv(key_states, self.num_kv_groups)?.contiguous()?;
191 let value_states =
192 crate::utils::repeat_kv(value_states, self.num_kv_groups)?.contiguous()?;
193
194 let attn_output = {
195 let scale = 1f64 / f64::sqrt(self.head_dim as f64);
196 let attn_weights = (query_states.matmul(&key_states.transpose(2, 3)?)? * scale)?;
197
198 let attn_weights = match attention_mask {
199 None => attn_weights,
200 Some(mask) => attn_weights.broadcast_add(mask)?,
201 };
202 let attn_weights = candle_nn::ops::softmax_last_dim(&attn_weights)?;
203 attn_weights.matmul(&value_states)?
204 };
205 attn_output
206 .transpose(1, 2)?
207 .reshape((b_sz, q_len, self.hidden_size))?
208 .apply(&self.o_proj)
209 }
210
211 fn clear_kv_cache(&mut self) {
212 self.kv_cache = None
213 }
214}
215
216#[derive(Debug, Clone)]
217struct DecoderLayer {
218 self_attn: Attention,
219 mlp: MLP,
220 post_attention_layernorm: RmsNorm,
221 post_feedforward_layernorm: RmsNorm,
222}
223
224impl DecoderLayer {
225 fn new(rotary_emb: Arc<RotaryEmbedding>, cfg: &Config, vb: VarBuilder) -> Result<Self> {
226 let self_attn = Attention::new(rotary_emb, cfg, vb.pp("self_attn"))?;
227 let mlp = MLP::new(cfg, vb.pp("mlp"))?;
228 let post_feedforward_layernorm = rms_norm(
229 cfg.hidden_size,
230 cfg.rms_norm_eps,
231 vb.pp("post_feedforward_layernorm"),
232 )?;
233 let post_attention_layernorm = rms_norm(
234 cfg.hidden_size,
235 cfg.rms_norm_eps,
236 vb.pp("post_attention_layernorm"),
237 )?;
238 Ok(Self {
239 self_attn,
240 mlp,
241 post_attention_layernorm,
242 post_feedforward_layernorm,
243 })
244 }
245
246 fn forward(
247 &mut self,
248 xs: &Tensor,
249 attention_mask: Option<&Tensor>,
250 seqlen_offset: usize,
251 ) -> Result<Tensor> {
252 let residual = xs;
253 let xs = self.self_attn.forward(xs, attention_mask, seqlen_offset)?;
254 let xs = self.post_attention_layernorm.forward(&xs)?;
255 let xs = (xs + residual)?;
256 let residual = &xs;
257 let xs = self.mlp.forward(&xs)?;
258 let xs = self.post_feedforward_layernorm.forward(&xs)?;
259 residual + xs
260 }
261
262 fn clear_kv_cache(&mut self) {
263 self.self_attn.clear_kv_cache()
264 }
265}
266
267#[derive(Debug, Clone)]
268pub struct Model {
269 embed_tokens: candle_nn::Embedding,
270 layers: Vec<DecoderLayer>,
271 norm: RmsNorm,
272 lm_head: Linear,
273 device: Device,
274 dtype: DType,
275}
276
277impl Model {
278 pub fn new(cfg: &Config, vb: VarBuilder) -> Result<Self> {
279 let vb_m = vb.pp("model");
280 let embed_tokens =
281 candle_nn::embedding(cfg.vocab_size, cfg.hidden_size, vb_m.pp("embed_tokens"))?;
282 let rotary_emb = Arc::new(RotaryEmbedding::new(vb.dtype(), cfg, vb_m.device())?);
283 let mut layers = Vec::with_capacity(cfg.num_hidden_layers);
284 let vb_l = vb_m.pp("layers");
285 for layer_idx in 0..cfg.num_hidden_layers {
286 let layer = DecoderLayer::new(rotary_emb.clone(), cfg, vb_l.pp(layer_idx))?;
287 layers.push(layer)
288 }
289 let norm = rms_norm(cfg.hidden_size, cfg.rms_norm_eps, vb_m.pp("norm"))?;
290 let lm_head = if cfg.tie_word_embeddings {
291 Linear::new(embed_tokens.embeddings().clone(), None)
292 } else {
293 linear_no_bias(cfg.hidden_size, cfg.vocab_size, vb.pp("lm_head"))?
294 };
295 Ok(Self {
296 embed_tokens,
297 layers,
298 norm,
299 lm_head,
300 device: vb.device().clone(),
301 dtype: vb.dtype(),
302 })
303 }
304
305 fn prepare_decoder_attention_mask(
306 &self,
307 b_size: usize,
308 tgt_len: usize,
309 seqlen_offset: usize,
310 ) -> Result<Tensor> {
311 let mask: Vec<_> = (0..tgt_len)
313 .flat_map(|i| (0..tgt_len).map(move |j| if i < j { f32::NEG_INFINITY } else { 0. }))
314 .collect();
315 let mask = Tensor::from_slice(&mask, (tgt_len, tgt_len), &self.device)?;
316 let mask = if seqlen_offset > 0 {
317 let mask0 = Tensor::zeros((tgt_len, seqlen_offset), self.dtype, &self.device)?;
318 Tensor::cat(&[&mask0, &mask], D::Minus1)?
319 } else {
320 mask
321 };
322 mask.expand((b_size, 1, tgt_len, tgt_len + seqlen_offset))?
323 .to_dtype(self.dtype)
324 }
325
326 pub fn forward(&mut self, input_ids: &Tensor, seqlen_offset: usize) -> Result<Tensor> {
327 let (b_size, seq_len) = input_ids.dims2()?;
328 let attention_mask = if seq_len <= 1 {
329 None
330 } else {
331 let mask = self.prepare_decoder_attention_mask(b_size, seq_len, seqlen_offset)?;
332 Some(mask)
333 };
334 let mut xs = self.embed_tokens.forward(input_ids)?;
335 for layer in self.layers.iter_mut() {
336 xs = layer.forward(&xs, attention_mask.as_ref(), seqlen_offset)?
337 }
338 xs.narrow(1, seq_len - 1, 1)?
339 .apply(&self.norm)?
340 .apply(&self.lm_head)
341 }
342
343 pub fn clear_kv_cache(&mut self) {
344 for layer in self.layers.iter_mut() {
345 layer.clear_kv_cache()
346 }
347 }
348}