Skip to main content

runtime/models_v2/
gpt2.rs

1//! GPT-2 Model V2 - Clean implementation using solid abstractions
2//!
3//! This implements the GPT-2 architecture which features:
4//! - Learned absolute position embeddings
5//! - Standard multi-head attention (no GQA)
6//! - Pre-norm or post-norm layer normalization
7//! - Uses unified Tensor type from tensor_core
8
9use crate::model_config;
10use super::traits::*;
11use anyhow::Result;
12use serde::{Serialize, Deserialize};
13
14/// GPT-2 model configuration
15model_config!(GPT2Config {
16    vocab_size: usize = 50257,
17    hidden_size: usize = 768,
18    intermediate_size: usize = 3072,
19    num_hidden_layers: usize = 12,
20    num_attention_heads: usize = 12,
21    num_key_value_heads: usize = 12,
22    hidden_act: String = "gelu".to_string(),
23    max_position_embeddings: usize = 1024,
24    initializer_range: f32 = 0.02,
25    layer_norm_eps: f32 = 1e-5,
26    use_cache: bool = true,
27    pad_token_id: i64 = 50256,
28    bos_token_id: i64 = 50256,
29    eos_token_id: i64 = 50256,
30    tie_word_embeddings: bool = true,
31    attention_dropout: f32 = 0.1,
32    residual_dropout: f32 = 0.1,
33});
34
35impl GPT2Config {
36    pub fn from_gguf_config(gguf: &crate::weight_loader_core::GGUFModelConfig) -> Self {
37        Self {
38            vocab_size: gguf.vocab_size,
39            hidden_size: gguf.hidden_size,
40            intermediate_size: gguf.intermediate_size,
41            num_hidden_layers: gguf.num_hidden_layers,
42            num_attention_heads: gguf.num_attention_heads,
43            num_key_value_heads: gguf.num_key_value_heads,
44            max_position_embeddings: gguf.max_position_embeddings,
45            ..Default::default()
46        }
47    }
48}
49
50/// Main GPT-2 model
51pub struct GPT2ModelV2 {
52    config: GPT2Config,
53    device: Device,
54
55    wte: Tensor,  // Token embeddings
56    wpe: Tensor,  // Position embeddings
57    layers: Vec<GPT2Layer>,
58    ln_f: Tensor, // Final layer norm
59    lm_head: Tensor,
60}
61
62/// GPT-2 transformer layer
63pub struct GPT2Layer {
64    attn: GPT2Attention,
65    mlp: GPT2MLP,
66    ln_1: Tensor,
67    ln_2: Tensor,
68}
69
70/// GPT-2 attention
71pub struct GPT2Attention {
72    c_attn: Tensor,  // Combined QKV projection
73    c_proj: Tensor,  // Output projection
74    num_heads: usize,
75    head_dim: usize,
76    scale: f32,
77}
78
79/// GPT-2 MLP
80pub struct GPT2MLP {
81    c_fc: Tensor,
82    c_proj: Tensor,
83}
84
85impl Model for GPT2ModelV2 {
86    type Config = GPT2Config;
87
88    fn new(config: GPT2Config) -> Result<Self> {
89        let device = Device::CPU;
90
91        let wte = ops_fn::zeros(&[config.vocab_size, config.hidden_size], DataType::Float32, &device)?;
92        let wpe = ops_fn::zeros(&[config.max_position_embeddings, config.hidden_size], DataType::Float32, &device)?;
93        let ln_f = ops_fn::zeros(&[config.hidden_size], DataType::Float32, &device)?;
94
95        let lm_head = if config.tie_word_embeddings {
96            wte.clone()
97        } else {
98            ops_fn::zeros(&[config.hidden_size, config.vocab_size], DataType::Float32, &device)?
99        };
100
101        let mut layers = Vec::with_capacity(config.num_hidden_layers);
102        for _ in 0..config.num_hidden_layers {
103            layers.push(GPT2Layer::new(&config, &device)?);
104        }
105
106        Ok(Self {
107            config,
108            device,
109            wte,
110            wpe,
111            layers,
112            ln_f,
113            lm_head,
114        })
115    }
116
117    fn from_weights(config: GPT2Config, weights: ModelWeights) -> Result<Self> {
118        let mut model = Self::new(config)?;
119
120        if let Some(wte) = weights.get("wte.weight") {
121            model.wte = wte.clone();
122        }
123        if let Some(wpe) = weights.get("wpe.weight") {
124            model.wpe = wpe.clone();
125        }
126        if let Some(ln_f) = weights.get("ln_f.weight") {
127            model.ln_f = ln_f.clone();
128        }
129
130        // lm_head is tied to wte in GPT-2
131        if model.config.tie_word_embeddings {
132            model.lm_head = model.wte.clone();
133        } else if let Some(lm_head) = weights.get("lm_head.weight") {
134            model.lm_head = ops_fn::transpose(lm_head)?;
135        }
136
137        for (i, layer) in model.layers.iter_mut().enumerate() {
138            layer.load_weights(&weights, i)?;
139        }
140
141        Ok(model)
142    }
143
144    fn forward(&self, inputs: &ModelInputs) -> Result<ModelOutputs> {
145        match inputs {
146            ModelInputs::Text { input_ids, .. } => {
147                let shape = input_ids.shape();
148                let seq_len = shape[1];
149
150                // Token embeddings
151                let mut hidden_states = ops_fn::embedding(input_ids, &self.wte)?;
152
153                // Position embeddings
154                let position_ids: Vec<i64> = (0..seq_len as i64).collect();
155                let position_tensor = Tensor::from_i64_slice(&position_ids, &[1, seq_len], &self.device)?;
156                let position_embeds = ops_fn::embedding(&position_tensor, &self.wpe)?;
157                hidden_states = ops_fn::add(&hidden_states, &position_embeds)?;
158
159                // Transformer layers
160                for layer in &self.layers {
161                    hidden_states = layer.forward(&hidden_states)?;
162                }
163
164                // Final layer norm
165                hidden_states = ops_fn::layer_norm(&hidden_states, &self.ln_f, None, self.config.layer_norm_eps)?;
166
167                // LM head (tied weights)
168                let logits = if self.config.tie_word_embeddings {
169                    // For tied weights, we need to do hidden @ wte.T
170                    // Flatten to 2D for matmul, then reshape back
171                    let wte_candle = self.wte.to_candle()?;
172                    let hidden_candle = hidden_states.to_candle()?.contiguous()?;
173                    let batch = hidden_candle.dims()[0];
174                    let seq = hidden_candle.dims()[1];
175                    let hidden_size = hidden_candle.dims()[2];
176                    let flat = hidden_candle.reshape(&[batch * seq, hidden_size])?;
177                    let logits_flat = flat.matmul(&wte_candle.t()?)?;
178                    let logits_candle = logits_flat.reshape(&[batch, seq, self.config.vocab_size])?;
179                    Tensor::from_candle(logits_candle)
180                } else {
181                    ops_fn::matmul(&hidden_states, &self.lm_head)?
182                };
183
184                Ok(ModelOutputs::Logits {
185                    logits,
186                    hidden_states: None,
187                })
188            }
189            _ => Err(anyhow::anyhow!("GPT-2 model only supports text inputs")),
190        }
191    }
192
193    fn generate(&self, prompt: &str, config: &GenerationConfig) -> Result<String> {
194        use crate::tokenizer::Tokenizer;
195        use rand::Rng;
196
197        let tokenizer = Tokenizer::new();
198        let mut tokens: Vec<u32> = tokenizer.encode(prompt);
199
200        for _ in 0..config.max_new_tokens {
201            // Truncate to max position embeddings
202            let start_idx = if tokens.len() > self.config.max_position_embeddings {
203                tokens.len() - self.config.max_position_embeddings
204            } else {
205                0
206            };
207            let context = &tokens[start_idx..];
208
209            let tokens_i64: Vec<i64> = context.iter().map(|&t| t as i64).collect();
210            let input_tensor = Tensor::from_i64_slice(&tokens_i64, &[1, context.len()], &self.device)?;
211
212            let inputs = ModelInputs::Text {
213                input_ids: input_tensor,
214                attention_mask: None,
215                position_ids: None,
216            };
217
218            let outputs = self.forward(&inputs)?;
219
220            let logits = match outputs {
221                ModelOutputs::Logits { logits, .. } => logits,
222                _ => return Err(anyhow::anyhow!("Expected logits output")),
223            };
224
225            let logits_candle = logits.to_candle()?;
226            let shape = logits_candle.dims();
227
228            let last_logits = if shape.len() == 3 {
229                let seq_len = shape[1];
230                logits_candle.narrow(1, seq_len - 1, 1)?.squeeze(1)?.squeeze(0)?
231            } else {
232                let seq_len = shape[0];
233                logits_candle.narrow(0, seq_len - 1, 1)?.squeeze(0)?
234            };
235
236            let logits_vec: Vec<f32> = last_logits.to_vec1()?;
237
238            let next_token = if config.do_sample && config.temperature > 0.0 {
239                let scaled: Vec<f32> = logits_vec.iter().map(|&x| x / config.temperature).collect();
240                let max_val = scaled.iter().cloned().fold(f32::NEG_INFINITY, f32::max);
241                let exp_sum: f32 = scaled.iter().map(|&x| (x - max_val).exp()).sum();
242                let probs: Vec<f32> = scaled.iter().map(|&x| (x - max_val).exp() / exp_sum).collect();
243
244                let mut rng = rand::thread_rng();
245                let random_val: f32 = rng.gen();
246                let mut cumulative = 0.0;
247                let mut sampled = 0u32;
248
249                for (idx, &prob) in probs.iter().enumerate() {
250                    cumulative += prob;
251                    if random_val <= cumulative {
252                        sampled = idx as u32;
253                        break;
254                    }
255                }
256                sampled
257            } else {
258                logits_vec.iter().enumerate().max_by(|a, b| a.1.partial_cmp(b.1).unwrap()).map(|(i, _)| i as u32).unwrap_or(0)
259            };
260
261            if next_token == config.eos_token_id {
262                break;
263            }
264
265            tokens.push(next_token);
266        }
267
268        Ok(tokenizer.decode(&tokens))
269    }
270
271    fn config(&self) -> &Self::Config {
272        &self.config
273    }
274
275    fn memory_requirements(&self) -> MemoryRequirements {
276        let param_size = self.config.vocab_size * self.config.hidden_size +
277                        self.config.max_position_embeddings * self.config.hidden_size +
278                        self.config.num_hidden_layers * (
279                            3 * self.config.hidden_size * self.config.hidden_size +
280                            2 * self.config.hidden_size * self.config.intermediate_size
281                        );
282
283        let param_bytes = param_size * 4;
284        let kv_cache_bytes = 2 * self.config.num_hidden_layers *
285                           self.config.max_position_embeddings *
286                           self.config.hidden_size * 4;
287
288        MemoryRequirements {
289            gpu_memory: param_bytes,
290            cpu_memory: param_bytes / 4,
291            kv_cache_memory: kv_cache_bytes,
292            peak_memory: param_bytes + kv_cache_bytes,
293        }
294    }
295
296    fn to_device(&mut self, device: &Device) -> Result<()> {
297        self.wte = self.wte.to_device(device)?;
298        self.wpe = self.wpe.to_device(device)?;
299        self.ln_f = self.ln_f.to_device(device)?;
300        self.lm_head = self.lm_head.to_device(device)?;
301
302        for layer in &mut self.layers {
303            layer.to_device(device)?;
304        }
305
306        self.device = device.clone();
307        Ok(())
308    }
309}
310
311impl GPT2Layer {
312    fn new(config: &GPT2Config, device: &Device) -> Result<Self> {
313        let attn = GPT2Attention::new(config, device)?;
314        let mlp = GPT2MLP::new(config, device)?;
315        let ln_1 = ops_fn::zeros(&[config.hidden_size], DataType::Float32, device)?;
316        let ln_2 = ops_fn::zeros(&[config.hidden_size], DataType::Float32, device)?;
317
318        Ok(Self { attn, mlp, ln_1, ln_2 })
319    }
320
321    fn forward(&self, hidden_states: &Tensor) -> Result<Tensor> {
322        // Pre-norm attention
323        let normed = ops_fn::layer_norm(hidden_states, &self.ln_1, None, 1e-5)?;
324        let attn_output = self.attn.forward(&normed)?;
325        let hidden_states = ops_fn::add(hidden_states, &attn_output)?;
326
327        // Pre-norm MLP
328        let normed = ops_fn::layer_norm(&hidden_states, &self.ln_2, None, 1e-5)?;
329        let mlp_output = self.mlp.forward(&normed)?;
330        let output = ops_fn::add(&hidden_states, &mlp_output)?;
331
332        Ok(output)
333    }
334
335    fn load_weights(&mut self, weights: &ModelWeights, layer_idx: usize) -> Result<()> {
336        let prefix = format!("h.{}", layer_idx);
337
338        if let Some(c_attn) = weights.get(&format!("{}.attn.c_attn.weight", prefix)) {
339            self.attn.c_attn = ops_fn::transpose(c_attn)?;
340        }
341        if let Some(c_proj) = weights.get(&format!("{}.attn.c_proj.weight", prefix)) {
342            self.attn.c_proj = ops_fn::transpose(c_proj)?;
343        }
344        if let Some(c_fc) = weights.get(&format!("{}.mlp.c_fc.weight", prefix)) {
345            self.mlp.c_fc = ops_fn::transpose(c_fc)?;
346        }
347        if let Some(c_proj) = weights.get(&format!("{}.mlp.c_proj.weight", prefix)) {
348            self.mlp.c_proj = ops_fn::transpose(c_proj)?;
349        }
350        if let Some(ln_1) = weights.get(&format!("{}.ln_1.weight", prefix)) {
351            self.ln_1 = ln_1.clone();
352        }
353        if let Some(ln_2) = weights.get(&format!("{}.ln_2.weight", prefix)) {
354            self.ln_2 = ln_2.clone();
355        }
356
357        Ok(())
358    }
359
360    fn to_device(&mut self, device: &Device) -> Result<()> {
361        self.attn.to_device(device)?;
362        self.mlp.to_device(device)?;
363        self.ln_1 = self.ln_1.to_device(device)?;
364        self.ln_2 = self.ln_2.to_device(device)?;
365        Ok(())
366    }
367}
368
369impl GPT2Attention {
370    fn new(config: &GPT2Config, device: &Device) -> Result<Self> {
371        let num_heads = config.num_attention_heads;
372        let head_dim = config.hidden_size / num_heads;
373        let scale = 1.0 / (head_dim as f32).sqrt();
374
375        // GPT-2 uses combined QKV projection
376        let c_attn = ops_fn::zeros(&[config.hidden_size, 3 * config.hidden_size], DataType::Float32, device)?;
377        let c_proj = ops_fn::zeros(&[config.hidden_size, config.hidden_size], DataType::Float32, device)?;
378
379        Ok(Self {
380            c_attn,
381            c_proj,
382            num_heads,
383            head_dim,
384            scale,
385        })
386    }
387
388    fn forward(&self, hidden_states: &Tensor) -> Result<Tensor> {
389        let shape = hidden_states.shape();
390        let (batch_size, seq_len, hidden_size) = (shape[0], shape[1], shape[2]);
391
392        // Combined QKV projection
393        let qkv = ops_fn::matmul(hidden_states, &self.c_attn)?;
394        let qkv_candle = qkv.to_candle()?;
395
396        // Split into Q, K, V
397        let q = qkv_candle.narrow(2, 0, hidden_size)?;
398        let k = qkv_candle.narrow(2, hidden_size, hidden_size)?;
399        let v = qkv_candle.narrow(2, 2 * hidden_size, hidden_size)?;
400
401        // Reshape for multi-head attention
402        let q = q.reshape(&[batch_size, seq_len, self.num_heads, self.head_dim])?.transpose(1, 2)?;
403        let k = k.reshape(&[batch_size, seq_len, self.num_heads, self.head_dim])?.transpose(1, 2)?;
404        let v = v.reshape(&[batch_size, seq_len, self.num_heads, self.head_dim])?.transpose(1, 2)?;
405
406        // Attention scores
407        let k_t = k.transpose(2, 3)?.contiguous()?;
408        let q = q.contiguous()?;
409        let scores = q.matmul(&k_t)?;
410        let scaled_scores = (scores * (self.scale as f64))?;
411
412        // Causal mask
413        let device = scaled_scores.device();
414        let causal_mask = {
415            let mut mask_data = vec![0.0f32; seq_len * seq_len];
416            for i in 0..seq_len {
417                for j in 0..seq_len {
418                    if j > i {
419                        mask_data[i * seq_len + j] = f32::NEG_INFINITY;
420                    }
421                }
422            }
423            candle_core::Tensor::from_vec(mask_data, &[1, 1, seq_len, seq_len], device)?
424        };
425
426        let masked_scores = scaled_scores.broadcast_add(&causal_mask)?;
427        let attention_weights = candle_nn::ops::softmax_last_dim(&masked_scores)?;
428        let v = v.contiguous()?;
429        let attn_output = attention_weights.matmul(&v)?;
430
431        // Reshape back
432        let attn_output = attn_output.transpose(1, 2)?.reshape(&[batch_size, seq_len, hidden_size])?;
433        let attn_output = Tensor::from_candle(attn_output);
434
435        // Output projection
436        ops_fn::matmul(&attn_output, &self.c_proj)
437    }
438
439    fn to_device(&mut self, device: &Device) -> Result<()> {
440        self.c_attn = self.c_attn.to_device(device)?;
441        self.c_proj = self.c_proj.to_device(device)?;
442        Ok(())
443    }
444}
445
446impl GPT2MLP {
447    fn new(config: &GPT2Config, device: &Device) -> Result<Self> {
448        let c_fc = ops_fn::zeros(&[config.hidden_size, config.intermediate_size], DataType::Float32, device)?;
449        let c_proj = ops_fn::zeros(&[config.intermediate_size, config.hidden_size], DataType::Float32, device)?;
450
451        Ok(Self { c_fc, c_proj })
452    }
453
454    fn forward(&self, hidden_states: &Tensor) -> Result<Tensor> {
455        let fc_output = ops_fn::matmul(hidden_states, &self.c_fc)?;
456        let activated = ops_fn::gelu(&fc_output)?;
457        ops_fn::matmul(&activated, &self.c_proj)
458    }
459
460    fn to_device(&mut self, device: &Device) -> Result<()> {
461        self.c_fc = self.c_fc.to_device(device)?;
462        self.c_proj = self.c_proj.to_device(device)?;
463        Ok(())
464    }
465}
466
467#[cfg(test)]
468mod tests {
469    use super::*;
470
471    #[test]
472    fn test_gpt2_model_creation() {
473        let config = GPT2Config {
474            vocab_size: 1000,
475            hidden_size: 128,
476            intermediate_size: 512,
477            num_hidden_layers: 2,
478            num_attention_heads: 4,
479            num_key_value_heads: 4,
480            max_position_embeddings: 256,
481            ..Default::default()
482        };
483
484        let model = GPT2ModelV2::new(config).unwrap();
485        assert_eq!(model.config().vocab_size(), 1000);
486    }
487
488    #[test]
489    fn test_gpt2_forward_pass() {
490        let config = GPT2Config {
491            vocab_size: 100,
492            hidden_size: 64,
493            intermediate_size: 256,
494            num_hidden_layers: 1,
495            num_attention_heads: 4,
496            num_key_value_heads: 4,
497            max_position_embeddings: 32,
498            ..Default::default()
499        };
500
501        let model = GPT2ModelV2::new(config).unwrap();
502        let input_ids = ops_fn::zeros(&[2, 8], DataType::Int64, &Device::CPU).unwrap();
503        let inputs = ModelInputs::text(input_ids);
504
505        let outputs = model.forward(&inputs).unwrap();
506        match outputs {
507            ModelOutputs::Logits { logits, .. } => {
508                assert_eq!(logits.shape(), &[2, 8, 100]);
509            }
510            _ => panic!("Expected logits output"),
511        }
512    }
513}