Skip to main content

runtime/models_v2/
chatglm.rs

1//! ChatGLM Model V2 - Clean implementation using solid abstractions
2//!
3//! This implements the ChatGLM architecture including:
4//! - ChatGLM-6B, ChatGLM2-6B, ChatGLM3-6B, GLM-4
5//!
6//! ChatGLM unique characteristics:
7//! - Structure: embedding -> encoder -> output_layer
8//! - Uses `transformer.embedding.word_embeddings.weight`
9//! - Uses `transformer.encoder.layers.{i}`
10//! - Uses `transformer.output_layer.weight`
11//! - Packed QKV attention (query_key_value combined)
12//! - SwiGLU activation in MLP
13
14use crate::model_config;
15use super::traits::*;
16use anyhow::Result;
17use serde::{Serialize, Deserialize};
18
19/// ChatGLM model configuration using the model_config macro
20model_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    // ChatGLM specific
39    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    /// Create ChatGLMConfig from GGUF model configuration
48    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
65/// Main ChatGLM model implementation
66pub struct ChatGLMModelV2 {
67    config: ChatGLMConfig,
68    device: Device,
69
70    // Model components using unified Tensor type
71    // ChatGLM uses: transformer.embedding.word_embeddings, transformer.encoder, transformer.output_layer
72    word_embeddings: Tensor,
73    layers: Vec<ChatGLMLayer>,
74    final_layernorm: Tensor,
75    output_layer: Tensor,
76}
77
78/// ChatGLM transformer layer
79pub struct ChatGLMLayer {
80    input_layernorm: Tensor,
81    self_attention: ChatGLMAttention,
82    post_attention_layernorm: Tensor,
83    mlp: ChatGLMMLP,
84}
85
86/// ChatGLM attention mechanism with packed QKV
87pub struct ChatGLMAttention {
88    // ChatGLM uses packed QKV: query_key_value combined weight
89    query_key_value: Tensor,
90    // Optional biases for QKV (ChatGLM often uses bias)
91    qkv_bias: Option<Tensor>,
92    // Output projection
93    dense: Tensor,
94    num_heads: usize,
95    num_key_value_heads: usize,
96    head_dim: usize,
97    scale: f32,
98}
99
100/// ChatGLM MLP with SwiGLU activation
101pub struct ChatGLMMLP {
102    // ChatGLM uses dense_h_to_4h (combined gate+up) and dense_4h_to_h
103    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        // Create embedding layer
115        let word_embeddings = ops_fn::zeros(
116            &[config.vocab_size, config.hidden_size],
117            DataType::Float32,
118            &device
119        )?;
120
121        // Final layer norm
122        let final_layernorm = ops_fn::zeros(
123            &[config.hidden_size],
124            DataType::Float32,
125            &device
126        )?;
127
128        // Output layer (lm_head equivalent)
129        let output_layer = ops_fn::zeros(
130            &[config.hidden_size, config.vocab_size],
131            DataType::Float32,
132            &device
133        )?;
134
135        // Create transformer layers
136        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        // Load embedding weights (ChatGLM uses transformer.embedding.word_embeddings.weight)
155        if let Some(embed_weights) = weights.get("transformer.embedding.word_embeddings.weight") {
156            model.word_embeddings = embed_weights.clone();
157        }
158
159        // Load final layer norm
160        if let Some(ln_weights) = weights.get("transformer.encoder.final_layernorm.weight") {
161            model.final_layernorm = ln_weights.clone();
162        }
163
164        // Load output layer (transpose for matmul: [vocab, hidden] -> [hidden, vocab])
165        if let Some(output_weights) = weights.get("transformer.output_layer.weight") {
166            model.output_layer = ops_fn::transpose(output_weights)?;
167        }
168
169        // Load layer weights
170        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                // 1. Token embedding
181                let mut hidden_states = ops_fn::embedding(input_ids, &self.word_embeddings)?;
182
183                // 2. Apply transformer layers (with RoPE)
184                for layer in &self.layers {
185                    hidden_states = layer.forward(&hidden_states, attention_mask.as_ref(), self.config.rope_theta)?;
186                }
187
188                // 3. Final layer norm
189                hidden_states = ops_fn::layer_norm(&hidden_states, &self.final_layernorm, None, self.config.rms_norm_eps)?;
190
191                // 4. Output layer (language modeling head)
192                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        // 1. Tokenize prompt
216        let tokenizer = Tokenizer::new();
217        let mut tokens: Vec<u32> = tokenizer.encode(prompt);
218
219        // 2. Generation loop
220        for _ in 0..config.max_new_tokens {
221            // Create input tensor from current tokens
222            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            // 3. Forward pass
232            let outputs = self.forward(&inputs)?;
233
234            // 4. Get logits and sample next token
235            let logits = match outputs {
236                ModelOutputs::Logits { logits, .. } => logits,
237                _ => return Err(anyhow::anyhow!("Expected logits output")),
238            };
239
240            // Get last token logits
241            let logits_candle = logits.to_candle()?;
242            let shape = logits_candle.dims();
243
244            // Extract last position logits [batch, seq, vocab] -> [vocab]
245            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            // Convert to probabilities and sample
259            let logits_vec: Vec<f32> = last_logits.to_vec1()?;
260
261            let next_token = if config.do_sample && config.temperature > 0.0 {
262                // Temperature sampling
263                let scaled: Vec<f32> = logits_vec.iter()
264                    .map(|&x| x / config.temperature)
265                    .collect();
266
267                // Softmax
268                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                // Sample from distribution
275                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                // Greedy sampling
290                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            // 5. Check for EOS
302            if next_token == config.eos_token_id {
303                break;
304            }
305
306            // 6. Append token
307            tokens.push(next_token);
308        }
309
310        // 7. Decode and return
311        Ok(tokenizer.decode(&tokens))
312    }
313
314    fn config(&self) -> &Self::Config {
315        &self.config
316    }
317
318    fn memory_requirements(&self) -> MemoryRequirements {
319        // Calculate approximate memory requirements
320        let param_size = self.config.vocab_size * self.config.hidden_size + // embeddings
321                        self.config.num_hidden_layers * (
322                            // Packed QKV + dense projection
323                            (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                            // MLP
327                            2 * self.config.hidden_size * self.config.intermediate_size
328                        );
329
330        let param_bytes = param_size * 4; // float32
331        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        // 1. Pre-attention layer norm
374        let normed = ops_fn::layer_norm(hidden_states, &self.input_layernorm, None, 1e-5)?;
375
376        // 2. Self attention (with RoPE and packed QKV)
377        let attn_output = self.self_attention.forward(&normed, attention_mask, rope_theta)?;
378
379        // 3. Residual connection
380        let hidden_states = ops_fn::add(hidden_states, &attn_output)?;
381
382        // 4. Pre-MLP layer norm
383        let normed = ops_fn::layer_norm(&hidden_states, &self.post_attention_layernorm, None, 1e-5)?;
384
385        // 5. MLP with SwiGLU
386        let mlp_output = self.mlp.forward(&normed)?;
387
388        // 6. Residual connection
389        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        // Load layer norms
398        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        // Load attention weights
406        self.self_attention.load_weights(weights, layer_idx)?;
407
408        // Load MLP weights
409        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
423/// Apply Rotary Position Embedding (RoPE) to Q and K tensors
424/// Input shape: [batch, heads, seq, head_dim]
425/// Returns tensors with same shape but with positional information encoded
426fn 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    // Compute inverse frequencies: 1 / (theta^(2i/d)) for i in [0, d/2)
436    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    // Create position indices [0, 1, 2, ..., seq_len-1]
442    let positions: Vec<f32> = (0..seq_len).map(|p| p as f32).collect();
443
444    // Compute angles: pos * inv_freq -> [seq_len, half_dim]
445    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    // Compute cos and sin
455    let cos = angles_tensor.cos()?;
456    let sin = angles_tensor.sin()?;
457
458    // Reshape for broadcasting: [1, 1, seq_len, half_dim]
459    let cos = cos.unsqueeze(0)?.unsqueeze(0)?;
460    let sin = sin.unsqueeze(0)?.unsqueeze(0)?;
461
462    // Apply RoPE rotation
463    // Split q and k into two halves along head_dim
464    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    // Apply rotation
470    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    // Concatenate rotated halves
476    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        // ChatGLM uses packed QKV: [hidden, (num_heads + 2*num_kv_heads) * head_dim]
490        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        // Optional QKV bias
498        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        // Output projection (dense)
505        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        // Get batch and sequence length from hidden_states shape
524        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        // 1. Compute packed QKV projection
534        let qkv = ops_fn::matmul(hidden_states, &self.query_key_value)?;
535
536        // Add bias if present
537        let qkv = if let Some(ref bias) = self.qkv_bias {
538            ops_fn::add(&qkv, bias)?
539        } else {
540            qkv
541        };
542
543        // 2. Split QKV into Q, K, V
544        let qkv_candle = qkv.to_candle()?;
545
546        // QKV layout: [batch, seq, (num_heads + 2*num_kv_heads) * head_dim]
547        // Split into: Q [batch, seq, num_heads * head_dim]
548        //             K [batch, seq, num_kv_heads * head_dim]
549        //             V [batch, seq, num_kv_heads * head_dim]
550        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        // 3. Reshape for multi-head attention
558        // Q: [batch, seq, num_heads * head_dim] -> [batch, num_heads, seq, head_dim]
559        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        // 4. Apply RoPE (Rotary Position Embedding) to Q and K
572        let (q_with_rope, k_with_rope) = apply_rope(&q_reshaped, &k_reshaped, seq_len, self.head_dim, rope_theta)?;
573
574        // 5. Handle GQA (Grouped Query Attention) - repeat K/V heads to match Q heads
575        let num_groups = self.num_heads / self.num_key_value_heads;
576        let (k_expanded, v_expanded) = if num_groups > 1 {
577            // Repeat K and V along the head dimension
578            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        // 6. Scaled dot-product attention
592        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        // Apply causal mask
601        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        // Softmax over last dimension
617        let attention_weights = candle_nn::ops::softmax_last_dim(&masked_scores)?;
618
619        // Apply attention to values
620        let v_contiguous = v_expanded.contiguous()?;
621        let attn_output = attention_weights.matmul(&v_contiguous)?;
622
623        // 7. Reshape back: [batch, heads, seq, head_dim] -> [batch, seq, heads * head_dim]
624        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        // 8. Output projection (dense)
631        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        // Load packed QKV weight (transpose for matmul: [out, in] -> [in, out])
640        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        // Load QKV bias if present
645        if let Some(qkv_bias) = weights.get(&format!("{}.query_key_value.bias", prefix)) {
646            self.qkv_bias = Some(qkv_bias.clone());
647        }
648
649        // Load output projection (dense) - transpose for matmul
650        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        // ChatGLM MLP uses SwiGLU, with combined gate+up projection
670        // dense_h_to_4h: [hidden, intermediate*2] (for gate and up combined)
671        // dense_4h_to_h: [intermediate, hidden]
672        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        // 1. Combined gate+up projection
692        let h_to_4h = ops_fn::matmul(hidden_states, &self.dense_h_to_4h)?;
693
694        // 2. Split into gate and up parts for SwiGLU
695        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        // 3. Apply SwiGLU: activation(gate) * up
703        // ChatGLM uses SwiGLU (SiLU/Swish for the gate)
704        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)?, // Default to SiLU for ChatGLM
711        };
712        let gated = ops_fn::mul(&gate_activated, &up_tensor)?;
713
714        // 4. Down projection
715        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        // Load MLP weights (transpose for matmul: [out, in] -> [in, out])
724        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]); // batch, seq, vocab
783            }
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}