1use 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 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
49pub struct YiModelV2 {
51 config: YiConfig,
52 device: Device,
53
54 embed_tokens: Tensor,
56 layers: Vec<YiLayer>,
57 norm: Tensor,
58 lm_head: Tensor,
59}
60
61pub struct YiLayer {
63 self_attn: YiAttention,
64 mlp: YiMLP,
65 input_layernorm: Tensor,
66 post_attention_layernorm: Tensor,
67}
68
69pub 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
81pub 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 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 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 if let Some(lm_head_weights) = weights.get("lm_head.weight") {
148 model.lm_head = ops_fn::transpose(lm_head_weights)?;
149 }
150
151 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 let mut hidden_states = ops_fn::embedding(input_ids, &self.embed_tokens)?;
164
165 for layer in &self.layers {
167 hidden_states = layer.forward(&hidden_states, attention_mask.as_ref(), self.config.rope_theta)?;
168 }
169
170 hidden_states = ops_fn::layer_norm(&hidden_states, &self.norm, None, self.config.rms_norm_eps)?;
172
173 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 let tokenizer = Tokenizer::new();
199 let mut tokens: Vec<u32> = tokenizer.encode(prompt);
200
201 for _ in 0..config.max_new_tokens {
203 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 let outputs = self.forward(&inputs)?;
215
216 let logits = match outputs {
218 ModelOutputs::Logits { logits, .. } => logits,
219 _ => return Err(anyhow::anyhow!("Expected logits output")),
220 };
221
222 let logits_candle = logits.to_candle()?;
224 let shape = logits_candle.dims();
225
226 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 let logits_vec: Vec<f32> = last_logits.to_vec1()?;
242
243 let next_token = if config.do_sample && config.temperature > 0.0 {
244 let scaled: Vec<f32> = logits_vec.iter()
246 .map(|&x| x / config.temperature)
247 .collect();
248
249 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 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 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 if next_token == config.eos_token_id {
285 break;
286 }
287
288 tokens.push(next_token);
290 }
291
292 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 let normed = ops_fn::layer_norm(hidden_states, &self.input_layernorm, None, 1e-5)?;
353
354 let attn_output = self.self_attn.forward(&normed, attention_mask, rope_theta)?;
356
357 let hidden_states = ops_fn::add(hidden_states, &attn_output)?;
359
360 let normed = ops_fn::layer_norm(&hidden_states, &self.post_attention_layernorm, None, 1e-5)?;
362
363 let mlp_output = self.mlp.forward(&normed)?;
365
366 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 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 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 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
420fn 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 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 let positions: Vec<f32> = (0..seq_len).map(|p| p as f32).collect();
438
439 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 let cos = angles_tensor.cos()?;
451 let sin = angles_tensor.sin()?;
452
453 let cos = cos.unsqueeze(0)?.unsqueeze(0)?;
455 let sin = sin.unsqueeze(0)?.unsqueeze(0)?;
456
457 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 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 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 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 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 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 let (q_with_rope, k_with_rope) = apply_rope(&q_reshaped, &k_reshaped, seq_len, self.head_dim, rope_theta)?;
535
536 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 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 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 let attention_weights = candle_nn::ops::softmax_last_dim(&masked_scores)?;
580
581 let v_contiguous = v_expanded.contiguous()?;
583 let attn_output = attention_weights.matmul(&v_contiguous)?;
584
585 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 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 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 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 let gated = ops_fn::mul(&gate_activated, &up_output)?;
635
636 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 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 assert_eq!(model.config().num_attention_heads / model.config().num_key_value_heads, 8);
736 }
737}