1use crate::model_config;
12use super::traits::*;
13use anyhow::Result;
14use serde::{Serialize, Deserialize};
15
16model_config!(MambaConfig {
18 vocab_size: usize = 50280,
19 hidden_size: usize = 768, num_hidden_layers: usize = 24, d_state: usize = 16, d_conv: usize = 4, expand: usize = 2, dt_rank: usize = 0, d_inner: usize = 0, dt_scale: f32 = 1.0,
27 dt_min: f32 = 0.001,
28 dt_max: f32 = 0.1,
29 dt_init_floor: f32 = 1e-4,
30 conv_bias: bool = true,
31 bias: bool = false,
32 layer_norm_epsilon: f32 = 1e-5,
33 rms_norm: bool = true,
34 initializer_range: f32 = 0.02,
35 tie_embeddings: bool = true,
36 pad_token_id: i64 = 0,
37 bos_token_id: i64 = 0,
38 eos_token_id: i64 = 0,
39});
40
41impl MambaConfig {
42 pub fn from_gguf_config(gguf: &crate::weight_loader_core::GGUFModelConfig) -> Self {
44 let hidden_size = gguf.hidden_size;
46 let expand = 2;
47 let d_inner = hidden_size * expand;
48 let dt_rank = ((hidden_size as f32 / 16.0).ceil() as usize).max(1);
49
50 Self {
51 vocab_size: gguf.vocab_size,
52 hidden_size,
53 num_hidden_layers: gguf.num_hidden_layers,
54 d_state: 16,
55 d_conv: 4,
56 expand,
57 dt_rank,
58 d_inner,
59 layer_norm_epsilon: gguf.rms_norm_eps,
60 ..Default::default()
61 }
62 }
63
64 pub fn effective_d_inner(&self) -> usize {
66 if self.d_inner > 0 {
67 self.d_inner
68 } else {
69 self.hidden_size * self.expand
70 }
71 }
72
73 pub fn effective_dt_rank(&self) -> usize {
75 if self.dt_rank > 0 {
76 self.dt_rank
77 } else {
78 ((self.hidden_size as f32 / 16.0).ceil() as usize).max(1)
79 }
80 }
81}
82
83pub struct MambaModelV2 {
85 config: MambaConfig,
86 device: Device,
87 backbone: MambaBackbone,
88 lm_head: Tensor,
89}
90
91pub struct MambaBackbone {
93 embeddings: Tensor,
94 layers: Vec<MambaBlock>,
95 norm_f: Tensor, config: MambaConfig,
97}
98
99pub struct MambaBlock {
101 mixer: MambaMixer,
102 norm: Tensor, config: MambaConfig,
104}
105
106pub struct MambaMixer {
108 in_proj: Tensor,
110
111 conv1d_weight: Tensor,
113 conv1d_bias: Option<Tensor>,
114
115 x_proj: Tensor, dt_proj: Tensor, dt_proj_bias: Option<Tensor>,
119
120 a_log: Tensor, d: Tensor, out_proj: Tensor, d_inner: usize,
129 d_state: usize,
130 d_conv: usize,
131 dt_rank: usize,
132}
133
134#[derive(Clone)]
136pub struct MambaState {
137 h: Tensor,
139 conv_cache: Tensor,
141}
142
143impl Model for MambaModelV2 {
144 type Config = MambaConfig;
145
146 fn new(config: MambaConfig) -> Result<Self> {
147 let device = Device::CPU;
148 let backbone = MambaBackbone::new(&config, &device)?;
149 let lm_head = if config.tie_embeddings {
150 ops_fn::zeros(&[config.vocab_size, config.hidden_size], DataType::Float32, &device)?
152 } else {
153 ops_fn::zeros(&[config.hidden_size, config.vocab_size], DataType::Float32, &device)?
154 };
155
156 Ok(Self { config, device, backbone, lm_head })
157 }
158
159 fn from_weights(config: MambaConfig, weights: ModelWeights) -> Result<Self> {
160 let mut model = Self::new(config)?;
161 model.backbone.load_weights(&weights)?;
162 if !model.config.tie_embeddings {
163 if let Some(w) = weights.get("lm_head.weight") {
164 model.lm_head = ops_fn::transpose(w)?;
165 }
166 }
167 Ok(model)
168 }
169
170 fn forward(&self, inputs: &ModelInputs) -> Result<ModelOutputs> {
171 let input_ids = match inputs {
172 ModelInputs::Text { input_ids, .. } => input_ids,
173 _ => return Err(anyhow::anyhow!("Mamba expects text input")),
174 };
175
176 let hidden_states = self.backbone.forward(input_ids, None)?;
177
178 let logits = if self.config.tie_embeddings {
180 let embed_t = ops_fn::transpose(&self.backbone.embeddings)?;
182 ops_fn::matmul(&hidden_states, &embed_t)?
183 } else {
184 ops_fn::matmul(&hidden_states, &self.lm_head)?
185 };
186
187 Ok(ModelOutputs::Logits {
188 logits,
189 hidden_states: Some(hidden_states)
190 })
191 }
192
193 fn generate(&self, prompt: &str, config: &GenerationConfig) -> Result<String> {
194 use crate::tokenizer::Tokenizer;
195 use rand::Rng;
196
197 let tokenizer = Tokenizer::new();
199 let mut tokens: Vec<u32> = tokenizer.encode(prompt);
200
201 let batch_size = 1;
203 let d_inner = self.config.effective_d_inner();
204 let d_state = self.config.d_state;
205 let d_conv = self.config.d_conv;
206
207 let mut layer_states: Vec<MambaState> = Vec::new();
208 for _ in 0..self.config.num_hidden_layers {
209 layer_states.push(MambaState {
210 h: ops_fn::zeros(&[batch_size, d_inner, d_state], DataType::Float32, &self.device)?,
211 conv_cache: ops_fn::zeros(&[batch_size, d_inner, d_conv - 1], DataType::Float32, &self.device)?,
212 });
213 }
214
215 for &token in &tokens[..tokens.len().saturating_sub(1)] {
218 let input_tensor = Tensor::from_i64_slice(&[token as i64], &[1, 1], &self.device)?;
219 self.backbone.forward_with_state(&input_tensor, &mut layer_states)?;
221 }
222
223 for _ in 0..config.max_new_tokens {
225 let last_token = *tokens.last().unwrap_or(&0);
227 let input_tensor = Tensor::from_i64_slice(&[last_token as i64], &[1, 1], &self.device)?;
228
229 let hidden_states = self.backbone.forward_with_state(&input_tensor, &mut layer_states)?;
231
232 let logits = if self.config.tie_embeddings {
234 let embed_t = ops_fn::transpose(&self.backbone.embeddings)?;
235 ops_fn::matmul(&hidden_states, &embed_t)?
236 } else {
237 ops_fn::matmul(&hidden_states, &self.lm_head)?
238 };
239
240 let logits_candle = logits.to_candle()?;
242 let shape = logits_candle.dims();
243
244 let last_logits = if shape.len() == 3 {
246 logits_candle.squeeze(1)?.squeeze(0)?
247 } else if shape.len() == 2 {
248 logits_candle.squeeze(0)?
249 } else {
250 logits_candle.clone()
251 };
252
253 let logits_vec: Vec<f32> = last_logits.to_vec1()?;
254
255 let next_token = if config.do_sample && config.temperature > 0.0 {
256 let scaled: Vec<f32> = logits_vec.iter()
258 .map(|&x| x / config.temperature)
259 .collect();
260
261 let max_val = scaled.iter().cloned().fold(f32::NEG_INFINITY, f32::max);
263 let exp_sum: f32 = scaled.iter().map(|&x| (x - max_val).exp()).sum();
264 let probs: Vec<f32> = scaled.iter()
265 .map(|&x| (x - max_val).exp() / exp_sum)
266 .collect();
267
268 let mut rng = rand::thread_rng();
270 let random_val: f32 = rng.gen();
271 let mut cumulative = 0.0;
272 let mut sampled = 0u32;
273
274 for (idx, &prob) in probs.iter().enumerate() {
275 cumulative += prob;
276 if random_val <= cumulative {
277 sampled = idx as u32;
278 break;
279 }
280 }
281 sampled
282 } else {
283 let mut max_idx = 0;
285 let mut max_val = logits_vec[0];
286 for (idx, &val) in logits_vec.iter().enumerate() {
287 if val > max_val {
288 max_val = val;
289 max_idx = idx;
290 }
291 }
292 max_idx as u32
293 };
294
295 if next_token == config.eos_token_id {
297 break;
298 }
299
300 tokens.push(next_token);
302 }
303
304 Ok(tokenizer.decode(&tokens))
306 }
307
308 fn config(&self) -> &Self::Config { &self.config }
309
310 fn memory_requirements(&self) -> MemoryRequirements {
311 let d_inner = self.config.effective_d_inner();
312 let param_size = (self.config.vocab_size * self.config.hidden_size +
313 self.config.num_hidden_layers * (
314 self.config.hidden_size * d_inner * 2 +
316 d_inner * self.config.d_conv +
318 d_inner * (self.config.effective_dt_rank() + self.config.d_state * 2) +
320 self.config.effective_dt_rank() * d_inner +
322 d_inner * self.config.d_state + d_inner +
324 d_inner * self.config.hidden_size
326 )) * 4;
327
328 let state_size = self.config.num_hidden_layers *
330 (d_inner * self.config.d_state + d_inner * (self.config.d_conv - 1)) * 4;
331
332 MemoryRequirements {
333 gpu_memory: param_size,
334 cpu_memory: param_size / 4,
335 kv_cache_memory: state_size, peak_memory: param_size + param_size / 2,
337 }
338 }
339
340 fn to_device(&mut self, device: &Device) -> Result<()> {
341 self.device = device.clone();
342 self.backbone.to_device(device)?;
343 if !self.config.tie_embeddings {
344 self.lm_head = self.lm_head.to_device(device)?;
345 }
346 Ok(())
347 }
348}
349
350impl MambaBackbone {
351 fn new(config: &MambaConfig, device: &Device) -> Result<Self> {
352 let embeddings = ops_fn::zeros(&[config.vocab_size, config.hidden_size], DataType::Float32, device)?;
353 let norm_f = ops_fn::zeros(&[config.hidden_size], DataType::Float32, device)?;
354
355 let mut layers = Vec::new();
356 for _ in 0..config.num_hidden_layers {
357 layers.push(MambaBlock::new(config, device)?);
358 }
359
360 Ok(Self { embeddings, layers, norm_f, config: config.clone() })
361 }
362
363 fn forward(&self, input_ids: &Tensor, states: Option<&mut Vec<MambaState>>) -> Result<Tensor> {
364 let mut hidden_states = ops_fn::embedding(input_ids, &self.embeddings)?;
366
367 match states {
369 Some(layer_states) => {
370 for (i, layer) in self.layers.iter().enumerate() {
371 hidden_states = layer.forward_with_state(&hidden_states, &mut layer_states[i])?;
372 }
373 }
374 None => {
375 for layer in &self.layers {
376 hidden_states = layer.forward(&hidden_states)?;
377 }
378 }
379 }
380
381 if self.config.rms_norm {
383 ops_fn::rms_norm(&hidden_states, &self.norm_f, self.config.layer_norm_epsilon)
384 } else {
385 ops_fn::layer_norm(&hidden_states, &self.norm_f, None, self.config.layer_norm_epsilon)
386 }
387 }
388
389 fn forward_with_state(&self, input_ids: &Tensor, states: &mut Vec<MambaState>) -> Result<Tensor> {
390 self.forward(input_ids, Some(states))
391 }
392
393 fn load_weights(&mut self, weights: &ModelWeights) -> Result<()> {
394 if let Some(w) = weights.get("backbone.embeddings.weight")
396 .or_else(|| weights.get("backbone.embedding.weight"))
397 .or_else(|| weights.get("model.embed_tokens.weight"))
398 {
399 self.embeddings = w.clone();
400 }
401
402 if let Some(w) = weights.get("backbone.norm_f.weight")
403 .or_else(|| weights.get("backbone.final_layernorm.weight"))
404 .or_else(|| weights.get("model.norm.weight"))
405 {
406 self.norm_f = w.clone();
407 }
408
409 for (i, layer) in self.layers.iter_mut().enumerate() {
410 layer.load_weights(weights, i)?;
411 }
412
413 Ok(())
414 }
415
416 fn to_device(&mut self, device: &Device) -> Result<()> {
417 self.embeddings = self.embeddings.to_device(device)?;
418 self.norm_f = self.norm_f.to_device(device)?;
419 for layer in &mut self.layers {
420 layer.to_device(device)?;
421 }
422 Ok(())
423 }
424}
425
426impl MambaBlock {
427 fn new(config: &MambaConfig, device: &Device) -> Result<Self> {
428 let norm = ops_fn::zeros(&[config.hidden_size], DataType::Float32, device)?;
429 let mixer = MambaMixer::new(config, device)?;
430
431 Ok(Self { mixer, norm, config: config.clone() })
432 }
433
434 fn forward(&self, hidden_states: &Tensor) -> Result<Tensor> {
435 let residual = hidden_states.clone();
436
437 let normalized = if self.config.rms_norm {
439 ops_fn::rms_norm(hidden_states, &self.norm, self.config.layer_norm_epsilon)?
440 } else {
441 ops_fn::layer_norm(hidden_states, &self.norm, None, self.config.layer_norm_epsilon)?
442 };
443
444 let mixed = self.mixer.forward(&normalized)?;
446
447 ops_fn::add(&residual, &mixed)
449 }
450
451 fn forward_with_state(&self, hidden_states: &Tensor, state: &mut MambaState) -> Result<Tensor> {
452 let residual = hidden_states.clone();
453
454 let normalized = if self.config.rms_norm {
456 ops_fn::rms_norm(hidden_states, &self.norm, self.config.layer_norm_epsilon)?
457 } else {
458 ops_fn::layer_norm(hidden_states, &self.norm, None, self.config.layer_norm_epsilon)?
459 };
460
461 let mixed = self.mixer.forward_with_state(&normalized, state)?;
463
464 ops_fn::add(&residual, &mixed)
466 }
467
468 fn load_weights(&mut self, weights: &ModelWeights, layer_idx: usize) -> Result<()> {
469 let prefix = format!("backbone.layers.{}", layer_idx);
470
471 if let Some(w) = weights.get(&format!("{}.norm.weight", prefix)) {
472 self.norm = w.clone();
473 }
474
475 self.mixer.load_weights(weights, layer_idx)?;
476 Ok(())
477 }
478
479 fn to_device(&mut self, device: &Device) -> Result<()> {
480 self.norm = self.norm.to_device(device)?;
481 self.mixer.to_device(device)?;
482 Ok(())
483 }
484}
485
486impl MambaMixer {
487 fn new(config: &MambaConfig, device: &Device) -> Result<Self> {
488 let d_inner = config.effective_d_inner();
489 let dt_rank = config.effective_dt_rank();
490 let d_state = config.d_state;
491 let d_conv = config.d_conv;
492
493 let in_proj = ops_fn::zeros(&[config.hidden_size, d_inner * 2], DataType::Float32, device)?;
495
496 let conv1d_weight = ops_fn::zeros(&[d_inner, d_conv], DataType::Float32, device)?;
498 let conv1d_bias = if config.conv_bias {
499 Some(ops_fn::zeros(&[d_inner], DataType::Float32, device)?)
500 } else {
501 None
502 };
503
504 let x_proj = ops_fn::zeros(&[d_inner, dt_rank + d_state * 2], DataType::Float32, device)?;
506
507 let dt_proj = ops_fn::zeros(&[dt_rank, d_inner], DataType::Float32, device)?;
509 let dt_proj_bias = if config.bias {
510 Some(ops_fn::zeros(&[d_inner], DataType::Float32, device)?)
511 } else {
512 None
513 };
514
515 let a_log = ops_fn::zeros(&[d_inner, d_state], DataType::Float32, device)?;
517
518 let d = ops_fn::zeros(&[d_inner], DataType::Float32, device)?;
520
521 let out_proj = ops_fn::zeros(&[d_inner, config.hidden_size], DataType::Float32, device)?;
523
524 Ok(Self {
525 in_proj,
526 conv1d_weight,
527 conv1d_bias,
528 x_proj,
529 dt_proj,
530 dt_proj_bias,
531 a_log,
532 d,
533 out_proj,
534 d_inner,
535 d_state,
536 d_conv,
537 dt_rank,
538 })
539 }
540
541 fn forward(&self, hidden_states: &Tensor) -> Result<Tensor> {
542 let shape = hidden_states.shape();
543 let (batch_size, seq_len, _d_model) = if shape.len() == 3 {
544 (shape[0], shape[1], shape[2])
545 } else if shape.len() == 2 {
546 (1, shape[0], shape[1])
547 } else {
548 return Err(anyhow::anyhow!("Invalid hidden_states shape: {:?}", shape));
549 };
550
551 let projected = ops_fn::matmul(hidden_states, &self.in_proj)?;
553
554 let (x, z) = self.split_xz(&projected)?;
556
557 let x_conv = self.apply_conv1d(&x, batch_size, seq_len)?;
559
560 let x_act = ops_fn::silu(&x_conv)?;
562
563 let y = self.selective_scan(&x_act, batch_size, seq_len)?;
565
566 let z_act = ops_fn::silu(&z)?;
568 let gated = ops_fn::mul(&y, &z_act)?;
569
570 ops_fn::matmul(&gated, &self.out_proj)
572 }
573
574 fn forward_with_state(&self, hidden_states: &Tensor, state: &mut MambaState) -> Result<Tensor> {
575 let shape = hidden_states.shape();
576 let (_batch_size, seq_len, _d_model) = if shape.len() == 3 {
577 (shape[0], shape[1], shape[2])
578 } else if shape.len() == 2 {
579 (1, shape[0], shape[1])
580 } else {
581 return Err(anyhow::anyhow!("Invalid hidden_states shape: {:?}", shape));
582 };
583
584 if seq_len == 1 {
586 return self.forward_step(hidden_states, state);
587 }
588
589 self.forward(hidden_states)
591 }
592
593 fn forward_step(&self, hidden_states: &Tensor, state: &mut MambaState) -> Result<Tensor> {
595 let projected = ops_fn::matmul(hidden_states, &self.in_proj)?;
597
598 let (x, z) = self.split_xz(&projected)?;
600
601 let x_conv = self.apply_conv1d_step(&x, state)?;
603
604 let x_act = ops_fn::silu(&x_conv)?;
606
607 let y = self.selective_scan_step(&x_act, state)?;
609
610 let z_act = ops_fn::silu(&z)?;
612 let gated = ops_fn::mul(&y, &z_act)?;
613
614 ops_fn::matmul(&gated, &self.out_proj)
616 }
617
618 fn split_xz(&self, projected: &Tensor) -> Result<(Tensor, Tensor)> {
620 let candle_tensor = projected.to_candle()?;
621 let dims = candle_tensor.dims();
622 let last_dim = dims.len() - 1;
623
624 let x_candle = candle_tensor.narrow(last_dim, 0, self.d_inner)?;
626 let z_candle = candle_tensor.narrow(last_dim, self.d_inner, self.d_inner)?;
627
628 Ok((Tensor::from_candle(x_candle), Tensor::from_candle(z_candle)))
629 }
630
631 fn apply_conv1d(&self, x: &Tensor, batch_size: usize, seq_len: usize) -> Result<Tensor> {
633 let x_candle = x.to_candle()?;
640 let w_candle = self.conv1d_weight.to_candle()?;
641
642 let pad_len = self.d_conv - 1;
644 let zeros_shape = [batch_size, pad_len, self.d_inner];
645 let zeros = candle_core::Tensor::zeros(&zeros_shape, x_candle.dtype(), x_candle.device())?;
646
647 let x_3d = if x_candle.dims().len() == 2 {
649 x_candle.unsqueeze(0)?
650 } else {
651 x_candle.clone()
652 };
653
654 let x_padded = candle_core::Tensor::cat(&[&zeros, &x_3d], 1)?;
656
657 let mut outputs = Vec::new();
660
661 for i in 0..seq_len {
662 let window = x_padded.narrow(1, i, self.d_conv)?;
664
665 let window_t = window.transpose(1, 2)?;
667
668 let conv_out = window_t.broadcast_mul(&w_candle)?;
670 let summed = conv_out.sum(2)?; outputs.push(summed);
673 }
674
675 let result = candle_core::Tensor::stack(&outputs, 1)?;
677
678 let result = if let Some(ref bias) = self.conv1d_bias {
680 let b_candle = bias.to_candle()?;
681 result.broadcast_add(&b_candle)?
682 } else {
683 result
684 };
685
686 Ok(Tensor::from_candle(result))
687 }
688
689 fn apply_conv1d_step(&self, x: &Tensor, state: &mut MambaState) -> Result<Tensor> {
691 let x_candle = x.to_candle()?;
693 let x_squeezed = x_candle.squeeze(1)?; let cache_candle = state.conv_cache.to_candle()?;
697
698 if self.d_conv > 1 {
701 let shifted = if self.d_conv > 2 {
702 cache_candle.narrow(2, 1, self.d_conv - 2)?
703 } else {
704 candle_core::Tensor::zeros(&[cache_candle.dims()[0], self.d_inner, 0], cache_candle.dtype(), cache_candle.device())?
706 };
707
708 let x_expanded = x_squeezed.unsqueeze(2)?;
710
711 state.conv_cache = if shifted.dims()[2] > 0 {
713 let new_cache = candle_core::Tensor::cat(&[&shifted, &x_expanded], 2)?;
714 Tensor::from_candle(new_cache)
715 } else {
716 Tensor::from_candle(x_expanded)
717 };
718 }
719
720 let w_candle = self.conv1d_weight.to_candle()?;
722
723 let cache_for_conv = state.conv_cache.to_candle()?;
725 let x_for_cat = x_squeezed.unsqueeze(2)?;
726 let full_window = candle_core::Tensor::cat(&[&cache_for_conv, &x_for_cat], 2)?;
727
728 let conv_out = full_window.broadcast_mul(&w_candle)?;
730 let result = conv_out.sum(2)?; let result = if let Some(ref bias) = self.conv1d_bias {
734 let b_candle = bias.to_candle()?;
735 result.broadcast_add(&b_candle)?
736 } else {
737 result
738 };
739
740 Ok(Tensor::from_candle(result.unsqueeze(1)?))
742 }
743
744 fn selective_scan(&self, x: &Tensor, batch_size: usize, seq_len: usize) -> Result<Tensor> {
746 let dbc = ops_fn::matmul(x, &self.x_proj)?;
750 let dbc_candle = dbc.to_candle()?;
751
752 let dt_raw = dbc_candle.narrow(2, 0, self.dt_rank)?;
754 let b = dbc_candle.narrow(2, self.dt_rank, self.d_state)?;
755 let c = dbc_candle.narrow(2, self.dt_rank + self.d_state, self.d_state)?;
756
757 let dt_proj_candle = self.dt_proj.to_candle()?;
760 let dt = dt_raw.broadcast_matmul(&dt_proj_candle)?;
761
762 let dt = if let Some(ref bias) = self.dt_proj_bias {
764 let b_candle = bias.to_candle()?;
765 dt.broadcast_add(&b_candle)?
766 } else {
767 dt
768 };
769
770 let dt = softplus(&dt)?;
772
773 let a_log_candle = self.a_log.to_candle()?;
775 let a = a_log_candle.exp()?.neg()?;
776
777 let x_candle = x.to_candle()?;
779 let d_candle = self.d.to_candle()?;
780
781 let mut h = candle_core::Tensor::zeros(&[batch_size, self.d_inner, self.d_state], candle_core::DType::F32, x_candle.device())?;
783
784 let mut outputs = Vec::new();
785
786 for t in 0..seq_len {
787 let x_t = x_candle.narrow(1, t, 1)?.squeeze(1)?; let dt_t = dt.narrow(1, t, 1)?.squeeze(1)?; let b_t = b.narrow(1, t, 1)?.squeeze(1)?; let c_t = c.narrow(1, t, 1)?.squeeze(1)?; let dt_expanded = dt_t.unsqueeze(2)?; let dt_a = dt_expanded.broadcast_mul(&a)?;
797 let a_bar = dt_a.exp()?; let b_expanded = b_t.unsqueeze(1)?; let dt_b = dt_expanded.broadcast_mul(&b_expanded)?; let x_expanded = x_t.unsqueeze(2)?; let ah = a_bar.mul(&h)?;
808 let bx = dt_b.mul(&x_expanded.broadcast_as(dt_b.dims())?)?;
809 h = ah.add(&bx)?;
810
811 let c_expanded = c_t.unsqueeze(1)?; let y_state = h.mul(&c_expanded.broadcast_as(h.dims())?)?.sum(2)?; let y_skip = x_t.broadcast_mul(&d_candle)?;
818 let y_t = y_state.add(&y_skip)?;
819
820 outputs.push(y_t);
821 }
822
823 let result = candle_core::Tensor::stack(&outputs, 1)?;
825
826 Ok(Tensor::from_candle(result))
827 }
828
829 fn selective_scan_step(&self, x: &Tensor, state: &mut MambaState) -> Result<Tensor> {
831 let x_candle = x.to_candle()?;
833 let x_t = x_candle.squeeze(1)?; let dbc = ops_fn::matmul(x, &self.x_proj)?;
837 let dbc_candle = dbc.to_candle()?.squeeze(1)?; let dt_raw = dbc_candle.narrow(1, 0, self.dt_rank)?;
840 let b_t = dbc_candle.narrow(1, self.dt_rank, self.d_state)?;
841 let c_t = dbc_candle.narrow(1, self.dt_rank + self.d_state, self.d_state)?;
842
843 let dt_proj_candle = self.dt_proj.to_candle()?;
845 let dt_t = dt_raw.matmul(&dt_proj_candle)?;
846
847 let dt_t = if let Some(ref bias) = self.dt_proj_bias {
848 let b_candle = bias.to_candle()?;
849 dt_t.broadcast_add(&b_candle)?
850 } else {
851 dt_t
852 };
853
854 let dt_t = softplus(&dt_t)?;
855
856 let a_log_candle = self.a_log.to_candle()?;
858 let a = a_log_candle.exp()?.neg()?;
859
860 let dt_expanded = dt_t.unsqueeze(2)?;
862 let dt_a = dt_expanded.broadcast_mul(&a)?;
863 let a_bar = dt_a.exp()?;
864
865 let b_expanded = b_t.unsqueeze(1)?;
866 let dt_b = dt_expanded.broadcast_mul(&b_expanded)?;
867
868 let h_candle = state.h.to_candle()?;
870 let x_expanded = x_t.unsqueeze(2)?;
871
872 let ah = a_bar.mul(&h_candle)?;
873 let bx = dt_b.mul(&x_expanded.broadcast_as(dt_b.dims())?)?;
874 let h_new = ah.add(&bx)?;
875
876 state.h = Tensor::from_candle(h_new.clone());
877
878 let c_expanded = c_t.unsqueeze(1)?;
880 let y_state = h_new.mul(&c_expanded.broadcast_as(h_new.dims())?)?.sum(2)?;
881
882 let d_candle = self.d.to_candle()?;
883 let y_skip = x_t.broadcast_mul(&d_candle)?;
884 let y_t = y_state.add(&y_skip)?;
885
886 Ok(Tensor::from_candle(y_t.unsqueeze(1)?))
888 }
889
890 fn load_weights(&mut self, weights: &ModelWeights, layer_idx: usize) -> Result<()> {
891 let prefix = format!("backbone.layers.{}.mixer", layer_idx);
892
893 if let Some(w) = weights.get(&format!("{}.in_proj.weight", prefix)) {
895 self.in_proj = ops_fn::transpose(w)?;
896 }
897
898 if let Some(w) = weights.get(&format!("{}.conv1d.weight", prefix)) {
900 let w_candle = w.to_candle()?;
902 let dims = w_candle.dims();
903 if dims.len() == 3 && dims[1] == 1 {
904 let reshaped = w_candle.squeeze(1)?;
905 self.conv1d_weight = Tensor::from_candle(reshaped);
906 } else {
907 self.conv1d_weight = w.clone();
908 }
909 }
910 if let Some(w) = weights.get(&format!("{}.conv1d.bias", prefix)) {
911 self.conv1d_bias = Some(w.clone());
912 }
913
914 if let Some(w) = weights.get(&format!("{}.x_proj.weight", prefix)) {
916 self.x_proj = ops_fn::transpose(w)?;
917 }
918
919 if let Some(w) = weights.get(&format!("{}.dt_proj.weight", prefix)) {
921 self.dt_proj = w.clone();
922 }
923 if let Some(w) = weights.get(&format!("{}.dt_proj.bias", prefix)) {
924 self.dt_proj_bias = Some(w.clone());
925 }
926
927 if let Some(w) = weights.get(&format!("{}.A_log", prefix)) {
929 self.a_log = w.clone();
930 }
931
932 if let Some(w) = weights.get(&format!("{}.D", prefix)) {
934 self.d = w.clone();
935 }
936
937 if let Some(w) = weights.get(&format!("{}.out_proj.weight", prefix)) {
939 self.out_proj = ops_fn::transpose(w)?;
940 }
941
942 Ok(())
943 }
944
945 fn to_device(&mut self, device: &Device) -> Result<()> {
946 self.in_proj = self.in_proj.to_device(device)?;
947 self.conv1d_weight = self.conv1d_weight.to_device(device)?;
948 if let Some(ref mut bias) = self.conv1d_bias {
949 *bias = bias.to_device(device)?;
950 }
951 self.x_proj = self.x_proj.to_device(device)?;
952 self.dt_proj = self.dt_proj.to_device(device)?;
953 if let Some(ref mut bias) = self.dt_proj_bias {
954 *bias = bias.to_device(device)?;
955 }
956 self.a_log = self.a_log.to_device(device)?;
957 self.d = self.d.to_device(device)?;
958 self.out_proj = self.out_proj.to_device(device)?;
959 Ok(())
960 }
961}
962
963fn softplus(x: &candle_core::Tensor) -> Result<candle_core::Tensor> {
965 let one = candle_core::Tensor::ones(x.dims(), x.dtype(), x.device())?;
968 let exp_x = x.exp()?;
969 let one_plus_exp = one.add(&exp_x)?;
970 Ok(one_plus_exp.log()?)
971}
972
973#[cfg(test)]
974mod tests {
975 use super::*;
976
977 #[test]
978 fn test_mamba_config() {
979 let config = MambaConfig::default();
980 assert_eq!(config.vocab_size, 50280);
981 assert_eq!(config.hidden_size, 768);
982 assert_eq!(config.effective_d_inner(), 768 * 2);
983 assert_eq!(config.effective_dt_rank(), 48); }
985
986 #[test]
987 fn test_mamba_model_creation() {
988 let config = MambaConfig {
989 vocab_size: 1000,
990 hidden_size: 128,
991 num_hidden_layers: 2,
992 d_state: 8,
993 d_conv: 4,
994 expand: 2,
995 ..Default::default()
996 };
997
998 let model = MambaModelV2::new(config).unwrap();
999 assert_eq!(model.config().vocab_size(), 1000);
1000 assert_eq!(model.config().hidden_size(), 128);
1001 assert_eq!(model.config().num_layers(), 2);
1002 }
1003
1004 #[test]
1005 fn test_mamba_forward_pass() {
1006 let config = MambaConfig {
1007 vocab_size: 100,
1008 hidden_size: 64,
1009 num_hidden_layers: 1,
1010 d_state: 8,
1011 d_conv: 4,
1012 expand: 2,
1013 ..Default::default()
1014 };
1015
1016 let model = MambaModelV2::new(config).unwrap();
1017 let input_ids = ops_fn::zeros(&[2, 8], DataType::Int64, &Device::CPU).unwrap();
1018 let inputs = ModelInputs::text(input_ids);
1019
1020 let outputs = model.forward(&inputs).unwrap();
1021 match outputs {
1022 ModelOutputs::Logits { logits, .. } => {
1023 assert_eq!(logits.shape(), &[2, 8, 100]);
1024 }
1025 _ => panic!("Expected logits output"),
1026 }
1027 }
1028
1029 #[test]
1030 fn test_mamba_generation() {
1031 let config = MambaConfig {
1032 vocab_size: 256,
1033 hidden_size: 64,
1034 num_hidden_layers: 1,
1035 d_state: 8,
1036 d_conv: 4,
1037 expand: 2,
1038 ..Default::default()
1039 };
1040
1041 let model = MambaModelV2::new(config).unwrap();
1042 let gen_config = GenerationConfig {
1043 max_new_tokens: 5,
1044 ..Default::default()
1045 };
1046
1047 let output = model.generate("Hello", &gen_config).unwrap();
1048 assert!(!output.is_empty());
1049 }
1050}