1use crate::models::glm4::EosTokenId;
2use crate::{
3 models::with_tracing::{linear_b, linear_no_bias, Linear, RmsNorm},
4 utils::repeat_kv,
5};
6use candle::{DType, Device, IndexOp, Module, Result, Tensor, D};
7use candle_nn::{kv_cache::KvCache, Activation, VarBuilder};
8use std::sync::Arc;
9
10#[derive(Debug, Clone, serde::Deserialize)]
11pub struct Config {
12 pub vocab_size: usize,
13 pub hidden_size: usize,
14 pub intermediate_size: usize,
15 pub num_hidden_layers: usize,
16 pub num_attention_heads: usize,
17 pub head_dim: Option<usize>,
18 pub partial_rotary_factor: Option<f32>,
19 pub attention_bias: Option<bool>,
20 pub num_key_value_heads: usize,
21 pub max_position_embeddings: usize,
22 pub sliding_window: Option<usize>,
23 pub tie_word_embeddings: bool,
24 pub rope_theta: f64,
25 pub rms_norm_eps: f64,
26 pub hidden_act: Activation,
27 pub eos_token_id: Option<EosTokenId>,
28}
29
30#[derive(Debug, Clone)]
31pub(crate) struct RotaryEmbedding {
32 sin: Tensor,
33 cos: Tensor,
34 rotary_dim: usize,
35}
36
37impl RotaryEmbedding {
38 pub(crate) fn new(dtype: DType, cfg: &Config, dev: &Device) -> Result<Self> {
39 let dim = cfg
40 .head_dim
41 .unwrap_or(cfg.hidden_size / cfg.num_attention_heads);
42 let rotary_dim = if let Some(factor) = cfg.partial_rotary_factor {
43 (factor * dim as f32) as usize
44 } else {
45 dim
46 };
47 let max_seq_len = cfg.max_position_embeddings;
48 let inv_freq: Vec<_> = (0..rotary_dim)
49 .step_by(2)
50 .map(|i| 1f32 / cfg.rope_theta.powf(i as f64 / rotary_dim as f64) as f32)
51 .collect();
52 let inv_freq_len = inv_freq.len();
53 let inv_freq = Tensor::from_vec(inv_freq, (1, inv_freq_len), dev)?.to_dtype(dtype)?;
54 let t = Tensor::arange(0u32, max_seq_len as u32, dev)?
55 .to_dtype(dtype)?
56 .reshape((max_seq_len, 1))?;
57 let freqs = t.matmul(&inv_freq)?;
58 Ok(Self {
59 sin: freqs.sin()?,
60 cos: freqs.cos()?,
61 rotary_dim,
62 })
63 }
64
65 pub(crate) fn apply(&self, xs: &Tensor, offset: usize) -> Result<Tensor> {
66 let (_, _, seq_len, _) = xs.dims4()?;
67 let (s, e) = (offset, offset + seq_len);
68 let cos = self.cos.i((s..e, ..))?.contiguous()?;
69 let sin = self.sin.i((s..e, ..))?.contiguous()?;
70 let xs_rot = xs
71 .i((0, .., .., ..self.rotary_dim))?
72 .unsqueeze(0)?
73 .contiguous()?;
74 let xs_pass = xs.i((0, .., .., self.rotary_dim..))?.unsqueeze(0)?;
75 let xs_rot = candle_nn::rotary_emb::rope_i(&xs_rot, &cos, &sin).unwrap();
76 Tensor::cat(&[&xs_rot, &xs_pass], D::Minus1)?.contiguous()
77 }
78}
79
80#[derive(Debug, Clone)]
81pub(crate) struct Mlp {
82 gate_up_proj: Linear,
83 down_proj: Linear,
84 act_fn: Activation,
85}
86
87impl Mlp {
88 pub(crate) fn new(cfg: &Config, vb: VarBuilder) -> Result<Self> {
89 Ok(Self {
90 gate_up_proj: linear_no_bias(
91 cfg.hidden_size,
92 cfg.intermediate_size * 2,
93 vb.pp("gate_up_proj"),
94 )?,
95 down_proj: linear_no_bias(cfg.intermediate_size, cfg.hidden_size, vb.pp("down_proj"))?,
96 act_fn: cfg.hidden_act,
97 })
98 }
99}
100
101impl Module for Mlp {
102 fn forward(&self, x: &Tensor) -> Result<Tensor> {
103 let w = self.gate_up_proj.forward(x)?;
104 let dim = w.dims().len() - 1;
105 let gate = w.narrow(dim, 0, w.dim(dim)? / 2)?.contiguous()?;
106 let gate = gate.apply(&self.act_fn)?;
107 let up_states = w
108 .narrow(dim, w.dim(dim)? / 2, w.dim(dim)? / 2)?
109 .contiguous()?;
110 self.down_proj.forward(&(gate * up_states)?)
111 }
112}
113
114#[derive(Debug, Clone)]
115pub(crate) struct Attention {
116 q_proj: Linear,
117 k_proj: Linear,
118 v_proj: Linear,
119 o_proj: Linear,
120 num_heads: usize,
121 num_kv_heads: usize,
122 num_kv_groups: usize,
123 head_dim: usize,
124 hidden_size: usize,
125 rotary_emb: Arc<RotaryEmbedding>,
126 kv_cache: KvCache,
127}
128
129impl Attention {
130 pub(crate) fn new(
131 cfg: &Config,
132 rotary_emb: Arc<RotaryEmbedding>,
133 vb: VarBuilder,
134 ) -> Result<Self> {
135 let head_dim = cfg
136 .head_dim
137 .unwrap_or(cfg.hidden_size / cfg.num_attention_heads);
138 let num_heads = cfg.num_attention_heads;
139 let num_kv_heads = cfg.num_key_value_heads;
140 let num_kv_groups = num_heads / num_kv_heads;
141
142 let q_proj = linear_b(
143 cfg.hidden_size,
144 num_heads * head_dim,
145 cfg.attention_bias.unwrap_or(false),
146 vb.pp("q_proj"),
147 )?;
148 let k_proj = linear_b(
149 cfg.hidden_size,
150 num_kv_heads * head_dim,
151 cfg.attention_bias.unwrap_or(false),
152 vb.pp("k_proj"),
153 )?;
154 let v_proj = linear_b(
155 cfg.hidden_size,
156 num_kv_heads * head_dim,
157 cfg.attention_bias.unwrap_or(false),
158 vb.pp("v_proj"),
159 )?;
160 let o_proj = linear_b(
161 num_heads * head_dim,
162 cfg.hidden_size,
163 false,
164 vb.pp("o_proj"),
165 )?;
166
167 let hidden_size = head_dim * cfg.num_attention_heads;
169
170 let kv_cache = KvCache::new(2, 512);
173
174 Ok(Self {
175 q_proj,
176 k_proj,
177 v_proj,
178 o_proj,
179 num_heads,
180 num_kv_heads,
181 num_kv_groups,
182 head_dim,
183 hidden_size,
184 rotary_emb,
185 kv_cache,
186 })
187 }
188
189 pub(crate) fn forward(
190 &mut self,
191 x: &Tensor,
192 attn_mask: Option<&Tensor>,
193 offset: usize,
194 ) -> Result<Tensor> {
195 let (b, l, _) = x.dims3()?;
196
197 let q = self.q_proj.forward(x)?;
198 let k = self.k_proj.forward(x)?;
199 let v = self.v_proj.forward(x)?;
200
201 let q = q
202 .reshape((b, l, self.num_heads, self.head_dim))?
203 .transpose(1, 2)?;
204 let k = k
205 .reshape((b, l, self.num_kv_heads, self.head_dim))?
206 .transpose(1, 2)?;
207 let v = v
208 .reshape((b, l, self.num_kv_heads, self.head_dim))?
209 .transpose(1, 2)?;
210
211 let q = self.rotary_emb.apply(&q, offset)?;
212 let k = self.rotary_emb.apply(&k, offset)?;
213
214 let (k, v) = self.kv_cache.append(&k.contiguous()?, &v.contiguous()?)?;
215
216 let k = repeat_kv(k, self.num_kv_groups)?;
217 let v = repeat_kv(v, self.num_kv_groups)?;
218
219 let scale = 1.0 / (self.head_dim as f64).sqrt();
220 let mut scores = (q.matmul(&k.transpose(2, 3)?)? * scale)?;
221 if let Some(m) = attn_mask {
222 scores = scores.broadcast_add(m)?;
223 }
224 let probs = candle_nn::ops::softmax_last_dim(&scores)?;
225 let ctx = probs.matmul(&v)?;
226
227 ctx.transpose(1, 2)?
228 .reshape((b, l, self.hidden_size))?
229 .apply(&self.o_proj)
230 }
231
232 pub(crate) fn clear_kv_cache(&mut self) {
233 self.kv_cache.reset();
234 }
235}
236
237#[derive(Debug, Clone)]
238struct DecoderLayer {
239 self_attn: Attention,
240 mlp: Mlp,
241 input_layernorm: RmsNorm,
242 post_attention_layernorm: RmsNorm,
243 post_mlp_layernorm: RmsNorm,
244 post_self_attn_layernorm: RmsNorm,
245}
246
247impl DecoderLayer {
248 fn new(cfg: &Config, rotary: Arc<RotaryEmbedding>, vb: VarBuilder) -> Result<Self> {
249 let self_attn = Attention::new(cfg, rotary, vb.pp("self_attn"))?;
250 let mlp = Mlp::new(cfg, vb.pp("mlp"))?;
251
252 let input_layernorm =
253 RmsNorm::new(cfg.hidden_size, cfg.rms_norm_eps, vb.pp("input_layernorm"))?;
254 let post_attention_layernorm = RmsNorm::new(
255 cfg.hidden_size,
256 cfg.rms_norm_eps,
257 vb.pp("post_attention_layernorm"),
258 )?;
259 let post_self_attn_layernorm = RmsNorm::new(
260 cfg.hidden_size,
261 cfg.rms_norm_eps,
262 vb.pp("post_self_attn_layernorm"),
263 )?;
264 let post_mlp_layernorm = RmsNorm::new(
265 cfg.hidden_size,
266 cfg.rms_norm_eps,
267 vb.pp("post_mlp_layernorm"),
268 )?;
269
270 Ok(Self {
271 self_attn,
272 mlp,
273 input_layernorm,
274 post_attention_layernorm,
275 post_self_attn_layernorm,
276 post_mlp_layernorm,
277 })
278 }
279
280 fn forward(&mut self, xs: &Tensor, mask: Option<&Tensor>, offset: usize) -> Result<Tensor> {
281 let residual = xs;
282 let hidden_states = self.input_layernorm.forward(xs)?;
283 let hidden_states = self.self_attn.forward(&hidden_states, mask, offset)?;
284 let hidden_states = self.post_self_attn_layernorm.forward(&hidden_states)?;
285 let hidden_states = (residual + hidden_states)?;
286 let residual = &hidden_states;
287 let hidden_states = self.post_attention_layernorm.forward(&hidden_states)?;
288 let hidden_states = self.mlp.forward(&hidden_states)?;
289 let hidden_states = self.post_mlp_layernorm.forward(&hidden_states)?;
290 residual + hidden_states
291 }
292
293 fn clear_kv_cache(&mut self) {
294 self.self_attn.clear_kv_cache();
295 }
296}
297
298#[derive(Debug, Clone)]
299pub struct Model {
300 embed_tokens: candle_nn::Embedding,
301 layers: Vec<DecoderLayer>,
302 norm: RmsNorm,
303 device: Device,
304 dtype: DType,
305}
306
307impl Model {
308 pub fn new(cfg: &Config, vb: VarBuilder) -> Result<Self> {
309 let embed_tokens =
310 candle_nn::embedding(cfg.vocab_size, cfg.hidden_size, vb.pp("model.embed_tokens"))?;
311 let rotary = Arc::new(RotaryEmbedding::new(vb.dtype(), cfg, vb.device())?);
312 let mut layers = Vec::with_capacity(cfg.num_hidden_layers);
313 let vb_l = vb.pp("model.layers");
314 for i in 0..cfg.num_hidden_layers {
315 layers.push(DecoderLayer::new(cfg, rotary.clone(), vb_l.pp(i))?);
316 }
317 Ok(Self {
318 embed_tokens,
319 layers,
320 norm: RmsNorm::new(cfg.hidden_size, cfg.rms_norm_eps, vb.pp("model.norm"))?,
321 device: vb.device().clone(),
322 dtype: vb.dtype(),
323 })
324 }
325
326 fn clear_kv_cache(&mut self) {
327 for l in &mut self.layers {
328 l.clear_kv_cache();
329 }
330 }
331
332 fn causal_mask(
333 &self,
334 b: usize,
335 tgt: usize,
336 offset: usize,
337 sw: Option<usize>,
338 ) -> Result<Tensor> {
339 let minf = f32::NEG_INFINITY;
340 let mask: Vec<_> = (0..tgt)
341 .flat_map(|i| {
342 (0..(tgt + offset)).map(move |j| {
343 let past_ok = j <= i + offset;
344 let sw_ok = match sw {
345 Some(w) => (i + offset) as i64 - j as i64 <= w as i64,
346 None => true,
347 };
348 if past_ok && sw_ok {
349 0.
350 } else {
351 minf
352 }
353 })
354 })
355 .collect();
356 Tensor::from_slice(&mask, (b, 1, tgt, tgt + offset), &self.device)?.to_dtype(self.dtype)
357 }
358
359 pub fn forward(&mut self, input: &Tensor, offset: usize) -> Result<Tensor> {
360 let (b, l) = input.dims2()?;
361 let mut h = self.embed_tokens.forward(input)?;
362
363 let causal = if l == 1 {
364 None
365 } else {
366 Some(self.causal_mask(b, l, offset, None)?)
367 };
368
369 for layer in &mut self.layers {
370 h = layer.forward(&h, causal.as_ref(), offset)?;
371 }
372 self.norm.forward(&h)
373 }
374}
375
376#[derive(Debug, Clone)]
377pub struct ModelForCausalLM {
378 base: Model,
379 lm_head: Linear,
380}
381
382impl ModelForCausalLM {
383 pub fn new(cfg: &Config, vb: VarBuilder) -> Result<Self> {
384 let base = Model::new(cfg, vb.clone())?;
385 let lm_head = if cfg.tie_word_embeddings {
386 Linear::from_weights(base.embed_tokens.embeddings().clone(), None)
387 } else {
388 linear_no_bias(cfg.hidden_size, cfg.vocab_size, vb.pp("lm_head"))?
389 };
390 Ok(Self { base, lm_head })
391 }
392
393 pub fn forward(&mut self, input: &Tensor, offset: usize) -> Result<Tensor> {
394 let (_, l) = input.dims2()?;
395 self.base
396 .forward(input, offset)?
397 .narrow(1, l - 1, 1)?
398 .apply(&self.lm_head)
399 }
400
401 pub fn clear_kv_cache(&mut self) {
402 self.base.clear_kv_cache();
403 }
404}