1use crate::model_config;
9use super::traits::*;
10use anyhow::Result;
11use serde::{Serialize, Deserialize};
12
13model_config!(GemmaConfig {
15 vocab_size: usize = 256000,
16 hidden_size: usize = 3072,
17 intermediate_size: usize = 24576,
18 num_hidden_layers: usize = 28,
19 num_attention_heads: usize = 16,
20 num_key_value_heads: usize = 16,
21 hidden_act: String = "gelu".to_string(),
22 max_position_embeddings: usize = 8192,
23 initializer_range: f32 = 0.02,
24 rms_norm_eps: f32 = 1e-6,
25 use_cache: bool = true,
26 pad_token_id: i64 = 0,
27 bos_token_id: i64 = 2,
28 eos_token_id: i64 = 1,
29 tie_word_embeddings: bool = true,
30 rope_theta: f32 = 10000.0,
31 attention_bias: bool = false,
32 head_dim: usize = 256,
33});
34
35impl GemmaConfig {
36 pub fn from_gguf_config(gguf: &crate::weight_loader_core::GGUFModelConfig) -> Self {
38 Self {
39 vocab_size: gguf.vocab_size,
40 hidden_size: gguf.hidden_size,
41 intermediate_size: gguf.intermediate_size,
42 num_hidden_layers: gguf.num_hidden_layers,
43 num_attention_heads: gguf.num_attention_heads,
44 num_key_value_heads: gguf.num_key_value_heads,
45 rms_norm_eps: gguf.rms_norm_eps,
46 rope_theta: gguf.rope_theta,
47 max_position_embeddings: gguf.max_position_embeddings,
48 head_dim: gguf.head_dim,
49 ..Default::default()
50 }
51 }
52}
53
54pub struct GemmaModelV2 {
56 config: GemmaConfig,
57 device: Device,
58
59 embed_tokens: Tensor,
61 layers: Vec<GemmaLayer>,
62 norm: Tensor,
63 lm_head: Tensor,
64}
65
66pub struct GemmaLayer {
68 self_attn: GemmaAttention,
69 mlp: GemmaMLP,
70 input_layernorm: Tensor,
71 post_attention_layernorm: Tensor,
72}
73
74pub struct GemmaAttention {
76 q_proj: Tensor,
77 k_proj: Tensor,
78 v_proj: Tensor,
79 o_proj: Tensor,
80 num_heads: usize,
81 num_key_value_heads: usize,
82 head_dim: usize,
83 scale: f32,
84}
85
86pub struct GemmaMLP {
88 gate_proj: Tensor,
89 up_proj: Tensor,
90 down_proj: Tensor,
91 hidden_act: String,
92}
93
94impl Model for GemmaModelV2 {
95 type Config = GemmaConfig;
96
97 fn new(config: GemmaConfig) -> Result<Self> {
98 let device = Device::CPU;
99
100 let embed_tokens = ops_fn::zeros(
101 &[config.vocab_size, config.hidden_size],
102 DataType::Float32,
103 &device
104 )?;
105
106 let norm = ops_fn::zeros(
107 &[config.hidden_size],
108 DataType::Float32,
109 &device
110 )?;
111
112 let lm_head = if config.tie_word_embeddings {
114 embed_tokens.clone()
115 } else {
116 ops_fn::zeros(
117 &[config.hidden_size, config.vocab_size],
118 DataType::Float32,
119 &device
120 )?
121 };
122
123 let mut layers = Vec::with_capacity(config.num_hidden_layers);
125 for _ in 0..config.num_hidden_layers {
126 layers.push(GemmaLayer::new(&config, &device)?);
127 }
128
129 Ok(Self {
130 config,
131 device,
132 embed_tokens,
133 layers,
134 norm,
135 lm_head,
136 })
137 }
138
139 fn from_weights(config: GemmaConfig, weights: ModelWeights) -> Result<Self> {
140 let mut model = Self::new(config)?;
141
142 if let Some(embed_weights) = weights.get("model.embed_tokens.weight") {
144 model.embed_tokens = embed_weights.clone();
145 }
146
147 if let Some(norm_weights) = weights.get("model.norm.weight") {
148 model.norm = norm_weights.clone();
149 }
150
151 if !model.config.tie_word_embeddings {
154 if let Some(lm_head_weights) = weights.get("lm_head.weight") {
155 model.lm_head = ops_fn::transpose(lm_head_weights)?;
156 }
157 }
158
159 for (i, layer) in model.layers.iter_mut().enumerate() {
161 layer.load_weights(&weights, i)?;
162 }
163
164 Ok(model)
165 }
166
167 fn forward(&self, inputs: &ModelInputs) -> Result<ModelOutputs> {
168 match inputs {
169 ModelInputs::Text { input_ids, attention_mask, .. } => {
170 let mut hidden_states = ops_fn::embedding(input_ids, &self.embed_tokens)?;
172
173 let scale = (self.config.hidden_size as f32).sqrt();
175 hidden_states = ops_fn::scale(&hidden_states, scale)?;
176
177 for layer in &self.layers {
179 hidden_states = layer.forward(&hidden_states, attention_mask.as_ref(), self.config.rope_theta, self.config.head_dim)?;
180 }
181
182 hidden_states = ops_fn::layer_norm(&hidden_states, &self.norm, None, self.config.rms_norm_eps)?;
184
185 let logits = if self.config.tie_word_embeddings {
188 let lm_head_t = ops_fn::transpose(&self.embed_tokens)?;
189 ops_fn::matmul(&hidden_states, &lm_head_t)?
190 } else {
191 ops_fn::matmul(&hidden_states, &self.lm_head)?
192 };
193
194 Ok(ModelOutputs::Logits {
195 logits,
196 hidden_states: None,
197 })
198 }
199 ModelInputs::Multimodal { input_ids, .. } => {
200 let text_inputs = ModelInputs::Text {
201 input_ids: input_ids.clone(),
202 attention_mask: None,
203 position_ids: None,
204 };
205 self.forward(&text_inputs)
206 }
207 _ => Err(anyhow::anyhow!("Gemma model only supports text and multimodal inputs")),
208 }
209 }
210
211 fn generate(&self, prompt: &str, config: &GenerationConfig) -> Result<String> {
212 use crate::tokenizer::Tokenizer;
213 use rand::Rng;
214
215 let tokenizer = Tokenizer::new();
216 let mut tokens: Vec<u32> = tokenizer.encode(prompt);
217
218 for _ in 0..config.max_new_tokens {
219 let tokens_i64: Vec<i64> = tokens.iter().map(|&t| t as i64).collect();
220 let input_tensor = Tensor::from_i64_slice(&tokens_i64, &[1, tokens.len()], &self.device)?;
221
222 let inputs = ModelInputs::Text {
223 input_ids: input_tensor,
224 attention_mask: None,
225 position_ids: None,
226 };
227
228 let outputs = self.forward(&inputs)?;
229
230 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()?;
236 let shape = logits_candle.dims();
237
238 let last_logits = if shape.len() == 3 {
239 let seq_len = shape[1];
240 logits_candle
241 .narrow(1, seq_len - 1, 1)?
242 .squeeze(1)?
243 .squeeze(0)?
244 } else {
245 let seq_len = shape[0];
246 logits_candle
247 .narrow(0, seq_len - 1, 1)?
248 .squeeze(0)?
249 };
250
251 let logits_vec: Vec<f32> = last_logits.to_vec1()?;
252
253 let next_token = if config.do_sample && config.temperature > 0.0 {
254 let scaled: Vec<f32> = logits_vec.iter()
255 .map(|&x| x / config.temperature)
256 .collect();
257
258 let max_val = scaled.iter().cloned().fold(f32::NEG_INFINITY, f32::max);
259 let exp_sum: f32 = scaled.iter().map(|&x| (x - max_val).exp()).sum();
260 let probs: Vec<f32> = scaled.iter()
261 .map(|&x| (x - max_val).exp() / exp_sum)
262 .collect();
263
264 let mut rng = rand::thread_rng();
265 let random_val: f32 = rng.gen();
266 let mut cumulative = 0.0;
267 let mut sampled = 0u32;
268
269 for (idx, &prob) in probs.iter().enumerate() {
270 cumulative += prob;
271 if random_val <= cumulative {
272 sampled = idx as u32;
273 break;
274 }
275 }
276 sampled
277 } else {
278 let mut max_idx = 0;
279 let mut max_val = logits_vec[0];
280 for (idx, &val) in logits_vec.iter().enumerate() {
281 if val > max_val {
282 max_val = val;
283 max_idx = idx;
284 }
285 }
286 max_idx as u32
287 };
288
289 if next_token == config.eos_token_id {
290 break;
291 }
292
293 tokens.push(next_token);
294 }
295
296 Ok(tokenizer.decode(&tokens))
297 }
298
299 fn config(&self) -> &Self::Config {
300 &self.config
301 }
302
303 fn memory_requirements(&self) -> MemoryRequirements {
304 let param_size = self.config.vocab_size * self.config.hidden_size +
305 self.config.num_hidden_layers * (
306 4 * self.config.hidden_size * self.config.hidden_size +
307 3 * self.config.hidden_size * self.config.intermediate_size
308 );
309
310 let param_bytes = param_size * 4;
311 let kv_cache_bytes = 2 * self.config.num_hidden_layers *
312 self.config.max_position_embeddings *
313 self.config.hidden_size * 4;
314
315 MemoryRequirements {
316 gpu_memory: param_bytes,
317 cpu_memory: param_bytes / 4,
318 kv_cache_memory: kv_cache_bytes,
319 peak_memory: param_bytes + kv_cache_bytes,
320 }
321 }
322
323 fn to_device(&mut self, device: &Device) -> Result<()> {
324 self.embed_tokens = self.embed_tokens.to_device(device)?;
325 self.norm = self.norm.to_device(device)?;
326 self.lm_head = self.lm_head.to_device(device)?;
327
328 for layer in &mut self.layers {
329 layer.to_device(device)?;
330 }
331
332 self.device = device.clone();
333 Ok(())
334 }
335}
336
337impl GemmaLayer {
338 fn new(config: &GemmaConfig, device: &Device) -> Result<Self> {
339 let self_attn = GemmaAttention::new(config, device)?;
340 let mlp = GemmaMLP::new(config, device)?;
341
342 let input_layernorm = ops_fn::zeros(&[config.hidden_size], DataType::Float32, device)?;
343 let post_attention_layernorm = ops_fn::zeros(&[config.hidden_size], DataType::Float32, device)?;
344
345 Ok(Self {
346 self_attn,
347 mlp,
348 input_layernorm,
349 post_attention_layernorm,
350 })
351 }
352
353 fn forward(&self, hidden_states: &Tensor, attention_mask: Option<&Tensor>, rope_theta: f32, head_dim: usize) -> Result<Tensor> {
354 let normed = ops_fn::layer_norm(hidden_states, &self.input_layernorm, None, 1e-6)?;
356
357 let attn_output = self.self_attn.forward(&normed, attention_mask, rope_theta, head_dim)?;
359
360 let hidden_states = ops_fn::add(hidden_states, &attn_output)?;
362
363 let normed = ops_fn::layer_norm(&hidden_states, &self.post_attention_layernorm, None, 1e-6)?;
365
366 let mlp_output = self.mlp.forward(&normed)?;
368
369 let output = ops_fn::add(&hidden_states, &mlp_output)?;
371
372 Ok(output)
373 }
374
375 fn load_weights(&mut self, weights: &ModelWeights, layer_idx: usize) -> Result<()> {
376 let prefix = format!("model.layers.{}", layer_idx);
377
378 if let Some(q_proj) = weights.get(&format!("{}.self_attn.q_proj.weight", prefix)) {
380 self.self_attn.q_proj = ops_fn::transpose(q_proj)?;
381 }
382 if let Some(k_proj) = weights.get(&format!("{}.self_attn.k_proj.weight", prefix)) {
383 self.self_attn.k_proj = ops_fn::transpose(k_proj)?;
384 }
385 if let Some(v_proj) = weights.get(&format!("{}.self_attn.v_proj.weight", prefix)) {
386 self.self_attn.v_proj = ops_fn::transpose(v_proj)?;
387 }
388 if let Some(o_proj) = weights.get(&format!("{}.self_attn.o_proj.weight", prefix)) {
389 self.self_attn.o_proj = ops_fn::transpose(o_proj)?;
390 }
391
392 if let Some(gate_proj) = weights.get(&format!("{}.mlp.gate_proj.weight", prefix)) {
394 self.mlp.gate_proj = ops_fn::transpose(gate_proj)?;
395 }
396 if let Some(up_proj) = weights.get(&format!("{}.mlp.up_proj.weight", prefix)) {
397 self.mlp.up_proj = ops_fn::transpose(up_proj)?;
398 }
399 if let Some(down_proj) = weights.get(&format!("{}.mlp.down_proj.weight", prefix)) {
400 self.mlp.down_proj = ops_fn::transpose(down_proj)?;
401 }
402
403 if let Some(input_ln) = weights.get(&format!("{}.input_layernorm.weight", prefix)) {
405 self.input_layernorm = input_ln.clone();
406 }
407 if let Some(post_ln) = weights.get(&format!("{}.post_attention_layernorm.weight", prefix)) {
408 self.post_attention_layernorm = post_ln.clone();
409 }
410
411 Ok(())
412 }
413
414 fn to_device(&mut self, device: &Device) -> Result<()> {
415 self.self_attn.to_device(device)?;
416 self.mlp.to_device(device)?;
417 self.input_layernorm = self.input_layernorm.to_device(device)?;
418 self.post_attention_layernorm = self.post_attention_layernorm.to_device(device)?;
419 Ok(())
420 }
421}
422
423fn apply_rope(
425 q: &candle_core::Tensor,
426 k: &candle_core::Tensor,
427 seq_len: usize,
428 head_dim: usize,
429 rope_theta: f32,
430) -> Result<(candle_core::Tensor, candle_core::Tensor)> {
431 let device = q.device();
432
433 let half_dim = head_dim / 2;
434 let inv_freq: Vec<f32> = (0..half_dim)
435 .map(|i| 1.0 / rope_theta.powf((2 * i) as f32 / head_dim as f32))
436 .collect();
437
438 let positions: Vec<f32> = (0..seq_len).map(|p| p as f32).collect();
439
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 let cos = angles_tensor.cos()?;
450 let sin = angles_tensor.sin()?;
451
452 let cos = cos.unsqueeze(0)?.unsqueeze(0)?;
453 let sin = sin.unsqueeze(0)?.unsqueeze(0)?;
454
455 let q_half1 = q.narrow(3, 0, half_dim)?;
456 let q_half2 = q.narrow(3, half_dim, half_dim)?;
457 let k_half1 = k.narrow(3, 0, half_dim)?;
458 let k_half2 = k.narrow(3, half_dim, half_dim)?;
459
460 let q_rot1 = (q_half1.broadcast_mul(&cos)? - q_half2.broadcast_mul(&sin)?)?;
461 let q_rot2 = (q_half1.broadcast_mul(&sin)? + q_half2.broadcast_mul(&cos)?)?;
462 let k_rot1 = (k_half1.broadcast_mul(&cos)? - k_half2.broadcast_mul(&sin)?)?;
463 let k_rot2 = (k_half1.broadcast_mul(&sin)? + k_half2.broadcast_mul(&cos)?)?;
464
465 let q_rotated = candle_core::Tensor::cat(&[&q_rot1, &q_rot2], 3)?;
466 let k_rotated = candle_core::Tensor::cat(&[&k_rot1, &k_rot2], 3)?;
467
468 Ok((q_rotated, k_rotated))
469}
470
471impl GemmaAttention {
472 fn new(config: &GemmaConfig, device: &Device) -> Result<Self> {
473 let num_heads = config.num_attention_heads;
474 let num_key_value_heads = config.num_key_value_heads;
475 let head_dim = config.head_dim;
476 let scale = 1.0 / (head_dim as f32).sqrt();
477
478 let q_proj = ops_fn::zeros(&[config.hidden_size, num_heads * head_dim], DataType::Float32, device)?;
479 let k_proj = ops_fn::zeros(&[config.hidden_size, num_key_value_heads * head_dim], DataType::Float32, device)?;
480 let v_proj = ops_fn::zeros(&[config.hidden_size, num_key_value_heads * head_dim], DataType::Float32, device)?;
481 let o_proj = ops_fn::zeros(&[num_heads * head_dim, config.hidden_size], DataType::Float32, device)?;
482
483 Ok(Self {
484 q_proj,
485 k_proj,
486 v_proj,
487 o_proj,
488 num_heads,
489 num_key_value_heads,
490 head_dim,
491 scale,
492 })
493 }
494
495 fn forward(&self, hidden_states: &Tensor, _attention_mask: Option<&Tensor>, rope_theta: f32, head_dim: usize) -> Result<Tensor> {
496 let shape = hidden_states.shape();
497 let (batch_size, seq_len, _hidden_size) = if shape.len() == 3 {
498 (shape[0], shape[1], shape[2])
499 } else if shape.len() == 2 {
500 (1, shape[0], shape[1])
501 } else {
502 return Err(anyhow::anyhow!("Invalid hidden_states shape: {:?}", shape));
503 };
504
505 let query_states = ops_fn::matmul(hidden_states, &self.q_proj)?;
506 let key_states = ops_fn::matmul(hidden_states, &self.k_proj)?;
507 let value_states = ops_fn::matmul(hidden_states, &self.v_proj)?;
508
509 let q_candle = query_states.to_candle()?;
510 let k_candle = key_states.to_candle()?;
511 let v_candle = value_states.to_candle()?;
512
513 let q_reshaped = q_candle
514 .reshape(&[batch_size, seq_len, self.num_heads, head_dim])?
515 .transpose(1, 2)?;
516
517 let k_reshaped = k_candle
518 .reshape(&[batch_size, seq_len, self.num_key_value_heads, head_dim])?
519 .transpose(1, 2)?;
520
521 let v_reshaped = v_candle
522 .reshape(&[batch_size, seq_len, self.num_key_value_heads, head_dim])?
523 .transpose(1, 2)?;
524
525 let (q_with_rope, k_with_rope) = apply_rope(&q_reshaped, &k_reshaped, seq_len, head_dim, rope_theta)?;
526
527 let num_groups = self.num_heads / self.num_key_value_heads;
528 let (k_expanded, v_expanded) = if num_groups > 1 {
529 let k_rep = k_with_rope
530 .unsqueeze(2)?
531 .broadcast_as(&[batch_size, self.num_key_value_heads, num_groups, seq_len, head_dim])?
532 .reshape(&[batch_size, self.num_heads, seq_len, head_dim])?;
533 let v_rep = v_reshaped
534 .unsqueeze(2)?
535 .broadcast_as(&[batch_size, self.num_key_value_heads, num_groups, seq_len, head_dim])?
536 .reshape(&[batch_size, self.num_heads, seq_len, head_dim])?;
537 (k_rep, v_rep)
538 } else {
539 (k_with_rope, v_reshaped)
540 };
541
542 let k_t = k_expanded.transpose(2, 3)?;
543
544 let q_contiguous = q_with_rope.contiguous()?;
545 let k_contiguous = k_t.contiguous()?;
546
547 let scores = q_contiguous.matmul(&k_contiguous)?;
548 let scaled_scores = (scores * (self.scale as f64))?;
549
550 let device = scaled_scores.device();
551 let causal_mask = {
552 let mut mask_data = vec![0.0f32; seq_len * seq_len];
553 for i in 0..seq_len {
554 for j in 0..seq_len {
555 if j > i {
556 mask_data[i * seq_len + j] = f32::NEG_INFINITY;
557 }
558 }
559 }
560 candle_core::Tensor::from_vec(mask_data, &[1, 1, seq_len, seq_len], device)?
561 };
562
563 let masked_scores = scaled_scores.broadcast_add(&causal_mask)?;
564
565 let attention_weights = candle_nn::ops::softmax_last_dim(&masked_scores)?;
566
567 let v_contiguous = v_expanded.contiguous()?;
568 let attn_output = attention_weights.matmul(&v_contiguous)?;
569
570 let attn_output = attn_output
571 .transpose(1, 2)?
572 .reshape(&[batch_size, seq_len, self.num_heads * head_dim])?;
573
574 let attn_output = Tensor::from_candle(attn_output);
575
576 let output = ops_fn::matmul(&attn_output, &self.o_proj)?;
577
578 Ok(output)
579 }
580
581 fn to_device(&mut self, device: &Device) -> Result<()> {
582 self.q_proj = self.q_proj.to_device(device)?;
583 self.k_proj = self.k_proj.to_device(device)?;
584 self.v_proj = self.v_proj.to_device(device)?;
585 self.o_proj = self.o_proj.to_device(device)?;
586 Ok(())
587 }
588}
589
590impl GemmaMLP {
591 fn new(config: &GemmaConfig, device: &Device) -> Result<Self> {
592 let gate_proj = ops_fn::zeros(&[config.hidden_size, config.intermediate_size], DataType::Float32, device)?;
593 let up_proj = ops_fn::zeros(&[config.hidden_size, config.intermediate_size], DataType::Float32, device)?;
594 let down_proj = ops_fn::zeros(&[config.intermediate_size, config.hidden_size], DataType::Float32, device)?;
595
596 Ok(Self {
597 gate_proj,
598 up_proj,
599 down_proj,
600 hidden_act: config.hidden_act.clone(),
601 })
602 }
603
604 fn forward(&self, hidden_states: &Tensor) -> Result<Tensor> {
605 let gate_output = ops_fn::matmul(hidden_states, &self.gate_proj)?;
606 let up_output = ops_fn::matmul(hidden_states, &self.up_proj)?;
607
608 let gate_activated = match self.hidden_act.as_str() {
609 "gelu" | "gelu_new" => ops_fn::gelu(&gate_output)?,
610 "silu" | "swish" => ops_fn::silu(&gate_output)?,
611 _ => return Err(anyhow::anyhow!("Unsupported activation: {}", self.hidden_act)),
612 };
613
614 let gated = ops_fn::mul(&gate_activated, &up_output)?;
615
616 let output = ops_fn::matmul(&gated, &self.down_proj)?;
617
618 Ok(output)
619 }
620
621 fn to_device(&mut self, device: &Device) -> Result<()> {
622 self.gate_proj = self.gate_proj.to_device(device)?;
623 self.up_proj = self.up_proj.to_device(device)?;
624 self.down_proj = self.down_proj.to_device(device)?;
625 Ok(())
626 }
627}
628
629#[cfg(test)]
630mod tests {
631 use super::*;
632
633 #[test]
634 fn test_gemma_model_creation() {
635 let config = GemmaConfig {
636 vocab_size: 1000,
637 hidden_size: 128,
638 intermediate_size: 512,
639 num_hidden_layers: 2,
640 num_attention_heads: 8,
641 num_key_value_heads: 8,
642 head_dim: 16,
643 ..Default::default()
644 };
645
646 let model = GemmaModelV2::new(config).unwrap();
647 assert_eq!(model.config().vocab_size(), 1000);
648 assert_eq!(model.config().hidden_size(), 128);
649 assert_eq!(model.config().num_layers(), 2);
650 }
651
652 #[test]
653 fn test_gemma_forward_pass() {
654 let config = GemmaConfig {
655 vocab_size: 100,
656 hidden_size: 64,
657 intermediate_size: 256,
658 num_hidden_layers: 1,
659 num_attention_heads: 4,
660 num_key_value_heads: 4,
661 head_dim: 16,
662 ..Default::default()
663 };
664
665 let model = GemmaModelV2::new(config).unwrap();
666 let input_ids = ops_fn::zeros(&[2, 8], DataType::Int64, &Device::CPU).unwrap();
667 let inputs = ModelInputs::text(input_ids);
668
669 let outputs = model.forward(&inputs).unwrap();
670 match outputs {
671 ModelOutputs::Logits { logits, .. } => {
672 assert_eq!(logits.shape(), &[2, 8, 100]);
673 }
674 _ => panic!("Expected logits output"),
675 }
676 }
677
678 #[test]
679 fn test_gemma_generation() {
680 let config = GemmaConfig {
681 vocab_size: 256,
682 hidden_size: 64,
683 intermediate_size: 256,
684 num_hidden_layers: 1,
685 num_attention_heads: 4,
686 num_key_value_heads: 4,
687 head_dim: 16,
688 ..Default::default()
689 };
690 let model = GemmaModelV2::new(config).unwrap();
691 let gen_config = GenerationConfig {
692 max_new_tokens: 5,
693 ..Default::default()
694 };
695
696 let output = model.generate("Hello", &gen_config).unwrap();
697 assert!(!output.is_empty());
698 }
699}