1use crate::model_config;
16use super::traits::*;
17use anyhow::Result;
18use serde::{Serialize, Deserialize};
19
20model_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, 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 parallel_attn: bool = true,
40 bias: bool = false,
41 new_decoder_architecture: bool = false,
42});
43
44impl FalconConfig {
45 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
62pub struct FalconModelV2 {
64 config: FalconConfig,
65 device: Device,
66
67 word_embeddings: Tensor,
69 h: Vec<FalconDecoderLayer>,
70 ln_f: Tensor,
71 lm_head: Tensor,
72}
73
74pub struct FalconDecoderLayer {
76 self_attention: FalconAttention,
77 mlp: FalconMLP,
78 input_layernorm: Tensor,
79 post_attention_layernorm: Option<Tensor>,
81 config: FalconConfig,
82}
83
84pub struct FalconAttention {
86 query_key_value: Tensor, dense: Tensor, num_heads: usize,
89 num_key_value_heads: usize,
90 head_dim: usize,
91 scale: f32,
92}
93
94pub 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 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 if let Some(embed_weights) = weights.get("transformer.word_embeddings.weight") {
151 model.word_embeddings = embed_weights.clone();
152 }
153
154 if let Some(norm_weights) = weights.get("transformer.ln_f.weight") {
156 model.ln_f = norm_weights.clone();
157 }
158
159 if let Some(lm_head_weights) = weights.get("lm_head.weight") {
161 model.lm_head = ops_fn::transpose(lm_head_weights)?;
162 }
163
164 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 let mut hidden_states = ops_fn::embedding(input_ids, &self.word_embeddings)?;
177
178 for layer in &self.h {
180 hidden_states = layer.forward(&hidden_states, attention_mask.as_ref(), self.config.rope_theta)?;
181 }
182
183 hidden_states = ops_fn::layer_norm(&hidden_states, &self.ln_f, None, self.config.rms_norm_eps)?;
185
186 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 let tokenizer = Tokenizer::new();
212 let mut tokens: Vec<u32> = tokenizer.encode(prompt);
213
214 for _ in 0..config.max_new_tokens {
216 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 let outputs = self.forward(&inputs)?;
228
229 let logits = match outputs {
231 ModelOutputs::Logits { logits, .. } => logits,
232 _ => return Err(anyhow::anyhow!("Expected logits output")),
233 };
234
235 let logits_candle = logits.to_candle()?;
237 let shape = logits_candle.dims();
238
239 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 let logits_vec: Vec<f32> = last_logits.to_vec1()?;
255
256 let next_token = if config.do_sample && config.temperature > 0.0 {
257 let scaled: Vec<f32> = logits_vec.iter()
259 .map(|&x| x / config.temperature)
260 .collect();
261
262 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 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 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 if next_token == config.eos_token_id {
298 break;
299 }
300
301 tokens.push(next_token);
303 }
304
305 Ok(tokenizer.decode(&tokens))
307 }
308
309 fn config(&self) -> &Self::Config {
310 &self.config
311 }
312
313 fn memory_requirements(&self) -> MemoryRequirements {
314 let param_size = self.config.vocab_size * self.config.hidden_size + self.config.num_hidden_layers * (
317 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 2 * self.config.hidden_size * self.config.intermediate_size
322 );
323
324 let param_bytes = param_size * 4; 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 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 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 let attn_output = self.self_attention.forward(&normed, attention_mask, rope_theta)?;
381 let mlp_output = self.mlp.forward(&normed)?;
382
383 let combined = ops_fn::add(&attn_output, &mlp_output)?;
385 ops_fn::add(&residual, &combined)
386 } else {
387 let attn_output = self.self_attention.forward(&normed, attention_mask, rope_theta)?;
389 let hidden_states = ops_fn::add(&residual, &attn_output)?;
390
391 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 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 if let Some(ln) = weights.get(&format!("{}.input_layernorm.weight", prefix)) {
409 self.input_layernorm = ln.clone();
410 }
411
412 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 self.self_attention.load_weights(weights, layer_idx)?;
419
420 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
437fn 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 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 let positions: Vec<f32> = (0..seq_len).map(|p| p as f32).collect();
457
458 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 let cos = angles_tensor.cos()?;
470 let sin = angles_tensor.sin()?;
471
472 let cos = cos.unsqueeze(0)?.unsqueeze(0)?;
474 let sin = sin.unsqueeze(0)?.unsqueeze(0)?;
475
476 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 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 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 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 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 let qkv = ops_fn::matmul(hidden_states, &self.query_key_value)?;
541 let qkv_candle = qkv.to_candle()?;
542
543 let q_size = self.num_heads * self.head_dim;
546 let kv_size = self.num_key_value_heads * self.head_dim;
547
548 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 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 let (q_with_rope, k_with_rope) = apply_rope(&q_reshaped, &k_reshaped, seq_len, self.head_dim, rope_theta)?;
570
571 let num_groups = self.num_heads / self.num_key_value_heads;
573 let (k_expanded, v_expanded) = if num_groups > 1 {
574 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 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 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 let attention_weights = candle_nn::ops::softmax_last_dim(&masked_scores)?;
616
617 let v_contiguous = v_expanded.contiguous()?;
619 let attn_output = attention_weights.matmul(&v_contiguous)?;
620
621 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 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 if let Some(qkv) = weights.get(&format!("{}.query_key_value.weight", prefix)) {
640 self.query_key_value = ops_fn::transpose(qkv)?;
641 }
642
643 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 let intermediate = ops_fn::matmul(hidden_states, &self.dense_h_to_4h)?;
681
682 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 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 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, ..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, ..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 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, ..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 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, ..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}