Skip to main content

runtime/models_v2/
falcon.rs

1//! Falcon Model V2 - Clean implementation using solid abstractions
2//!
3//! This implements the Falcon architecture which includes:
4//! - Falcon-7B, Falcon-40B, Falcon-180B models
5//! - Uses unified Tensor type from tensor_core
6//! - Implements Model trait from model_core
7//!
8//! Falcon-specific characteristics:
9//! - Uses `transformer.h.{i}` layer prefix instead of `model.layers.{i}`
10//! - Uses `word_embeddings` instead of `embed_tokens`
11//! - Uses `ln_f` instead of `norm` for final layer norm
12//! - Parallel attention + MLP (like Phi)
13//! - Packed QKV attention (query_key_value)
14
15use crate::model_config;
16use super::traits::*;
17use anyhow::Result;
18use serde::{Serialize, Deserialize};
19
20/// Falcon model configuration using the model_config macro
21model_config!(FalconConfig {
22    vocab_size: usize = 65024,
23    hidden_size: usize = 4544,
24    intermediate_size: usize = 18176,
25    num_hidden_layers: usize = 32,
26    num_attention_heads: usize = 71,
27    num_key_value_heads: usize = 1,  // Falcon typically uses MQA (1 kv head) or GQA
28    hidden_act: String = "gelu".to_string(),
29    max_position_embeddings: usize = 2048,
30    initializer_range: f32 = 0.02,
31    rms_norm_eps: f32 = 1e-5,
32    use_cache: bool = true,
33    pad_token_id: i64 = 11,
34    bos_token_id: i64 = 11,
35    eos_token_id: i64 = 11,
36    tie_word_embeddings: bool = false,
37    rope_theta: f32 = 10000.0,
38    // Falcon specific
39    parallel_attn: bool = true,
40    bias: bool = false,
41    new_decoder_architecture: bool = false,
42});
43
44impl FalconConfig {
45    /// Create FalconConfig from GGUF model configuration
46    pub fn from_gguf_config(gguf: &crate::weight_loader_core::GGUFModelConfig) -> Self {
47        Self {
48            vocab_size: gguf.vocab_size,
49            hidden_size: gguf.hidden_size,
50            intermediate_size: gguf.intermediate_size,
51            num_hidden_layers: gguf.num_hidden_layers,
52            num_attention_heads: gguf.num_attention_heads,
53            num_key_value_heads: gguf.num_key_value_heads,
54            rms_norm_eps: gguf.rms_norm_eps,
55            rope_theta: gguf.rope_theta,
56            max_position_embeddings: gguf.max_position_embeddings,
57            ..Default::default()
58        }
59    }
60}
61
62/// Main Falcon model implementation
63pub struct FalconModelV2 {
64    config: FalconConfig,
65    device: Device,
66
67    // Model components using unified Tensor type
68    word_embeddings: Tensor,
69    h: Vec<FalconDecoderLayer>,
70    ln_f: Tensor,
71    lm_head: Tensor,
72}
73
74/// Falcon transformer layer
75pub struct FalconDecoderLayer {
76    self_attention: FalconAttention,
77    mlp: FalconMLP,
78    input_layernorm: Tensor,
79    // For non-parallel architectures
80    post_attention_layernorm: Option<Tensor>,
81    config: FalconConfig,
82}
83
84/// Falcon attention mechanism with packed QKV
85pub struct FalconAttention {
86    query_key_value: Tensor,  // Packed QKV projection
87    dense: Tensor,            // Output projection
88    num_heads: usize,
89    num_key_value_heads: usize,
90    head_dim: usize,
91    scale: f32,
92}
93
94/// Falcon MLP (feed-forward network) - standard 2-layer MLP
95pub struct FalconMLP {
96    dense_h_to_4h: Tensor,
97    dense_4h_to_h: Tensor,
98    hidden_act: String,
99}
100
101impl Model for FalconModelV2 {
102    type Config = FalconConfig;
103
104    fn new(config: FalconConfig) -> Result<Self> {
105        let device = Device::CPU;
106
107        let word_embeddings = ops_fn::zeros(
108            &[config.vocab_size, config.hidden_size],
109            DataType::Float32,
110            &device
111        )?;
112
113        let ln_f = ops_fn::zeros(
114            &[config.hidden_size],
115            DataType::Float32,
116            &device
117        )?;
118
119        let lm_head = if config.tie_word_embeddings {
120            word_embeddings.clone()
121        } else {
122            ops_fn::zeros(
123                &[config.hidden_size, config.vocab_size],
124                DataType::Float32,
125                &device
126            )?
127        };
128
129        // Create transformer layers
130        let mut h = Vec::with_capacity(config.num_hidden_layers);
131        for _ in 0..config.num_hidden_layers {
132            h.push(FalconDecoderLayer::new(&config, &device)?);
133        }
134
135        Ok(Self {
136            config,
137            device,
138            word_embeddings,
139            h,
140            ln_f,
141            lm_head,
142        })
143    }
144
145    fn from_weights(config: FalconConfig, weights: ModelWeights) -> Result<Self> {
146        let mut model = Self::new(config)?;
147
148        // Load weights from the unified weight container
149        // Falcon uses transformer.word_embeddings.weight
150        if let Some(embed_weights) = weights.get("transformer.word_embeddings.weight") {
151            model.word_embeddings = embed_weights.clone();
152        }
153
154        // Falcon uses transformer.ln_f.weight for final layer norm
155        if let Some(norm_weights) = weights.get("transformer.ln_f.weight") {
156            model.ln_f = norm_weights.clone();
157        }
158
159        // lm_head weight needs transpose for matmul: [vocab, hidden] -> [hidden, vocab]
160        if let Some(lm_head_weights) = weights.get("lm_head.weight") {
161            model.lm_head = ops_fn::transpose(lm_head_weights)?;
162        }
163
164        // Load layer weights
165        for (i, layer) in model.h.iter_mut().enumerate() {
166            layer.load_weights(&weights, i)?;
167        }
168
169        Ok(model)
170    }
171
172    fn forward(&self, inputs: &ModelInputs) -> Result<ModelOutputs> {
173        match inputs {
174            ModelInputs::Text { input_ids, attention_mask, .. } => {
175                // 1. Token embedding
176                let mut hidden_states = ops_fn::embedding(input_ids, &self.word_embeddings)?;
177
178                // 2. Apply transformer layers (with RoPE and parallel attention+MLP)
179                for layer in &self.h {
180                    hidden_states = layer.forward(&hidden_states, attention_mask.as_ref(), self.config.rope_theta)?;
181                }
182
183                // 3. Final layer norm
184                hidden_states = ops_fn::layer_norm(&hidden_states, &self.ln_f, None, self.config.rms_norm_eps)?;
185
186                // 4. Language modeling head
187                let logits = ops_fn::matmul(&hidden_states, &self.lm_head)?;
188
189                Ok(ModelOutputs::Logits {
190                    logits,
191                    hidden_states: None,
192                })
193            }
194            ModelInputs::Multimodal { input_ids, .. } => {
195                let text_inputs = ModelInputs::Text {
196                    input_ids: input_ids.clone(),
197                    attention_mask: None,
198                    position_ids: None,
199                };
200                self.forward(&text_inputs)
201            }
202            _ => Err(anyhow::anyhow!("Falcon model only supports text and multimodal inputs")),
203        }
204    }
205
206    fn generate(&self, prompt: &str, config: &GenerationConfig) -> Result<String> {
207        use crate::tokenizer::Tokenizer;
208        use rand::Rng;
209
210        // 1. Tokenize prompt
211        let tokenizer = Tokenizer::new();
212        let mut tokens: Vec<u32> = tokenizer.encode(prompt);
213
214        // 2. Generation loop
215        for _ in 0..config.max_new_tokens {
216            // Create input tensor from current tokens
217            let tokens_i64: Vec<i64> = tokens.iter().map(|&t| t as i64).collect();
218            let input_tensor = Tensor::from_i64_slice(&tokens_i64, &[1, tokens.len()], &self.device)?;
219
220            let inputs = ModelInputs::Text {
221                input_ids: input_tensor,
222                attention_mask: None,
223                position_ids: None,
224            };
225
226            // 3. Forward pass
227            let outputs = self.forward(&inputs)?;
228
229            // 4. Get logits and sample next token
230            let logits = match outputs {
231                ModelOutputs::Logits { logits, .. } => logits,
232                _ => return Err(anyhow::anyhow!("Expected logits output")),
233            };
234
235            // Get last token logits
236            let logits_candle = logits.to_candle()?;
237            let shape = logits_candle.dims();
238
239            // Extract last position logits [batch, seq, vocab] -> [vocab]
240            let last_logits = if shape.len() == 3 {
241                let seq_len = shape[1];
242                logits_candle
243                    .narrow(1, seq_len - 1, 1)?
244                    .squeeze(1)?
245                    .squeeze(0)?
246            } else {
247                let seq_len = shape[0];
248                logits_candle
249                    .narrow(0, seq_len - 1, 1)?
250                    .squeeze(0)?
251            };
252
253            // Convert to probabilities and sample
254            let logits_vec: Vec<f32> = last_logits.to_vec1()?;
255
256            let next_token = if config.do_sample && config.temperature > 0.0 {
257                // Temperature sampling
258                let scaled: Vec<f32> = logits_vec.iter()
259                    .map(|&x| x / config.temperature)
260                    .collect();
261
262                // Softmax
263                let max_val = scaled.iter().cloned().fold(f32::NEG_INFINITY, f32::max);
264                let exp_sum: f32 = scaled.iter().map(|&x| (x - max_val).exp()).sum();
265                let probs: Vec<f32> = scaled.iter()
266                    .map(|&x| (x - max_val).exp() / exp_sum)
267                    .collect();
268
269                // Sample from distribution
270                let mut rng = rand::thread_rng();
271                let random_val: f32 = rng.gen();
272                let mut cumulative = 0.0;
273                let mut sampled = 0u32;
274
275                for (idx, &prob) in probs.iter().enumerate() {
276                    cumulative += prob;
277                    if random_val <= cumulative {
278                        sampled = idx as u32;
279                        break;
280                    }
281                }
282                sampled
283            } else {
284                // Greedy sampling
285                let mut max_idx = 0;
286                let mut max_val = logits_vec[0];
287                for (idx, &val) in logits_vec.iter().enumerate() {
288                    if val > max_val {
289                        max_val = val;
290                        max_idx = idx;
291                    }
292                }
293                max_idx as u32
294            };
295
296            // 5. Check for EOS
297            if next_token == config.eos_token_id {
298                break;
299            }
300
301            // 6. Append token
302            tokens.push(next_token);
303        }
304
305        // 7. Decode and return
306        Ok(tokenizer.decode(&tokens))
307    }
308
309    fn config(&self) -> &Self::Config {
310        &self.config
311    }
312
313    fn memory_requirements(&self) -> MemoryRequirements {
314        // Calculate approximate memory requirements
315        let param_size = self.config.vocab_size * self.config.hidden_size + // embeddings
316                        self.config.num_hidden_layers * (
317                            // Packed QKV + output projection
318                            self.config.hidden_size * (self.config.num_attention_heads + 2 * self.config.num_key_value_heads) * (self.config.hidden_size / self.config.num_attention_heads) +
319                            self.config.hidden_size * self.config.hidden_size +
320                            // MLP
321                            2 * self.config.hidden_size * self.config.intermediate_size
322                        );
323
324        let param_bytes = param_size * 4; // float32
325        let kv_cache_bytes = 2 * self.config.num_hidden_layers *
326                           self.config.max_position_embeddings *
327                           self.config.num_key_value_heads *
328                           (self.config.hidden_size / self.config.num_attention_heads) * 4;
329
330        MemoryRequirements {
331            gpu_memory: param_bytes,
332            cpu_memory: param_bytes / 4,
333            kv_cache_memory: kv_cache_bytes,
334            peak_memory: param_bytes + kv_cache_bytes,
335        }
336    }
337
338    fn to_device(&mut self, device: &Device) -> Result<()> {
339        self.word_embeddings = self.word_embeddings.to_device(device)?;
340        self.ln_f = self.ln_f.to_device(device)?;
341        self.lm_head = self.lm_head.to_device(device)?;
342
343        for layer in &mut self.h {
344            layer.to_device(device)?;
345        }
346
347        self.device = device.clone();
348        Ok(())
349    }
350}
351
352impl FalconDecoderLayer {
353    fn new(config: &FalconConfig, device: &Device) -> Result<Self> {
354        let input_layernorm = ops_fn::zeros(&[config.hidden_size], DataType::Float32, device)?;
355
356        // Non-parallel architectures have a post-attention layernorm
357        let post_attention_layernorm = if !config.parallel_attn {
358            Some(ops_fn::zeros(&[config.hidden_size], DataType::Float32, device)?)
359        } else {
360            None
361        };
362
363        Ok(Self {
364            self_attention: FalconAttention::new(config, device)?,
365            mlp: FalconMLP::new(config, device)?,
366            input_layernorm,
367            post_attention_layernorm,
368            config: config.clone(),
369        })
370    }
371
372    fn forward(&self, hidden_states: &Tensor, attention_mask: Option<&Tensor>, rope_theta: f32) -> Result<Tensor> {
373        let residual = hidden_states.clone();
374
375        // Pre-attention layer norm
376        let normed = ops_fn::layer_norm(hidden_states, &self.input_layernorm, None, self.config.rms_norm_eps)?;
377
378        if self.config.parallel_attn {
379            // Falcon parallel architecture: attention and MLP computed in parallel
380            let attn_output = self.self_attention.forward(&normed, attention_mask, rope_theta)?;
381            let mlp_output = self.mlp.forward(&normed)?;
382
383            // Add both outputs to residual
384            let combined = ops_fn::add(&attn_output, &mlp_output)?;
385            ops_fn::add(&residual, &combined)
386        } else {
387            // Sequential architecture (Falcon-40B uses this)
388            let attn_output = self.self_attention.forward(&normed, attention_mask, rope_theta)?;
389            let hidden_states = ops_fn::add(&residual, &attn_output)?;
390
391            // Post-attention layer norm
392            if let Some(post_ln) = &self.post_attention_layernorm {
393                let normed = ops_fn::layer_norm(&hidden_states, post_ln, None, self.config.rms_norm_eps)?;
394                let mlp_output = self.mlp.forward(&normed)?;
395                ops_fn::add(&hidden_states, &mlp_output)
396            } else {
397                // Fallback if post_attention_layernorm is not available
398                let mlp_output = self.mlp.forward(&hidden_states)?;
399                ops_fn::add(&hidden_states, &mlp_output)
400            }
401        }
402    }
403
404    fn load_weights(&mut self, weights: &ModelWeights, layer_idx: usize) -> Result<()> {
405        let prefix = format!("transformer.h.{}", layer_idx);
406
407        // Load input layernorm
408        if let Some(ln) = weights.get(&format!("{}.input_layernorm.weight", prefix)) {
409            self.input_layernorm = ln.clone();
410        }
411
412        // Load post-attention layernorm for sequential architectures
413        if let Some(post_ln) = weights.get(&format!("{}.post_attention_layernorm.weight", prefix)) {
414            self.post_attention_layernorm = Some(post_ln.clone());
415        }
416
417        // Load attention weights
418        self.self_attention.load_weights(weights, layer_idx)?;
419
420        // Load MLP weights
421        self.mlp.load_weights(weights, layer_idx)?;
422
423        Ok(())
424    }
425
426    fn to_device(&mut self, device: &Device) -> Result<()> {
427        self.input_layernorm = self.input_layernorm.to_device(device)?;
428        if let Some(post_ln) = &self.post_attention_layernorm {
429            self.post_attention_layernorm = Some(post_ln.to_device(device)?);
430        }
431        self.self_attention.to_device(device)?;
432        self.mlp.to_device(device)?;
433        Ok(())
434    }
435}
436
437/// Apply Rotary Position Embedding (RoPE) to Q and K tensors
438/// Input shape: [batch, heads, seq, head_dim]
439/// Returns tensors with same shape but with positional information encoded
440fn apply_rope(
441    q: &candle_core::Tensor,
442    k: &candle_core::Tensor,
443    seq_len: usize,
444    head_dim: usize,
445    rope_theta: f32,
446) -> Result<(candle_core::Tensor, candle_core::Tensor)> {
447    let device = q.device();
448
449    // Compute inverse frequencies: 1 / (theta^(2i/d)) for i in [0, d/2)
450    let half_dim = head_dim / 2;
451    let inv_freq: Vec<f32> = (0..half_dim)
452        .map(|i| 1.0 / rope_theta.powf((2 * i) as f32 / head_dim as f32))
453        .collect();
454
455    // Create position indices [0, 1, 2, ..., seq_len-1]
456    let positions: Vec<f32> = (0..seq_len).map(|p| p as f32).collect();
457
458    // Compute angles: pos * inv_freq -> [seq_len, half_dim]
459    let mut angles = Vec::with_capacity(seq_len * half_dim);
460    for pos in &positions {
461        for freq in &inv_freq {
462            angles.push(pos * freq);
463        }
464    }
465
466    let angles_tensor = candle_core::Tensor::from_vec(angles, &[seq_len, half_dim], device)?;
467
468    // Compute cos and sin
469    let cos = angles_tensor.cos()?;
470    let sin = angles_tensor.sin()?;
471
472    // Reshape for broadcasting: [1, 1, seq_len, half_dim]
473    let cos = cos.unsqueeze(0)?.unsqueeze(0)?;
474    let sin = sin.unsqueeze(0)?.unsqueeze(0)?;
475
476    // Apply RoPE rotation
477    // Split q and k into two halves along head_dim
478    let q_half1 = q.narrow(3, 0, half_dim)?;
479    let q_half2 = q.narrow(3, half_dim, half_dim)?;
480    let k_half1 = k.narrow(3, 0, half_dim)?;
481    let k_half2 = k.narrow(3, half_dim, half_dim)?;
482
483    // Apply rotation
484    let q_rot1 = (q_half1.broadcast_mul(&cos)? - q_half2.broadcast_mul(&sin)?)?;
485    let q_rot2 = (q_half1.broadcast_mul(&sin)? + q_half2.broadcast_mul(&cos)?)?;
486    let k_rot1 = (k_half1.broadcast_mul(&cos)? - k_half2.broadcast_mul(&sin)?)?;
487    let k_rot2 = (k_half1.broadcast_mul(&sin)? + k_half2.broadcast_mul(&cos)?)?;
488
489    // Concatenate rotated halves
490    let q_rotated = candle_core::Tensor::cat(&[&q_rot1, &q_rot2], 3)?;
491    let k_rotated = candle_core::Tensor::cat(&[&k_rot1, &k_rot2], 3)?;
492
493    Ok((q_rotated, k_rotated))
494}
495
496impl FalconAttention {
497    fn new(config: &FalconConfig, device: &Device) -> Result<Self> {
498        let num_heads = config.num_attention_heads;
499        let num_key_value_heads = config.num_key_value_heads;
500        let head_dim = config.hidden_size / num_heads;
501        let scale = 1.0 / (head_dim as f32).sqrt();
502
503        // Packed QKV: [hidden_size, (num_heads + 2*num_kv_heads) * head_dim]
504        let qkv_size = (num_heads + 2 * num_key_value_heads) * head_dim;
505        let query_key_value = ops_fn::zeros(
506            &[config.hidden_size, qkv_size],
507            DataType::Float32,
508            device
509        )?;
510
511        let dense = ops_fn::zeros(
512            &[num_heads * head_dim, config.hidden_size],
513            DataType::Float32,
514            device
515        )?;
516
517        Ok(Self {
518            query_key_value,
519            dense,
520            num_heads,
521            num_key_value_heads,
522            head_dim,
523            scale,
524        })
525    }
526
527    fn forward(&self, hidden_states: &Tensor, _attention_mask: Option<&Tensor>, rope_theta: f32) -> Result<Tensor> {
528        // Get batch and sequence length from hidden_states shape
529        let shape = hidden_states.shape();
530        let (batch_size, seq_len, _hidden_size) = if shape.len() == 3 {
531            (shape[0], shape[1], shape[2])
532        } else if shape.len() == 2 {
533            (1, shape[0], shape[1])
534        } else {
535            return Err(anyhow::anyhow!("Invalid hidden_states shape: {:?}", shape));
536        };
537
538        // 1. Compute packed QKV projection
539        // [batch, seq, hidden] @ [hidden, qkv_size] -> [batch, seq, qkv_size]
540        let qkv = ops_fn::matmul(hidden_states, &self.query_key_value)?;
541        let qkv_candle = qkv.to_candle()?;
542
543        // 2. Split QKV
544        // qkv_size = (num_heads + 2*num_kv_heads) * head_dim
545        let q_size = self.num_heads * self.head_dim;
546        let kv_size = self.num_key_value_heads * self.head_dim;
547
548        // Split along the last dimension
549        let query_states = qkv_candle.narrow(2, 0, q_size)?;
550        let key_states = qkv_candle.narrow(2, q_size, kv_size)?;
551        let value_states = qkv_candle.narrow(2, q_size + kv_size, kv_size)?;
552
553        // 3. Reshape for multi-head attention
554        // Q: [batch, seq, num_heads * head_dim] -> [batch, num_heads, seq, head_dim]
555        // K: [batch, seq, num_kv_heads * head_dim] -> [batch, num_kv_heads, seq, head_dim]
556        let q_reshaped = query_states
557            .reshape(&[batch_size, seq_len, self.num_heads, self.head_dim])?
558            .transpose(1, 2)?;
559
560        let k_reshaped = key_states
561            .reshape(&[batch_size, seq_len, self.num_key_value_heads, self.head_dim])?
562            .transpose(1, 2)?;
563
564        let v_reshaped = value_states
565            .reshape(&[batch_size, seq_len, self.num_key_value_heads, self.head_dim])?
566            .transpose(1, 2)?;
567
568        // 4. Apply RoPE (Rotary Position Embedding) to Q and K
569        let (q_with_rope, k_with_rope) = apply_rope(&q_reshaped, &k_reshaped, seq_len, self.head_dim, rope_theta)?;
570
571        // 5. Handle GQA (Grouped Query Attention) / MQA - repeat K/V heads to match Q heads
572        let num_groups = self.num_heads / self.num_key_value_heads;
573        let (k_expanded, v_expanded) = if num_groups > 1 {
574            // Repeat K and V along the head dimension
575            let k_rep = k_with_rope
576                .unsqueeze(2)?
577                .broadcast_as(&[batch_size, self.num_key_value_heads, num_groups, seq_len, self.head_dim])?
578                .reshape(&[batch_size, self.num_heads, seq_len, self.head_dim])?;
579            let v_rep = v_reshaped
580                .unsqueeze(2)?
581                .broadcast_as(&[batch_size, self.num_key_value_heads, num_groups, seq_len, self.head_dim])?
582                .reshape(&[batch_size, self.num_heads, seq_len, self.head_dim])?;
583            (k_rep, v_rep)
584        } else {
585            (k_with_rope, v_reshaped)
586        };
587
588        // 6. Scaled dot-product attention
589        // scores = Q @ K^T / sqrt(head_dim)
590        let k_t = k_expanded.transpose(2, 3)?;
591
592        let q_contiguous = q_with_rope.contiguous()?;
593        let k_contiguous = k_t.contiguous()?;
594
595        let scores = q_contiguous.matmul(&k_contiguous)?;
596        let scaled_scores = (scores * (self.scale as f64))?;
597
598        // Apply causal mask: prevent attending to future positions
599        let device = scaled_scores.device();
600        let causal_mask = {
601            let mut mask_data = vec![0.0f32; seq_len * seq_len];
602            for i in 0..seq_len {
603                for j in 0..seq_len {
604                    if j > i {
605                        mask_data[i * seq_len + j] = f32::NEG_INFINITY;
606                    }
607                }
608            }
609            candle_core::Tensor::from_vec(mask_data, &[1, 1, seq_len, seq_len], device)?
610        };
611
612        let masked_scores = scaled_scores.broadcast_add(&causal_mask)?;
613
614        // Softmax over last dimension
615        let attention_weights = candle_nn::ops::softmax_last_dim(&masked_scores)?;
616
617        // Apply attention to values
618        let v_contiguous = v_expanded.contiguous()?;
619        let attn_output = attention_weights.matmul(&v_contiguous)?;
620
621        // 7. Reshape back: [batch, heads, seq, head_dim] -> [batch, seq, heads * head_dim]
622        let attn_output = attn_output
623            .transpose(1, 2)?
624            .reshape(&[batch_size, seq_len, self.num_heads * self.head_dim])?;
625
626        let attn_output = Tensor::from_candle(attn_output);
627
628        // 8. Output projection
629        let output = ops_fn::matmul(&attn_output, &self.dense)?;
630
631        Ok(output)
632    }
633
634    fn load_weights(&mut self, weights: &ModelWeights, layer_idx: usize) -> Result<()> {
635        let prefix = format!("transformer.h.{}.self_attention", layer_idx);
636
637        // Falcon uses packed query_key_value weight
638        // Transpose for matmul: [out, in] -> [in, out]
639        if let Some(qkv) = weights.get(&format!("{}.query_key_value.weight", prefix)) {
640            self.query_key_value = ops_fn::transpose(qkv)?;
641        }
642
643        // Output projection (dense)
644        if let Some(dense) = weights.get(&format!("{}.dense.weight", prefix)) {
645            self.dense = ops_fn::transpose(dense)?;
646        }
647
648        Ok(())
649    }
650
651    fn to_device(&mut self, device: &Device) -> Result<()> {
652        self.query_key_value = self.query_key_value.to_device(device)?;
653        self.dense = self.dense.to_device(device)?;
654        Ok(())
655    }
656}
657
658impl FalconMLP {
659    fn new(config: &FalconConfig, device: &Device) -> Result<Self> {
660        let dense_h_to_4h = ops_fn::zeros(
661            &[config.hidden_size, config.intermediate_size],
662            DataType::Float32,
663            device
664        )?;
665        let dense_4h_to_h = ops_fn::zeros(
666            &[config.intermediate_size, config.hidden_size],
667            DataType::Float32,
668            device
669        )?;
670
671        Ok(Self {
672            dense_h_to_4h,
673            dense_4h_to_h,
674            hidden_act: config.hidden_act.clone(),
675        })
676    }
677
678    fn forward(&self, hidden_states: &Tensor) -> Result<Tensor> {
679        // First projection with activation
680        let intermediate = ops_fn::matmul(hidden_states, &self.dense_h_to_4h)?;
681
682        // Apply activation (Falcon uses GELU)
683        let activated = match self.hidden_act.as_str() {
684            "gelu" | "gelu_new" => ops_fn::gelu(&intermediate)?,
685            "silu" | "swish" => ops_fn::silu(&intermediate)?,
686            _ => return Err(anyhow::anyhow!("Unsupported activation: {}", self.hidden_act)),
687        };
688
689        // Second projection
690        let output = ops_fn::matmul(&activated, &self.dense_4h_to_h)?;
691
692        Ok(output)
693    }
694
695    fn load_weights(&mut self, weights: &ModelWeights, layer_idx: usize) -> Result<()> {
696        let prefix = format!("transformer.h.{}.mlp", layer_idx);
697
698        // Transpose for matmul: [out, in] -> [in, out]
699        if let Some(w) = weights.get(&format!("{}.dense_h_to_4h.weight", prefix)) {
700            self.dense_h_to_4h = ops_fn::transpose(w)?;
701        }
702        if let Some(w) = weights.get(&format!("{}.dense_4h_to_h.weight", prefix)) {
703            self.dense_4h_to_h = ops_fn::transpose(w)?;
704        }
705
706        Ok(())
707    }
708
709    fn to_device(&mut self, device: &Device) -> Result<()> {
710        self.dense_h_to_4h = self.dense_h_to_4h.to_device(device)?;
711        self.dense_4h_to_h = self.dense_4h_to_h.to_device(device)?;
712        Ok(())
713    }
714}
715
716#[cfg(test)]
717mod tests {
718    use super::*;
719
720    #[test]
721    fn test_falcon_model_creation() {
722        let config = FalconConfig {
723            vocab_size: 1000,
724            hidden_size: 128,
725            intermediate_size: 512,
726            num_hidden_layers: 2,
727            num_attention_heads: 8,
728            num_key_value_heads: 1, // MQA
729            ..Default::default()
730        };
731
732        let model = FalconModelV2::new(config).unwrap();
733        assert_eq!(model.config().vocab_size(), 1000);
734        assert_eq!(model.config().hidden_size(), 128);
735        assert_eq!(model.config().num_layers(), 2);
736    }
737
738    #[test]
739    fn test_falcon_forward_pass() {
740        let config = FalconConfig {
741            vocab_size: 100,
742            hidden_size: 64,
743            intermediate_size: 256,
744            num_hidden_layers: 1,
745            num_attention_heads: 8,
746            num_key_value_heads: 1, // MQA
747            ..Default::default()
748        };
749
750        let model = FalconModelV2::new(config).unwrap();
751        let input_ids = ops_fn::zeros(&[2, 8], DataType::Int64, &Device::CPU).unwrap();
752        let inputs = ModelInputs::text(input_ids);
753
754        let outputs = model.forward(&inputs).unwrap();
755        match outputs {
756            ModelOutputs::Logits { logits, .. } => {
757                assert_eq!(logits.shape(), &[2, 8, 100]);
758            }
759            _ => panic!("Expected logits output"),
760        }
761    }
762
763    #[test]
764    fn test_falcon_gqa() {
765        // Test with Grouped Query Attention (GQA)
766        let config = FalconConfig {
767            vocab_size: 100,
768            hidden_size: 64,
769            intermediate_size: 256,
770            num_hidden_layers: 1,
771            num_attention_heads: 8,
772            num_key_value_heads: 2, // GQA with 4 groups
773            ..Default::default()
774        };
775
776        let model = FalconModelV2::new(config).unwrap();
777        let input_ids = ops_fn::zeros(&[1, 4], DataType::Int64, &Device::CPU).unwrap();
778        let inputs = ModelInputs::text(input_ids);
779
780        let outputs = model.forward(&inputs).unwrap();
781        match outputs {
782            ModelOutputs::Logits { logits, .. } => {
783                assert_eq!(logits.shape(), &[1, 4, 100]);
784            }
785            _ => panic!("Expected logits output"),
786        }
787    }
788
789    #[test]
790    fn test_falcon_sequential_architecture() {
791        // Test with sequential (non-parallel) attention+MLP
792        let config = FalconConfig {
793            vocab_size: 100,
794            hidden_size: 64,
795            intermediate_size: 256,
796            num_hidden_layers: 1,
797            num_attention_heads: 8,
798            num_key_value_heads: 1,
799            parallel_attn: false, // Sequential architecture
800            ..Default::default()
801        };
802
803        let model = FalconModelV2::new(config).unwrap();
804        let input_ids = ops_fn::zeros(&[1, 4], DataType::Int64, &Device::CPU).unwrap();
805        let inputs = ModelInputs::text(input_ids);
806
807        let outputs = model.forward(&inputs).unwrap();
808        match outputs {
809            ModelOutputs::Logits { logits, .. } => {
810                assert_eq!(logits.shape(), &[1, 4, 100]);
811            }
812            _ => panic!("Expected logits output"),
813        }
814    }
815
816    #[test]
817    fn test_falcon_generation() {
818        let config = FalconConfig {
819            vocab_size: 256,
820            hidden_size: 64,
821            intermediate_size: 256,
822            num_hidden_layers: 1,
823            num_attention_heads: 8,
824            num_key_value_heads: 1,
825            ..Default::default()
826        };
827        let model = FalconModelV2::new(config).unwrap();
828        let gen_config = GenerationConfig {
829            max_new_tokens: 5,
830            ..Default::default()
831        };
832
833        let output = model.generate("Hello", &gen_config).unwrap();
834        assert!(!output.is_empty());
835    }
836}