1use 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}