Skip to main content

runtime/models_v2/
mamba.rs

1//! Mamba Model V2 - State-Space Model implementation
2//!
3//! This implements the Mamba architecture which is a state-space model (SSM):
4//! - Uses selective state-space layers instead of attention
5//! - Has Conv1D in the mixer block for local context
6//! - Uses state-space operations (selective scan)
7//! - Different memory characteristics (state vs KV cache)
8//!
9//! Supports: Mamba-130M, Mamba-370M, Mamba-790M, Mamba-1.4B, Mamba-2.8B
10
11use crate::model_config;
12use super::traits::*;
13use anyhow::Result;
14use serde::{Serialize, Deserialize};
15
16/// Mamba model configuration using the model_config macro
17model_config!(MambaConfig {
18    vocab_size: usize = 50280,
19    hidden_size: usize = 768,       // d_model
20    num_hidden_layers: usize = 24,  // n_layer
21    d_state: usize = 16,            // State dimension for SSM
22    d_conv: usize = 4,              // Convolution kernel size
23    expand: usize = 2,              // Expansion factor for d_inner
24    dt_rank: usize = 0,             // Delta time rank (0 = auto: ceil(d_model/16))
25    d_inner: usize = 0,             // Inner dimension (0 = auto: d_model * expand)
26    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    /// Create MambaConfig from GGUF model configuration
43    pub fn from_gguf_config(gguf: &crate::weight_loader_core::GGUFModelConfig) -> Self {
44        // GGUF may not have all Mamba-specific fields, use sensible defaults
45        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    /// Get the effective d_inner (inner dimension)
65    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    /// Get the effective dt_rank
74    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
83/// Main Mamba model implementation
84pub struct MambaModelV2 {
85    config: MambaConfig,
86    device: Device,
87    backbone: MambaBackbone,
88    lm_head: Tensor,
89}
90
91/// Mamba backbone containing embeddings and layers
92pub struct MambaBackbone {
93    embeddings: Tensor,
94    layers: Vec<MambaBlock>,
95    norm_f: Tensor,  // Final layer norm
96    config: MambaConfig,
97}
98
99/// Single Mamba block (pre-norm + mixer + residual)
100pub struct MambaBlock {
101    mixer: MambaMixer,
102    norm: Tensor,  // Pre-norm weights
103    config: MambaConfig,
104}
105
106/// Mamba mixer - the core SSM block
107pub struct MambaMixer {
108    // Input projection: projects to 2 * d_inner (for x and z branches)
109    in_proj: Tensor,
110
111    // Causal Conv1D for local context
112    conv1d_weight: Tensor,
113    conv1d_bias: Option<Tensor>,
114
115    // State-space parameter projections
116    x_proj: Tensor,      // Projects to dt_rank + 2*d_state (for delta, B, C)
117    dt_proj: Tensor,     // Projects delta from dt_rank to d_inner
118    dt_proj_bias: Option<Tensor>,
119
120    // SSM parameters
121    a_log: Tensor,       // Log of state transition matrix A [d_inner, d_state]
122    d: Tensor,           // Skip connection parameter [d_inner]
123
124    // Output projection
125    out_proj: Tensor,    // Projects d_inner back to d_model
126
127    // Dimensions
128    d_inner: usize,
129    d_state: usize,
130    d_conv: usize,
131    dt_rank: usize,
132}
133
134/// SSM state for generation (per layer)
135#[derive(Clone)]
136pub struct MambaState {
137    /// Hidden state [batch, d_inner, d_state]
138    h: Tensor,
139    /// Conv1D cache [batch, d_inner, d_conv-1]
140    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            // Will share with backbone.embeddings during forward
151            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        // Apply lm_head
179        let logits = if self.config.tie_embeddings {
180            // Use embeddings transposed for output projection
181            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        // 1. Tokenize prompt
198        let tokenizer = Tokenizer::new();
199        let mut tokens: Vec<u32> = tokenizer.encode(prompt);
200
201        // 2. Initialize states for all layers
202        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        // 3. Process prompt tokens to build up state
216        // For efficiency, we process the entire prompt first
217        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            // Process through backbone with state update
220            self.backbone.forward_with_state(&input_tensor, &mut layer_states)?;
221        }
222
223        // 4. Generation loop
224        for _ in 0..config.max_new_tokens {
225            // Get last token
226            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            // Forward with state
230            let hidden_states = self.backbone.forward_with_state(&input_tensor, &mut layer_states)?;
231
232            // Apply lm_head
233            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            // Get logits and sample next token
241            let logits_candle = logits.to_candle()?;
242            let shape = logits_candle.dims();
243
244            // Extract last position logits
245            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                // Temperature sampling
257                let scaled: Vec<f32> = logits_vec.iter()
258                    .map(|&x| x / config.temperature)
259                    .collect();
260
261                // Softmax
262                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                // Sample from distribution
269                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                // Greedy sampling
284                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            // Check for EOS
296            if next_token == config.eos_token_id {
297                break;
298            }
299
300            // Append token
301            tokens.push(next_token);
302        }
303
304        // 5. Decode and return
305        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                             // in_proj
315                             self.config.hidden_size * d_inner * 2 +
316                             // conv1d
317                             d_inner * self.config.d_conv +
318                             // x_proj
319                             d_inner * (self.config.effective_dt_rank() + self.config.d_state * 2) +
320                             // dt_proj
321                             self.config.effective_dt_rank() * d_inner +
322                             // A_log, D
323                             d_inner * self.config.d_state + d_inner +
324                             // out_proj
325                             d_inner * self.config.hidden_size
326                         )) * 4;
327
328        // State memory: h [batch, d_inner, d_state] + conv_cache [batch, d_inner, d_conv-1]
329        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, // Mamba uses state instead of KV cache
336            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        // Embedding lookup
365        let mut hidden_states = ops_fn::embedding(input_ids, &self.embeddings)?;
366
367        // Pass through layers
368        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        // Final normalization
382        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        // Try different naming conventions
395        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        // Pre-norm
438        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        // Mixer
445        let mixed = self.mixer.forward(&normalized)?;
446
447        // Residual
448        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        // Pre-norm
455        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        // Mixer with state
462        let mixed = self.mixer.forward_with_state(&normalized, state)?;
463
464        // Residual
465        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        // in_proj: [d_model, 2 * d_inner] - projects to x and z branches
494        let in_proj = ops_fn::zeros(&[config.hidden_size, d_inner * 2], DataType::Float32, device)?;
495
496        // conv1d: [d_inner, d_conv] - depthwise causal convolution
497        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        // x_proj: [d_inner, dt_rank + 2*d_state] - projects to (dt, B, C)
505        let x_proj = ops_fn::zeros(&[d_inner, dt_rank + d_state * 2], DataType::Float32, device)?;
506
507        // dt_proj: [dt_rank, d_inner] - projects dt from rank to d_inner
508        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        // A_log: [d_inner, d_state] - log of state transition matrix
516        let a_log = ops_fn::zeros(&[d_inner, d_state], DataType::Float32, device)?;
517
518        // D: [d_inner] - skip connection
519        let d = ops_fn::zeros(&[d_inner], DataType::Float32, device)?;
520
521        // out_proj: [d_inner, d_model]
522        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        // 1. Input projection: [B, L, D] -> [B, L, 2*d_inner]
552        let projected = ops_fn::matmul(hidden_states, &self.in_proj)?;
553
554        // 2. Split into x and z: [B, L, d_inner] each
555        let (x, z) = self.split_xz(&projected)?;
556
557        // 3. Apply causal Conv1D to x: [B, L, d_inner] -> [B, L, d_inner]
558        let x_conv = self.apply_conv1d(&x, batch_size, seq_len)?;
559
560        // 4. Apply SiLU activation to conv output
561        let x_act = ops_fn::silu(&x_conv)?;
562
563        // 5. Selective scan (SSM)
564        let y = self.selective_scan(&x_act, batch_size, seq_len)?;
565
566        // 6. Gate with z (SiLU(z) * y)
567        let z_act = ops_fn::silu(&z)?;
568        let gated = ops_fn::mul(&y, &z_act)?;
569
570        // 7. Output projection: [B, L, d_inner] -> [B, L, D]
571        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        // For single token generation (seq_len = 1), use stateful computation
585        if seq_len == 1 {
586            return self.forward_step(hidden_states, state);
587        }
588
589        // For longer sequences, use full forward and update state
590        self.forward(hidden_states)
591    }
592
593    /// Single step forward for generation with state
594    fn forward_step(&self, hidden_states: &Tensor, state: &mut MambaState) -> Result<Tensor> {
595        // 1. Input projection
596        let projected = ops_fn::matmul(hidden_states, &self.in_proj)?;
597
598        // 2. Split into x and z
599        let (x, z) = self.split_xz(&projected)?;
600
601        // 3. Apply causal Conv1D with cache update
602        let x_conv = self.apply_conv1d_step(&x, state)?;
603
604        // 4. Apply SiLU
605        let x_act = ops_fn::silu(&x_conv)?;
606
607        // 5. Selective scan step with state update
608        let y = self.selective_scan_step(&x_act, state)?;
609
610        // 6. Gate with z
611        let z_act = ops_fn::silu(&z)?;
612        let gated = ops_fn::mul(&y, &z_act)?;
613
614        // 7. Output projection
615        ops_fn::matmul(&gated, &self.out_proj)
616    }
617
618    /// Split projected tensor into x and z branches
619    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        // Split along last dimension: first half is x, second half is z
625        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    /// Apply causal Conv1D
632    fn apply_conv1d(&self, x: &Tensor, batch_size: usize, seq_len: usize) -> Result<Tensor> {
633        // x shape: [B, L, d_inner]
634        // conv1d_weight shape: [d_inner, d_conv]
635
636        // For simplicity, implement as a sliding window matmul
637        // This is equivalent to depthwise separable causal convolution
638
639        let x_candle = x.to_candle()?;
640        let w_candle = self.conv1d_weight.to_candle()?;
641
642        // Pad input with zeros on the left for causal convolution
643        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        // Reshape x to [B, L, d_inner] if needed
648        let x_3d = if x_candle.dims().len() == 2 {
649            x_candle.unsqueeze(0)?
650        } else {
651            x_candle.clone()
652        };
653
654        // Concatenate: [B, pad_len + L, d_inner]
655        let x_padded = candle_core::Tensor::cat(&[&zeros, &x_3d], 1)?;
656
657        // Apply convolution by gathering windows and multiplying
658        // For each position i, gather [i:i+d_conv] and dot with weights
659        let mut outputs = Vec::new();
660
661        for i in 0..seq_len {
662            // Extract window [B, d_conv, d_inner]
663            let window = x_padded.narrow(1, i, self.d_conv)?;
664
665            // Transpose to [B, d_inner, d_conv]
666            let window_t = window.transpose(1, 2)?;
667
668            // Element-wise multiply with weights [d_inner, d_conv] and sum over d_conv
669            let conv_out = window_t.broadcast_mul(&w_candle)?;
670            let summed = conv_out.sum(2)?;  // [B, d_inner]
671
672            outputs.push(summed);
673        }
674
675        // Stack outputs: [B, L, d_inner]
676        let result = candle_core::Tensor::stack(&outputs, 1)?;
677
678        // Add bias if present
679        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    /// Apply Conv1D step with cache
690    fn apply_conv1d_step(&self, x: &Tensor, state: &mut MambaState) -> Result<Tensor> {
691        // x shape: [B, 1, d_inner]
692        let x_candle = x.to_candle()?;
693        let x_squeezed = x_candle.squeeze(1)?;  // [B, d_inner]
694
695        // Update conv cache: shift and append new x
696        let cache_candle = state.conv_cache.to_candle()?;
697
698        // cache shape: [B, d_inner, d_conv-1]
699        // Shift left (drop oldest) and append new
700        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                // d_conv == 2, cache is [B, d_inner, 1], drop everything
705                candle_core::Tensor::zeros(&[cache_candle.dims()[0], self.d_inner, 0], cache_candle.dtype(), cache_candle.device())?
706            };
707
708            // Expand x to [B, d_inner, 1]
709            let x_expanded = x_squeezed.unsqueeze(2)?;
710
711            // Concatenate: [B, d_inner, d_conv-1]
712            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        // Apply convolution: gather cache + current x, multiply with weights
721        let w_candle = self.conv1d_weight.to_candle()?;
722
723        // Get full conv window [B, d_inner, d_conv]
724        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        // Element-wise multiply and sum
729        let conv_out = full_window.broadcast_mul(&w_candle)?;
730        let result = conv_out.sum(2)?;  // [B, d_inner]
731
732        // Add bias
733        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        // Return [B, 1, d_inner]
741        Ok(Tensor::from_candle(result.unsqueeze(1)?))
742    }
743
744    /// Selective scan (SSM) operation
745    fn selective_scan(&self, x: &Tensor, batch_size: usize, seq_len: usize) -> Result<Tensor> {
746        // x shape: [B, L, d_inner]
747
748        // 1. Project to get delta, B, C: [B, L, dt_rank + 2*d_state]
749        let dbc = ops_fn::matmul(x, &self.x_proj)?;
750        let dbc_candle = dbc.to_candle()?;
751
752        // Split into dt (delta), B, C
753        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        // 2. Project dt: [B, L, dt_rank] @ [dt_rank, d_inner] -> [B, L, d_inner]
758        // Use broadcast_matmul for 3D @ 2D
759        let dt_proj_candle = self.dt_proj.to_candle()?;
760        let dt = dt_raw.broadcast_matmul(&dt_proj_candle)?;
761
762        // Add bias and apply softplus
763        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        // Softplus: log(1 + exp(x))
771        let dt = softplus(&dt)?;
772
773        // 3. Get A from A_log: A = -exp(A_log)
774        let a_log_candle = self.a_log.to_candle()?;
775        let a = a_log_candle.exp()?.neg()?;
776
777        // 4. Selective scan loop
778        let x_candle = x.to_candle()?;
779        let d_candle = self.d.to_candle()?;
780
781        // Initialize hidden state h: [B, d_inner, d_state]
782        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            // Get current timestep values
788            let x_t = x_candle.narrow(1, t, 1)?.squeeze(1)?;  // [B, d_inner]
789            let dt_t = dt.narrow(1, t, 1)?.squeeze(1)?;        // [B, d_inner]
790            let b_t = b.narrow(1, t, 1)?.squeeze(1)?;          // [B, d_state]
791            let c_t = c.narrow(1, t, 1)?.squeeze(1)?;          // [B, d_state]
792
793            // Discretize: A_bar = exp(dt * A), B_bar = dt * B
794            // dt: [B, d_inner], A: [d_inner, d_state] -> dt_A: [B, d_inner, d_state]
795            let dt_expanded = dt_t.unsqueeze(2)?;  // [B, d_inner, 1]
796            let dt_a = dt_expanded.broadcast_mul(&a)?;
797            let a_bar = dt_a.exp()?;  // [B, d_inner, d_state]
798
799            // B_bar = dt[:, :, None] * B[:, None, :]
800            let b_expanded = b_t.unsqueeze(1)?;  // [B, 1, d_state]
801            let dt_b = dt_expanded.broadcast_mul(&b_expanded)?;  // [B, d_inner, d_state]
802
803            // x expanded for state update
804            let x_expanded = x_t.unsqueeze(2)?;  // [B, d_inner, 1]
805
806            // State update: h = A_bar * h + B_bar * x
807            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            // Output: y = (C @ h) + D * x
812            // h: [B, d_inner, d_state], C: [B, d_state] -> y: [B, d_inner]
813            let c_expanded = c_t.unsqueeze(1)?;  // [B, 1, d_state]
814            let y_state = h.mul(&c_expanded.broadcast_as(h.dims())?)?.sum(2)?;  // [B, d_inner]
815
816            // Add skip connection (broadcast multiply with D)
817            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        // Stack outputs: [B, L, d_inner]
824        let result = candle_core::Tensor::stack(&outputs, 1)?;
825
826        Ok(Tensor::from_candle(result))
827    }
828
829    /// Single step selective scan with state
830    fn selective_scan_step(&self, x: &Tensor, state: &mut MambaState) -> Result<Tensor> {
831        // x shape: [B, 1, d_inner]
832        let x_candle = x.to_candle()?;
833        let x_t = x_candle.squeeze(1)?;  // [B, d_inner]
834
835        // 1. Project to get delta, B, C
836        let dbc = ops_fn::matmul(x, &self.x_proj)?;
837        let dbc_candle = dbc.to_candle()?.squeeze(1)?;  // [B, dt_rank + 2*d_state]
838
839        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        // 2. Project dt: [B, dt_rank] @ [dt_rank, d_inner] -> [B, d_inner]
844        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        // 3. Get A
857        let a_log_candle = self.a_log.to_candle()?;
858        let a = a_log_candle.exp()?.neg()?;
859
860        // 4. Compute discretized matrices
861        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        // 5. State update
869        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        // 6. Compute output
879        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        // Return [B, 1, d_inner]
887        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        // in_proj
894        if let Some(w) = weights.get(&format!("{}.in_proj.weight", prefix)) {
895            self.in_proj = ops_fn::transpose(w)?;
896        }
897
898        // Conv1D
899        if let Some(w) = weights.get(&format!("{}.conv1d.weight", prefix)) {
900            // Conv weight might need reshaping from [d_inner, 1, d_conv] to [d_inner, d_conv]
901            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        // x_proj
915        if let Some(w) = weights.get(&format!("{}.x_proj.weight", prefix)) {
916            self.x_proj = ops_fn::transpose(w)?;
917        }
918
919        // dt_proj
920        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        // A_log
928        if let Some(w) = weights.get(&format!("{}.A_log", prefix)) {
929            self.a_log = w.clone();
930        }
931
932        // D
933        if let Some(w) = weights.get(&format!("{}.D", prefix)) {
934            self.d = w.clone();
935        }
936
937        // out_proj
938        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
963/// Softplus activation: log(1 + exp(x))
964fn softplus(x: &candle_core::Tensor) -> Result<candle_core::Tensor> {
965    // For numerical stability: softplus(x) = x + log(1 + exp(-|x|)) - min(0, x)
966    // Simplified: log(1 + exp(x))
967    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); // ceil(768/16)
984    }
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}