1use crate::model_config;
12use super::traits::*;
13use anyhow::Result;
14use serde::{Serialize, Deserialize};
15
16model_config!(Rwkv4Config {
18 vocab_size: usize = 50277,
19 hidden_size: usize = 768, num_hidden_layers: usize = 12,
21 intermediate_size: usize = 0, layer_norm_epsilon: f32 = 1e-5,
23 rescale_every: usize = 6, 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
51pub struct Rwkv4ModelV2 {
53 config: Rwkv4Config,
54 device: Device,
55 embeddings: Tensor,
56 blocks: Vec<Rwkv4Block>,
57 ln_out: Tensor,
58 head: Tensor,
59}
60
61pub 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
71pub struct Rwkv4TimeMixing {
73 time_decay: Tensor, time_first: Tensor, time_mix_k: Tensor, time_mix_v: Tensor, time_mix_r: Tensor, key: Tensor, value: Tensor, receptance: Tensor, output: Tensor, hidden_size: usize,
87}
88
89pub struct Rwkv4ChannelMixing {
91 time_mix_k: Tensor, time_mix_r: Tensor, key: Tensor, value: Tensor, receptance: Tensor, hidden_size: usize,
99 intermediate_size: usize,
100}
101
102#[derive(Clone)]
104pub struct Rwkv4State {
105 pub num: Tensor,
107 pub den: Tensor,
109 pub prev_x_tm: Tensor,
111 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 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 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 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 + self.config.hidden_size * inter_size * 2 + self.config.hidden_size * 10 )
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 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 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 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 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 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 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 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 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 let x_candle = x.to_candle()?;
422 let zeros = candle_core::Tensor::zeros(x_candle.dims(), x_candle.dtype(), x_candle.device())?;
423
424 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 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 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 let r_sigmoid = candle_nn::ops::sigmoid(&r)?;
449
450 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 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 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 state.prev_x_tm = Tensor::from_candle(x_candle.clone());
487
488 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 let r_sigmoid = candle_nn::ops::sigmoid(&r)?;
500
501 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 let ek = (u.broadcast_add(&k)?).exp()?;
510
511 let numerator = ek.broadcast_mul(&v)?.add(&w.broadcast_mul(&num_prev)?)?;
513
514 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 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 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 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 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 let k_relu = k.relu()?;
618 let k_squared = k_relu.sqr()?;
619
620 let v = k_squared.matmul(&v_proj)?;
622
623 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 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 state.prev_x_cm = Tensor::from_candle(x_candle.clone());
646
647 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 let k_relu = k.relu()?;
657 let k_squared = k_relu.sqr()?;
658
659 let v = k_squared.matmul(&v_proj)?;
661
662 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}