Skip to main content

runtime/models_v2/
gptj.rs

1//! GPT-J Model V2 - Clean implementation using solid abstractions
2//!
3//! This implements the GPT-J architecture which features:
4//! - Rotary Position Embeddings (RoPE)
5//! - Parallel attention + MLP computation
6//! - No bias in attention projections
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-J model configuration
15model_config!(GPTJConfig {
16    vocab_size: usize = 50400,
17    hidden_size: usize = 4096,
18    intermediate_size: usize = 16384,
19    num_hidden_layers: usize = 28,
20    num_attention_heads: usize = 16,
21    num_key_value_heads: usize = 16,
22    hidden_act: String = "gelu".to_string(),
23    max_position_embeddings: usize = 2048,
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    rope_theta: f32 = 10000.0,
32    rotary_dim: usize = 64,
33});
34
35impl GPTJConfig {
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            rope_theta: gguf.rope_theta,
45            max_position_embeddings: gguf.max_position_embeddings,
46            ..Default::default()
47        }
48    }
49}
50
51/// Main GPT-J model
52pub struct GPTJModelV2 {
53    config: GPTJConfig,
54    device: Device,
55    wte: Tensor,
56    layers: Vec<GPTJLayer>,
57    ln_f: Tensor,
58    lm_head: Tensor,
59}
60
61pub struct GPTJLayer {
62    attn: GPTJAttention,
63    mlp: GPTJMLP,
64    ln_1: Tensor,
65}
66
67pub struct GPTJAttention {
68    q_proj: Tensor,
69    k_proj: Tensor,
70    v_proj: Tensor,
71    o_proj: Tensor,
72    num_heads: usize,
73    head_dim: usize,
74    rotary_dim: usize,
75    scale: f32,
76}
77
78pub struct GPTJMLP {
79    fc_in: Tensor,
80    fc_out: Tensor,
81}
82
83impl Model for GPTJModelV2 {
84    type Config = GPTJConfig;
85
86    fn new(config: GPTJConfig) -> Result<Self> {
87        let device = Device::CPU;
88
89        let wte = ops_fn::zeros(&[config.vocab_size, config.hidden_size], DataType::Float32, &device)?;
90        let ln_f = ops_fn::zeros(&[config.hidden_size], DataType::Float32, &device)?;
91        let lm_head = if config.tie_word_embeddings {
92            wte.clone()
93        } else {
94            ops_fn::zeros(&[config.hidden_size, config.vocab_size], DataType::Float32, &device)?
95        };
96
97        let mut layers = Vec::with_capacity(config.num_hidden_layers);
98        for _ in 0..config.num_hidden_layers {
99            layers.push(GPTJLayer::new(&config, &device)?);
100        }
101
102        Ok(Self { config, device, wte, layers, ln_f, lm_head })
103    }
104
105    fn from_weights(config: GPTJConfig, weights: ModelWeights) -> Result<Self> {
106        let mut model = Self::new(config)?;
107
108        if let Some(wte) = weights.get("transformer.wte.weight") {
109            model.wte = wte.clone();
110        }
111        if let Some(ln_f) = weights.get("transformer.ln_f.weight") {
112            model.ln_f = ln_f.clone();
113        }
114        if !model.config.tie_word_embeddings {
115            if let Some(lm_head) = weights.get("lm_head.weight") {
116                model.lm_head = ops_fn::transpose(lm_head)?;
117            }
118        }
119
120        for (i, layer) in model.layers.iter_mut().enumerate() {
121            layer.load_weights(&weights, i)?;
122        }
123
124        Ok(model)
125    }
126
127    fn forward(&self, inputs: &ModelInputs) -> Result<ModelOutputs> {
128        match inputs {
129            ModelInputs::Text { input_ids, .. } => {
130                let mut hidden_states = ops_fn::embedding(input_ids, &self.wte)?;
131
132                for layer in &self.layers {
133                    hidden_states = layer.forward(&hidden_states, self.config.rope_theta)?;
134                }
135
136                hidden_states = ops_fn::layer_norm(&hidden_states, &self.ln_f, None, self.config.layer_norm_eps)?;
137
138                let logits = if self.config.tie_word_embeddings {
139                    // Flatten to 2D for matmul, then reshape back
140                    let wte_candle = self.wte.to_candle()?;
141                    let hidden_candle = hidden_states.to_candle()?.contiguous()?;
142                    let batch = hidden_candle.dims()[0];
143                    let seq = hidden_candle.dims()[1];
144                    let hidden_size = hidden_candle.dims()[2];
145                    let flat = hidden_candle.reshape(&[batch * seq, hidden_size])?;
146                    let logits_flat = flat.matmul(&wte_candle.t()?)?;
147                    let logits_candle = logits_flat.reshape(&[batch, seq, self.config.vocab_size])?;
148                    Tensor::from_candle(logits_candle)
149                } else {
150                    ops_fn::matmul(&hidden_states, &self.lm_head)?
151                };
152
153                Ok(ModelOutputs::Logits { logits, hidden_states: None })
154            }
155            _ => Err(anyhow::anyhow!("GPT-J model only supports text inputs")),
156        }
157    }
158
159    fn generate(&self, prompt: &str, config: &GenerationConfig) -> Result<String> {
160        use crate::tokenizer::Tokenizer;
161        use rand::Rng;
162
163        let tokenizer = Tokenizer::new();
164        let mut tokens: Vec<u32> = tokenizer.encode(prompt);
165
166        for _ in 0..config.max_new_tokens {
167            let tokens_i64: Vec<i64> = tokens.iter().map(|&t| t as i64).collect();
168            let input_tensor = Tensor::from_i64_slice(&tokens_i64, &[1, tokens.len()], &self.device)?;
169            let inputs = ModelInputs::text(input_tensor);
170            let outputs = self.forward(&inputs)?;
171
172            let logits = match outputs {
173                ModelOutputs::Logits { logits, .. } => logits,
174                _ => return Err(anyhow::anyhow!("Expected logits")),
175            };
176
177            let logits_candle = logits.to_candle()?;
178            let shape = logits_candle.dims();
179            let last_logits = logits_candle.narrow(1, shape[1] - 1, 1)?.squeeze(1)?.squeeze(0)?;
180            let logits_vec: Vec<f32> = last_logits.to_vec1()?;
181
182            let next_token = if config.do_sample && config.temperature > 0.0 {
183                let scaled: Vec<f32> = logits_vec.iter().map(|&x| x / config.temperature).collect();
184                let max_val = scaled.iter().cloned().fold(f32::NEG_INFINITY, f32::max);
185                let exp_sum: f32 = scaled.iter().map(|&x| (x - max_val).exp()).sum();
186                let probs: Vec<f32> = scaled.iter().map(|&x| (x - max_val).exp() / exp_sum).collect();
187                let mut rng = rand::thread_rng();
188                let r: f32 = rng.gen();
189                let mut cum = 0.0;
190                let mut sampled = 0u32;
191                for (i, &p) in probs.iter().enumerate() {
192                    cum += p;
193                    if r <= cum { sampled = i as u32; break; }
194                }
195                sampled
196            } else {
197                logits_vec.iter().enumerate().max_by(|a, b| a.1.partial_cmp(b.1).unwrap()).map(|(i, _)| i as u32).unwrap_or(0)
198            };
199
200            if next_token == config.eos_token_id { break; }
201            tokens.push(next_token);
202        }
203
204        Ok(tokenizer.decode(&tokens))
205    }
206
207    fn config(&self) -> &Self::Config { &self.config }
208
209    fn memory_requirements(&self) -> MemoryRequirements {
210        let param_size = self.config.vocab_size * self.config.hidden_size +
211                        self.config.num_hidden_layers * (4 * self.config.hidden_size.pow(2) + 2 * self.config.hidden_size * self.config.intermediate_size);
212        MemoryRequirements {
213            gpu_memory: param_size * 4,
214            cpu_memory: param_size,
215            kv_cache_memory: 2 * self.config.num_hidden_layers * self.config.max_position_embeddings * self.config.hidden_size * 4,
216            peak_memory: param_size * 5,
217        }
218    }
219
220    fn to_device(&mut self, device: &Device) -> Result<()> {
221        self.wte = self.wte.to_device(device)?;
222        self.ln_f = self.ln_f.to_device(device)?;
223        self.lm_head = self.lm_head.to_device(device)?;
224        for layer in &mut self.layers { layer.to_device(device)?; }
225        self.device = device.clone();
226        Ok(())
227    }
228}
229
230impl GPTJLayer {
231    fn new(config: &GPTJConfig, device: &Device) -> Result<Self> {
232        Ok(Self {
233            attn: GPTJAttention::new(config, device)?,
234            mlp: GPTJMLP::new(config, device)?,
235            ln_1: ops_fn::zeros(&[config.hidden_size], DataType::Float32, device)?,
236        })
237    }
238
239    fn forward(&self, hidden_states: &Tensor, rope_theta: f32) -> Result<Tensor> {
240        let normed = ops_fn::layer_norm(hidden_states, &self.ln_1, None, 1e-5)?;
241
242        // Parallel attention + MLP (GPT-J specific)
243        let attn_output = self.attn.forward(&normed, rope_theta)?;
244        let mlp_output = self.mlp.forward(&normed)?;
245
246        let combined = ops_fn::add(&attn_output, &mlp_output)?;
247        ops_fn::add(hidden_states, &combined)
248    }
249
250    fn load_weights(&mut self, weights: &ModelWeights, idx: usize) -> Result<()> {
251        let p = format!("transformer.h.{}", idx);
252        if let Some(w) = weights.get(&format!("{}.attn.q_proj.weight", p)) { self.attn.q_proj = ops_fn::transpose(w)?; }
253        if let Some(w) = weights.get(&format!("{}.attn.k_proj.weight", p)) { self.attn.k_proj = ops_fn::transpose(w)?; }
254        if let Some(w) = weights.get(&format!("{}.attn.v_proj.weight", p)) { self.attn.v_proj = ops_fn::transpose(w)?; }
255        if let Some(w) = weights.get(&format!("{}.attn.out_proj.weight", p)) { self.attn.o_proj = ops_fn::transpose(w)?; }
256        if let Some(w) = weights.get(&format!("{}.mlp.fc_in.weight", p)) { self.mlp.fc_in = ops_fn::transpose(w)?; }
257        if let Some(w) = weights.get(&format!("{}.mlp.fc_out.weight", p)) { self.mlp.fc_out = ops_fn::transpose(w)?; }
258        if let Some(w) = weights.get(&format!("{}.ln_1.weight", p)) { self.ln_1 = w.clone(); }
259        Ok(())
260    }
261
262    fn to_device(&mut self, device: &Device) -> Result<()> {
263        self.attn.to_device(device)?;
264        self.mlp.to_device(device)?;
265        self.ln_1 = self.ln_1.to_device(device)?;
266        Ok(())
267    }
268}
269
270fn apply_partial_rope(
271    q: &candle_core::Tensor,
272    k: &candle_core::Tensor,
273    seq_len: usize,
274    rotary_dim: usize,
275    rope_theta: f32,
276) -> Result<(candle_core::Tensor, candle_core::Tensor)> {
277    let device = q.device();
278    let half_rotary = rotary_dim / 2;
279
280    let inv_freq: Vec<f32> = (0..half_rotary)
281        .map(|i| 1.0 / rope_theta.powf((2 * i) as f32 / rotary_dim as f32))
282        .collect();
283    let positions: Vec<f32> = (0..seq_len).map(|p| p as f32).collect();
284
285    let mut angles = Vec::with_capacity(seq_len * half_rotary);
286    for pos in &positions {
287        for freq in &inv_freq { angles.push(pos * freq); }
288    }
289
290    let angles_t = candle_core::Tensor::from_vec(angles, &[seq_len, half_rotary], device)?;
291    let cos = angles_t.cos()?.unsqueeze(0)?.unsqueeze(0)?;
292    let sin = angles_t.sin()?.unsqueeze(0)?.unsqueeze(0)?;
293
294    // Apply RoPE only to first rotary_dim dimensions
295    let q_rot = q.narrow(3, 0, rotary_dim)?;
296    let q_pass = q.narrow(3, rotary_dim, q.dims()[3] - rotary_dim)?;
297    let k_rot = k.narrow(3, 0, rotary_dim)?;
298    let k_pass = k.narrow(3, rotary_dim, k.dims()[3] - rotary_dim)?;
299
300    let q1 = q_rot.narrow(3, 0, half_rotary)?;
301    let q2 = q_rot.narrow(3, half_rotary, half_rotary)?;
302    let k1 = k_rot.narrow(3, 0, half_rotary)?;
303    let k2 = k_rot.narrow(3, half_rotary, half_rotary)?;
304
305    let q_r1 = (q1.broadcast_mul(&cos)? - q2.broadcast_mul(&sin)?)?;
306    let q_r2 = (q1.broadcast_mul(&sin)? + q2.broadcast_mul(&cos)?)?;
307    let k_r1 = (k1.broadcast_mul(&cos)? - k2.broadcast_mul(&sin)?)?;
308    let k_r2 = (k1.broadcast_mul(&sin)? + k2.broadcast_mul(&cos)?)?;
309
310    let q_rotated = candle_core::Tensor::cat(&[&q_r1, &q_r2], 3)?;
311    let k_rotated = candle_core::Tensor::cat(&[&k_r1, &k_r2], 3)?;
312
313    let q_out = candle_core::Tensor::cat(&[&q_rotated, &q_pass], 3)?;
314    let k_out = candle_core::Tensor::cat(&[&k_rotated, &k_pass], 3)?;
315
316    Ok((q_out, k_out))
317}
318
319impl GPTJAttention {
320    fn new(config: &GPTJConfig, device: &Device) -> Result<Self> {
321        let head_dim = config.hidden_size / config.num_attention_heads;
322        Ok(Self {
323            q_proj: ops_fn::zeros(&[config.hidden_size, config.hidden_size], DataType::Float32, device)?,
324            k_proj: ops_fn::zeros(&[config.hidden_size, config.hidden_size], DataType::Float32, device)?,
325            v_proj: ops_fn::zeros(&[config.hidden_size, config.hidden_size], DataType::Float32, device)?,
326            o_proj: ops_fn::zeros(&[config.hidden_size, config.hidden_size], DataType::Float32, device)?,
327            num_heads: config.num_attention_heads,
328            head_dim,
329            rotary_dim: config.rotary_dim,
330            scale: 1.0 / (head_dim as f32).sqrt(),
331        })
332    }
333
334    fn forward(&self, hidden_states: &Tensor, rope_theta: f32) -> Result<Tensor> {
335        let shape = hidden_states.shape();
336        let (batch, seq_len, _) = (shape[0], shape[1], shape[2]);
337
338        let q = ops_fn::matmul(hidden_states, &self.q_proj)?.to_candle()?.reshape(&[batch, seq_len, self.num_heads, self.head_dim])?.transpose(1, 2)?;
339        let k = ops_fn::matmul(hidden_states, &self.k_proj)?.to_candle()?.reshape(&[batch, seq_len, self.num_heads, self.head_dim])?.transpose(1, 2)?;
340        let v = ops_fn::matmul(hidden_states, &self.v_proj)?.to_candle()?.reshape(&[batch, seq_len, self.num_heads, self.head_dim])?.transpose(1, 2)?;
341
342        let (q, k) = apply_partial_rope(&q, &k, seq_len, self.rotary_dim, rope_theta)?;
343
344        let q = q.contiguous()?;
345        let k_t = k.transpose(2, 3)?.contiguous()?;
346        let scores = (q.matmul(&k_t)? * (self.scale as f64))?;
347        let device = scores.device();
348        let mask = {
349            let mut m = vec![0.0f32; seq_len * seq_len];
350            for i in 0..seq_len { for j in (i+1)..seq_len { m[i*seq_len+j] = f32::NEG_INFINITY; } }
351            candle_core::Tensor::from_vec(m, &[1, 1, seq_len, seq_len], device)?
352        };
353        let masked = scores.broadcast_add(&mask)?;
354        let v = v.contiguous()?;
355        let attn = candle_nn::ops::softmax_last_dim(&masked)?.matmul(&v)?;
356        let out = attn.transpose(1, 2)?.reshape(&[batch, seq_len, self.num_heads * self.head_dim])?;
357        ops_fn::matmul(&Tensor::from_candle(out), &self.o_proj)
358    }
359
360    fn to_device(&mut self, device: &Device) -> Result<()> {
361        self.q_proj = self.q_proj.to_device(device)?;
362        self.k_proj = self.k_proj.to_device(device)?;
363        self.v_proj = self.v_proj.to_device(device)?;
364        self.o_proj = self.o_proj.to_device(device)?;
365        Ok(())
366    }
367}
368
369impl GPTJMLP {
370    fn new(config: &GPTJConfig, device: &Device) -> Result<Self> {
371        Ok(Self {
372            fc_in: ops_fn::zeros(&[config.hidden_size, config.intermediate_size], DataType::Float32, device)?,
373            fc_out: ops_fn::zeros(&[config.intermediate_size, config.hidden_size], DataType::Float32, device)?,
374        })
375    }
376
377    fn forward(&self, hidden_states: &Tensor) -> Result<Tensor> {
378        let h = ops_fn::matmul(hidden_states, &self.fc_in)?;
379        let h = ops_fn::gelu(&h)?;
380        ops_fn::matmul(&h, &self.fc_out)
381    }
382
383    fn to_device(&mut self, device: &Device) -> Result<()> {
384        self.fc_in = self.fc_in.to_device(device)?;
385        self.fc_out = self.fc_out.to_device(device)?;
386        Ok(())
387    }
388}
389
390#[cfg(test)]
391mod tests {
392    use super::*;
393
394    #[test]
395    fn test_gptj_creation() {
396        let config = GPTJConfig {
397            vocab_size: 1000, hidden_size: 128, intermediate_size: 512,
398            num_hidden_layers: 2, num_attention_heads: 4, num_key_value_heads: 4,
399            rotary_dim: 32, ..Default::default()
400        };
401        let model = GPTJModelV2::new(config).unwrap();
402        assert_eq!(model.config().vocab_size(), 1000);
403    }
404
405    #[test]
406    fn test_gptj_forward() {
407        let config = GPTJConfig {
408            vocab_size: 100, hidden_size: 64, intermediate_size: 256,
409            num_hidden_layers: 1, num_attention_heads: 4, num_key_value_heads: 4,
410            rotary_dim: 16, ..Default::default()
411        };
412        let model = GPTJModelV2::new(config).unwrap();
413        let inputs = ModelInputs::text(ops_fn::zeros(&[2, 8], DataType::Int64, &Device::CPU).unwrap());
414        let outputs = model.forward(&inputs).unwrap();
415        match outputs {
416            ModelOutputs::Logits { logits, .. } => assert_eq!(logits.shape(), &[2, 8, 100]),
417            _ => panic!("Expected logits"),
418        }
419    }
420}