Skip to main content

candle_transformers/models/
rwkv_v7.rs

1//! RWKV v7 "Goose" (x070) model implementation.
2//!
3//! The [RWKV model](https://wiki.rwkv.com/) is a recurrent neural network model
4//! with performance on par with transformer architectures. This implements the v7
5//! architecture (codenamed "Goose"), which introduces:
6//!
7//! - Delta-rule state update with in-context learning
8//! - Value residual stream across layers
9//! - LoRA-style projections for decay, gate, and ICL parameters
10//!
11//! Three variants are supported:
12//! - **v7**: Base architecture with linear attention + squared ReLU FFN
13//! - **v7a**: Adds DeepEmbed token-dependent gating to the FFN
14//! - **v7b**: Adds Deep Embedding Attention (DEA) — a full quadratic attention alongside RWKV
15//!
16//! # References
17//!
18//! - [RWKV-7 reference code](https://github.com/BlinkDL/RWKV-LM/tree/main/RWKV-v7)
19
20use candle::{DType, Device, IndexOp, Result, Tensor};
21use candle_nn::{embedding, Embedding, VarBuilder};
22
23// ─── Config ──────────────────────────────────────────────────────────────────
24
25/// Which RWKV v7 variant to use.
26#[derive(Debug, Clone, Copy, PartialEq, Eq, serde::Deserialize)]
27pub enum ModelVersion {
28    V7,
29    V7a,
30    V7b,
31}
32
33/// Configuration for RWKV v7 models.
34#[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
65/// Infer LoRA dimensions from actual weight shapes in the first block.
66/// This is more robust than computing from a formula, as different
67/// model sizes may use different LoRA dimensions.
68fn 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
77// ─── State ───────────────────────────────────────────────────────────────────
78
79/// Per-layer persistent state for RWKV v7 inference.
80pub struct StatePerLayer {
81    /// Previous token embedding for time-mix shifting. Shape: `(hidden_size,)`.
82    pub att_x_prev: Tensor,
83    /// WKV state matrix. Shape: `(n_heads, head_size, head_size)` in f32.
84    pub att_kv: Tensor,
85    /// Previous token embedding for channel-mix shifting. Shape: `(hidden_size,)`.
86    pub ffn_x_prev: Tensor,
87}
88
89/// KV cache state for DEA (v7b only).
90pub struct DeaState {
91    /// Token IDs seen so far (growing).
92    pub token_ids: Vec<u32>,
93    /// Per-layer K projections cache. Each entry: `(seq_len, 32)`.
94    pub k_cache: Vec<Tensor>,
95    /// Per-layer V projections cache. Each entry: `(seq_len, 32)`.
96    pub v_cache: Vec<Tensor>,
97    /// Per-layer previous Q for token-shifting. Each entry: `(256,)`.
98    pub q_prev: Vec<Tensor>,
99}
100
101/// Full inference state for RWKV v7.
102pub struct State {
103    pub per_layer: Vec<StatePerLayer>,
104    pub dea: Option<DeaState>,
105    pub pos: usize,
106}
107
108impl State {
109    /// Create state with F32 precision (default, most compatible).
110    pub fn new(cfg: &Config, dev: &Device) -> Result<Self> {
111        Self::new_with_dtype(cfg, dev, DType::F32)
112    }
113
114    /// Create state with specified dtype (F16/BF16 for faster inference).
115    ///
116    /// Note: The KV state (`att_kv`) always uses F32 for numerical stability
117    /// in the delta-rule accumulation. Other state tensors use the specified dtype.
118    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                // KV state stays F32 for numerical stability in accumulation
125                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
155// ─── Tokenizer ───────────────────────────────────────────────────────────────
156
157pub use crate::models::rwkv_v5::Tokenizer;
158
159// ─── Helpers ─────────────────────────────────────────────────────────────────
160
161/// Layer normalization that preserves input dtype when possible.
162/// All internal computation happens in F32 for numerical stability,
163/// then converts back to the original dtype.
164fn 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    // Convert to F32 for all internal computation (numerical stability)
169    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    // Convert back to original dtype if needed
182    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// ─── TimeMix (Attention) ─────────────────────────────────────────────────────
192
193#[derive(Debug, Clone)]
194struct TimeMix {
195    // Token-shift lerp mixes (pre-squeezed to 1D for efficiency)
196    x_r: Tensor,
197    x_w: Tensor,
198    x_k: Tensor,
199    x_v: Tensor,
200    x_a: Tensor,
201    x_g: Tensor,
202    // Decay LoRA (w0 pre-squeezed)
203    w0: Tensor,
204    w1: Tensor,
205    w2: Tensor,
206    // ICL rate LoRA (a0 pre-squeezed)
207    a0: Tensor,
208    a1: Tensor,
209    a2: Tensor,
210    // Value residual LoRA (None for layer 0, v0 pre-squeezed)
211    v0: Option<Tensor>,
212    v1: Option<Tensor>,
213    v2: Option<Tensor>,
214    // Gate LoRA
215    g1: Tensor,
216    g2: Tensor,
217    // Key processing (pre-squeezed)
218    k_k: Tensor,
219    k_a: Tensor,
220    // Bonus term (pre-flattened to 1D)
221    r_k: Tensor,
222    // Linear projections (pre-transposed for efficiency)
223    receptance_t: Tensor,
224    key_t: Tensor,
225    value_t: Tensor,
226    output_t: Tensor,
227    // GroupNorm weights
228    ln_x_weight: Tensor,
229    ln_x_bias: Tensor,
230    // Metadata
231    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        // Pre-squeeze (1,1,C) -> (C,) at load time to avoid per-token squeeze calls
249        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        // v0/v1/v2 exist for all layers in the weights file, but are only used for layers > 0
265        // (layer 0 stores v_first instead of blending toward it).
266        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            // Load and discard — these tensors exist in the file but are ignored at layer 0
274            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        // Pre-flatten r_k to (H*N,) to avoid reshape in forward
286        let r_k = vb
287            .get((n_heads, head_size), "r_k")?
288            .reshape(n_heads * head_size)?;
289
290        // Linear projections — pre-transpose and make contiguous for optimal memory access
291        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    /// Forward pass for a single token (RNN mode).
333    /// Input `x` shape: `[C]` (1D). Returns `(output [C], v_first [C])`.
334    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        // Helper: matrix multiply for 1D vec @ 2D weight: unsqueeze, matmul, squeeze
344        macro_rules! mm {
345            ($x:expr, $w:expr) => {
346                $x.unsqueeze(0)?.matmul($w)?.squeeze(0)?
347            };
348        }
349
350        // 1. Token shift: lerp between current and previous token
351        // (x_r, x_w, etc. are pre-squeezed at load time)
352        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        // 2. Linear projections (weights pre-transposed at load time)
362        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        // 3. Decay: w = exp(-0.606531 * sigmoid(w0 + tanh(xw @ w1) @ w2))
367        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()?; // sigmoid
370        let w = (w * (-0.606531))?.exp()?;
371
372        // 4. Value residual
373        let (v, v_first) = if self.layer_id == 0 {
374            // Layer 0: v_first = v (only one clone needed, v is moved)
375            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        // 5. ICL rate: a = sigmoid(a0 + (xa @ a1) @ a2)
389        let a = candle_nn::ops::sigmoid(&(&self.a0 + mm!(mm!(xa, &self.a1), &self.a2))?)?;
390
391        // 6. Gate: g = sigmoid(xg @ g1) @ g2
392        let g = mm!(candle_nn::ops::sigmoid(&mm!(xg, &self.g1))?, &self.g2);
393
394        // 7. Key processing (k_k, k_a pre-squeezed)
395        // kk = L2_normalize(k * k_k, per_head)
396        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        // k = k * (1 + (a - 1) * k_a)
403        let k = (&k * (1.0 + (&a - 1.0)?.broadcast_mul(&self.k_a)?)?)?;
404
405        // 8. State update (delta-rule core)
406        // vk = v.view(H,N,1) @ k.view(H,1,N)  — outer product
407        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        // ab = (-kk).view(H,N,1) @ (kk*a).view(H,1,N)  — ICL correction
412        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        // state = state * w.view(H,1,N) + state @ ab + vk
419        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        // out = state @ r.view(H,N,1)
429        let r_hn = r.reshape((h, n, 1))?;
430        let out = state.att_kv.to_dtype(r.dtype())?.matmul(&r_hn)?;
431
432        // 9. GroupNorm (H groups, eps=64e-5)
433        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        // 10. Bonus term: (r * k * r_k).sum_per_head * v (r_k pre-flattened)
444        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        // 11. Output (weight pre-transposed)
452        let out = mm!((out * g)?, &self.output_t);
453
454        Ok((out, v_first))
455    }
456}
457
458// ─── ChannelMix (FFN) ────────────────────────────────────────────────────────
459
460#[derive(Debug, Clone)]
461struct ChannelMix {
462    x_k: Tensor,     // Pre-squeezed to 1D
463    key_t: Tensor,   // Pre-transposed
464    value_t: Tensor, // Pre-transposed
465    // DeepEmbed (v7a, v7b only)
466    deep_embed: Option<DeepEmbed>,
467}
468
469#[derive(Debug, Clone)]
470struct DeepEmbed {
471    s_emb: Tensor, // (vocab_size, 1024) — pre-merged with emb @ s_emb_x^T
472    s0: Tensor,    // (dim_ffn,)
473    s1: Tensor,    // (hidden_size, 32)
474    s2: Tensor,    // (32, dim_ffn)
475}
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        // Pre-squeeze and pre-transpose at load time
483        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            // Load s_emb — the pre-merged embedding is computed in Model::new()
489            let s_emb = vb.get((cfg.vocab_size, 1024), "s_emb.weight")?;
490            // s0 stored as (1, 1, dim_ffn) in weights, squeeze to 1D for efficiency
491            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    /// Forward pass for a single token. Input `x` shape: `[C]`.
508    /// `token_ids` is needed for DeepEmbed (v7a/v7b).
509    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        // Token shift (x_k pre-squeezed)
522        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        // Squared ReLU: relu(key(k))^2 (key pre-transposed)
527        let mut k = mm!(k, &self.key_t).relu()?.sqr()?;
528
529        // DeepEmbed gating (v7a/v7b)
530        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            // ss = (x @ s1) @ s_emb[token_id].view(32, 32)
534            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            // k = k * ((ss @ s2) + s0)
540            let gate = (mm!(ss, &de.s2) + &de.s0)?;
541            k = (k * gate)?;
542        }
543
544        // Down-projection (value pre-transposed)
545        Ok(mm!(k, &self.value_t))
546    }
547}
548
549// ─── DeaAttention (v7b only) ─────────────────────────────────────────────────
550
551#[derive(Debug, Clone)]
552struct DeaAttention {
553    qq_weight: Tensor,  // (hidden_size, 256)
554    k1: Tensor,         // (hidden_size, 32)
555    k2: Tensor,         // (32, 256)
556    k_emb: Tensor,      // (vocab_size, 256) — pre-merged
557    v1: Tensor,         // (hidden_size, 32)
558    v2: Tensor,         // (32, hidden_size)
559    v_emb: Tensor,      // (vocab_size, hidden_size) — pre-merged
560    x_q: Tensor,        // (256,)
561    x_k: Tensor,        // (256,)
562    x_v: Tensor,        // (hidden_size,)
563    lnq_weight: Tensor, // (256,)
564    lnq_bias: Tensor,   // (256,)
565    lnk_weight: Tensor, // (256,)
566    lnk_bias: Tensor,   // (256,)
567    lnv_weight: Tensor, // (hidden_size,)
568    lnv_bias: Tensor,   // (hidden_size,)
569    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        // qq.weight stored as (256, hidden_size) in PyTorch format, transpose for matmul
577        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        // Token-shift params stored as (1, 1, dim), squeeze to 1D
585        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    /// Forward pass for DEA attention. Updates the KV cache in `dea_state`.
618    fn forward(&self, x: &Tensor, dea_state: &mut DeaState, token_ids: &[u32]) -> Result<Tensor> {
619        let dev = x.device();
620
621        // Helper for 1D vector @ 2D matrix multiplication
622        macro_rules! mm {
623            ($x:expr, $w:expr) => {
624                $x.unsqueeze(0)?.matmul($w)?.squeeze(0)?
625            };
626        }
627
628        // Q projection
629        let q = mm!(x, &self.qq_weight);
630
631        // K: project down, cache, project up, multiply by token embedding
632        let k_proj = mm!(x, &self.k1); // (32,)
633        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        // V: project down, cache, project up (with tanh), multiply by token embedding
652        let v_proj = mm!(x, &self.v1); // (32,)
653        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        // Token shifting on Q (using previous Q state)
665        // Important: save ORIGINAL q before shifting (reference line 160)
666        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(); // Save original, not shifted!
669        let q = q_shifted;
670
671        // Token shifting on K and V (pad left by 1)
672        // For seq_len=1: F.pad(k, (0,0,1,-1)) produces zeros, so k = k * (1 - x_k)
673        // For seq_len>1: shifted = [zeros, k[:-1]], so k = k + (shifted - k) * x_k
674        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            // Single token: shifted is zeros, so k = k + (0 - k) * x_k = k * (1 - x_k)
687            // Note: Candle doesn't support scalar - tensor directly, use neg + scalar
688            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            // Single token: v = v * (1 - x_v)
703            let scale = (1.0 - &self.x_v)?;
704            v_full.broadcast_mul(&scale)?
705        };
706
707        // LayerNorm on Q, K, V
708        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        // Soft-capped causal attention: 64 * tanh(q @ k^T / 1024)
713        let scores = q.unsqueeze(0)?.matmul(&k_full.t()?)?;
714        let scores = ((scores * (1.0 / 1024.0))?.tanh()? * 64.0)?;
715
716        // Attention output
717        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// ─── Block ───────────────────────────────────────────────────────────────────
725
726#[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        // Pre-norm (block 0 only) - store owned tensor if ln0 applied
791        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        // DEA attention (v7b only) — computed on x BEFORE ln1
800        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        // Time mixing (RWKV linear attention)
808        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        // Residual: x + att_out + dea_out (clone only when needed for addition)
814        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        // Channel mixing (FFN)
821        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// ─── Model ───────────────────────────────────────────────────────────────────
832
833#[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, // Pre-transposed for efficiency
840    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        // Pre-transpose head weight at load time
859        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        // Load-time merges for DeepEmbed (v7a/v7b) and DEA (v7b)
874        // IMPORTANT: Reference pre-normalizes emb.weight with ln0 BEFORE merging!
875        // See rwkv_v7b_demo.py line 103:
876        //   z['emb.weight'] = F.layer_norm(z['emb.weight'], ..., weight=z['blocks.0.ln0.weight'], ...)
877        if cfg.version == ModelVersion::V7a || cfg.version == ModelVersion::V7b {
878            // Get ln0 weights from block 0 to normalize embeddings
879            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            // Normalize embeddings with ln0 (applied to each row independently)
889            let emb_raw = model.embeddings.embeddings();
890            let emb_normalized = layer_norm(emb_raw, ln0_weight, ln0_bias, 1e-5)?;
891
892            // DeepEmbed merges (FFN s_emb)
893            for i in 0..cfg.num_hidden_layers {
894                if let Some(de) = &mut model.blocks[i].ffn.deep_embed {
895                    // s_emb += normalized_emb @ s_emb_x^T
896                    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            // DEA merges (v7b only)
902            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    /// Run a forward pass for a single token (RNN-style inference).
919    ///
920    /// `token_ids` should contain the token ID(s) being processed.
921    /// For v7a/v7b, these are used for DeepEmbed and DEA token-embedding lookups.
922    pub fn forward(&self, xs: &Tensor, state: &mut State, token_ids: &[u32]) -> Result<Tensor> {
923        let mut xs = xs.apply(&self.embeddings)?;
924        // xs shape: (1, 1, hidden_size) for single token; squeeze to (hidden_size,)
925        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        // Update DEA token ID cache after all blocks processed
941        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        // head_t is pre-transposed, no .t() needed
947        let xs = xs.unsqueeze(0)?.matmul(&self.head_t)?.squeeze(0)?;
948        state.pos += 1;
949        Ok(xs)
950    }
951
952    /// Process a sequence of tokens efficiently (batch prompt processing).
953    ///
954    /// This is significantly faster than calling `forward` token-by-token because:
955    /// - Embeddings are computed in one batch
956    /// - Linear projections are batched where possible
957    ///
958    /// Returns the logits for the last token only (for next-token prediction).
959    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        // For short sequences, fall back to single-token processing
965        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        // Batch embed all tokens at once: (seq_len,) -> (seq_len, hidden_size)
974        let input_ids = Tensor::new(token_ids, dev)?;
975        let xs = input_ids.apply(&self.embeddings)?;
976
977        // Process each token through all layers
978        // Note: RWKV state updates are sequential, but we batch the embedding lookup
979        let seq_len = token_ids.len();
980        let mut last_logits = None;
981
982        for t in 0..seq_len {
983            // Extract single token embedding: (hidden_size,)
984            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            // Update DEA token ID cache
1002            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            // Only compute logits for the last token
1009            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}