1use crate::model_config;
15use super::traits::*;
16use anyhow::Result;
17use serde::{Serialize, Deserialize};
18
19model_config!(ChatGLMConfig {
21 vocab_size: usize = 65024,
22 hidden_size: usize = 4096,
23 intermediate_size: usize = 13696,
24 num_hidden_layers: usize = 28,
25 num_attention_heads: usize = 32,
26 num_key_value_heads: usize = 2,
27 hidden_act: String = "swiglu".to_string(),
28 max_position_embeddings: usize = 8192,
29 initializer_range: f32 = 0.02,
30 rms_norm_eps: f32 = 1e-5,
31 use_cache: bool = true,
32 pad_token_id: i64 = 0,
33 bos_token_id: i64 = 1,
34 eos_token_id: i64 = 2,
35 tie_word_embeddings: bool = false,
36 rope_theta: f32 = 10000.0,
37 attention_dropout: f32 = 0.0,
38 add_bias_linear: bool = false,
40 add_qkv_bias: bool = true,
41 apply_residual_connection_post_layernorm: bool = false,
42 kv_channels: usize = 128,
43 multi_query_attention: bool = true,
44});
45
46impl ChatGLMConfig {
47 pub fn from_gguf_config(gguf: &crate::weight_loader_core::GGUFModelConfig) -> Self {
49 Self {
50 vocab_size: gguf.vocab_size,
51 hidden_size: gguf.hidden_size,
52 intermediate_size: gguf.intermediate_size,
53 num_hidden_layers: gguf.num_hidden_layers,
54 num_attention_heads: gguf.num_attention_heads,
55 num_key_value_heads: gguf.num_key_value_heads,
56 rms_norm_eps: gguf.rms_norm_eps,
57 rope_theta: gguf.rope_theta,
58 max_position_embeddings: gguf.max_position_embeddings,
59 kv_channels: gguf.head_dim,
60 ..Default::default()
61 }
62 }
63}
64
65pub struct ChatGLMModelV2 {
67 config: ChatGLMConfig,
68 device: Device,
69
70 word_embeddings: Tensor,
73 layers: Vec<ChatGLMLayer>,
74 final_layernorm: Tensor,
75 output_layer: Tensor,
76}
77
78pub struct ChatGLMLayer {
80 input_layernorm: Tensor,
81 self_attention: ChatGLMAttention,
82 post_attention_layernorm: Tensor,
83 mlp: ChatGLMMLP,
84}
85
86pub struct ChatGLMAttention {
88 query_key_value: Tensor,
90 qkv_bias: Option<Tensor>,
92 dense: Tensor,
94 num_heads: usize,
95 num_key_value_heads: usize,
96 head_dim: usize,
97 scale: f32,
98}
99
100pub struct ChatGLMMLP {
102 dense_h_to_4h: Tensor,
104 dense_4h_to_h: Tensor,
105 hidden_act: String,
106}
107
108impl Model for ChatGLMModelV2 {
109 type Config = ChatGLMConfig;
110
111 fn new(config: ChatGLMConfig) -> Result<Self> {
112 let device = Device::CPU;
113
114 let word_embeddings = ops_fn::zeros(
116 &[config.vocab_size, config.hidden_size],
117 DataType::Float32,
118 &device
119 )?;
120
121 let final_layernorm = ops_fn::zeros(
123 &[config.hidden_size],
124 DataType::Float32,
125 &device
126 )?;
127
128 let output_layer = ops_fn::zeros(
130 &[config.hidden_size, config.vocab_size],
131 DataType::Float32,
132 &device
133 )?;
134
135 let mut layers = Vec::with_capacity(config.num_hidden_layers);
137 for _ in 0..config.num_hidden_layers {
138 layers.push(ChatGLMLayer::new(&config, &device)?);
139 }
140
141 Ok(Self {
142 config,
143 device,
144 word_embeddings,
145 layers,
146 final_layernorm,
147 output_layer,
148 })
149 }
150
151 fn from_weights(config: ChatGLMConfig, weights: ModelWeights) -> Result<Self> {
152 let mut model = Self::new(config)?;
153
154 if let Some(embed_weights) = weights.get("transformer.embedding.word_embeddings.weight") {
156 model.word_embeddings = embed_weights.clone();
157 }
158
159 if let Some(ln_weights) = weights.get("transformer.encoder.final_layernorm.weight") {
161 model.final_layernorm = ln_weights.clone();
162 }
163
164 if let Some(output_weights) = weights.get("transformer.output_layer.weight") {
166 model.output_layer = ops_fn::transpose(output_weights)?;
167 }
168
169 for (i, layer) in model.layers.iter_mut().enumerate() {
171 layer.load_weights(&weights, i)?;
172 }
173
174 Ok(model)
175 }
176
177 fn forward(&self, inputs: &ModelInputs) -> Result<ModelOutputs> {
178 match inputs {
179 ModelInputs::Text { input_ids, attention_mask, .. } => {
180 let mut hidden_states = ops_fn::embedding(input_ids, &self.word_embeddings)?;
182
183 for layer in &self.layers {
185 hidden_states = layer.forward(&hidden_states, attention_mask.as_ref(), self.config.rope_theta)?;
186 }
187
188 hidden_states = ops_fn::layer_norm(&hidden_states, &self.final_layernorm, None, self.config.rms_norm_eps)?;
190
191 let logits = ops_fn::matmul(&hidden_states, &self.output_layer)?;
193
194 Ok(ModelOutputs::Logits {
195 logits,
196 hidden_states: None,
197 })
198 }
199 ModelInputs::Multimodal { input_ids, .. } => {
200 let text_inputs = ModelInputs::Text {
201 input_ids: input_ids.clone(),
202 attention_mask: None,
203 position_ids: None,
204 };
205 self.forward(&text_inputs)
206 }
207 _ => Err(anyhow::anyhow!("ChatGLM model only supports text and multimodal inputs")),
208 }
209 }
210
211 fn generate(&self, prompt: &str, config: &GenerationConfig) -> Result<String> {
212 use crate::tokenizer::Tokenizer;
213 use rand::Rng;
214
215 let tokenizer = Tokenizer::new();
217 let mut tokens: Vec<u32> = tokenizer.encode(prompt);
218
219 for _ in 0..config.max_new_tokens {
221 let tokens_i64: Vec<i64> = tokens.iter().map(|&t| t as i64).collect();
223 let input_tensor = Tensor::from_i64_slice(&tokens_i64, &[1, tokens.len()], &self.device)?;
224
225 let inputs = ModelInputs::Text {
226 input_ids: input_tensor,
227 attention_mask: None,
228 position_ids: None,
229 };
230
231 let outputs = self.forward(&inputs)?;
233
234 let logits = match outputs {
236 ModelOutputs::Logits { logits, .. } => logits,
237 _ => return Err(anyhow::anyhow!("Expected logits output")),
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 let seq_len = shape[1];
247 logits_candle
248 .narrow(1, seq_len - 1, 1)?
249 .squeeze(1)?
250 .squeeze(0)?
251 } else {
252 let seq_len = shape[0];
253 logits_candle
254 .narrow(0, seq_len - 1, 1)?
255 .squeeze(0)?
256 };
257
258 let logits_vec: Vec<f32> = last_logits.to_vec1()?;
260
261 let next_token = if config.do_sample && config.temperature > 0.0 {
262 let scaled: Vec<f32> = logits_vec.iter()
264 .map(|&x| x / config.temperature)
265 .collect();
266
267 let max_val = scaled.iter().cloned().fold(f32::NEG_INFINITY, f32::max);
269 let exp_sum: f32 = scaled.iter().map(|&x| (x - max_val).exp()).sum();
270 let probs: Vec<f32> = scaled.iter()
271 .map(|&x| (x - max_val).exp() / exp_sum)
272 .collect();
273
274 let mut rng = rand::thread_rng();
276 let random_val: f32 = rng.gen();
277 let mut cumulative = 0.0;
278 let mut sampled = 0u32;
279
280 for (idx, &prob) in probs.iter().enumerate() {
281 cumulative += prob;
282 if random_val <= cumulative {
283 sampled = idx as u32;
284 break;
285 }
286 }
287 sampled
288 } else {
289 let mut max_idx = 0;
291 let mut max_val = logits_vec[0];
292 for (idx, &val) in logits_vec.iter().enumerate() {
293 if val > max_val {
294 max_val = val;
295 max_idx = idx;
296 }
297 }
298 max_idx as u32
299 };
300
301 if next_token == config.eos_token_id {
303 break;
304 }
305
306 tokens.push(next_token);
308 }
309
310 Ok(tokenizer.decode(&tokens))
312 }
313
314 fn config(&self) -> &Self::Config {
315 &self.config
316 }
317
318 fn memory_requirements(&self) -> MemoryRequirements {
319 let param_size = self.config.vocab_size * self.config.hidden_size + self.config.num_hidden_layers * (
322 (self.config.num_attention_heads + 2 * self.config.num_key_value_heads) *
324 (self.config.hidden_size / self.config.num_attention_heads) * self.config.hidden_size +
325 self.config.hidden_size * self.config.hidden_size +
326 2 * self.config.hidden_size * self.config.intermediate_size
328 );
329
330 let param_bytes = param_size * 4; let kv_cache_bytes = 2 * self.config.num_hidden_layers *
332 self.config.max_position_embeddings *
333 self.config.hidden_size * 4;
334
335 MemoryRequirements {
336 gpu_memory: param_bytes,
337 cpu_memory: param_bytes / 4,
338 kv_cache_memory: kv_cache_bytes,
339 peak_memory: param_bytes + kv_cache_bytes,
340 }
341 }
342
343 fn to_device(&mut self, device: &Device) -> Result<()> {
344 self.word_embeddings = self.word_embeddings.to_device(device)?;
345 self.final_layernorm = self.final_layernorm.to_device(device)?;
346 self.output_layer = self.output_layer.to_device(device)?;
347
348 for layer in &mut self.layers {
349 layer.to_device(device)?;
350 }
351
352 self.device = device.clone();
353 Ok(())
354 }
355}
356
357impl ChatGLMLayer {
358 fn new(config: &ChatGLMConfig, device: &Device) -> Result<Self> {
359 let input_layernorm = ops_fn::zeros(&[config.hidden_size], DataType::Float32, device)?;
360 let post_attention_layernorm = ops_fn::zeros(&[config.hidden_size], DataType::Float32, device)?;
361 let self_attention = ChatGLMAttention::new(config, device)?;
362 let mlp = ChatGLMMLP::new(config, device)?;
363
364 Ok(Self {
365 input_layernorm,
366 self_attention,
367 post_attention_layernorm,
368 mlp,
369 })
370 }
371
372 fn forward(&self, hidden_states: &Tensor, attention_mask: Option<&Tensor>, rope_theta: f32) -> Result<Tensor> {
373 let normed = ops_fn::layer_norm(hidden_states, &self.input_layernorm, None, 1e-5)?;
375
376 let attn_output = self.self_attention.forward(&normed, attention_mask, rope_theta)?;
378
379 let hidden_states = ops_fn::add(hidden_states, &attn_output)?;
381
382 let normed = ops_fn::layer_norm(&hidden_states, &self.post_attention_layernorm, None, 1e-5)?;
384
385 let mlp_output = self.mlp.forward(&normed)?;
387
388 let output = ops_fn::add(&hidden_states, &mlp_output)?;
390
391 Ok(output)
392 }
393
394 fn load_weights(&mut self, weights: &ModelWeights, layer_idx: usize) -> Result<()> {
395 let prefix = format!("transformer.encoder.layers.{}", layer_idx);
396
397 if let Some(w) = weights.get(&format!("{}.input_layernorm.weight", prefix)) {
399 self.input_layernorm = w.clone();
400 }
401 if let Some(w) = weights.get(&format!("{}.post_attention_layernorm.weight", prefix)) {
402 self.post_attention_layernorm = w.clone();
403 }
404
405 self.self_attention.load_weights(weights, layer_idx)?;
407
408 self.mlp.load_weights(weights, layer_idx)?;
410
411 Ok(())
412 }
413
414 fn to_device(&mut self, device: &Device) -> Result<()> {
415 self.input_layernorm = self.input_layernorm.to_device(device)?;
416 self.post_attention_layernorm = self.post_attention_layernorm.to_device(device)?;
417 self.self_attention.to_device(device)?;
418 self.mlp.to_device(device)?;
419 Ok(())
420 }
421}
422
423fn apply_rope(
427 q: &candle_core::Tensor,
428 k: &candle_core::Tensor,
429 seq_len: usize,
430 head_dim: usize,
431 rope_theta: f32,
432) -> Result<(candle_core::Tensor, candle_core::Tensor)> {
433 let device = q.device();
434
435 let half_dim = head_dim / 2;
437 let inv_freq: Vec<f32> = (0..half_dim)
438 .map(|i| 1.0 / rope_theta.powf((2 * i) as f32 / head_dim as f32))
439 .collect();
440
441 let positions: Vec<f32> = (0..seq_len).map(|p| p as f32).collect();
443
444 let mut angles = Vec::with_capacity(seq_len * half_dim);
446 for pos in &positions {
447 for freq in &inv_freq {
448 angles.push(pos * freq);
449 }
450 }
451
452 let angles_tensor = candle_core::Tensor::from_vec(angles, &[seq_len, half_dim], device)?;
453
454 let cos = angles_tensor.cos()?;
456 let sin = angles_tensor.sin()?;
457
458 let cos = cos.unsqueeze(0)?.unsqueeze(0)?;
460 let sin = sin.unsqueeze(0)?.unsqueeze(0)?;
461
462 let q_half1 = q.narrow(3, 0, half_dim)?;
465 let q_half2 = q.narrow(3, half_dim, half_dim)?;
466 let k_half1 = k.narrow(3, 0, half_dim)?;
467 let k_half2 = k.narrow(3, half_dim, half_dim)?;
468
469 let q_rot1 = (q_half1.broadcast_mul(&cos)? - q_half2.broadcast_mul(&sin)?)?;
471 let q_rot2 = (q_half1.broadcast_mul(&sin)? + q_half2.broadcast_mul(&cos)?)?;
472 let k_rot1 = (k_half1.broadcast_mul(&cos)? - k_half2.broadcast_mul(&sin)?)?;
473 let k_rot2 = (k_half1.broadcast_mul(&sin)? + k_half2.broadcast_mul(&cos)?)?;
474
475 let q_rotated = candle_core::Tensor::cat(&[&q_rot1, &q_rot2], 3)?;
477 let k_rotated = candle_core::Tensor::cat(&[&k_rot1, &k_rot2], 3)?;
478
479 Ok((q_rotated, k_rotated))
480}
481
482impl ChatGLMAttention {
483 fn new(config: &ChatGLMConfig, device: &Device) -> Result<Self> {
484 let num_heads = config.num_attention_heads;
485 let num_key_value_heads = config.num_key_value_heads;
486 let head_dim = config.hidden_size / num_heads;
487 let scale = 1.0 / (head_dim as f32).sqrt();
488
489 let qkv_size = (num_heads + 2 * num_key_value_heads) * head_dim;
491 let query_key_value = ops_fn::zeros(
492 &[config.hidden_size, qkv_size],
493 DataType::Float32,
494 device
495 )?;
496
497 let qkv_bias = if config.add_qkv_bias {
499 Some(ops_fn::zeros(&[qkv_size], DataType::Float32, device)?)
500 } else {
501 None
502 };
503
504 let dense = ops_fn::zeros(
506 &[num_heads * head_dim, config.hidden_size],
507 DataType::Float32,
508 device
509 )?;
510
511 Ok(Self {
512 query_key_value,
513 qkv_bias,
514 dense,
515 num_heads,
516 num_key_value_heads,
517 head_dim,
518 scale,
519 })
520 }
521
522 fn forward(&self, hidden_states: &Tensor, _attention_mask: Option<&Tensor>, rope_theta: f32) -> Result<Tensor> {
523 let shape = hidden_states.shape();
525 let (batch_size, seq_len, _hidden_size) = if shape.len() == 3 {
526 (shape[0], shape[1], shape[2])
527 } else if shape.len() == 2 {
528 (1, shape[0], shape[1])
529 } else {
530 return Err(anyhow::anyhow!("Invalid hidden_states shape: {:?}", shape));
531 };
532
533 let qkv = ops_fn::matmul(hidden_states, &self.query_key_value)?;
535
536 let qkv = if let Some(ref bias) = self.qkv_bias {
538 ops_fn::add(&qkv, bias)?
539 } else {
540 qkv
541 };
542
543 let qkv_candle = qkv.to_candle()?;
545
546 let q_size = self.num_heads * self.head_dim;
551 let kv_size = self.num_key_value_heads * self.head_dim;
552
553 let q = qkv_candle.narrow(2, 0, q_size)?;
554 let k = qkv_candle.narrow(2, q_size, kv_size)?;
555 let v = qkv_candle.narrow(2, q_size + kv_size, kv_size)?;
556
557 let q_reshaped = q
560 .reshape(&[batch_size, seq_len, self.num_heads, self.head_dim])?
561 .transpose(1, 2)?;
562
563 let k_reshaped = k
564 .reshape(&[batch_size, seq_len, self.num_key_value_heads, self.head_dim])?
565 .transpose(1, 2)?;
566
567 let v_reshaped = v
568 .reshape(&[batch_size, seq_len, self.num_key_value_heads, self.head_dim])?
569 .transpose(1, 2)?;
570
571 let (q_with_rope, k_with_rope) = apply_rope(&q_reshaped, &k_reshaped, seq_len, self.head_dim, rope_theta)?;
573
574 let num_groups = self.num_heads / self.num_key_value_heads;
576 let (k_expanded, v_expanded) = if num_groups > 1 {
577 let k_rep = k_with_rope
579 .unsqueeze(2)?
580 .broadcast_as(&[batch_size, self.num_key_value_heads, num_groups, seq_len, self.head_dim])?
581 .reshape(&[batch_size, self.num_heads, seq_len, self.head_dim])?;
582 let v_rep = v_reshaped
583 .unsqueeze(2)?
584 .broadcast_as(&[batch_size, self.num_key_value_heads, num_groups, seq_len, self.head_dim])?
585 .reshape(&[batch_size, self.num_heads, seq_len, self.head_dim])?;
586 (k_rep, v_rep)
587 } else {
588 (k_with_rope, v_reshaped)
589 };
590
591 let k_t = k_expanded.transpose(2, 3)?;
593
594 let q_contiguous = q_with_rope.contiguous()?;
595 let k_contiguous = k_t.contiguous()?;
596
597 let scores = q_contiguous.matmul(&k_contiguous)?;
598 let scaled_scores = (scores * (self.scale as f64))?;
599
600 let device = scaled_scores.device();
602 let causal_mask = {
603 let mut mask_data = vec![0.0f32; seq_len * seq_len];
604 for i in 0..seq_len {
605 for j in 0..seq_len {
606 if j > i {
607 mask_data[i * seq_len + j] = f32::NEG_INFINITY;
608 }
609 }
610 }
611 candle_core::Tensor::from_vec(mask_data, &[1, 1, seq_len, seq_len], device)?
612 };
613
614 let masked_scores = scaled_scores.broadcast_add(&causal_mask)?;
615
616 let attention_weights = candle_nn::ops::softmax_last_dim(&masked_scores)?;
618
619 let v_contiguous = v_expanded.contiguous()?;
621 let attn_output = attention_weights.matmul(&v_contiguous)?;
622
623 let attn_output = attn_output
625 .transpose(1, 2)?
626 .reshape(&[batch_size, seq_len, self.num_heads * self.head_dim])?;
627
628 let attn_output = Tensor::from_candle(attn_output);
629
630 let output = ops_fn::matmul(&attn_output, &self.dense)?;
632
633 Ok(output)
634 }
635
636 fn load_weights(&mut self, weights: &ModelWeights, layer_idx: usize) -> Result<()> {
637 let prefix = format!("transformer.encoder.layers.{}.self_attention", layer_idx);
638
639 if let Some(qkv_weight) = weights.get(&format!("{}.query_key_value.weight", prefix)) {
641 self.query_key_value = ops_fn::transpose(qkv_weight)?;
642 }
643
644 if let Some(qkv_bias) = weights.get(&format!("{}.query_key_value.bias", prefix)) {
646 self.qkv_bias = Some(qkv_bias.clone());
647 }
648
649 if let Some(dense_weight) = weights.get(&format!("{}.dense.weight", prefix)) {
651 self.dense = ops_fn::transpose(dense_weight)?;
652 }
653
654 Ok(())
655 }
656
657 fn to_device(&mut self, device: &Device) -> Result<()> {
658 self.query_key_value = self.query_key_value.to_device(device)?;
659 if let Some(ref mut bias) = self.qkv_bias {
660 *bias = bias.to_device(device)?;
661 }
662 self.dense = self.dense.to_device(device)?;
663 Ok(())
664 }
665}
666
667impl ChatGLMMLP {
668 fn new(config: &ChatGLMConfig, device: &Device) -> Result<Self> {
669 let dense_h_to_4h = ops_fn::zeros(
673 &[config.hidden_size, config.intermediate_size * 2],
674 DataType::Float32,
675 device
676 )?;
677 let dense_4h_to_h = ops_fn::zeros(
678 &[config.intermediate_size, config.hidden_size],
679 DataType::Float32,
680 device
681 )?;
682
683 Ok(Self {
684 dense_h_to_4h,
685 dense_4h_to_h,
686 hidden_act: config.hidden_act.clone(),
687 })
688 }
689
690 fn forward(&self, hidden_states: &Tensor) -> Result<Tensor> {
691 let h_to_4h = ops_fn::matmul(hidden_states, &self.dense_h_to_4h)?;
693
694 let h_to_4h_candle = h_to_4h.to_candle()?;
696 let shape = h_to_4h_candle.dims();
697 let half_size = shape[shape.len() - 1] / 2;
698
699 let gate = h_to_4h_candle.narrow(shape.len() - 1, 0, half_size)?;
700 let up = h_to_4h_candle.narrow(shape.len() - 1, half_size, half_size)?;
701
702 let gate_tensor = Tensor::from_candle(gate);
705 let up_tensor = Tensor::from_candle(up);
706
707 let gate_activated = match self.hidden_act.as_str() {
708 "swiglu" | "silu" | "swish" => ops_fn::silu(&gate_tensor)?,
709 "gelu" => ops_fn::gelu(&gate_tensor)?,
710 _ => ops_fn::silu(&gate_tensor)?, };
712 let gated = ops_fn::mul(&gate_activated, &up_tensor)?;
713
714 let output = ops_fn::matmul(&gated, &self.dense_4h_to_h)?;
716
717 Ok(output)
718 }
719
720 fn load_weights(&mut self, weights: &ModelWeights, layer_idx: usize) -> Result<()> {
721 let prefix = format!("transformer.encoder.layers.{}.mlp", layer_idx);
722
723 if let Some(h_to_4h) = weights.get(&format!("{}.dense_h_to_4h.weight", prefix)) {
725 self.dense_h_to_4h = ops_fn::transpose(h_to_4h)?;
726 }
727 if let Some(h4_to_h) = weights.get(&format!("{}.dense_4h_to_h.weight", prefix)) {
728 self.dense_4h_to_h = ops_fn::transpose(h4_to_h)?;
729 }
730
731 Ok(())
732 }
733
734 fn to_device(&mut self, device: &Device) -> Result<()> {
735 self.dense_h_to_4h = self.dense_h_to_4h.to_device(device)?;
736 self.dense_4h_to_h = self.dense_4h_to_h.to_device(device)?;
737 Ok(())
738 }
739}
740
741#[cfg(test)]
742mod tests {
743 use super::*;
744
745 #[test]
746 fn test_chatglm_model_creation() {
747 let config = ChatGLMConfig {
748 vocab_size: 1000,
749 hidden_size: 128,
750 intermediate_size: 512,
751 num_hidden_layers: 2,
752 num_attention_heads: 8,
753 num_key_value_heads: 2,
754 ..Default::default()
755 };
756
757 let model = ChatGLMModelV2::new(config).unwrap();
758 assert_eq!(model.config().vocab_size(), 1000);
759 assert_eq!(model.config().hidden_size(), 128);
760 assert_eq!(model.config().num_layers(), 2);
761 }
762
763 #[test]
764 fn test_chatglm_forward_pass() {
765 let config = ChatGLMConfig {
766 vocab_size: 100,
767 hidden_size: 64,
768 intermediate_size: 256,
769 num_hidden_layers: 1,
770 num_attention_heads: 4,
771 num_key_value_heads: 2,
772 ..Default::default()
773 };
774
775 let model = ChatGLMModelV2::new(config).unwrap();
776 let input_ids = ops_fn::zeros(&[2, 8], DataType::Int64, &Device::CPU).unwrap();
777 let inputs = ModelInputs::text(input_ids);
778
779 let outputs = model.forward(&inputs).unwrap();
780 match outputs {
781 ModelOutputs::Logits { logits, .. } => {
782 assert_eq!(logits.shape(), &[2, 8, 100]); }
784 _ => panic!("Expected logits output"),
785 }
786 }
787
788 #[test]
789 fn test_chatglm_generation() {
790 let config = ChatGLMConfig {
791 vocab_size: 256,
792 hidden_size: 64,
793 intermediate_size: 256,
794 num_hidden_layers: 1,
795 num_attention_heads: 4,
796 num_key_value_heads: 2,
797 ..Default::default()
798 };
799 let model = ChatGLMModelV2::new(config).unwrap();
800 let gen_config = GenerationConfig {
801 max_new_tokens: 5,
802 ..Default::default()
803 };
804
805 let output = model.generate("Hello", &gen_config).unwrap();
806 assert!(!output.is_empty());
807 }
808
809 #[test]
810 fn test_chatglm_from_gguf_config() {
811 let gguf_config = crate::weight_loader_core::GGUFModelConfig {
812 architecture: "chatglm".to_string(),
813 vocab_size: 65024,
814 hidden_size: 4096,
815 intermediate_size: 13696,
816 num_hidden_layers: 28,
817 num_attention_heads: 32,
818 num_key_value_heads: 2,
819 head_dim: 128,
820 rms_norm_eps: 1e-5,
821 rope_theta: 10000.0,
822 max_position_embeddings: 8192,
823 };
824
825 let config = ChatGLMConfig::from_gguf_config(&gguf_config);
826 assert_eq!(config.vocab_size, 65024);
827 assert_eq!(config.hidden_size, 4096);
828 assert_eq!(config.num_key_value_heads, 2);
829 assert_eq!(config.kv_channels, 128);
830 }
831}