Skip to main content

runtime/models_v2/
minicpm.rs

1//! MiniCPM Model V2 - Clean implementation using solid abstractions
2//!
3//! This implements the MiniCPM architecture including:
4//! - MiniCPM-2B, MiniCPM-1B, MiniCPM-V (vision variant)
5//!
6//! MiniCPM unique features:
7//! - scale_emb scaling on embeddings
8//! - scale_depth / sqrt(num_layers) scaling per layer
9//! - tied embeddings with transpose for lm_head
10
11use crate::model_config;
12use super::traits::*;
13use anyhow::Result;
14use serde::{Serialize, Deserialize};
15
16model_config!(MiniCPMConfig {
17    vocab_size: usize = 122753,
18    hidden_size: usize = 2304,
19    intermediate_size: usize = 5760,
20    num_hidden_layers: usize = 40,
21    num_attention_heads: usize = 36,
22    num_key_value_heads: usize = 36,
23    hidden_act: String = "silu".to_string(),
24    max_position_embeddings: usize = 4096,
25    initializer_range: f32 = 0.02,
26    rms_norm_eps: f32 = 1e-5,
27    use_cache: bool = true,
28    pad_token_id: i64 = 0,
29    bos_token_id: i64 = 1,
30    eos_token_id: i64 = 2,
31    tie_word_embeddings: bool = true,
32    rope_theta: f32 = 10000.0,
33    rope_scaling: Option<String> = None,
34    attention_bias: bool = false,
35    attention_dropout: f32 = 0.0,
36    // MiniCPM specific
37    scale_emb: f32 = 12.0,
38    scale_depth: f32 = 1.4,
39    dim_model_base: usize = 256,
40    scale_width: f32 = 2.0,
41});
42
43impl MiniCPMConfig {
44    /// Create MiniCPMConfig from GGUF model configuration
45    pub fn from_gguf_config(gguf: &crate::weight_loader_core::GGUFModelConfig) -> Self {
46        Self {
47            vocab_size: gguf.vocab_size,
48            hidden_size: gguf.hidden_size,
49            intermediate_size: gguf.intermediate_size,
50            num_hidden_layers: gguf.num_hidden_layers,
51            num_attention_heads: gguf.num_attention_heads,
52            num_key_value_heads: gguf.num_key_value_heads,
53            rms_norm_eps: gguf.rms_norm_eps,
54            rope_theta: gguf.rope_theta,
55            max_position_embeddings: gguf.max_position_embeddings,
56            ..Default::default()
57        }
58    }
59}
60
61pub struct MiniCPMModelV2 {
62    config: MiniCPMConfig,
63    device: Device,
64    embed_tokens: Tensor,
65    layers: Vec<MiniCPMDecoderLayer>,
66    norm: Tensor,
67    lm_head: Option<Tensor>, // None if tie_word_embeddings is true
68}
69
70pub struct MiniCPMDecoderLayer {
71    self_attn: MiniCPMAttention,
72    mlp: MiniCPMMLP,
73    input_layernorm: Tensor,
74    post_attention_layernorm: Tensor,
75    scale: f32,
76}
77
78pub struct MiniCPMAttention {
79    q_proj: Tensor,
80    k_proj: Tensor,
81    v_proj: Tensor,
82    o_proj: Tensor,
83    num_heads: usize,
84    num_key_value_heads: usize,
85    head_dim: usize,
86    scale: f32,
87}
88
89pub struct MiniCPMMLP {
90    gate_proj: Tensor,
91    up_proj: Tensor,
92    down_proj: Tensor,
93    hidden_act: String,
94}
95
96/// Apply Rotary Position Embedding (RoPE) to Q and K tensors
97/// Input shape: [batch, heads, seq, head_dim]
98/// Returns tensors with same shape but with positional information encoded
99fn apply_rope(
100    q: &candle_core::Tensor,
101    k: &candle_core::Tensor,
102    seq_len: usize,
103    head_dim: usize,
104    rope_theta: f32,
105) -> Result<(candle_core::Tensor, candle_core::Tensor)> {
106    let device = q.device();
107
108    // Compute inverse frequencies: 1 / (theta^(2i/d)) for i in [0, d/2)
109    let half_dim = head_dim / 2;
110    let inv_freq: Vec<f32> = (0..half_dim)
111        .map(|i| 1.0 / rope_theta.powf((2 * i) as f32 / head_dim as f32))
112        .collect();
113
114    // Create position indices [0, 1, 2, ..., seq_len-1]
115    let positions: Vec<f32> = (0..seq_len).map(|p| p as f32).collect();
116
117    // Compute angles: pos * inv_freq -> [seq_len, half_dim]
118    let mut angles = Vec::with_capacity(seq_len * half_dim);
119    for pos in &positions {
120        for freq in &inv_freq {
121            angles.push(pos * freq);
122        }
123    }
124
125    let angles_tensor = candle_core::Tensor::from_vec(angles, &[seq_len, half_dim], device)?;
126
127    // Compute cos and sin
128    let cos = angles_tensor.cos()?;
129    let sin = angles_tensor.sin()?;
130
131    // Reshape for broadcasting: [1, 1, seq_len, half_dim]
132    let cos = cos.unsqueeze(0)?.unsqueeze(0)?;
133    let sin = sin.unsqueeze(0)?.unsqueeze(0)?;
134
135    // Apply RoPE rotation
136    // Split q and k into two halves along head_dim
137    // q = [q1, q2], k = [k1, k2] where each half has shape [..., half_dim]
138    // rotated_q = [q1*cos - q2*sin, q1*sin + q2*cos]
139    // rotated_k = [k1*cos - k2*sin, k1*sin + k2*cos]
140    let q_half1 = q.narrow(3, 0, half_dim)?;
141    let q_half2 = q.narrow(3, half_dim, half_dim)?;
142    let k_half1 = k.narrow(3, 0, half_dim)?;
143    let k_half2 = k.narrow(3, half_dim, half_dim)?;
144
145    // Apply rotation
146    let q_rot1 = (q_half1.broadcast_mul(&cos)? - q_half2.broadcast_mul(&sin)?)?;
147    let q_rot2 = (q_half1.broadcast_mul(&sin)? + q_half2.broadcast_mul(&cos)?)?;
148    let k_rot1 = (k_half1.broadcast_mul(&cos)? - k_half2.broadcast_mul(&sin)?)?;
149    let k_rot2 = (k_half1.broadcast_mul(&sin)? + k_half2.broadcast_mul(&cos)?)?;
150
151    // Concatenate rotated halves
152    let q_rotated = candle_core::Tensor::cat(&[&q_rot1, &q_rot2], 3)?;
153    let k_rotated = candle_core::Tensor::cat(&[&k_rot1, &k_rot2], 3)?;
154
155    Ok((q_rotated, k_rotated))
156}
157
158impl Model for MiniCPMModelV2 {
159    type Config = MiniCPMConfig;
160
161    fn new(config: MiniCPMConfig) -> Result<Self> {
162        let device = Device::CPU;
163        let embed_tokens = ops_fn::zeros(&[config.vocab_size, config.hidden_size], DataType::Float32, &device)?;
164        let norm = ops_fn::zeros(&[config.hidden_size], DataType::Float32, &device)?;
165
166        let lm_head = if config.tie_word_embeddings {
167            None
168        } else {
169            Some(ops_fn::zeros(&[config.hidden_size, config.vocab_size], DataType::Float32, &device)?)
170        };
171
172        let mut layers = Vec::new();
173        for layer_idx in 0..config.num_hidden_layers {
174            layers.push(MiniCPMDecoderLayer::new(&config, &device, layer_idx)?);
175        }
176
177        Ok(Self { config, device, embed_tokens, layers, norm, lm_head })
178    }
179
180    fn from_weights(config: MiniCPMConfig, weights: ModelWeights) -> Result<Self> {
181        let mut model = Self::new(config)?;
182
183        // Load embedding weights (no transpose - used for index lookup)
184        if let Some(w) = weights.get("model.embed_tokens.weight") {
185            model.embed_tokens = w.clone();
186        }
187
188        if let Some(w) = weights.get("model.norm.weight") {
189            model.norm = w.clone();
190        }
191
192        // For tied embeddings, lm_head is None and we use embed_tokens transposed at inference
193        // For separate lm_head, transpose: [vocab, hidden] -> [hidden, vocab] for matmul
194        if !model.config.tie_word_embeddings {
195            if let Some(w) = weights.get("lm_head.weight") {
196                model.lm_head = Some(ops_fn::transpose(w)?);
197            }
198        }
199
200        for (i, layer) in model.layers.iter_mut().enumerate() {
201            layer.load_weights(&weights, i)?;
202        }
203
204        Ok(model)
205    }
206
207    fn forward(&self, inputs: &ModelInputs) -> Result<ModelOutputs> {
208        let input_ids = match inputs {
209            ModelInputs::Text { input_ids, .. } => input_ids,
210            ModelInputs::Multimodal { input_ids, .. } => input_ids,
211            _ => return Err(anyhow::anyhow!("MiniCPM expects text input")),
212        };
213
214        // Embedding with MiniCPM scaling
215        let mut hidden_states = ops_fn::embedding(input_ids, &self.embed_tokens)?;
216        hidden_states = ops_fn::scale(&hidden_states, self.config.scale_emb)?;
217
218        // Forward through layers
219        for layer in &self.layers {
220            hidden_states = layer.forward(&hidden_states, None, self.config.rope_theta)?;
221        }
222
223        // Final normalization
224        hidden_states = ops_fn::layer_norm(&hidden_states, &self.norm, None, self.config.rms_norm_eps)?;
225
226        // Language modeling head
227        let logits = if let Some(ref lm_head) = self.lm_head {
228            // Separate lm_head (already transposed)
229            ops_fn::matmul(&hidden_states, lm_head)?
230        } else {
231            // Tied embeddings - need to transpose embed_tokens for matmul
232            // embed_tokens: [vocab, hidden] -> need [hidden, vocab]
233            let lm_head_t = ops_fn::transpose(&self.embed_tokens)?;
234            ops_fn::matmul(&hidden_states, &lm_head_t)?
235        };
236
237        // Apply final scaling (MiniCPM specific - inverse of embedding scale)
238        let logits = ops_fn::scale(&logits, 1.0 / self.config.scale_emb)?;
239
240        Ok(ModelOutputs::Logits { logits, hidden_states: Some(hidden_states) })
241    }
242
243    fn generate(&self, prompt: &str, config: &GenerationConfig) -> Result<String> {
244        use crate::tokenizer::Tokenizer;
245        use rand::Rng;
246
247        // 1. Tokenize prompt
248        let tokenizer = Tokenizer::new();
249        let mut tokens: Vec<u32> = tokenizer.encode(prompt);
250
251        // 2. Generation loop
252        for _ in 0..config.max_new_tokens {
253            // Create input tensor from current tokens
254            let tokens_i64: Vec<i64> = tokens.iter().map(|&t| t as i64).collect();
255            let input_tensor = Tensor::from_i64_slice(&tokens_i64, &[1, tokens.len()], &self.device)?;
256
257            let inputs = ModelInputs::Text {
258                input_ids: input_tensor,
259                attention_mask: None,
260                position_ids: None,
261            };
262
263            // 3. Forward pass
264            let outputs = self.forward(&inputs)?;
265
266            // 4. Get logits and sample next token
267            let logits = match outputs {
268                ModelOutputs::Logits { logits, .. } => logits,
269                _ => return Err(anyhow::anyhow!("Expected logits output")),
270            };
271
272            // Get last token logits
273            let logits_candle = logits.to_candle()?;
274            let shape = logits_candle.dims();
275
276            // Extract last position logits [batch, seq, vocab] -> [vocab]
277            let last_logits = if shape.len() == 3 {
278                let seq_len = shape[1];
279                logits_candle
280                    .narrow(1, seq_len - 1, 1)?
281                    .squeeze(1)?
282                    .squeeze(0)?
283            } else {
284                let seq_len = shape[0];
285                logits_candle
286                    .narrow(0, seq_len - 1, 1)?
287                    .squeeze(0)?
288            };
289
290            // Convert to probabilities and sample
291            let logits_vec: Vec<f32> = last_logits.to_vec1()?;
292
293            let next_token = if config.do_sample && config.temperature > 0.0 {
294                // Temperature sampling
295                let scaled: Vec<f32> = logits_vec.iter()
296                    .map(|&x| x / config.temperature)
297                    .collect();
298
299                // Softmax
300                let max_val = scaled.iter().cloned().fold(f32::NEG_INFINITY, f32::max);
301                let exp_sum: f32 = scaled.iter().map(|&x| (x - max_val).exp()).sum();
302                let probs: Vec<f32> = scaled.iter()
303                    .map(|&x| (x - max_val).exp() / exp_sum)
304                    .collect();
305
306                // Sample from distribution
307                let mut rng = rand::thread_rng();
308                let random_val: f32 = rng.gen();
309                let mut cumulative = 0.0;
310                let mut sampled = 0u32;
311
312                for (idx, &prob) in probs.iter().enumerate() {
313                    cumulative += prob;
314                    if random_val <= cumulative {
315                        sampled = idx as u32;
316                        break;
317                    }
318                }
319                sampled
320            } else {
321                // Greedy sampling
322                let mut max_idx = 0;
323                let mut max_val = logits_vec[0];
324                for (idx, &val) in logits_vec.iter().enumerate() {
325                    if val > max_val {
326                        max_val = val;
327                        max_idx = idx;
328                    }
329                }
330                max_idx as u32
331            };
332
333            // 5. Check for EOS
334            if next_token == config.eos_token_id {
335                break;
336            }
337
338            // 6. Append token
339            tokens.push(next_token);
340        }
341
342        // 7. Decode and return
343        Ok(tokenizer.decode(&tokens))
344    }
345
346    fn config(&self) -> &Self::Config { &self.config }
347
348    fn memory_requirements(&self) -> MemoryRequirements {
349        let param_size = (self.config.vocab_size * self.config.hidden_size +
350                         self.config.num_hidden_layers * self.config.hidden_size * self.config.hidden_size * 4) * 4;
351        MemoryRequirements {
352            gpu_memory: param_size, cpu_memory: param_size / 4,
353            kv_cache_memory: self.config.max_position_embeddings * self.config.hidden_size * 2 * 4,
354            peak_memory: param_size + param_size / 2,
355        }
356    }
357
358    fn to_device(&mut self, device: &Device) -> Result<()> {
359        self.device = device.clone();
360        self.embed_tokens = self.embed_tokens.to_device(device)?;
361        self.norm = self.norm.to_device(device)?;
362        if let Some(ref mut lm_head) = self.lm_head {
363            *lm_head = lm_head.to_device(device)?;
364        }
365        for layer in &mut self.layers {
366            layer.to_device(device)?;
367        }
368        Ok(())
369    }
370}
371
372impl MiniCPMDecoderLayer {
373    fn new(config: &MiniCPMConfig, device: &Device, _layer_idx: usize) -> Result<Self> {
374        // MiniCPM uses layer-wise scaling based on depth
375        let scale = config.scale_depth / (config.num_hidden_layers as f32).sqrt();
376
377        Ok(Self {
378            self_attn: MiniCPMAttention::new(config, device)?,
379            mlp: MiniCPMMLP::new(config, device)?,
380            input_layernorm: ops_fn::zeros(&[config.hidden_size], DataType::Float32, device)?,
381            post_attention_layernorm: ops_fn::zeros(&[config.hidden_size], DataType::Float32, device)?,
382            scale,
383        })
384    }
385
386    fn forward(&self, hidden_states: &Tensor, _attention_mask: Option<&Tensor>, rope_theta: f32) -> Result<Tensor> {
387        let residual = hidden_states.clone();
388
389        // Pre-attention normalization
390        let normed = ops_fn::layer_norm(hidden_states, &self.input_layernorm, None, 1e-5)?;
391
392        // Self-attention with RoPE
393        let attn_output = self.self_attn.forward(&normed, None, rope_theta)?;
394
395        // Apply MiniCPM scaling and residual connection
396        let attn_scaled = ops_fn::scale(&attn_output, self.scale)?;
397        let hidden_states = ops_fn::add(&residual, &attn_scaled)?;
398
399        // Pre-MLP normalization
400        let residual = hidden_states.clone();
401        let normed = ops_fn::layer_norm(&hidden_states, &self.post_attention_layernorm, None, 1e-5)?;
402
403        // MLP
404        let mlp_output = self.mlp.forward(&normed)?;
405
406        // Apply MiniCPM scaling and residual connection
407        let mlp_scaled = ops_fn::scale(&mlp_output, self.scale)?;
408        ops_fn::add(&residual, &mlp_scaled)
409    }
410
411    fn load_weights(&mut self, weights: &ModelWeights, layer_idx: usize) -> Result<()> {
412        let prefix = format!("model.layers.{}", layer_idx);
413
414        // Load layer norm weights (no transpose needed - 1D tensors)
415        if let Some(w) = weights.get(&format!("{}.input_layernorm.weight", prefix)) {
416            self.input_layernorm = w.clone();
417        }
418        if let Some(w) = weights.get(&format!("{}.post_attention_layernorm.weight", prefix)) {
419            self.post_attention_layernorm = w.clone();
420        }
421
422        self.self_attn.load_weights(weights, layer_idx)?;
423        self.mlp.load_weights(weights, layer_idx)?;
424        Ok(())
425    }
426
427    fn to_device(&mut self, device: &Device) -> Result<()> {
428        self.input_layernorm = self.input_layernorm.to_device(device)?;
429        self.post_attention_layernorm = self.post_attention_layernorm.to_device(device)?;
430        self.self_attn.to_device(device)?;
431        self.mlp.to_device(device)?;
432        Ok(())
433    }
434}
435
436impl MiniCPMAttention {
437    fn new(config: &MiniCPMConfig, device: &Device) -> Result<Self> {
438        let num_heads = config.num_attention_heads;
439        let num_key_value_heads = config.num_key_value_heads;
440        let head_dim = config.hidden_size / num_heads;
441        let scale = 1.0 / (head_dim as f32).sqrt();
442
443        Ok(Self {
444            q_proj: ops_fn::zeros(&[config.hidden_size, num_heads * head_dim], DataType::Float32, device)?,
445            k_proj: ops_fn::zeros(&[config.hidden_size, num_key_value_heads * head_dim], DataType::Float32, device)?,
446            v_proj: ops_fn::zeros(&[config.hidden_size, num_key_value_heads * head_dim], DataType::Float32, device)?,
447            o_proj: ops_fn::zeros(&[num_heads * head_dim, config.hidden_size], DataType::Float32, device)?,
448            num_heads,
449            num_key_value_heads,
450            head_dim,
451            scale,
452        })
453    }
454
455    fn forward(&self, hidden_states: &Tensor, _attention_mask: Option<&Tensor>, rope_theta: f32) -> Result<Tensor> {
456        // Get batch and sequence length from hidden_states shape
457        let shape = hidden_states.shape();
458        let (batch_size, seq_len, _hidden_size) = if shape.len() == 3 {
459            (shape[0], shape[1], shape[2])
460        } else if shape.len() == 2 {
461            (1, shape[0], shape[1])
462        } else {
463            return Err(anyhow::anyhow!("Invalid hidden_states shape: {:?}", shape));
464        };
465
466        // 1. Project to Q, K, V
467        let query_states = ops_fn::matmul(hidden_states, &self.q_proj)?;
468        let key_states = ops_fn::matmul(hidden_states, &self.k_proj)?;
469        let value_states = ops_fn::matmul(hidden_states, &self.v_proj)?;
470
471        // 2. Reshape for multi-head attention
472        let q_candle = query_states.to_candle()?;
473        let k_candle = key_states.to_candle()?;
474        let v_candle = value_states.to_candle()?;
475
476        let q_reshaped = q_candle
477            .reshape(&[batch_size, seq_len, self.num_heads, self.head_dim])?
478            .transpose(1, 2)?;
479
480        let k_reshaped = k_candle
481            .reshape(&[batch_size, seq_len, self.num_key_value_heads, self.head_dim])?
482            .transpose(1, 2)?;
483
484        let v_reshaped = v_candle
485            .reshape(&[batch_size, seq_len, self.num_key_value_heads, self.head_dim])?
486            .transpose(1, 2)?;
487
488        // 3. Apply RoPE (Rotary Position Embedding) to Q and K
489        let (q_with_rope, k_with_rope) = apply_rope(&q_reshaped, &k_reshaped, seq_len, self.head_dim, rope_theta)?;
490
491        // 4. Handle GQA (Grouped Query Attention) - repeat K/V heads to match Q heads
492        let num_groups = self.num_heads / self.num_key_value_heads;
493        let (k_expanded, v_expanded) = if num_groups > 1 {
494            let k_rep = k_with_rope
495                .unsqueeze(2)?
496                .broadcast_as(&[batch_size, self.num_key_value_heads, num_groups, seq_len, self.head_dim])?
497                .reshape(&[batch_size, self.num_heads, seq_len, self.head_dim])?;
498            let v_rep = v_reshaped
499                .unsqueeze(2)?
500                .broadcast_as(&[batch_size, self.num_key_value_heads, num_groups, seq_len, self.head_dim])?
501                .reshape(&[batch_size, self.num_heads, seq_len, self.head_dim])?;
502            (k_rep, v_rep)
503        } else {
504            (k_with_rope, v_reshaped)
505        };
506
507        // 5. Scaled dot-product attention
508        let k_t = k_expanded.transpose(2, 3)?;
509
510        let q_contiguous = q_with_rope.contiguous()?;
511        let k_contiguous = k_t.contiguous()?;
512
513        let scores = q_contiguous.matmul(&k_contiguous)?;
514        let scaled_scores = (scores * (self.scale as f64))?;
515
516        // Apply causal mask
517        let device = scaled_scores.device();
518        let causal_mask = {
519            let mut mask_data = vec![0.0f32; seq_len * seq_len];
520            for i in 0..seq_len {
521                for j in 0..seq_len {
522                    if j > i {
523                        mask_data[i * seq_len + j] = f32::NEG_INFINITY;
524                    }
525                }
526            }
527            candle_core::Tensor::from_vec(mask_data, &[1, 1, seq_len, seq_len], device)?
528        };
529
530        let masked_scores = scaled_scores.broadcast_add(&causal_mask)?;
531
532        // Softmax over last dimension
533        let attention_weights = candle_nn::ops::softmax_last_dim(&masked_scores)?;
534
535        // Apply attention to values
536        let v_contiguous = v_expanded.contiguous()?;
537        let attn_output = attention_weights.matmul(&v_contiguous)?;
538
539        // 6. Reshape back
540        let attn_output = attn_output
541            .transpose(1, 2)?
542            .reshape(&[batch_size, seq_len, self.num_heads * self.head_dim])?;
543
544        let attn_output = Tensor::from_candle(attn_output);
545
546        // 7. Output projection
547        let output = ops_fn::matmul(&attn_output, &self.o_proj)?;
548
549        Ok(output)
550    }
551
552    fn load_weights(&mut self, weights: &ModelWeights, layer_idx: usize) -> Result<()> {
553        let prefix = format!("model.layers.{}.self_attn", layer_idx);
554
555        // Load attention weights (transpose for matmul: [out, in] -> [in, out])
556        if let Some(w) = weights.get(&format!("{}.q_proj.weight", prefix)) {
557            self.q_proj = ops_fn::transpose(w)?;
558        }
559        if let Some(w) = weights.get(&format!("{}.k_proj.weight", prefix)) {
560            self.k_proj = ops_fn::transpose(w)?;
561        }
562        if let Some(w) = weights.get(&format!("{}.v_proj.weight", prefix)) {
563            self.v_proj = ops_fn::transpose(w)?;
564        }
565        if let Some(w) = weights.get(&format!("{}.o_proj.weight", prefix)) {
566            self.o_proj = ops_fn::transpose(w)?;
567        }
568        Ok(())
569    }
570
571    fn to_device(&mut self, device: &Device) -> Result<()> {
572        self.q_proj = self.q_proj.to_device(device)?;
573        self.k_proj = self.k_proj.to_device(device)?;
574        self.v_proj = self.v_proj.to_device(device)?;
575        self.o_proj = self.o_proj.to_device(device)?;
576        Ok(())
577    }
578}
579
580impl MiniCPMMLP {
581    fn new(config: &MiniCPMConfig, device: &Device) -> Result<Self> {
582        Ok(Self {
583            gate_proj: ops_fn::zeros(&[config.hidden_size, config.intermediate_size], DataType::Float32, device)?,
584            up_proj: ops_fn::zeros(&[config.hidden_size, config.intermediate_size], DataType::Float32, device)?,
585            down_proj: ops_fn::zeros(&[config.intermediate_size, config.hidden_size], DataType::Float32, device)?,
586            hidden_act: config.hidden_act.clone(),
587        })
588    }
589
590    fn forward(&self, hidden_states: &Tensor) -> Result<Tensor> {
591        // 1. Gate and up projections
592        let gate_output = ops_fn::matmul(hidden_states, &self.gate_proj)?;
593        let up_output = ops_fn::matmul(hidden_states, &self.up_proj)?;
594
595        // 2. Apply activation (SiLU for MiniCPM)
596        let gate_activated = match self.hidden_act.as_str() {
597            "silu" | "swish" => ops_fn::silu(&gate_output)?,
598            "gelu" => ops_fn::gelu(&gate_output)?,
599            _ => return Err(anyhow::anyhow!("Unsupported activation: {}", self.hidden_act)),
600        };
601
602        // 3. Element-wise multiplication (gating)
603        let gated = ops_fn::mul(&gate_activated, &up_output)?;
604
605        // 4. Down projection
606        ops_fn::matmul(&gated, &self.down_proj)
607    }
608
609    fn load_weights(&mut self, weights: &ModelWeights, layer_idx: usize) -> Result<()> {
610        let prefix = format!("model.layers.{}.mlp", layer_idx);
611
612        // Load MLP weights (transpose for matmul: [out, in] -> [in, out])
613        if let Some(w) = weights.get(&format!("{}.gate_proj.weight", prefix)) {
614            self.gate_proj = ops_fn::transpose(w)?;
615        }
616        if let Some(w) = weights.get(&format!("{}.up_proj.weight", prefix)) {
617            self.up_proj = ops_fn::transpose(w)?;
618        }
619        if let Some(w) = weights.get(&format!("{}.down_proj.weight", prefix)) {
620            self.down_proj = ops_fn::transpose(w)?;
621        }
622        Ok(())
623    }
624
625    fn to_device(&mut self, device: &Device) -> Result<()> {
626        self.gate_proj = self.gate_proj.to_device(device)?;
627        self.up_proj = self.up_proj.to_device(device)?;
628        self.down_proj = self.down_proj.to_device(device)?;
629        Ok(())
630    }
631}
632
633#[cfg(test)]
634mod tests {
635    use super::*;
636
637    #[test]
638    fn test_minicpm_model_creation() {
639        let config = MiniCPMConfig {
640            vocab_size: 1000,
641            hidden_size: 128,
642            intermediate_size: 512,
643            num_hidden_layers: 2,
644            num_attention_heads: 8,
645            num_key_value_heads: 8,
646            ..Default::default()
647        };
648
649        let model = MiniCPMModelV2::new(config).unwrap();
650        assert_eq!(model.config().vocab_size(), 1000);
651        assert_eq!(model.config().hidden_size(), 128);
652        assert_eq!(model.config().num_layers(), 2);
653    }
654
655    #[test]
656    fn test_minicpm_forward_pass() {
657        let config = MiniCPMConfig {
658            vocab_size: 100,
659            hidden_size: 64,
660            intermediate_size: 256,
661            num_hidden_layers: 1,
662            num_attention_heads: 4,
663            num_key_value_heads: 4,
664            ..Default::default()
665        };
666
667        let model = MiniCPMModelV2::new(config).unwrap();
668        let input_ids = ops_fn::zeros(&[2, 8], DataType::Int64, &Device::CPU).unwrap();
669        let inputs = ModelInputs::text(input_ids);
670
671        let outputs = model.forward(&inputs).unwrap();
672        match outputs {
673            ModelOutputs::Logits { logits, .. } => {
674                assert_eq!(logits.shape(), &[2, 8, 100]);
675            }
676            _ => panic!("Expected logits output"),
677        }
678    }
679
680    #[test]
681    fn test_minicpm_generation() {
682        let config = MiniCPMConfig {
683            vocab_size: 256,
684            hidden_size: 64,
685            intermediate_size: 256,
686            num_hidden_layers: 1,
687            num_attention_heads: 4,
688            num_key_value_heads: 4,
689            ..Default::default()
690        };
691        let model = MiniCPMModelV2::new(config).unwrap();
692        let gen_config = GenerationConfig {
693            max_new_tokens: 5,
694            ..Default::default()
695        };
696
697        let output = model.generate("Hello", &gen_config).unwrap();
698        assert!(!output.is_empty());
699    }
700
701    #[test]
702    fn test_minicpm_gqa() {
703        // Test with GQA (num_key_value_heads < num_attention_heads)
704        let config = MiniCPMConfig {
705            vocab_size: 100,
706            hidden_size: 64,
707            intermediate_size: 256,
708            num_hidden_layers: 1,
709            num_attention_heads: 8,
710            num_key_value_heads: 2, // 4 groups
711            ..Default::default()
712        };
713
714        let model = MiniCPMModelV2::new(config).unwrap();
715        let input_ids = ops_fn::zeros(&[1, 4], DataType::Int64, &Device::CPU).unwrap();
716        let inputs = ModelInputs::text(input_ids);
717
718        let outputs = model.forward(&inputs).unwrap();
719        match outputs {
720            ModelOutputs::Logits { logits, .. } => {
721                assert_eq!(logits.shape(), &[1, 4, 100]);
722            }
723            _ => panic!("Expected logits output"),
724        }
725    }
726}