Skip to main content

runtime/models_v2/
yi.rs

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