Skip to main content

runtime/models_v2/
rwkv4.rs

1//! RWKV-4 Model V2 - Linear Attention with Time and Channel Mixing
2//!
3//! This implements the RWKV-4 architecture which features:
4//! - Time mixing: Linear attention with exponential decay (WKV computation)
5//! - Channel mixing: FFN with receptance, key, value gates
6//! - Running state instead of attention/KV cache
7//! - O(1) memory per token during generation
8//!
9//! Supports: RWKV-4 models (169M to 14B)
10
11use crate::model_config;
12use super::traits::*;
13use anyhow::Result;
14use serde::{Serialize, Deserialize};
15
16/// RWKV-4 model configuration
17model_config!(Rwkv4Config {
18    vocab_size: usize = 50277,
19    hidden_size: usize = 768,      // d_model
20    num_hidden_layers: usize = 12,
21    intermediate_size: usize = 0,  // 0 = auto: hidden_size * 4
22    layer_norm_epsilon: f32 = 1e-5,
23    rescale_every: usize = 6,      // Rescale every N layers for numerical stability
24    tie_word_embeddings: bool = false,
25    pad_token_id: i64 = 0,
26    bos_token_id: i64 = 0,
27    eos_token_id: i64 = 0,
28});
29
30impl Rwkv4Config {
31    pub fn from_gguf_config(gguf: &crate::weight_loader_core::GGUFModelConfig) -> Self {
32        Self {
33            vocab_size: gguf.vocab_size,
34            hidden_size: gguf.hidden_size,
35            num_hidden_layers: gguf.num_hidden_layers,
36            intermediate_size: gguf.intermediate_size,
37            layer_norm_epsilon: gguf.rms_norm_eps,
38            ..Default::default()
39        }
40    }
41
42    pub fn effective_intermediate_size(&self) -> usize {
43        if self.intermediate_size > 0 {
44            self.intermediate_size
45        } else {
46            self.hidden_size * 4
47        }
48    }
49}
50
51/// Main RWKV-4 model
52pub struct Rwkv4ModelV2 {
53    config: Rwkv4Config,
54    device: Device,
55    embeddings: Tensor,
56    blocks: Vec<Rwkv4Block>,
57    ln_out: Tensor,
58    head: Tensor,
59}
60
61/// RWKV-4 block with time mixing and channel mixing
62pub struct Rwkv4Block {
63    ln1: Tensor,
64    ln2: Tensor,
65    time_mixing: Rwkv4TimeMixing,
66    channel_mixing: Rwkv4ChannelMixing,
67    layer_idx: usize,
68    rescale_every: usize,
69}
70
71/// RWKV-4 Time Mixing (attention-like with linear complexity)
72pub struct Rwkv4TimeMixing {
73    // Learnable parameters
74    time_decay: Tensor,      // w [hidden_size]
75    time_first: Tensor,      // u [hidden_size]
76    time_mix_k: Tensor,      // [hidden_size]
77    time_mix_v: Tensor,      // [hidden_size]
78    time_mix_r: Tensor,      // [hidden_size]
79
80    // Projections
81    key: Tensor,             // W_k [hidden_size, hidden_size]
82    value: Tensor,           // W_v [hidden_size, hidden_size]
83    receptance: Tensor,      // W_r [hidden_size, hidden_size]
84    output: Tensor,          // W_o [hidden_size, hidden_size]
85
86    hidden_size: usize,
87}
88
89/// RWKV-4 Channel Mixing (FFN-like)
90pub struct Rwkv4ChannelMixing {
91    time_mix_k: Tensor,      // [hidden_size]
92    time_mix_r: Tensor,      // [hidden_size]
93
94    key: Tensor,             // [hidden_size, intermediate_size]
95    value: Tensor,           // [intermediate_size, hidden_size]
96    receptance: Tensor,      // [hidden_size, hidden_size]
97
98    hidden_size: usize,
99    intermediate_size: usize,
100}
101
102/// RWKV state for generation
103#[derive(Clone)]
104pub struct Rwkv4State {
105    /// Numerator state [batch, hidden_size]
106    pub num: Tensor,
107    /// Denominator state [batch, hidden_size]
108    pub den: Tensor,
109    /// Previous x for time mixing [batch, hidden_size]
110    pub prev_x_tm: Tensor,
111    /// Previous x for channel mixing [batch, hidden_size]
112    pub prev_x_cm: Tensor,
113}
114
115impl Model for Rwkv4ModelV2 {
116    type Config = Rwkv4Config;
117
118    fn new(config: Rwkv4Config) -> Result<Self> {
119        let device = Device::CPU;
120
121        let embeddings = ops_fn::zeros(&[config.vocab_size, config.hidden_size], DataType::Float32, &device)?;
122        let ln_out = ops_fn::zeros(&[config.hidden_size], DataType::Float32, &device)?;
123        let head = ops_fn::zeros(&[config.hidden_size, config.vocab_size], DataType::Float32, &device)?;
124
125        let mut blocks = Vec::with_capacity(config.num_hidden_layers);
126        for i in 0..config.num_hidden_layers {
127            blocks.push(Rwkv4Block::new(&config, i, &device)?);
128        }
129
130        Ok(Self {
131            config,
132            device,
133            embeddings,
134            blocks,
135            ln_out,
136            head,
137        })
138    }
139
140    fn from_weights(config: Rwkv4Config, weights: ModelWeights) -> Result<Self> {
141        let mut model = Self::new(config)?;
142
143        if let Some(w) = weights.get("emb.weight").or_else(|| weights.get("rwkv.embeddings.weight")) {
144            model.embeddings = w.clone();
145        }
146
147        if let Some(w) = weights.get("ln_out.weight").or_else(|| weights.get("rwkv.ln_out.weight")) {
148            model.ln_out = w.clone();
149        }
150
151        if let Some(w) = weights.get("head.weight") {
152            model.head = ops_fn::transpose(w)?;
153        }
154
155        for (i, block) in model.blocks.iter_mut().enumerate() {
156            block.load_weights(&weights, i)?;
157        }
158
159        Ok(model)
160    }
161
162    fn forward(&self, inputs: &ModelInputs) -> Result<ModelOutputs> {
163        match inputs {
164            ModelInputs::Text { input_ids, .. } => {
165                let mut hidden_states = ops_fn::embedding(input_ids, &self.embeddings)?;
166
167                for block in &self.blocks {
168                    hidden_states = block.forward(&hidden_states)?;
169                }
170
171                hidden_states = ops_fn::layer_norm(&hidden_states, &self.ln_out, None, self.config.layer_norm_epsilon)?;
172                let logits = ops_fn::matmul(&hidden_states, &self.head)?;
173
174                Ok(ModelOutputs::Logits {
175                    logits,
176                    hidden_states: None,
177                })
178            }
179            _ => Err(anyhow::anyhow!("RWKV-4 only supports text inputs")),
180        }
181    }
182
183    fn generate(&self, prompt: &str, config: &GenerationConfig) -> Result<String> {
184        use crate::tokenizer::Tokenizer;
185        use rand::Rng;
186
187        let tokenizer = Tokenizer::new();
188        let mut tokens: Vec<u32> = tokenizer.encode(prompt);
189
190        // Initialize states for all layers
191        let batch_size = 1;
192        let mut layer_states: Vec<Rwkv4State> = Vec::new();
193        for _ in 0..self.config.num_hidden_layers {
194            layer_states.push(Rwkv4State {
195                num: ops_fn::zeros(&[batch_size, self.config.hidden_size], DataType::Float32, &self.device)?,
196                den: ops_fn::zeros(&[batch_size, self.config.hidden_size], DataType::Float32, &self.device)?,
197                prev_x_tm: ops_fn::zeros(&[batch_size, self.config.hidden_size], DataType::Float32, &self.device)?,
198                prev_x_cm: ops_fn::zeros(&[batch_size, self.config.hidden_size], DataType::Float32, &self.device)?,
199            });
200        }
201
202        // Process prompt
203        for &token in &tokens[..tokens.len().saturating_sub(1)] {
204            let input_tensor = Tensor::from_i64_slice(&[token as i64], &[1, 1], &self.device)?;
205            let mut hidden = ops_fn::embedding(&input_tensor, &self.embeddings)?;
206            hidden = hidden.reshape(&[1, self.config.hidden_size])?;
207
208            for (i, block) in self.blocks.iter().enumerate() {
209                hidden = block.forward_with_state(&hidden, &mut layer_states[i])?;
210            }
211        }
212
213        // Generation loop
214        for _ in 0..config.max_new_tokens {
215            let last_token = *tokens.last().unwrap_or(&0);
216            let input_tensor = Tensor::from_i64_slice(&[last_token as i64], &[1, 1], &self.device)?;
217            let mut hidden = ops_fn::embedding(&input_tensor, &self.embeddings)?;
218            hidden = hidden.reshape(&[1, self.config.hidden_size])?;
219
220            for (i, block) in self.blocks.iter().enumerate() {
221                hidden = block.forward_with_state(&hidden, &mut layer_states[i])?;
222            }
223
224            hidden = ops_fn::layer_norm(&hidden, &self.ln_out, None, self.config.layer_norm_epsilon)?;
225            let logits = ops_fn::matmul(&hidden, &self.head)?;
226
227            let logits_vec: Vec<f32> = logits.to_candle()?.flatten_all()?.to_vec1()?;
228
229            let next_token = if config.do_sample && config.temperature > 0.0 {
230                let scaled: Vec<f32> = logits_vec.iter().map(|&x| x / config.temperature).collect();
231                let max_val = scaled.iter().cloned().fold(f32::NEG_INFINITY, f32::max);
232                let exp_sum: f32 = scaled.iter().map(|&x| (x - max_val).exp()).sum();
233                let probs: Vec<f32> = scaled.iter().map(|&x| (x - max_val).exp() / exp_sum).collect();
234
235                let mut rng = rand::thread_rng();
236                let random_val: f32 = rng.gen();
237                let mut cumulative = 0.0;
238                let mut sampled = 0u32;
239
240                for (idx, &prob) in probs.iter().enumerate() {
241                    cumulative += prob;
242                    if random_val <= cumulative {
243                        sampled = idx as u32;
244                        break;
245                    }
246                }
247                sampled
248            } else {
249                logits_vec.iter()
250                    .enumerate()
251                    .max_by(|(_, a), (_, b)| a.partial_cmp(b).unwrap())
252                    .map(|(idx, _)| idx as u32)
253                    .unwrap_or(0)
254            };
255
256            if next_token == config.eos_token_id {
257                break;
258            }
259
260            tokens.push(next_token);
261        }
262
263        Ok(tokenizer.decode(&tokens))
264    }
265
266    fn config(&self) -> &Self::Config { &self.config }
267
268    fn memory_requirements(&self) -> MemoryRequirements {
269        let inter_size = self.config.effective_intermediate_size();
270        let param_size = (
271            self.config.vocab_size * self.config.hidden_size +
272            self.config.num_hidden_layers * (
273                4 * self.config.hidden_size * self.config.hidden_size +  // time mixing projections
274                self.config.hidden_size * inter_size * 2 +                // channel mixing
275                self.config.hidden_size * 10                              // various parameters
276            )
277        ) * 4;
278
279        let state_size = self.config.num_hidden_layers * self.config.hidden_size * 4 * 4;
280
281        MemoryRequirements {
282            gpu_memory: param_size,
283            cpu_memory: param_size / 4,
284            kv_cache_memory: state_size,
285            peak_memory: param_size + param_size / 2,
286        }
287    }
288
289    fn to_device(&mut self, device: &Device) -> Result<()> {
290        self.device = device.clone();
291        self.embeddings = self.embeddings.to_device(device)?;
292        self.ln_out = self.ln_out.to_device(device)?;
293        self.head = self.head.to_device(device)?;
294        for block in &mut self.blocks {
295            block.to_device(device)?;
296        }
297        Ok(())
298    }
299}
300
301impl Rwkv4Block {
302    fn new(config: &Rwkv4Config, layer_idx: usize, device: &Device) -> Result<Self> {
303        let ln1 = ops_fn::zeros(&[config.hidden_size], DataType::Float32, device)?;
304        let ln2 = ops_fn::zeros(&[config.hidden_size], DataType::Float32, device)?;
305        let time_mixing = Rwkv4TimeMixing::new(config, device)?;
306        let channel_mixing = Rwkv4ChannelMixing::new(config, device)?;
307
308        Ok(Self {
309            ln1,
310            ln2,
311            time_mixing,
312            channel_mixing,
313            layer_idx,
314            rescale_every: config.rescale_every,
315        })
316    }
317
318    fn forward(&self, hidden_states: &Tensor) -> Result<Tensor> {
319        let shape = hidden_states.shape();
320        let (batch_size, seq_len, _) = (shape[0], shape[1], shape[2]);
321
322        let mut output = hidden_states.clone();
323
324        // Process each position sequentially (RWKV is autoregressive by nature)
325        for t in 0..seq_len {
326            let x = hidden_states.to_candle()?.narrow(1, t, 1)?.squeeze(1)?;
327            let x = Tensor::from_candle(x);
328
329            // Time mixing
330            let ln_x = ops_fn::layer_norm(&x, &self.ln1, None, 1e-5)?;
331            let tm_out = self.time_mixing.forward(&ln_x, t)?;
332            let x = ops_fn::add(&x, &tm_out)?;
333
334            // Channel mixing
335            let ln_x = ops_fn::layer_norm(&x, &self.ln2, None, 1e-5)?;
336            let cm_out = self.channel_mixing.forward(&ln_x, t)?;
337            let x = ops_fn::add(&x, &cm_out)?;
338
339            // Rescale for numerical stability
340            let x = if self.rescale_every > 0 && (self.layer_idx + 1) % self.rescale_every == 0 {
341                ops_fn::scale(&x, 0.5)?
342            } else {
343                x
344            };
345
346            // Update output at position t
347            let x_expanded = x.to_candle()?.unsqueeze(1)?;
348            let output_candle = output.to_candle()?;
349            output = Tensor::from_candle(output_candle.slice_assign(&[0..batch_size, t..t+1, 0..shape[2]], &x_expanded)?);
350        }
351
352        Ok(output)
353    }
354
355    fn forward_with_state(&self, hidden_states: &Tensor, state: &mut Rwkv4State) -> Result<Tensor> {
356        // hidden_states: [batch, hidden_size]
357
358        // Time mixing
359        let ln_x = ops_fn::layer_norm(hidden_states, &self.ln1, None, 1e-5)?;
360        let tm_out = self.time_mixing.forward_with_state(&ln_x, state)?;
361        let x = ops_fn::add(hidden_states, &tm_out)?;
362
363        // Channel mixing
364        let ln_x = ops_fn::layer_norm(&x, &self.ln2, None, 1e-5)?;
365        let cm_out = self.channel_mixing.forward_with_state(&ln_x, state)?;
366        let x = ops_fn::add(&x, &cm_out)?;
367
368        // Rescale
369        if self.rescale_every > 0 && (self.layer_idx + 1) % self.rescale_every == 0 {
370            ops_fn::scale(&x, 0.5)
371        } else {
372            Ok(x)
373        }
374    }
375
376    fn load_weights(&mut self, weights: &ModelWeights, layer_idx: usize) -> Result<()> {
377        let prefix = format!("blocks.{}", layer_idx);
378
379        if let Some(w) = weights.get(&format!("{}.ln1.weight", prefix)) {
380            self.ln1 = w.clone();
381        }
382        if let Some(w) = weights.get(&format!("{}.ln2.weight", prefix)) {
383            self.ln2 = w.clone();
384        }
385
386        self.time_mixing.load_weights(weights, layer_idx)?;
387        self.channel_mixing.load_weights(weights, layer_idx)?;
388
389        Ok(())
390    }
391
392    fn to_device(&mut self, device: &Device) -> Result<()> {
393        self.ln1 = self.ln1.to_device(device)?;
394        self.ln2 = self.ln2.to_device(device)?;
395        self.time_mixing.to_device(device)?;
396        self.channel_mixing.to_device(device)?;
397        Ok(())
398    }
399}
400
401impl Rwkv4TimeMixing {
402    fn new(config: &Rwkv4Config, device: &Device) -> Result<Self> {
403        let hidden_size = config.hidden_size;
404
405        Ok(Self {
406            time_decay: ops_fn::zeros(&[hidden_size], DataType::Float32, device)?,
407            time_first: ops_fn::zeros(&[hidden_size], DataType::Float32, device)?,
408            time_mix_k: ops_fn::zeros(&[hidden_size], DataType::Float32, device)?,
409            time_mix_v: ops_fn::zeros(&[hidden_size], DataType::Float32, device)?,
410            time_mix_r: ops_fn::zeros(&[hidden_size], DataType::Float32, device)?,
411            key: ops_fn::zeros(&[hidden_size, hidden_size], DataType::Float32, device)?,
412            value: ops_fn::zeros(&[hidden_size, hidden_size], DataType::Float32, device)?,
413            receptance: ops_fn::zeros(&[hidden_size, hidden_size], DataType::Float32, device)?,
414            output: ops_fn::zeros(&[hidden_size, hidden_size], DataType::Float32, device)?,
415            hidden_size,
416        })
417    }
418
419    fn forward(&self, x: &Tensor, _position: usize) -> Result<Tensor> {
420        // Simplified forward without state (uses zeros for prev_x)
421        let x_candle = x.to_candle()?;
422        let zeros = candle_core::Tensor::zeros(x_candle.dims(), x_candle.dtype(), x_candle.device())?;
423
424        // Time mixing
425        let mix_k = self.time_mix_k.to_candle()?;
426        let mix_v = self.time_mix_v.to_candle()?;
427        let mix_r = self.time_mix_r.to_candle()?;
428
429        // xk = x * time_mix_k + prev_x * (1 - time_mix_k)
430        let one_minus_k = candle_core::Tensor::ones_like(&mix_k)?.sub(&mix_k)?;
431        let one_minus_v = candle_core::Tensor::ones_like(&mix_v)?.sub(&mix_v)?;
432        let one_minus_r = candle_core::Tensor::ones_like(&mix_r)?.sub(&mix_r)?;
433        let xk = x_candle.broadcast_mul(&mix_k)?.add(&zeros.broadcast_mul(&one_minus_k)?)?;
434        let xv = x_candle.broadcast_mul(&mix_v)?.add(&zeros.broadcast_mul(&one_minus_v)?)?;
435        let xr = x_candle.broadcast_mul(&mix_r)?.add(&zeros.broadcast_mul(&one_minus_r)?)?;
436
437        // Project
438        let k_proj = self.key.to_candle()?;
439        let v_proj = self.value.to_candle()?;
440        let r_proj = self.receptance.to_candle()?;
441        let o_proj = self.output.to_candle()?;
442
443        let k = xk.matmul(&k_proj)?;
444        let v = xv.matmul(&v_proj)?;
445        let r = xr.matmul(&r_proj)?;
446
447        // Receptance gating
448        let r_sigmoid = candle_nn::ops::sigmoid(&r)?;
449
450        // WKV computation (simplified - no state)
451        let w = self.time_decay.to_candle()?.neg()?.exp()?;
452        let u = self.time_first.to_candle()?;
453
454        let wkv = {
455            let ek = (u.broadcast_add(&k)?).exp()?;
456            let numerator = ek.broadcast_mul(&v)?;
457            let denominator = ek;
458            numerator.broadcast_div(&denominator.broadcast_add(&candle_core::Tensor::ones_like(&denominator)?)?)?
459        };
460
461        // Output
462        let output = r_sigmoid.broadcast_mul(&wkv)?;
463        let output = output.matmul(&o_proj)?;
464
465        Ok(Tensor::from_candle(output))
466    }
467
468    fn forward_with_state(&self, x: &Tensor, state: &mut Rwkv4State) -> Result<Tensor> {
469        let x_candle = x.to_candle()?;
470        let prev_x = state.prev_x_tm.to_candle()?;
471
472        // Time mixing with previous x
473        let mix_k = self.time_mix_k.to_candle()?;
474        let mix_v = self.time_mix_v.to_candle()?;
475        let mix_r = self.time_mix_r.to_candle()?;
476
477        let one_minus_k = candle_core::Tensor::ones_like(&mix_k)?.sub(&mix_k)?;
478        let one_minus_v = candle_core::Tensor::ones_like(&mix_v)?.sub(&mix_v)?;
479        let one_minus_r = candle_core::Tensor::ones_like(&mix_r)?.sub(&mix_r)?;
480
481        let xk = x_candle.broadcast_mul(&mix_k)?.add(&prev_x.broadcast_mul(&one_minus_k)?)?;
482        let xv = x_candle.broadcast_mul(&mix_v)?.add(&prev_x.broadcast_mul(&one_minus_v)?)?;
483        let xr = x_candle.broadcast_mul(&mix_r)?.add(&prev_x.broadcast_mul(&one_minus_r)?)?;
484
485        // Update state
486        state.prev_x_tm = Tensor::from_candle(x_candle.clone());
487
488        // Project
489        let k_proj = self.key.to_candle()?;
490        let v_proj = self.value.to_candle()?;
491        let r_proj = self.receptance.to_candle()?;
492        let o_proj = self.output.to_candle()?;
493
494        let k = xk.matmul(&k_proj)?;
495        let v = xv.matmul(&v_proj)?;
496        let r = xr.matmul(&r_proj)?;
497
498        // Receptance gating
499        let r_sigmoid = candle_nn::ops::sigmoid(&r)?;
500
501        // WKV computation with state
502        let w = self.time_decay.to_candle()?.neg()?.exp()?;
503        let u = self.time_first.to_candle()?;
504
505        let num_prev = state.num.to_candle()?;
506        let den_prev = state.den.to_candle()?;
507
508        // e^(u + k)
509        let ek = (u.broadcast_add(&k)?).exp()?;
510
511        // numerator = e^(u+k) * v + w * num_prev
512        let numerator = ek.broadcast_mul(&v)?.add(&w.broadcast_mul(&num_prev)?)?;
513
514        // denominator = e^(u+k) + w * den_prev
515        let denominator = ek.add(&w.broadcast_mul(&den_prev)?)?;
516
517        let wkv = numerator.broadcast_div(&denominator.broadcast_add(&candle_core::Tensor::full(1e-8f32, denominator.dims(), denominator.device())?)?)?;
518
519        // Update WKV state
520        let ek_simple = k.exp()?;
521        state.num = Tensor::from_candle(w.broadcast_mul(&num_prev)?.add(&ek_simple.broadcast_mul(&v)?)?);
522        state.den = Tensor::from_candle(w.broadcast_mul(&den_prev)?.add(&ek_simple)?);
523
524        // Output
525        let output = r_sigmoid.broadcast_mul(&wkv)?;
526        let output = output.matmul(&o_proj)?;
527
528        Ok(Tensor::from_candle(output))
529    }
530
531    fn load_weights(&mut self, weights: &ModelWeights, layer_idx: usize) -> Result<()> {
532        let prefix = format!("blocks.{}.att", layer_idx);
533
534        if let Some(w) = weights.get(&format!("{}.time_decay", prefix)) {
535            self.time_decay = w.clone();
536        }
537        if let Some(w) = weights.get(&format!("{}.time_first", prefix)) {
538            self.time_first = w.clone();
539        }
540        if let Some(w) = weights.get(&format!("{}.time_mix_k", prefix)) {
541            self.time_mix_k = w.clone();
542        }
543        if let Some(w) = weights.get(&format!("{}.time_mix_v", prefix)) {
544            self.time_mix_v = w.clone();
545        }
546        if let Some(w) = weights.get(&format!("{}.time_mix_r", prefix)) {
547            self.time_mix_r = w.clone();
548        }
549        if let Some(w) = weights.get(&format!("{}.key.weight", prefix)) {
550            self.key = ops_fn::transpose(w)?;
551        }
552        if let Some(w) = weights.get(&format!("{}.value.weight", prefix)) {
553            self.value = ops_fn::transpose(w)?;
554        }
555        if let Some(w) = weights.get(&format!("{}.receptance.weight", prefix)) {
556            self.receptance = ops_fn::transpose(w)?;
557        }
558        if let Some(w) = weights.get(&format!("{}.output.weight", prefix)) {
559            self.output = ops_fn::transpose(w)?;
560        }
561
562        Ok(())
563    }
564
565    fn to_device(&mut self, device: &Device) -> Result<()> {
566        self.time_decay = self.time_decay.to_device(device)?;
567        self.time_first = self.time_first.to_device(device)?;
568        self.time_mix_k = self.time_mix_k.to_device(device)?;
569        self.time_mix_v = self.time_mix_v.to_device(device)?;
570        self.time_mix_r = self.time_mix_r.to_device(device)?;
571        self.key = self.key.to_device(device)?;
572        self.value = self.value.to_device(device)?;
573        self.receptance = self.receptance.to_device(device)?;
574        self.output = self.output.to_device(device)?;
575        Ok(())
576    }
577}
578
579impl Rwkv4ChannelMixing {
580    fn new(config: &Rwkv4Config, device: &Device) -> Result<Self> {
581        let hidden_size = config.hidden_size;
582        let intermediate_size = config.effective_intermediate_size();
583
584        Ok(Self {
585            time_mix_k: ops_fn::zeros(&[hidden_size], DataType::Float32, device)?,
586            time_mix_r: ops_fn::zeros(&[hidden_size], DataType::Float32, device)?,
587            key: ops_fn::zeros(&[hidden_size, intermediate_size], DataType::Float32, device)?,
588            value: ops_fn::zeros(&[intermediate_size, hidden_size], DataType::Float32, device)?,
589            receptance: ops_fn::zeros(&[hidden_size, hidden_size], DataType::Float32, device)?,
590            hidden_size,
591            intermediate_size,
592        })
593    }
594
595    fn forward(&self, x: &Tensor, _position: usize) -> Result<Tensor> {
596        let x_candle = x.to_candle()?;
597        let zeros = candle_core::Tensor::zeros(x_candle.dims(), x_candle.dtype(), x_candle.device())?;
598
599        // Channel mixing with zeros for prev_x
600        let mix_k = self.time_mix_k.to_candle()?;
601        let mix_r = self.time_mix_r.to_candle()?;
602
603        let one_minus_k = candle_core::Tensor::ones_like(&mix_k)?.sub(&mix_k)?;
604        let one_minus_r = candle_core::Tensor::ones_like(&mix_r)?.sub(&mix_r)?;
605        let xk = x_candle.broadcast_mul(&mix_k)?.add(&zeros.broadcast_mul(&one_minus_k)?)?;
606        let xr = x_candle.broadcast_mul(&mix_r)?.add(&zeros.broadcast_mul(&one_minus_r)?)?;
607
608        // Project
609        let k_proj = self.key.to_candle()?;
610        let v_proj = self.value.to_candle()?;
611        let r_proj = self.receptance.to_candle()?;
612
613        let k = xk.matmul(&k_proj)?;
614        let r = xr.matmul(&r_proj)?;
615
616        // Squared ReLU for k
617        let k_relu = k.relu()?;
618        let k_squared = k_relu.sqr()?;
619
620        // Value projection
621        let v = k_squared.matmul(&v_proj)?;
622
623        // Receptance gating
624        let r_sigmoid = candle_nn::ops::sigmoid(&r)?;
625        let output = r_sigmoid.broadcast_mul(&v)?;
626
627        Ok(Tensor::from_candle(output))
628    }
629
630    fn forward_with_state(&self, x: &Tensor, state: &mut Rwkv4State) -> Result<Tensor> {
631        let x_candle = x.to_candle()?;
632        let prev_x = state.prev_x_cm.to_candle()?;
633
634        // Channel mixing with state
635        let mix_k = self.time_mix_k.to_candle()?;
636        let mix_r = self.time_mix_r.to_candle()?;
637
638        let one_minus_k = candle_core::Tensor::ones_like(&mix_k)?.sub(&mix_k)?;
639        let one_minus_r = candle_core::Tensor::ones_like(&mix_r)?.sub(&mix_r)?;
640
641        let xk = x_candle.broadcast_mul(&mix_k)?.add(&prev_x.broadcast_mul(&one_minus_k)?)?;
642        let xr = x_candle.broadcast_mul(&mix_r)?.add(&prev_x.broadcast_mul(&one_minus_r)?)?;
643
644        // Update state
645        state.prev_x_cm = Tensor::from_candle(x_candle.clone());
646
647        // Project
648        let k_proj = self.key.to_candle()?;
649        let v_proj = self.value.to_candle()?;
650        let r_proj = self.receptance.to_candle()?;
651
652        let k = xk.matmul(&k_proj)?;
653        let r = xr.matmul(&r_proj)?;
654
655        // Squared ReLU
656        let k_relu = k.relu()?;
657        let k_squared = k_relu.sqr()?;
658
659        // Value projection
660        let v = k_squared.matmul(&v_proj)?;
661
662        // Receptance gating
663        let r_sigmoid = candle_nn::ops::sigmoid(&r)?;
664        let output = r_sigmoid.broadcast_mul(&v)?;
665
666        Ok(Tensor::from_candle(output))
667    }
668
669    fn load_weights(&mut self, weights: &ModelWeights, layer_idx: usize) -> Result<()> {
670        let prefix = format!("blocks.{}.ffn", layer_idx);
671
672        if let Some(w) = weights.get(&format!("{}.time_mix_k", prefix)) {
673            self.time_mix_k = w.clone();
674        }
675        if let Some(w) = weights.get(&format!("{}.time_mix_r", prefix)) {
676            self.time_mix_r = w.clone();
677        }
678        if let Some(w) = weights.get(&format!("{}.key.weight", prefix)) {
679            self.key = ops_fn::transpose(w)?;
680        }
681        if let Some(w) = weights.get(&format!("{}.value.weight", prefix)) {
682            self.value = ops_fn::transpose(w)?;
683        }
684        if let Some(w) = weights.get(&format!("{}.receptance.weight", prefix)) {
685            self.receptance = ops_fn::transpose(w)?;
686        }
687
688        Ok(())
689    }
690
691    fn to_device(&mut self, device: &Device) -> Result<()> {
692        self.time_mix_k = self.time_mix_k.to_device(device)?;
693        self.time_mix_r = self.time_mix_r.to_device(device)?;
694        self.key = self.key.to_device(device)?;
695        self.value = self.value.to_device(device)?;
696        self.receptance = self.receptance.to_device(device)?;
697        Ok(())
698    }
699}
700
701#[cfg(test)]
702mod tests {
703    use super::*;
704
705    #[test]
706    fn test_rwkv4_config() {
707        let config = Rwkv4Config::default();
708        assert_eq!(config.vocab_size, 50277);
709        assert_eq!(config.hidden_size, 768);
710        assert_eq!(config.effective_intermediate_size(), 768 * 4);
711    }
712
713    #[test]
714    fn test_rwkv4_model_creation() {
715        let config = Rwkv4Config {
716            vocab_size: 1000,
717            hidden_size: 64,
718            num_hidden_layers: 2,
719            ..Default::default()
720        };
721
722        let model = Rwkv4ModelV2::new(config).unwrap();
723        assert_eq!(model.config().vocab_size(), 1000);
724        assert_eq!(model.config().hidden_size(), 64);
725        assert_eq!(model.config().num_layers(), 2);
726    }
727
728    #[test]
729    fn test_rwkv4_forward_pass() {
730        let config = Rwkv4Config {
731            vocab_size: 100,
732            hidden_size: 32,
733            num_hidden_layers: 1,
734            ..Default::default()
735        };
736
737        let model = Rwkv4ModelV2::new(config).unwrap();
738        let input_ids = ops_fn::zeros(&[1, 4], DataType::Int64, &Device::CPU).unwrap();
739        let inputs = ModelInputs::text(input_ids);
740
741        let outputs = model.forward(&inputs).unwrap();
742        match outputs {
743            ModelOutputs::Logits { logits, .. } => {
744                assert_eq!(logits.shape(), &[1, 4, 100]);
745            }
746            _ => panic!("Expected logits output"),
747        }
748    }
749}