Skip to main content

runtime/models_v2/
gptneox.rs

1//! GPT-NeoX Model V2 - Clean implementation
2//!
3//! GPT-NeoX architecture features:
4//! - Rotary Position Embeddings (RoPE)
5//! - Parallel attention + MLP
6//! - Pre-norm layer normalization
7
8use crate::model_config;
9use super::traits::*;
10use anyhow::Result;
11use serde::{Serialize, Deserialize};
12
13model_config!(GPTNeoXConfig {
14    vocab_size: usize = 50432,
15    hidden_size: usize = 6144,
16    intermediate_size: usize = 24576,
17    num_hidden_layers: usize = 44,
18    num_attention_heads: usize = 64,
19    num_key_value_heads: usize = 64,
20    hidden_act: String = "gelu".to_string(),
21    max_position_embeddings: usize = 2048,
22    initializer_range: f32 = 0.02,
23    layer_norm_eps: f32 = 1e-5,
24    use_cache: bool = true,
25    pad_token_id: i64 = 0,
26    bos_token_id: i64 = 0,
27    eos_token_id: i64 = 0,
28    tie_word_embeddings: bool = false,
29    rope_theta: f32 = 10000.0,
30    rotary_pct: f32 = 0.25,
31    use_parallel_residual: bool = true,
32});
33
34impl GPTNeoXConfig {
35    pub fn from_gguf_config(gguf: &crate::weight_loader_core::GGUFModelConfig) -> Self {
36        Self {
37            vocab_size: gguf.vocab_size,
38            hidden_size: gguf.hidden_size,
39            intermediate_size: gguf.intermediate_size,
40            num_hidden_layers: gguf.num_hidden_layers,
41            num_attention_heads: gguf.num_attention_heads,
42            num_key_value_heads: gguf.num_key_value_heads,
43            rope_theta: gguf.rope_theta,
44            max_position_embeddings: gguf.max_position_embeddings,
45            ..Default::default()
46        }
47    }
48}
49
50pub struct GPTNeoXModelV2 {
51    config: GPTNeoXConfig,
52    device: Device,
53    embed_in: Tensor,
54    layers: Vec<GPTNeoXLayer>,
55    final_layer_norm: Tensor,
56    embed_out: Tensor,
57}
58
59pub struct GPTNeoXLayer {
60    attention: GPTNeoXAttention,
61    mlp: GPTNeoXMLP,
62    input_layernorm: Tensor,
63    post_attention_layernorm: Tensor,
64    use_parallel_residual: bool,
65}
66
67pub struct GPTNeoXAttention {
68    query_key_value: Tensor,
69    dense: Tensor,
70    num_heads: usize,
71    head_dim: usize,
72    rotary_dim: usize,
73    scale: f32,
74}
75
76pub struct GPTNeoXMLP {
77    dense_h_to_4h: Tensor,
78    dense_4h_to_h: Tensor,
79}
80
81impl Model for GPTNeoXModelV2 {
82    type Config = GPTNeoXConfig;
83
84    fn new(config: GPTNeoXConfig) -> Result<Self> {
85        let device = Device::CPU;
86        let embed_in = ops_fn::zeros(&[config.vocab_size, config.hidden_size], DataType::Float32, &device)?;
87        let final_layer_norm = ops_fn::zeros(&[config.hidden_size], DataType::Float32, &device)?;
88        let embed_out = ops_fn::zeros(&[config.hidden_size, config.vocab_size], DataType::Float32, &device)?;
89
90        let mut layers = Vec::with_capacity(config.num_hidden_layers);
91        for _ in 0..config.num_hidden_layers {
92            layers.push(GPTNeoXLayer::new(&config, &device)?);
93        }
94
95        Ok(Self { config, device, embed_in, layers, final_layer_norm, embed_out })
96    }
97
98    fn from_weights(config: GPTNeoXConfig, weights: ModelWeights) -> Result<Self> {
99        let mut model = Self::new(config)?;
100        if let Some(w) = weights.get("gpt_neox.embed_in.weight") { model.embed_in = w.clone(); }
101        if let Some(w) = weights.get("gpt_neox.final_layer_norm.weight") { model.final_layer_norm = w.clone(); }
102        if let Some(w) = weights.get("embed_out.weight") { model.embed_out = ops_fn::transpose(w)?; }
103        for (i, layer) in model.layers.iter_mut().enumerate() { layer.load_weights(&weights, i)?; }
104        Ok(model)
105    }
106
107    fn forward(&self, inputs: &ModelInputs) -> Result<ModelOutputs> {
108        match inputs {
109            ModelInputs::Text { input_ids, .. } => {
110                let mut hidden_states = ops_fn::embedding(input_ids, &self.embed_in)?;
111                for layer in &self.layers {
112                    hidden_states = layer.forward(&hidden_states, self.config.rope_theta)?;
113                }
114                hidden_states = ops_fn::layer_norm(&hidden_states, &self.final_layer_norm, None, self.config.layer_norm_eps)?;
115                let logits = ops_fn::matmul(&hidden_states, &self.embed_out)?;
116                Ok(ModelOutputs::Logits { logits, hidden_states: None })
117            }
118            _ => Err(anyhow::anyhow!("GPT-NeoX only supports text inputs")),
119        }
120    }
121
122    fn generate(&self, prompt: &str, config: &GenerationConfig) -> Result<String> {
123        use crate::tokenizer::Tokenizer;
124        use rand::Rng;
125
126        let tokenizer = Tokenizer::new();
127        let mut tokens: Vec<u32> = tokenizer.encode(prompt);
128
129        for _ in 0..config.max_new_tokens {
130            let tokens_i64: Vec<i64> = tokens.iter().map(|&t| t as i64).collect();
131            let input = Tensor::from_i64_slice(&tokens_i64, &[1, tokens.len()], &self.device)?;
132            let outputs = self.forward(&ModelInputs::text(input))?;
133            let logits = match outputs { ModelOutputs::Logits { logits, .. } => logits, _ => return Err(anyhow::anyhow!("Expected logits")) };
134            let logits_candle = logits.to_candle()?;
135            let last = logits_candle.narrow(1, logits_candle.dims()[1] - 1, 1)?.squeeze(1)?.squeeze(0)?;
136            let logits_vec: Vec<f32> = last.to_vec1()?;
137
138            let next = if config.do_sample && config.temperature > 0.0 {
139                let scaled: Vec<f32> = logits_vec.iter().map(|&x| x / config.temperature).collect();
140                let max_v = scaled.iter().cloned().fold(f32::NEG_INFINITY, f32::max);
141                let exp_sum: f32 = scaled.iter().map(|&x| (x - max_v).exp()).sum();
142                let probs: Vec<f32> = scaled.iter().map(|&x| (x - max_v).exp() / exp_sum).collect();
143                let mut rng = rand::thread_rng();
144                let r: f32 = rng.gen();
145                let mut cum = 0.0;
146                let mut s = 0u32;
147                for (i, &p) in probs.iter().enumerate() { cum += p; if r <= cum { s = i as u32; break; } }
148                s
149            } else {
150                logits_vec.iter().enumerate().max_by(|a, b| a.1.partial_cmp(b.1).unwrap()).map(|(i, _)| i as u32).unwrap_or(0)
151            };
152            if next == config.eos_token_id { break; }
153            tokens.push(next);
154        }
155        Ok(tokenizer.decode(&tokens))
156    }
157
158    fn config(&self) -> &Self::Config { &self.config }
159    fn memory_requirements(&self) -> MemoryRequirements {
160        let p = self.config.vocab_size * self.config.hidden_size + self.config.num_hidden_layers * 8 * self.config.hidden_size.pow(2);
161        MemoryRequirements { gpu_memory: p * 4, cpu_memory: p, kv_cache_memory: 2 * self.config.num_hidden_layers * self.config.max_position_embeddings * self.config.hidden_size * 4, peak_memory: p * 5 }
162    }
163    fn to_device(&mut self, device: &Device) -> Result<()> {
164        self.embed_in = self.embed_in.to_device(device)?;
165        self.final_layer_norm = self.final_layer_norm.to_device(device)?;
166        self.embed_out = self.embed_out.to_device(device)?;
167        for l in &mut self.layers { l.to_device(device)?; }
168        self.device = device.clone();
169        Ok(())
170    }
171}
172
173impl GPTNeoXLayer {
174    fn new(config: &GPTNeoXConfig, device: &Device) -> Result<Self> {
175        Ok(Self {
176            attention: GPTNeoXAttention::new(config, device)?,
177            mlp: GPTNeoXMLP::new(config, device)?,
178            input_layernorm: ops_fn::zeros(&[config.hidden_size], DataType::Float32, device)?,
179            post_attention_layernorm: ops_fn::zeros(&[config.hidden_size], DataType::Float32, device)?,
180            use_parallel_residual: config.use_parallel_residual,
181        })
182    }
183
184    fn forward(&self, hidden_states: &Tensor, rope_theta: f32) -> Result<Tensor> {
185        let ln1 = ops_fn::layer_norm(hidden_states, &self.input_layernorm, None, 1e-5)?;
186        let attn_out = self.attention.forward(&ln1, rope_theta)?;
187
188        if self.use_parallel_residual {
189            let ln2 = ops_fn::layer_norm(hidden_states, &self.post_attention_layernorm, None, 1e-5)?;
190            let mlp_out = self.mlp.forward(&ln2)?;
191            let combined = ops_fn::add(&attn_out, &mlp_out)?;
192            ops_fn::add(hidden_states, &combined)
193        } else {
194            let h = ops_fn::add(hidden_states, &attn_out)?;
195            let ln2 = ops_fn::layer_norm(&h, &self.post_attention_layernorm, None, 1e-5)?;
196            let mlp_out = self.mlp.forward(&ln2)?;
197            ops_fn::add(&h, &mlp_out)
198        }
199    }
200
201    fn load_weights(&mut self, weights: &ModelWeights, idx: usize) -> Result<()> {
202        let p = format!("gpt_neox.layers.{}", idx);
203        if let Some(w) = weights.get(&format!("{}.attention.query_key_value.weight", p)) { self.attention.query_key_value = ops_fn::transpose(w)?; }
204        if let Some(w) = weights.get(&format!("{}.attention.dense.weight", p)) { self.attention.dense = ops_fn::transpose(w)?; }
205        if let Some(w) = weights.get(&format!("{}.mlp.dense_h_to_4h.weight", p)) { self.mlp.dense_h_to_4h = ops_fn::transpose(w)?; }
206        if let Some(w) = weights.get(&format!("{}.mlp.dense_4h_to_h.weight", p)) { self.mlp.dense_4h_to_h = ops_fn::transpose(w)?; }
207        if let Some(w) = weights.get(&format!("{}.input_layernorm.weight", p)) { self.input_layernorm = w.clone(); }
208        if let Some(w) = weights.get(&format!("{}.post_attention_layernorm.weight", p)) { self.post_attention_layernorm = w.clone(); }
209        Ok(())
210    }
211
212    fn to_device(&mut self, device: &Device) -> Result<()> {
213        self.attention.to_device(device)?;
214        self.mlp.to_device(device)?;
215        self.input_layernorm = self.input_layernorm.to_device(device)?;
216        self.post_attention_layernorm = self.post_attention_layernorm.to_device(device)?;
217        Ok(())
218    }
219}
220
221impl GPTNeoXAttention {
222    fn new(config: &GPTNeoXConfig, device: &Device) -> Result<Self> {
223        let head_dim = config.hidden_size / config.num_attention_heads;
224        let rotary_dim = (head_dim as f32 * config.rotary_pct) as usize;
225        Ok(Self {
226            query_key_value: ops_fn::zeros(&[config.hidden_size, 3 * config.hidden_size], DataType::Float32, device)?,
227            dense: ops_fn::zeros(&[config.hidden_size, config.hidden_size], DataType::Float32, device)?,
228            num_heads: config.num_attention_heads,
229            head_dim,
230            rotary_dim,
231            scale: 1.0 / (head_dim as f32).sqrt(),
232        })
233    }
234
235    fn forward(&self, hidden_states: &Tensor, rope_theta: f32) -> Result<Tensor> {
236        let shape = hidden_states.shape();
237        let (batch, seq_len, hidden_size) = (shape[0], shape[1], shape[2]);
238
239        let qkv = ops_fn::matmul(hidden_states, &self.query_key_value)?.to_candle()?;
240        let q = qkv.narrow(2, 0, hidden_size)?.reshape(&[batch, seq_len, self.num_heads, self.head_dim])?.transpose(1, 2)?;
241        let k = qkv.narrow(2, hidden_size, hidden_size)?.reshape(&[batch, seq_len, self.num_heads, self.head_dim])?.transpose(1, 2)?;
242        let v = qkv.narrow(2, 2 * hidden_size, hidden_size)?.reshape(&[batch, seq_len, self.num_heads, self.head_dim])?.transpose(1, 2)?;
243
244        let (q, k) = apply_partial_rope(&q, &k, seq_len, self.rotary_dim, rope_theta)?;
245
246        let q = q.contiguous()?;
247        let k_t = k.transpose(2, 3)?.contiguous()?;
248        let scores = (q.matmul(&k_t)? * (self.scale as f64))?;
249        let device = scores.device();
250        let mask = {
251            let mut m = vec![0.0f32; seq_len * seq_len];
252            for i in 0..seq_len { for j in (i+1)..seq_len { m[i*seq_len+j] = f32::NEG_INFINITY; } }
253            candle_core::Tensor::from_vec(m, &[1, 1, seq_len, seq_len], device)?
254        };
255        let v = v.contiguous()?;
256        let attn = candle_nn::ops::softmax_last_dim(&scores.broadcast_add(&mask)?)?.matmul(&v)?;
257        let out = attn.transpose(1, 2)?.reshape(&[batch, seq_len, hidden_size])?;
258        ops_fn::matmul(&Tensor::from_candle(out), &self.dense)
259    }
260
261    fn to_device(&mut self, device: &Device) -> Result<()> {
262        self.query_key_value = self.query_key_value.to_device(device)?;
263        self.dense = self.dense.to_device(device)?;
264        Ok(())
265    }
266}
267
268fn apply_partial_rope(q: &candle_core::Tensor, k: &candle_core::Tensor, seq_len: usize, rotary_dim: usize, rope_theta: f32) -> Result<(candle_core::Tensor, candle_core::Tensor)> {
269    let device = q.device();
270    let half = rotary_dim / 2;
271    let inv_freq: Vec<f32> = (0..half).map(|i| 1.0 / rope_theta.powf((2 * i) as f32 / rotary_dim as f32)).collect();
272    let positions: Vec<f32> = (0..seq_len).map(|p| p as f32).collect();
273    let mut angles = Vec::with_capacity(seq_len * half);
274    for pos in &positions { for freq in &inv_freq { angles.push(pos * freq); } }
275    let angles_t = candle_core::Tensor::from_vec(angles, &[seq_len, half], device)?;
276    let cos = angles_t.cos()?.unsqueeze(0)?.unsqueeze(0)?;
277    let sin = angles_t.sin()?.unsqueeze(0)?.unsqueeze(0)?;
278
279    let head_dim = q.dims()[3];
280    if rotary_dim >= head_dim {
281        let q1 = q.narrow(3, 0, half)?;
282        let q2 = q.narrow(3, half, half)?;
283        let k1 = k.narrow(3, 0, half)?;
284        let k2 = k.narrow(3, half, half)?;
285        let qr1 = (q1.broadcast_mul(&cos)? - q2.broadcast_mul(&sin)?)?;
286        let qr2 = (q1.broadcast_mul(&sin)? + q2.broadcast_mul(&cos)?)?;
287        let kr1 = (k1.broadcast_mul(&cos)? - k2.broadcast_mul(&sin)?)?;
288        let kr2 = (k1.broadcast_mul(&sin)? + k2.broadcast_mul(&cos)?)?;
289        Ok((candle_core::Tensor::cat(&[&qr1, &qr2], 3)?, candle_core::Tensor::cat(&[&kr1, &kr2], 3)?))
290    } else {
291        let q_rot = q.narrow(3, 0, rotary_dim)?;
292        let q_pass = q.narrow(3, rotary_dim, head_dim - rotary_dim)?;
293        let k_rot = k.narrow(3, 0, rotary_dim)?;
294        let k_pass = k.narrow(3, rotary_dim, head_dim - rotary_dim)?;
295        let q1 = q_rot.narrow(3, 0, half)?;
296        let q2 = q_rot.narrow(3, half, half)?;
297        let k1 = k_rot.narrow(3, 0, half)?;
298        let k2 = k_rot.narrow(3, half, half)?;
299        let qr = candle_core::Tensor::cat(&[&(q1.broadcast_mul(&cos)? - q2.broadcast_mul(&sin)?)?, &(q1.broadcast_mul(&sin)? + q2.broadcast_mul(&cos)?)?], 3)?;
300        let kr = candle_core::Tensor::cat(&[&(k1.broadcast_mul(&cos)? - k2.broadcast_mul(&sin)?)?, &(k1.broadcast_mul(&sin)? + k2.broadcast_mul(&cos)?)?], 3)?;
301        Ok((candle_core::Tensor::cat(&[&qr, &q_pass], 3)?, candle_core::Tensor::cat(&[&kr, &k_pass], 3)?))
302    }
303}
304
305impl GPTNeoXMLP {
306    fn new(config: &GPTNeoXConfig, device: &Device) -> Result<Self> {
307        Ok(Self {
308            dense_h_to_4h: ops_fn::zeros(&[config.hidden_size, config.intermediate_size], DataType::Float32, device)?,
309            dense_4h_to_h: ops_fn::zeros(&[config.intermediate_size, config.hidden_size], DataType::Float32, device)?,
310        })
311    }
312    fn forward(&self, x: &Tensor) -> Result<Tensor> {
313        let h = ops_fn::matmul(x, &self.dense_h_to_4h)?;
314        let h = ops_fn::gelu(&h)?;
315        ops_fn::matmul(&h, &self.dense_4h_to_h)
316    }
317    fn to_device(&mut self, device: &Device) -> Result<()> {
318        self.dense_h_to_4h = self.dense_h_to_4h.to_device(device)?;
319        self.dense_4h_to_h = self.dense_4h_to_h.to_device(device)?;
320        Ok(())
321    }
322}
323
324#[cfg(test)]
325mod tests {
326    use super::*;
327    #[test]
328    fn test_gptneox_creation() {
329        let config = GPTNeoXConfig { vocab_size: 1000, hidden_size: 128, intermediate_size: 512, num_hidden_layers: 2, num_attention_heads: 4, num_key_value_heads: 4, ..Default::default() };
330        let model = GPTNeoXModelV2::new(config).unwrap();
331        assert_eq!(model.config().vocab_size(), 1000);
332    }
333    #[test]
334    fn test_gptneox_forward() {
335        let config = GPTNeoXConfig { vocab_size: 100, hidden_size: 64, intermediate_size: 256, num_hidden_layers: 1, num_attention_heads: 4, num_key_value_heads: 4, ..Default::default() };
336        let model = GPTNeoXModelV2::new(config).unwrap();
337        let inputs = ModelInputs::text(ops_fn::zeros(&[2, 8], DataType::Int64, &Device::CPU).unwrap());
338        match model.forward(&inputs).unwrap() { ModelOutputs::Logits { logits, .. } => assert_eq!(logits.shape(), &[2, 8, 100]), _ => panic!() }
339    }
340}