1use crate::model_config;
12use super::traits::*;
13use anyhow::Result;
14use serde::{Serialize, Deserialize};
15
16model_config!(Rwkv6Config {
18 vocab_size: usize = 65536,
19 hidden_size: usize = 2048,
20 num_hidden_layers: usize = 24,
21 intermediate_size: usize = 0, head_size: usize = 64, num_heads: usize = 0, 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
61pub struct Rwkv6ModelV2 {
63 config: Rwkv6Config,
64 device: Device,
65 embeddings: Tensor,
66 blocks: Vec<Rwkv6Block>,
67 ln_out: Tensor,
68 head: Tensor,
69}
70
71pub 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
81pub struct Rwkv6TimeMixing {
83 time_maa_x: Tensor, time_maa_w: Tensor, time_maa_k: Tensor, time_maa_v: Tensor, time_maa_r: Tensor, time_maa_g: Tensor, time_decay: Tensor, time_decay_w1: Tensor, time_decay_w2: Tensor, time_first: Tensor, receptance: Tensor,
99 key: Tensor,
100 value: Tensor,
101 output: Tensor,
102 gate: Tensor, ln_x: Tensor,
106
107 hidden_size: usize,
108 head_size: usize,
109 num_heads: usize,
110}
111
112pub 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#[derive(Clone)]
127pub struct Rwkv6State {
128 pub wkv_state: Tensor,
130 pub prev_x_tm: Tensor,
132 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 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 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 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 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; 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 let x_candle = x.to_candle()?;
451
452 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 let r_sigmoid = candle_nn::ops::sigmoid(&r)?;
466 let g_silu = candle_nn::ops::silu(&g)?;
467
468 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 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 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 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 let r_sigmoid = candle_nn::ops::sigmoid(&r)?;
508 let g_silu = candle_nn::ops::silu(&g)?;
509
510 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 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 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 let ek = k_heads.exp()?;
528 let kv_outer = ek.unsqueeze(3)?.matmul(&v_heads.unsqueeze(2)?)?; 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 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 let wkv_flat = wkv_out.reshape(&[batch_size, self.hidden_size])?;
543
544 let ln_x = self.ln_x.to_candle()?;
546 let wkv_normed = wkv_flat.broadcast_mul(&ln_x)?;
547
548 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}