1use candle::{DType, Device, IndexOp, Result, Tensor};
21use candle_nn::{embedding, Embedding, VarBuilder};
22
23#[derive(Debug, Clone, Copy, PartialEq, Eq, serde::Deserialize)]
27pub enum ModelVersion {
28 V7,
29 V7a,
30 V7b,
31}
32
33#[derive(Debug, Clone, serde::Deserialize)]
35pub struct Config {
36 pub version: ModelVersion,
37 pub vocab_size: usize,
38 pub hidden_size: usize,
39 pub num_hidden_layers: usize,
40 #[serde(default = "default_head_size")]
41 pub head_size: usize,
42 pub intermediate_size: Option<usize>,
43 #[serde(default = "default_rescale_every")]
44 pub rescale_every: usize,
45}
46
47fn default_head_size() -> usize {
48 64
49}
50
51fn default_rescale_every() -> usize {
52 0
53}
54
55impl Config {
56 fn n_heads(&self) -> usize {
57 self.hidden_size / self.head_size
58 }
59
60 fn dim_ffn(&self) -> usize {
61 self.intermediate_size.unwrap_or(self.hidden_size * 4)
62 }
63}
64
65fn infer_lora_dims(vb: &VarBuilder) -> Result<(usize, usize, usize, usize)> {
69 let att = vb.pp("blocks").pp(0).pp("att");
70 let d_decay = att.get_unchecked("w1")?.dim(1)?;
71 let d_aaa = att.get_unchecked("a1")?.dim(1)?;
72 let d_mv = att.get_unchecked("v1")?.dim(1)?;
73 let d_gate = att.get_unchecked("g1")?.dim(1)?;
74 Ok((d_decay, d_aaa, d_mv, d_gate))
75}
76
77pub struct StatePerLayer {
81 pub att_x_prev: Tensor,
83 pub att_kv: Tensor,
85 pub ffn_x_prev: Tensor,
87}
88
89pub struct DeaState {
91 pub token_ids: Vec<u32>,
93 pub k_cache: Vec<Tensor>,
95 pub v_cache: Vec<Tensor>,
97 pub q_prev: Vec<Tensor>,
99}
100
101pub struct State {
103 pub per_layer: Vec<StatePerLayer>,
104 pub dea: Option<DeaState>,
105 pub pos: usize,
106}
107
108impl State {
109 pub fn new(cfg: &Config, dev: &Device) -> Result<Self> {
111 Self::new_with_dtype(cfg, dev, DType::F32)
112 }
113
114 pub fn new_with_dtype(cfg: &Config, dev: &Device, dtype: DType) -> Result<Self> {
119 let n_heads = cfg.n_heads();
120 let mut per_layer = Vec::with_capacity(cfg.num_hidden_layers);
121 for _layer_idx in 0..cfg.num_hidden_layers {
122 per_layer.push(StatePerLayer {
123 att_x_prev: Tensor::zeros(cfg.hidden_size, dtype, dev)?,
124 att_kv: Tensor::zeros((n_heads, cfg.head_size, cfg.head_size), DType::F32, dev)?,
126 ffn_x_prev: Tensor::zeros(cfg.hidden_size, dtype, dev)?,
127 });
128 }
129 let dea = if cfg.version == ModelVersion::V7b {
130 let mut k_cache = Vec::with_capacity(cfg.num_hidden_layers);
131 let mut v_cache = Vec::with_capacity(cfg.num_hidden_layers);
132 let mut q_prev = Vec::with_capacity(cfg.num_hidden_layers);
133 for _ in 0..cfg.num_hidden_layers {
134 k_cache.push(Tensor::zeros((0, 32), dtype, dev)?);
135 v_cache.push(Tensor::zeros((0, 32), dtype, dev)?);
136 q_prev.push(Tensor::zeros(256, dtype, dev)?);
137 }
138 Some(DeaState {
139 token_ids: Vec::new(),
140 k_cache,
141 v_cache,
142 q_prev,
143 })
144 } else {
145 None
146 };
147 Ok(Self {
148 per_layer,
149 dea,
150 pos: 0,
151 })
152 }
153}
154
155pub use crate::models::rwkv_v5::Tokenizer;
158
159fn layer_norm(xs: &Tensor, weight: &Tensor, bias: &Tensor, eps: f64) -> Result<Tensor> {
165 let xs_dtype = xs.dtype();
166 let needs_conversion = xs_dtype != DType::F32;
167
168 let xs_f32 = if needs_conversion {
170 xs.to_dtype(DType::F32)?
171 } else {
172 xs.clone()
173 };
174
175 let dim = xs_f32.dim(candle::D::Minus1)?;
176 let mean = (xs_f32.sum_keepdim(candle::D::Minus1)? / dim as f64)?;
177 let centered = xs_f32.broadcast_sub(&mean)?;
178 let var = (centered.sqr()?.sum_keepdim(candle::D::Minus1)? / dim as f64)?;
179 let xs = centered.broadcast_div(&(var + eps)?.sqrt()?)?;
180
181 let xs = if needs_conversion {
183 xs.to_dtype(xs_dtype)?
184 } else {
185 xs
186 };
187 let xs = xs.broadcast_mul(weight)?.broadcast_add(bias)?;
188 Ok(xs)
189}
190
191#[derive(Debug, Clone)]
194struct TimeMix {
195 x_r: Tensor,
197 x_w: Tensor,
198 x_k: Tensor,
199 x_v: Tensor,
200 x_a: Tensor,
201 x_g: Tensor,
202 w0: Tensor,
204 w1: Tensor,
205 w2: Tensor,
206 a0: Tensor,
208 a1: Tensor,
209 a2: Tensor,
210 v0: Option<Tensor>,
212 v1: Option<Tensor>,
213 v2: Option<Tensor>,
214 g1: Tensor,
216 g2: Tensor,
217 k_k: Tensor,
219 k_a: Tensor,
220 r_k: Tensor,
222 receptance_t: Tensor,
224 key_t: Tensor,
225 value_t: Tensor,
226 output_t: Tensor,
227 ln_x_weight: Tensor,
229 ln_x_bias: Tensor,
230 layer_id: usize,
232 n_heads: usize,
233 head_size: usize,
234}
235
236impl TimeMix {
237 fn new(
238 layer_id: usize,
239 cfg: &Config,
240 lora: (usize, usize, usize, usize),
241 vb: VarBuilder,
242 ) -> Result<Self> {
243 let c = cfg.hidden_size;
244 let (d_decay, d_aaa, d_mv, d_gate) = lora;
245 let n_heads = cfg.n_heads();
246 let head_size = cfg.head_size;
247
248 let x_r = vb.get((1, 1, c), "x_r")?.squeeze(0)?.squeeze(0)?;
250 let x_w = vb.get((1, 1, c), "x_w")?.squeeze(0)?.squeeze(0)?;
251 let x_k = vb.get((1, 1, c), "x_k")?.squeeze(0)?.squeeze(0)?;
252 let x_v = vb.get((1, 1, c), "x_v")?.squeeze(0)?.squeeze(0)?;
253 let x_a = vb.get((1, 1, c), "x_a")?.squeeze(0)?.squeeze(0)?;
254 let x_g = vb.get((1, 1, c), "x_g")?.squeeze(0)?.squeeze(0)?;
255
256 let w0 = vb.get((1, 1, c), "w0")?.squeeze(0)?.squeeze(0)?;
257 let w1 = vb.get((c, d_decay), "w1")?;
258 let w2 = vb.get((d_decay, c), "w2")?;
259
260 let a0 = vb.get((1, 1, c), "a0")?.squeeze(0)?.squeeze(0)?;
261 let a1 = vb.get((c, d_aaa), "a1")?;
262 let a2 = vb.get((d_aaa, c), "a2")?;
263
264 let (v0, v1, v2) = if layer_id > 0 {
267 (
268 Some(vb.get((1, 1, c), "v0")?.squeeze(0)?.squeeze(0)?),
269 Some(vb.get((c, d_mv), "v1")?),
270 Some(vb.get((d_mv, c), "v2")?),
271 )
272 } else {
273 let _ = vb.get((1, 1, c), "v0");
275 let _ = vb.get((c, d_mv), "v1");
276 let _ = vb.get((d_mv, c), "v2");
277 (None, None, None)
278 };
279
280 let g1 = vb.get((c, d_gate), "g1")?;
281 let g2 = vb.get((d_gate, c), "g2")?;
282
283 let k_k = vb.get((1, 1, c), "k_k")?.squeeze(0)?.squeeze(0)?;
284 let k_a = vb.get((1, 1, c), "k_a")?.squeeze(0)?.squeeze(0)?;
285 let r_k = vb
287 .get((n_heads, head_size), "r_k")?
288 .reshape(n_heads * head_size)?;
289
290 let receptance_t = vb.get((c, c), "receptance.weight")?.t()?.contiguous()?;
292 let key_t = vb.get((c, c), "key.weight")?.t()?.contiguous()?;
293 let value_t = vb.get((c, c), "value.weight")?.t()?.contiguous()?;
294 let output_t = vb.get((c, c), "output.weight")?.t()?.contiguous()?;
295
296 let ln_x_weight = vb.get(c, "ln_x.weight")?;
297 let ln_x_bias = vb.get(c, "ln_x.bias")?;
298
299 Ok(Self {
300 x_r,
301 x_w,
302 x_k,
303 x_v,
304 x_a,
305 x_g,
306 w0,
307 w1,
308 w2,
309 a0,
310 a1,
311 a2,
312 v0,
313 v1,
314 v2,
315 g1,
316 g2,
317 k_k,
318 k_a,
319 r_k,
320 receptance_t,
321 key_t,
322 value_t,
323 output_t,
324 ln_x_weight,
325 ln_x_bias,
326 layer_id,
327 n_heads,
328 head_size,
329 })
330 }
331
332 fn forward(
335 &self,
336 x: &Tensor,
337 state: &mut StatePerLayer,
338 v_first: Option<Tensor>,
339 ) -> Result<(Tensor, Tensor)> {
340 let h = self.n_heads;
341 let n = self.head_size;
342
343 macro_rules! mm {
345 ($x:expr, $w:expr) => {
346 $x.unsqueeze(0)?.matmul($w)?.squeeze(0)?
347 };
348 }
349
350 let xx = (&state.att_x_prev - x)?;
353 let xr = (x + xx.broadcast_mul(&self.x_r)?)?;
354 let xw = (x + xx.broadcast_mul(&self.x_w)?)?;
355 let xk = (x + xx.broadcast_mul(&self.x_k)?)?;
356 let xv = (x + xx.broadcast_mul(&self.x_v)?)?;
357 let xa = (x + xx.broadcast_mul(&self.x_a)?)?;
358 let xg = (x + xx.broadcast_mul(&self.x_g)?)?;
359 state.att_x_prev = x.clone();
360
361 let r = mm!(xr, &self.receptance_t);
363 let k = mm!(xk, &self.key_t);
364 let v = mm!(xv, &self.value_t);
365
366 let w = mm!(mm!(xw, &self.w1).tanh()?, &self.w2);
368 let w = (&self.w0 + &w)?.to_dtype(DType::F32)?;
369 let w = (w.neg()?.exp()? + 1.0)?.recip()?; let w = (w * (-0.606531))?.exp()?;
371
372 let (v, v_first) = if self.layer_id == 0 {
374 let v_first = v.clone();
376 (v, v_first)
377 } else {
378 let v_first = v_first.unwrap();
379 if let (Some(v0), Some(v1), Some(v2)) = (&self.v0, &self.v1, &self.v2) {
380 let gate = candle_nn::ops::sigmoid(&(v0 + mm!(mm!(xv, v1), v2))?)?;
381 let v = (&v + (&v_first - &v)?.broadcast_mul(&gate)?)?;
382 (v, v_first)
383 } else {
384 (v, v_first)
385 }
386 };
387
388 let a = candle_nn::ops::sigmoid(&(&self.a0 + mm!(mm!(xa, &self.a1), &self.a2))?)?;
390
391 let g = mm!(candle_nn::ops::sigmoid(&mm!(xg, &self.g1))?, &self.g2);
393
394 let kk = (&k * &self.k_k)?;
397 let kk = kk.reshape((h, n))?;
398 let kk_norm = (kk.sqr()?.sum_keepdim(1)?.sqrt()? + 1e-12)?;
399 let kk = kk.broadcast_div(&kk_norm)?;
400 let kk = kk.reshape(h * n)?;
401
402 let k = (&k * (1.0 + (&a - 1.0)?.broadcast_mul(&self.k_a)?)?)?;
404
405 let v_hn = v.reshape((h, n, 1))?;
408 let k_hn = k.reshape((h, 1, n))?;
409 let vk = v_hn.matmul(&k_hn)?;
410
411 let kk_h = kk.reshape((h, n))?;
413 let a_h = a.reshape((h, n))?;
414 let neg_kk = kk_h.neg()?.reshape((h, n, 1))?;
415 let kk_a = (&kk_h * &a_h)?.reshape((h, 1, n))?;
416 let ab = neg_kk.matmul(&kk_a)?;
417
418 let w_h = w.reshape((h, 1, n))?;
420 let att_kv = &state.att_kv;
421 let new_state = (att_kv.broadcast_mul(&w_h)?
422 + att_kv
423 .to_dtype(DType::F32)?
424 .matmul(&ab.to_dtype(DType::F32)?)?
425 + vk.to_dtype(DType::F32)?)?;
426 state.att_kv = new_state;
427
428 let r_hn = r.reshape((h, n, 1))?;
430 let out = state.att_kv.to_dtype(r.dtype())?.matmul(&r_hn)?;
431
432 let out = {
434 let reshaped = out.reshape((h, n))?;
435 let mean = reshaped.mean_keepdim(1)?;
436 let centered = reshaped.broadcast_sub(&mean)?;
437 let var = centered.sqr()?.mean_keepdim(1)?;
438 let normed = centered.broadcast_div(&(var + 64e-5)?.sqrt()?)?;
439 normed.reshape(h * n)?
440 };
441 let out = (out.broadcast_mul(&self.ln_x_weight)? + &self.ln_x_bias)?;
442
443 let bonus = (&r * &k * &self.r_k)?
445 .reshape((h, n))?
446 .sum_keepdim(1)?
447 .broadcast_mul(&v.reshape((h, n))?)?
448 .reshape(h * n)?;
449 let out = (out + bonus)?;
450
451 let out = mm!((out * g)?, &self.output_t);
453
454 Ok((out, v_first))
455 }
456}
457
458#[derive(Debug, Clone)]
461struct ChannelMix {
462 x_k: Tensor, key_t: Tensor, value_t: Tensor, deep_embed: Option<DeepEmbed>,
467}
468
469#[derive(Debug, Clone)]
470struct DeepEmbed {
471 s_emb: Tensor, s0: Tensor, s1: Tensor, s2: Tensor, }
476
477impl ChannelMix {
478 fn new(_layer_id: usize, cfg: &Config, vb: VarBuilder) -> Result<Self> {
479 let c = cfg.hidden_size;
480 let dim_ffn = cfg.dim_ffn();
481
482 let x_k = vb.get((1, 1, c), "x_k")?.squeeze(0)?.squeeze(0)?;
484 let key_t = vb.get((dim_ffn, c), "key.weight")?.t()?.contiguous()?;
485 let value_t = vb.get((c, dim_ffn), "value.weight")?.t()?.contiguous()?;
486
487 let deep_embed = if cfg.version == ModelVersion::V7a || cfg.version == ModelVersion::V7b {
488 let s_emb = vb.get((cfg.vocab_size, 1024), "s_emb.weight")?;
490 let s0 = vb.get((1, 1, dim_ffn), "s0")?.squeeze(0)?.squeeze(0)?;
492 let s1 = vb.get((c, 32), "s1")?;
493 let s2 = vb.get((32, dim_ffn), "s2")?;
494 Some(DeepEmbed { s_emb, s0, s1, s2 })
495 } else {
496 None
497 };
498
499 Ok(Self {
500 x_k,
501 key_t,
502 value_t,
503 deep_embed,
504 })
505 }
506
507 fn forward(
510 &self,
511 x: &Tensor,
512 state: &mut StatePerLayer,
513 token_ids: Option<&[u32]>,
514 ) -> Result<Tensor> {
515 macro_rules! mm {
516 ($x:expr, $w:expr) => {
517 $x.unsqueeze(0)?.matmul($w)?.squeeze(0)?
518 };
519 }
520
521 let xx = (&state.ffn_x_prev - x)?;
523 let k = (x + xx.broadcast_mul(&self.x_k)?)?;
524 state.ffn_x_prev = x.clone();
525
526 let mut k = mm!(k, &self.key_t).relu()?.sqr()?;
528
529 if let Some(de) = &self.deep_embed {
531 let token_ids = token_ids.expect("v7a/v7b requires token_ids in forward");
532 let token_id = token_ids[0] as usize;
533 let semb = de.s_emb.i(token_id)?;
535 let ss = mm!(x, &de.s1)
536 .unsqueeze(0)?
537 .matmul(&semb.reshape((32, 32))?)?
538 .squeeze(0)?;
539 let gate = (mm!(ss, &de.s2) + &de.s0)?;
541 k = (k * gate)?;
542 }
543
544 Ok(mm!(k, &self.value_t))
546 }
547}
548
549#[derive(Debug, Clone)]
552struct DeaAttention {
553 qq_weight: Tensor, k1: Tensor, k2: Tensor, k_emb: Tensor, v1: Tensor, v2: Tensor, v_emb: Tensor, x_q: Tensor, x_k: Tensor, x_v: Tensor, lnq_weight: Tensor, lnq_bias: Tensor, lnk_weight: Tensor, lnk_bias: Tensor, lnv_weight: Tensor, lnv_bias: Tensor, layer_id: usize,
570 hidden_size: usize,
571}
572
573impl DeaAttention {
574 fn new(layer_id: usize, cfg: &Config, vb: VarBuilder) -> Result<Self> {
575 let c = cfg.hidden_size;
576 let qq_weight = vb.get((256, c), "qq.weight")?.t()?.contiguous()?;
578 let k1 = vb.get((c, 32), "k1")?;
579 let k2 = vb.get((32, 256), "k2")?;
580 let k_emb = vb.get((cfg.vocab_size, 256), "k_emb.weight")?;
581 let v1 = vb.get((c, 32), "v1")?;
582 let v2 = vb.get((32, c), "v2")?;
583 let v_emb = vb.get((cfg.vocab_size, c), "v_emb.weight")?;
584 let x_q = vb.get((1, 1, 256), "x_q")?.squeeze(0)?.squeeze(0)?;
586 let x_k = vb.get((1, 1, 256), "x_k")?.squeeze(0)?.squeeze(0)?;
587 let x_v = vb.get((1, 1, c), "x_v")?.squeeze(0)?.squeeze(0)?;
588
589 let lnq_weight = vb.get(256, "lnq.weight")?;
590 let lnq_bias = vb.get(256, "lnq.bias")?;
591 let lnk_weight = vb.get(256, "lnk.weight")?;
592 let lnk_bias = vb.get(256, "lnk.bias")?;
593 let lnv_weight = vb.get(c, "lnv.weight")?;
594 let lnv_bias = vb.get(c, "lnv.bias")?;
595 Ok(Self {
596 qq_weight,
597 k1,
598 k2,
599 k_emb,
600 v1,
601 v2,
602 v_emb,
603 x_q,
604 x_k,
605 x_v,
606 lnq_weight,
607 lnq_bias,
608 lnk_weight,
609 lnk_bias,
610 lnv_weight,
611 lnv_bias,
612 layer_id,
613 hidden_size: c,
614 })
615 }
616
617 fn forward(&self, x: &Tensor, dea_state: &mut DeaState, token_ids: &[u32]) -> Result<Tensor> {
619 let dev = x.device();
620
621 macro_rules! mm {
623 ($x:expr, $w:expr) => {
624 $x.unsqueeze(0)?.matmul($w)?.squeeze(0)?
625 };
626 }
627
628 let q = mm!(x, &self.qq_weight);
630
631 let k_proj = mm!(x, &self.k1); let k_proj_2d = k_proj.reshape((1, 32))?;
634 let old_k = &dea_state.k_cache[self.layer_id];
635 dea_state.k_cache[self.layer_id] = if old_k.dim(0)? == 0 {
636 k_proj_2d.clone()
637 } else {
638 Tensor::cat(&[old_k, &k_proj_2d], 0)?
639 };
640 let all_token_ids: Vec<u32> = dea_state
641 .token_ids
642 .iter()
643 .copied()
644 .chain(token_ids.iter().copied())
645 .collect();
646 let ctx_tensor = Tensor::new(&all_token_ids[..], dev)?;
647 let k_full = dea_state.k_cache[self.layer_id].matmul(&self.k2)?;
648 let k_emb_sel = self.k_emb.index_select(&ctx_tensor, 0)?;
649 let k_full = (k_full * k_emb_sel)?;
650
651 let v_proj = mm!(x, &self.v1); let v_proj_2d = v_proj.reshape((1, 32))?;
654 let old_v = &dea_state.v_cache[self.layer_id];
655 dea_state.v_cache[self.layer_id] = if old_v.dim(0)? == 0 {
656 v_proj_2d.clone()
657 } else {
658 Tensor::cat(&[old_v, &v_proj_2d], 0)?
659 };
660 let v_full = dea_state.v_cache[self.layer_id].matmul(&self.v2)?.tanh()?;
661 let v_emb_sel = self.v_emb.index_select(&ctx_tensor, 0)?;
662 let v_full = (v_full * v_emb_sel)?;
663
664 let q_prev = &dea_state.q_prev[self.layer_id];
667 let q_shifted = (&q + (q_prev - &q)?.broadcast_mul(&self.x_q)?)?;
668 dea_state.q_prev[self.layer_id] = q.clone(); let q = q_shifted;
670
671 let seq_len = k_full.dim(0)?;
675
676 let k_full = if seq_len > 1 {
677 let k_shifted = Tensor::cat(
678 &[
679 &Tensor::zeros((1, 256), k_full.dtype(), dev)?,
680 &k_full.i(..seq_len - 1)?,
681 ],
682 0,
683 )?;
684 (&k_full + (&k_shifted - &k_full)?.broadcast_mul(&self.x_k)?)?
685 } else {
686 let scale = (self.x_k.neg()? + 1.0)?;
689
690 k_full.broadcast_mul(&scale)?
691 };
692 let v_full = if seq_len > 1 {
693 let v_shifted = Tensor::cat(
694 &[
695 &Tensor::zeros((1, self.hidden_size), v_full.dtype(), dev)?,
696 &v_full.i(..seq_len - 1)?,
697 ],
698 0,
699 )?;
700 (&v_full + (&v_shifted - &v_full)?.broadcast_mul(&self.x_v)?)?
701 } else {
702 let scale = (1.0 - &self.x_v)?;
704 v_full.broadcast_mul(&scale)?
705 };
706
707 let q = layer_norm(&q.unsqueeze(0)?, &self.lnq_weight, &self.lnq_bias, 1e-5)?.squeeze(0)?;
709 let k_full = layer_norm(&k_full, &self.lnk_weight, &self.lnk_bias, 1e-5)?;
710 let v_full = layer_norm(&v_full, &self.lnv_weight, &self.lnv_bias, 1e-5)?;
711
712 let scores = q.unsqueeze(0)?.matmul(&k_full.t()?)?;
714 let scores = ((scores * (1.0 / 1024.0))?.tanh()? * 64.0)?;
715
716 let attn_weights = candle_nn::ops::softmax_last_dim(&scores)?;
718 let out = attn_weights.matmul(&v_full)?.squeeze(0)?;
719
720 Ok(out)
721 }
722}
723
724#[derive(Debug, Clone)]
727struct Block {
728 ln0_weight: Option<Tensor>,
729 ln0_bias: Option<Tensor>,
730 ln1_weight: Tensor,
731 ln1_bias: Tensor,
732 ln2_weight: Tensor,
733 ln2_bias: Tensor,
734 att: TimeMix,
735 ffn: ChannelMix,
736 dea: Option<DeaAttention>,
737 layer_id: usize,
738}
739
740impl Block {
741 fn new(
742 layer_id: usize,
743 cfg: &Config,
744 lora: (usize, usize, usize, usize),
745 vb: VarBuilder,
746 ) -> Result<Self> {
747 let c = cfg.hidden_size;
748
749 let (ln0_weight, ln0_bias) = if layer_id == 0 {
750 (Some(vb.get(c, "ln0.weight")?), Some(vb.get(c, "ln0.bias")?))
751 } else {
752 (None, None)
753 };
754
755 let ln1_weight = vb.get(c, "ln1.weight")?;
756 let ln1_bias = vb.get(c, "ln1.bias")?;
757 let ln2_weight = vb.get(c, "ln2.weight")?;
758 let ln2_bias = vb.get(c, "ln2.bias")?;
759
760 let att = TimeMix::new(layer_id, cfg, lora, vb.pp("att"))?;
761 let ffn = ChannelMix::new(layer_id, cfg, vb.pp("ffn"))?;
762
763 let dea = if cfg.version == ModelVersion::V7b {
764 Some(DeaAttention::new(layer_id, cfg, vb.pp("qkv"))?)
765 } else {
766 None
767 };
768
769 Ok(Self {
770 ln0_weight,
771 ln0_bias,
772 ln1_weight,
773 ln1_bias,
774 ln2_weight,
775 ln2_bias,
776 att,
777 ffn,
778 dea,
779 layer_id,
780 })
781 }
782
783 fn forward(
784 &self,
785 x: &Tensor,
786 state: &mut State,
787 v_first: Option<Tensor>,
788 token_ids: Option<&[u32]>,
789 ) -> Result<(Tensor, Tensor)> {
790 let x_owned: Option<Tensor> = if let (Some(w), Some(b)) = (&self.ln0_weight, &self.ln0_bias)
792 {
793 Some(layer_norm(x, w, b, 1e-5)?)
794 } else {
795 None
796 };
797 let x_ref: &Tensor = x_owned.as_ref().unwrap_or(x);
798
799 let dea_out = if let Some(dea) = &self.dea {
801 let dea_state = state.dea.as_mut().expect("v7b requires DeaState");
802 Some(dea.forward(x_ref, dea_state, token_ids.unwrap())?)
803 } else {
804 None
805 };
806
807 let x_ln1 = layer_norm(x_ref, &self.ln1_weight, &self.ln1_bias, 1e-5)?;
809 let (att_out, v_first) =
810 self.att
811 .forward(&x_ln1, &mut state.per_layer[self.layer_id], v_first)?;
812
813 let x = if let Some(dea_out) = dea_out {
815 (x_ref + &att_out + dea_out)?
816 } else {
817 (x_ref + att_out)?
818 };
819
820 let x_ln2 = layer_norm(&x, &self.ln2_weight, &self.ln2_bias, 1e-5)?;
822 let ffn_out = self
823 .ffn
824 .forward(&x_ln2, &mut state.per_layer[self.layer_id], token_ids)?;
825 let x = (x + ffn_out)?;
826
827 Ok((x, v_first))
828 }
829}
830
831#[derive(Debug, Clone)]
834pub struct Model {
835 embeddings: Embedding,
836 blocks: Vec<Block>,
837 ln_out_weight: Tensor,
838 ln_out_bias: Tensor,
839 head_t: Tensor, pub version: ModelVersion,
841}
842
843impl Model {
844 pub fn new(cfg: &Config, vb: VarBuilder) -> Result<Self> {
845 let c = cfg.hidden_size;
846 let lora = infer_lora_dims(&vb)?;
847
848 let embeddings = embedding(cfg.vocab_size, c, vb.pp("emb"))?;
849
850 let mut blocks = Vec::with_capacity(cfg.num_hidden_layers);
851 let vb_b = vb.pp("blocks");
852 for layer_id in 0..cfg.num_hidden_layers {
853 blocks.push(Block::new(layer_id, cfg, lora, vb_b.pp(layer_id))?);
854 }
855
856 let ln_out_weight = vb.get(c, "ln_out.weight")?;
857 let ln_out_bias = vb.get(c, "ln_out.bias")?;
858 let head_t = vb
860 .get((cfg.vocab_size, c), "head.weight")?
861 .t()?
862 .contiguous()?;
863
864 let mut model = Self {
865 embeddings,
866 blocks,
867 ln_out_weight,
868 ln_out_bias,
869 head_t,
870 version: cfg.version,
871 };
872
873 if cfg.version == ModelVersion::V7a || cfg.version == ModelVersion::V7b {
878 let ln0_weight = &model.blocks[0]
880 .ln0_weight
881 .as_ref()
882 .expect("v7a/v7b requires ln0");
883 let ln0_bias = &model.blocks[0]
884 .ln0_bias
885 .as_ref()
886 .expect("v7a/v7b requires ln0");
887
888 let emb_raw = model.embeddings.embeddings();
890 let emb_normalized = layer_norm(emb_raw, ln0_weight, ln0_bias, 1e-5)?;
891
892 for i in 0..cfg.num_hidden_layers {
894 if let Some(de) = &mut model.blocks[i].ffn.deep_embed {
895 let s_emb_x = vb_b.pp(i).pp("ffn").get((1024, c), "s_emb_x.weight")?;
897 de.s_emb = (&de.s_emb + emb_normalized.matmul(&s_emb_x.t()?)?)?;
898 }
899 }
900
901 if cfg.version == ModelVersion::V7b {
903 for i in 0..cfg.num_hidden_layers {
904 if let Some(dea) = &mut model.blocks[i].dea {
905 let k_emb_x = vb_b.pp(i).pp("qkv").get((256, c), "k_emb_x.weight")?;
906 dea.k_emb = (&dea.k_emb + emb_normalized.matmul(&k_emb_x.t()?)?)?;
907
908 let v_emb_x = vb_b.pp(i).pp("qkv").get((c, c), "v_emb_x.weight")?;
909 dea.v_emb = (&dea.v_emb + emb_normalized.matmul(&v_emb_x.t()?)?)?;
910 }
911 }
912 }
913 }
914
915 Ok(model)
916 }
917
918 pub fn forward(&self, xs: &Tensor, state: &mut State, token_ids: &[u32]) -> Result<Tensor> {
923 let mut xs = xs.apply(&self.embeddings)?;
924 xs = xs.squeeze(0)?.squeeze(0)?;
926
927 let token_ids_opt = if self.version == ModelVersion::V7 {
928 None
929 } else {
930 Some(token_ids)
931 };
932
933 let mut v_first: Option<Tensor> = None;
934 for block in &self.blocks {
935 let (new_xs, new_v_first) = block.forward(&xs, state, v_first, token_ids_opt)?;
936 xs = new_xs;
937 v_first = Some(new_v_first);
938 }
939
940 if let Some(dea_state) = &mut state.dea {
942 dea_state.token_ids.extend_from_slice(token_ids);
943 }
944
945 let xs = layer_norm(&xs, &self.ln_out_weight, &self.ln_out_bias, 1e-5)?;
946 let xs = xs.unsqueeze(0)?.matmul(&self.head_t)?.squeeze(0)?;
948 state.pos += 1;
949 Ok(xs)
950 }
951
952 pub fn forward_seq(&self, token_ids: &[u32], state: &mut State) -> Result<Tensor> {
960 if token_ids.is_empty() {
961 candle::bail!("token_ids cannot be empty");
962 }
963
964 if token_ids.len() == 1 {
966 let dev = state.per_layer[0].att_x_prev.device();
967 let input = Tensor::new(&[token_ids[0]], dev)?.unsqueeze(0)?;
968 return self.forward(&input, state, token_ids);
969 }
970
971 let dev = state.per_layer[0].att_x_prev.device();
972
973 let input_ids = Tensor::new(token_ids, dev)?;
975 let xs = input_ids.apply(&self.embeddings)?;
976
977 let seq_len = token_ids.len();
980 let mut last_logits = None;
981
982 for t in 0..seq_len {
983 let x = xs.i(t)?;
985
986 let token_ids_opt = if self.version == ModelVersion::V7 {
987 None
988 } else {
989 Some(&token_ids[t..t + 1])
990 };
991
992 let mut x_out = x;
993 let mut v_first: Option<Tensor> = None;
994
995 for block in &self.blocks {
996 let (new_x, new_v_first) = block.forward(&x_out, state, v_first, token_ids_opt)?;
997 x_out = new_x;
998 v_first = Some(new_v_first);
999 }
1000
1001 if let Some(dea_state) = &mut state.dea {
1003 dea_state.token_ids.push(token_ids[t]);
1004 }
1005
1006 state.pos += 1;
1007
1008 if t == seq_len - 1 {
1010 let x_norm = layer_norm(&x_out, &self.ln_out_weight, &self.ln_out_bias, 1e-5)?;
1011 last_logits = Some(x_norm.unsqueeze(0)?.matmul(&self.head_t)?.squeeze(0)?);
1012 }
1013 }
1014
1015 last_logits.ok_or_else(|| candle::Error::Msg("No tokens processed".to_string()))
1016 }
1017}