Skip to main content

runtime/models_v2/
starcoder.rs

1//! StarCoder Model V2 - Clean implementation
2//!
3//! StarCoder architecture features:
4//! - Multi-Query Attention (MQA) - single KV head
5//! - GPT-2 style architecture with learned position embeddings
6//! - Fill-in-the-Middle (FIM) token support
7//! - Larger context windows
8
9use crate::model_config;
10use super::traits::*;
11use anyhow::Result;
12use serde::{Serialize, Deserialize};
13
14/// StarCoder model configuration
15model_config!(StarCoderConfig {
16    vocab_size: usize = 49152,
17    hidden_size: usize = 6144,
18    intermediate_size: usize = 24576,
19    num_hidden_layers: usize = 40,
20    num_attention_heads: usize = 48,
21    num_key_value_heads: usize = 1,  // MQA: single KV head
22    hidden_act: String = "gelu_new".to_string(),
23    max_position_embeddings: usize = 8192,
24    initializer_range: f32 = 0.02,
25    layer_norm_eps: f32 = 1e-5,
26    use_cache: bool = true,
27    pad_token_id: i64 = 49152,
28    bos_token_id: i64 = 49152,
29    eos_token_id: i64 = 0,
30    tie_word_embeddings: bool = true,
31});
32
33impl StarCoderConfig {
34    pub fn from_gguf_config(gguf: &crate::weight_loader_core::GGUFModelConfig) -> Self {
35        Self {
36            vocab_size: gguf.vocab_size,
37            hidden_size: gguf.hidden_size,
38            intermediate_size: gguf.intermediate_size,
39            num_hidden_layers: gguf.num_hidden_layers,
40            num_attention_heads: gguf.num_attention_heads,
41            num_key_value_heads: gguf.num_key_value_heads,
42            max_position_embeddings: gguf.max_position_embeddings,
43            ..Default::default()
44        }
45    }
46}
47
48pub struct StarCoderModelV2 {
49    config: StarCoderConfig,
50    device: Device,
51    wte: Tensor,
52    wpe: Tensor,
53    layers: Vec<StarCoderLayer>,
54    ln_f: Tensor,
55    lm_head: Tensor,
56}
57
58pub struct StarCoderLayer {
59    attn: StarCoderAttention,
60    mlp: StarCoderMLP,
61    ln_1: Tensor,
62    ln_2: Tensor,
63}
64
65pub struct StarCoderAttention {
66    c_attn: Tensor,  // Combined QKV projection
67    c_proj: Tensor,
68    num_heads: usize,
69    head_dim: usize,
70    scale: f32,
71}
72
73pub struct StarCoderMLP {
74    c_fc: Tensor,
75    c_proj: Tensor,
76}
77
78impl Model for StarCoderModelV2 {
79    type Config = StarCoderConfig;
80
81    fn new(config: StarCoderConfig) -> Result<Self> {
82        let device = Device::CPU;
83        let wte = ops_fn::zeros(&[config.vocab_size, config.hidden_size], DataType::Float32, &device)?;
84        let wpe = ops_fn::zeros(&[config.max_position_embeddings, config.hidden_size], DataType::Float32, &device)?;
85        let ln_f = ops_fn::zeros(&[config.hidden_size], DataType::Float32, &device)?;
86        let lm_head = wte.clone();
87
88        let mut layers = Vec::with_capacity(config.num_hidden_layers);
89        for _ in 0..config.num_hidden_layers {
90            layers.push(StarCoderLayer::new(&config, &device)?);
91        }
92
93        Ok(Self { config, device, wte, wpe, layers, ln_f, lm_head })
94    }
95
96    fn from_weights(config: StarCoderConfig, weights: ModelWeights) -> Result<Self> {
97        let mut model = Self::new(config)?;
98        if let Some(w) = weights.get("transformer.wte.weight") { model.wte = w.clone(); }
99        if let Some(w) = weights.get("transformer.wpe.weight") { model.wpe = w.clone(); }
100        if let Some(w) = weights.get("transformer.ln_f.weight") { model.ln_f = w.clone(); }
101        if model.config.tie_word_embeddings {
102            model.lm_head = model.wte.clone();
103        } else if let Some(w) = weights.get("lm_head.weight") {
104            model.lm_head = w.clone();
105        }
106        for (i, layer) in model.layers.iter_mut().enumerate() { layer.load_weights(&weights, i)?; }
107        Ok(model)
108    }
109
110    fn forward(&self, inputs: &ModelInputs) -> Result<ModelOutputs> {
111        match inputs {
112            ModelInputs::Text { input_ids, .. } => {
113                let shape = input_ids.shape();
114                let seq_len = shape[1];
115
116                let mut hidden = ops_fn::embedding(input_ids, &self.wte)?;
117                let positions: Vec<i64> = (0..seq_len as i64).collect();
118                let pos_tensor = Tensor::from_i64_slice(&positions, &[1, seq_len], &self.device)?;
119                let pos_embeds = ops_fn::embedding(&pos_tensor, &self.wpe)?;
120                hidden = ops_fn::add(&hidden, &pos_embeds)?;
121
122                for layer in &self.layers {
123                    hidden = layer.forward(&hidden)?;
124                }
125
126                hidden = ops_fn::layer_norm(&hidden, &self.ln_f, None, self.config.layer_norm_eps)?;
127
128                // Tied embeddings - flatten to 2D, matmul, reshape back
129                let lm_head_candle = self.lm_head.to_candle()?;
130                let hidden_candle = hidden.to_candle()?.contiguous()?;
131                let batch = hidden_candle.dims()[0];
132                let seq = hidden_candle.dims()[1];
133                let hidden_size = hidden_candle.dims()[2];
134                let flat = hidden_candle.reshape(&[batch * seq, hidden_size])?;
135                let logits_flat = flat.matmul(&lm_head_candle.t()?)?;
136                let logits_candle = logits_flat.reshape(&[batch, seq, self.config.vocab_size])?;
137                let logits = Tensor::from_candle(logits_candle);
138
139                Ok(ModelOutputs::Logits { logits, hidden_states: None })
140            }
141            _ => Err(anyhow::anyhow!("StarCoder only supports text inputs")),
142        }
143    }
144
145    fn generate(&self, prompt: &str, config: &GenerationConfig) -> Result<String> {
146        use crate::tokenizer::Tokenizer;
147        use rand::Rng;
148        let tokenizer = Tokenizer::new();
149        let mut tokens: Vec<u32> = tokenizer.encode(prompt);
150        for _ in 0..config.max_new_tokens {
151            let tokens_i64: Vec<i64> = tokens.iter().map(|&t| t as i64).collect();
152            let input = Tensor::from_i64_slice(&tokens_i64, &[1, tokens.len()], &self.device)?;
153            let outputs = self.forward(&ModelInputs::text(input))?;
154            let logits = match outputs { ModelOutputs::Logits { logits, .. } => logits, _ => return Err(anyhow::anyhow!("Expected logits")) };
155            let logits_candle = logits.to_candle()?;
156            let last = logits_candle.narrow(1, logits_candle.dims()[1] - 1, 1)?.squeeze(1)?.squeeze(0)?;
157            let logits_vec: Vec<f32> = last.to_vec1()?;
158            let next = if config.do_sample && config.temperature > 0.0 {
159                let scaled: Vec<f32> = logits_vec.iter().map(|&x| x / config.temperature).collect();
160                let max_v = scaled.iter().cloned().fold(f32::NEG_INFINITY, f32::max);
161                let exp_sum: f32 = scaled.iter().map(|&x| (x - max_v).exp()).sum();
162                let probs: Vec<f32> = scaled.iter().map(|&x| (x - max_v).exp() / exp_sum).collect();
163                let mut rng = rand::thread_rng();
164                let r: f32 = rng.gen();
165                let mut cum = 0.0;
166                let mut s = 0u32;
167                for (i, &p) in probs.iter().enumerate() { cum += p; if r <= cum { s = i as u32; break; } }
168                s
169            } else {
170                logits_vec.iter().enumerate().max_by(|a, b| a.1.partial_cmp(b.1).unwrap()).map(|(i, _)| i as u32).unwrap_or(0)
171            };
172            if next == config.eos_token_id { break; }
173            tokens.push(next);
174        }
175        Ok(tokenizer.decode(&tokens))
176    }
177
178    fn config(&self) -> &Self::Config { &self.config }
179    fn memory_requirements(&self) -> MemoryRequirements {
180        let p = self.config.vocab_size * self.config.hidden_size + self.config.num_hidden_layers * 8 * self.config.hidden_size.pow(2);
181        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 }
182    }
183    fn to_device(&mut self, device: &Device) -> Result<()> {
184        self.wte = self.wte.to_device(device)?;
185        self.wpe = self.wpe.to_device(device)?;
186        self.ln_f = self.ln_f.to_device(device)?;
187        self.lm_head = self.lm_head.to_device(device)?;
188        for layer in &mut self.layers { layer.to_device(device)?; }
189        self.device = device.clone();
190        Ok(())
191    }
192}
193
194impl StarCoderLayer {
195    fn new(config: &StarCoderConfig, device: &Device) -> Result<Self> {
196        Ok(Self {
197            attn: StarCoderAttention::new(config, device)?,
198            mlp: StarCoderMLP::new(config, device)?,
199            ln_1: ops_fn::zeros(&[config.hidden_size], DataType::Float32, device)?,
200            ln_2: ops_fn::zeros(&[config.hidden_size], DataType::Float32, device)?,
201        })
202    }
203
204    fn forward(&self, hidden_states: &Tensor) -> Result<Tensor> {
205        let residual = hidden_states.clone();
206        let h = ops_fn::layer_norm(hidden_states, &self.ln_1, None, 1e-5)?;
207        let attn_out = self.attn.forward(&h)?;
208        let h = ops_fn::add(&residual, &attn_out)?;
209
210        let residual = h.clone();
211        let h = ops_fn::layer_norm(&h, &self.ln_2, None, 1e-5)?;
212        let mlp_out = self.mlp.forward(&h)?;
213        ops_fn::add(&residual, &mlp_out)
214    }
215
216    fn load_weights(&mut self, weights: &ModelWeights, idx: usize) -> Result<()> {
217        let p = format!("transformer.h.{}", idx);
218        if let Some(w) = weights.get(&format!("{}.attn.c_attn.weight", p)) { self.attn.c_attn = ops_fn::transpose(w)?; }
219        if let Some(w) = weights.get(&format!("{}.attn.c_proj.weight", p)) { self.attn.c_proj = ops_fn::transpose(w)?; }
220        if let Some(w) = weights.get(&format!("{}.mlp.c_fc.weight", p)) { self.mlp.c_fc = ops_fn::transpose(w)?; }
221        if let Some(w) = weights.get(&format!("{}.mlp.c_proj.weight", p)) { self.mlp.c_proj = ops_fn::transpose(w)?; }
222        if let Some(w) = weights.get(&format!("{}.ln_1.weight", p)) { self.ln_1 = w.clone(); }
223        if let Some(w) = weights.get(&format!("{}.ln_2.weight", p)) { self.ln_2 = w.clone(); }
224        Ok(())
225    }
226
227    fn to_device(&mut self, device: &Device) -> Result<()> {
228        self.attn.to_device(device)?;
229        self.mlp.to_device(device)?;
230        self.ln_1 = self.ln_1.to_device(device)?;
231        self.ln_2 = self.ln_2.to_device(device)?;
232        Ok(())
233    }
234}
235
236impl StarCoderAttention {
237    fn new(config: &StarCoderConfig, device: &Device) -> Result<Self> {
238        let head_dim = config.hidden_size / config.num_attention_heads;
239        // MQA: Q has num_heads, K/V have 1 head each
240        let qkv_size = config.hidden_size + 2 * head_dim;  // Q (all heads) + K (1 head) + V (1 head)
241        Ok(Self {
242            c_attn: ops_fn::zeros(&[config.hidden_size, qkv_size], DataType::Float32, device)?,
243            c_proj: ops_fn::zeros(&[config.hidden_size, config.hidden_size], DataType::Float32, device)?,
244            num_heads: config.num_attention_heads,
245            head_dim,
246            scale: 1.0 / (head_dim as f32).sqrt(),
247        })
248    }
249
250    fn forward(&self, hidden_states: &Tensor) -> Result<Tensor> {
251        let shape = hidden_states.shape();
252        let (batch, seq, hidden_size) = (shape[0], shape[1], shape[2]);
253
254        let qkv = ops_fn::matmul(hidden_states, &self.c_attn)?.to_candle()?;
255
256        // Split QKV - Q has full hidden_size, K and V each have head_dim
257        let q = qkv.narrow(2, 0, hidden_size)?;
258        let k = qkv.narrow(2, hidden_size, self.head_dim)?;
259        let v = qkv.narrow(2, hidden_size + self.head_dim, self.head_dim)?;
260
261        let q = q.reshape(&[batch, seq, self.num_heads, self.head_dim])?.transpose(1, 2)?;
262        let k = k.reshape(&[batch, seq, 1, self.head_dim])?.transpose(1, 2)?;
263        let v = v.reshape(&[batch, seq, 1, self.head_dim])?.transpose(1, 2)?;
264
265        // Broadcast K, V to all heads
266        let k = k.broadcast_as(&[batch, self.num_heads, seq, self.head_dim])?.contiguous()?;
267        let v = v.broadcast_as(&[batch, self.num_heads, seq, self.head_dim])?.contiguous()?;
268
269        let q = q.contiguous()?;
270        let k_t = k.transpose(2, 3)?.contiguous()?;
271        let scores = (q.matmul(&k_t)? * (self.scale as f64))?;
272
273        // Causal mask
274        let device = scores.device();
275        let mask = {
276            let mut m = vec![0.0f32; seq * seq];
277            for i in 0..seq { for j in (i+1)..seq { m[i*seq+j] = f32::NEG_INFINITY; } }
278            candle_core::Tensor::from_vec(m, &[1, 1, seq, seq], device)?
279        };
280        let scores = scores.broadcast_add(&mask)?;
281
282        let attn = candle_nn::ops::softmax_last_dim(&scores)?.matmul(&v)?;
283        let out = attn.transpose(1, 2)?.reshape(&[batch, seq, hidden_size])?;
284        ops_fn::matmul(&Tensor::from_candle(out), &self.c_proj)
285    }
286
287    fn to_device(&mut self, device: &Device) -> Result<()> {
288        self.c_attn = self.c_attn.to_device(device)?;
289        self.c_proj = self.c_proj.to_device(device)?;
290        Ok(())
291    }
292}
293
294impl StarCoderMLP {
295    fn new(config: &StarCoderConfig, device: &Device) -> Result<Self> {
296        Ok(Self {
297            c_fc: ops_fn::zeros(&[config.hidden_size, config.intermediate_size], DataType::Float32, device)?,
298            c_proj: ops_fn::zeros(&[config.intermediate_size, config.hidden_size], DataType::Float32, device)?,
299        })
300    }
301    fn forward(&self, x: &Tensor) -> Result<Tensor> {
302        let h = ops_fn::matmul(x, &self.c_fc)?;
303        let h = ops_fn::gelu(&h)?;
304        ops_fn::matmul(&h, &self.c_proj)
305    }
306    fn to_device(&mut self, device: &Device) -> Result<()> {
307        self.c_fc = self.c_fc.to_device(device)?;
308        self.c_proj = self.c_proj.to_device(device)?;
309        Ok(())
310    }
311}
312
313#[cfg(test)]
314mod tests {
315    use super::*;
316    #[test]
317    fn test_starcoder_creation() {
318        let config = StarCoderConfig { vocab_size: 1000, hidden_size: 128, intermediate_size: 512, num_hidden_layers: 2, num_attention_heads: 4, num_key_value_heads: 1, ..Default::default() };
319        let model = StarCoderModelV2::new(config).unwrap();
320        assert_eq!(model.config().vocab_size(), 1000);
321    }
322    #[test]
323    fn test_starcoder_forward() {
324        let config = StarCoderConfig { vocab_size: 100, hidden_size: 64, intermediate_size: 256, num_hidden_layers: 1, num_attention_heads: 4, num_key_value_heads: 1, ..Default::default() };
325        let model = StarCoderModelV2::new(config).unwrap();
326        let inputs = ModelInputs::text(ops_fn::zeros(&[2, 8], DataType::Int64, &Device::CPU).unwrap());
327        match model.forward(&inputs).unwrap() { ModelOutputs::Logits { logits, .. } => assert_eq!(logits.shape(), &[2, 8, 100]), _ => panic!() }
328    }
329}