1use crate::model_config;
10use super::traits::*;
11use anyhow::Result;
12use serde::{Serialize, Deserialize};
13
14model_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
50pub struct GPT2ModelV2 {
52 config: GPT2Config,
53 device: Device,
54
55 wte: Tensor, wpe: Tensor, layers: Vec<GPT2Layer>,
58 ln_f: Tensor, lm_head: Tensor,
60}
61
62pub struct GPT2Layer {
64 attn: GPT2Attention,
65 mlp: GPT2MLP,
66 ln_1: Tensor,
67 ln_2: Tensor,
68}
69
70pub struct GPT2Attention {
72 c_attn: Tensor, c_proj: Tensor, num_heads: usize,
75 head_dim: usize,
76 scale: f32,
77}
78
79pub 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 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 let mut hidden_states = ops_fn::embedding(input_ids, &self.wte)?;
152
153 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 for layer in &self.layers {
161 hidden_states = layer.forward(&hidden_states)?;
162 }
163
164 hidden_states = ops_fn::layer_norm(&hidden_states, &self.ln_f, None, self.config.layer_norm_eps)?;
166
167 let logits = if self.config.tie_word_embeddings {
169 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 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 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 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 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 let qkv = ops_fn::matmul(hidden_states, &self.c_attn)?;
394 let qkv_candle = qkv.to_candle()?;
395
396 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 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 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 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 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 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}