1use 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 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 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>, }
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
96fn 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 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 let positions: Vec<f32> = (0..seq_len).map(|p| p as f32).collect();
116
117 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 let cos = angles_tensor.cos()?;
129 let sin = angles_tensor.sin()?;
130
131 let cos = cos.unsqueeze(0)?.unsqueeze(0)?;
133 let sin = sin.unsqueeze(0)?.unsqueeze(0)?;
134
135 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 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 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 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 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 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 for layer in &self.layers {
220 hidden_states = layer.forward(&hidden_states, None, self.config.rope_theta)?;
221 }
222
223 hidden_states = ops_fn::layer_norm(&hidden_states, &self.norm, None, self.config.rms_norm_eps)?;
225
226 let logits = if let Some(ref lm_head) = self.lm_head {
228 ops_fn::matmul(&hidden_states, lm_head)?
230 } else {
231 let lm_head_t = ops_fn::transpose(&self.embed_tokens)?;
234 ops_fn::matmul(&hidden_states, &lm_head_t)?
235 };
236
237 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 let tokenizer = Tokenizer::new();
249 let mut tokens: Vec<u32> = tokenizer.encode(prompt);
250
251 for _ in 0..config.max_new_tokens {
253 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 let outputs = self.forward(&inputs)?;
265
266 let logits = match outputs {
268 ModelOutputs::Logits { logits, .. } => logits,
269 _ => return Err(anyhow::anyhow!("Expected logits output")),
270 };
271
272 let logits_candle = logits.to_candle()?;
274 let shape = logits_candle.dims();
275
276 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 let logits_vec: Vec<f32> = last_logits.to_vec1()?;
292
293 let next_token = if config.do_sample && config.temperature > 0.0 {
294 let scaled: Vec<f32> = logits_vec.iter()
296 .map(|&x| x / config.temperature)
297 .collect();
298
299 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 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 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 if next_token == config.eos_token_id {
335 break;
336 }
337
338 tokens.push(next_token);
340 }
341
342 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 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 let normed = ops_fn::layer_norm(hidden_states, &self.input_layernorm, None, 1e-5)?;
391
392 let attn_output = self.self_attn.forward(&normed, None, rope_theta)?;
394
395 let attn_scaled = ops_fn::scale(&attn_output, self.scale)?;
397 let hidden_states = ops_fn::add(&residual, &attn_scaled)?;
398
399 let residual = hidden_states.clone();
401 let normed = ops_fn::layer_norm(&hidden_states, &self.post_attention_layernorm, None, 1e-5)?;
402
403 let mlp_output = self.mlp.forward(&normed)?;
405
406 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 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 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 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 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 let (q_with_rope, k_with_rope) = apply_rope(&q_reshaped, &k_reshaped, seq_len, self.head_dim, rope_theta)?;
490
491 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 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 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 let attention_weights = candle_nn::ops::softmax_last_dim(&masked_scores)?;
534
535 let v_contiguous = v_expanded.contiguous()?;
537 let attn_output = attention_weights.matmul(&v_contiguous)?;
538
539 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 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 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 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 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 let gated = ops_fn::mul(&gate_activated, &up_output)?;
604
605 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 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 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, ..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}