Skip to main content

runtime/models_v2/
rwkv6.rs

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