1use crate::model_config;
10use super::traits::*;
11use anyhow::Result;
12use serde::{Serialize, Deserialize};
13
14model_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
51pub 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 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 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 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}