1use crate::model_config;
9use super::traits::*;
10use anyhow::Result;
11use serde::{Serialize, Deserialize};
12
13model_config!(T5Config {
14 vocab_size: usize = 32128,
15 d_model: usize = 512,
16 d_kv: usize = 64,
17 d_ff: usize = 2048,
18 num_layers: usize = 6,
19 num_decoder_layers: usize = 6,
20 num_heads: usize = 8,
21 relative_attention_num_buckets: usize = 32,
22 relative_attention_max_distance: usize = 128,
23 dropout_rate: f32 = 0.1,
24 layer_norm_epsilon: f32 = 1e-6,
25 initializer_factor: f32 = 1.0,
26 feed_forward_proj: String = "relu".to_string(),
27 is_encoder_decoder: bool = true,
28 use_cache: bool = true,
29 pad_token_id: i64 = 0,
30 eos_token_id: i64 = 1,
31 decoder_start_token_id: i64 = 0,
32 tie_word_embeddings: bool = true,
33 is_gated_act: bool = false,
34 hidden_size: usize = 512,
36 num_hidden_layers: usize = 6,
37});
38
39impl T5Config {
40 pub fn from_gguf_config(gguf: &crate::weight_loader_core::GGUFModelConfig) -> Self {
42 Self {
43 vocab_size: gguf.vocab_size,
44 d_model: gguf.hidden_size,
45 hidden_size: gguf.hidden_size,
46 d_kv: gguf.head_dim,
47 d_ff: gguf.intermediate_size,
48 num_layers: gguf.num_hidden_layers,
49 num_decoder_layers: gguf.num_hidden_layers,
50 num_hidden_layers: gguf.num_hidden_layers,
51 num_heads: gguf.num_attention_heads,
52 layer_norm_epsilon: gguf.rms_norm_eps,
53 ..Default::default()
54 }
55 }
56}
57
58pub struct T5ModelV2 {
59 config: T5Config,
60 device: Device,
61 shared: Tensor, encoder: T5Stack,
63 decoder: T5Stack,
64 lm_head: Option<Tensor>,
65}
66
67pub struct T5Stack {
68 block: Vec<T5Block>,
69 final_layer_norm: Tensor,
70 config: T5Config,
71 is_decoder: bool,
72}
73
74pub struct T5Block {
75 self_attention: T5LayerSelfAttention,
76 cross_attention: Option<T5LayerCrossAttention>,
77 ff: T5LayerFF,
78 config: T5Config,
79 is_decoder: bool,
80}
81
82pub struct T5LayerSelfAttention {
83 attention: T5Attention,
84 layer_norm: Tensor,
85}
86
87pub struct T5LayerCrossAttention {
88 attention: T5Attention,
89 layer_norm: Tensor,
90}
91
92pub struct T5LayerFF {
93 wi: Tensor, wo: Tensor, wi_1: Option<Tensor>, layer_norm: Tensor,
97 is_gated: bool,
98 activation: String,
99}
100
101pub struct T5Attention {
102 q: Tensor,
103 k: Tensor,
104 v: Tensor,
105 o: Tensor,
106 relative_attention_bias: Option<Tensor>,
107 num_heads: usize,
108 d_kv: usize,
109 is_decoder: bool,
110 has_relative_attention_bias: bool,
111}
112
113impl Model for T5ModelV2 {
114 type Config = T5Config;
115
116 fn new(config: T5Config) -> Result<Self> {
117 let device = Device::CPU;
118 let shared = ops_fn::zeros(&[config.vocab_size, config.d_model], DataType::Float32, &device)?;
119
120 let encoder = T5Stack::new(&config, &device, false)?;
121 let decoder = T5Stack::new(&config, &device, true)?;
122
123 let lm_head = if config.tie_word_embeddings {
124 None } else {
126 Some(ops_fn::zeros(&[config.d_model, config.vocab_size], DataType::Float32, &device)?)
127 };
128
129 Ok(Self { config, device, shared, encoder, decoder, lm_head })
130 }
131
132 fn from_weights(config: T5Config, weights: ModelWeights) -> Result<Self> {
133 let mut model = Self::new(config)?;
134
135 if let Some(w) = weights.get("shared.weight") {
137 model.shared = w.clone();
138 } else if let Some(w) = weights.get("encoder.embed_tokens.weight") {
139 model.shared = w.clone();
140 }
141
142 if let Some(w) = weights.get("lm_head.weight") {
144 if model.lm_head.is_some() {
145 model.lm_head = Some(ops_fn::transpose(w)?);
146 }
147 }
148
149 model.encoder.load_weights(&weights, "encoder")?;
151 model.decoder.load_weights(&weights, "decoder")?;
152
153 Ok(model)
154 }
155
156 fn forward(&self, inputs: &ModelInputs) -> Result<ModelOutputs> {
157 let input_ids = match inputs {
158 ModelInputs::Text { input_ids, .. } => input_ids,
159 _ => return Err(anyhow::anyhow!("T5 expects text input")),
160 };
161
162 let encoder_hidden_states = self.encoder.forward(input_ids, &self.shared, None, None)?;
164
165 let decoder_input = input_ids;
167
168 let decoder_hidden_states = self.decoder.forward(
170 decoder_input,
171 &self.shared,
172 Some(&encoder_hidden_states),
173 None,
174 )?;
175
176 let logits = if let Some(ref lm_head) = self.lm_head {
178 ops_fn::matmul(&decoder_hidden_states, lm_head)?
179 } else {
180 let shared_t = ops_fn::transpose(&self.shared)?;
182 ops_fn::matmul(&decoder_hidden_states, &shared_t)?
183 };
184
185 Ok(ModelOutputs::Sequence {
186 logits,
187 encoder_hidden_states: Some(encoder_hidden_states),
188 decoder_hidden_states: Some(decoder_hidden_states),
189 })
190 }
191
192 fn generate(&self, prompt: &str, config: &GenerationConfig) -> Result<String> {
193 use crate::tokenizer::Tokenizer;
194
195 let tokenizer = Tokenizer::new();
197 let input_tokens: Vec<u32> = tokenizer.encode(prompt);
198
199 let input_i64: Vec<i64> = input_tokens.iter().map(|&t| t as i64).collect();
201 let input_tensor = Tensor::from_i64_slice(&input_i64, &[1, input_tokens.len()], &self.device)?;
202
203 let encoder_hidden_states = self.encoder.forward(&input_tensor, &self.shared, None, None)?;
205
206 let mut decoder_tokens: Vec<i64> = vec![self.config.decoder_start_token_id];
208
209 for _ in 0..config.max_new_tokens {
211 let decoder_tensor = Tensor::from_i64_slice(
213 &decoder_tokens,
214 &[1, decoder_tokens.len()],
215 &self.device,
216 )?;
217
218 let decoder_hidden_states = self.decoder.forward(
220 &decoder_tensor,
221 &self.shared,
222 Some(&encoder_hidden_states),
223 None,
224 )?;
225
226 let logits = if let Some(ref lm_head) = self.lm_head {
228 ops_fn::matmul(&decoder_hidden_states, lm_head)?
229 } else {
230 let shared_t = ops_fn::transpose(&self.shared)?;
231 ops_fn::matmul(&decoder_hidden_states, &shared_t)?
232 };
233
234 let logits_candle = logits.to_candle()?;
236 let shape = logits_candle.dims();
237 let seq_len = if shape.len() == 3 { shape[1] } else { shape[0] };
238
239 let last_logits = if shape.len() == 3 {
240 logits_candle.narrow(1, seq_len - 1, 1)?.squeeze(1)?.squeeze(0)?
241 } else {
242 logits_candle.narrow(0, seq_len - 1, 1)?.squeeze(0)?
243 };
244
245 let logits_vec: Vec<f32> = last_logits.to_vec1()?;
246
247 let next_token = if config.do_sample && config.temperature > 0.0 {
249 let scaled: Vec<f32> = logits_vec.iter()
251 .map(|&x| x / config.temperature)
252 .collect();
253
254 let max_val = scaled.iter().cloned().fold(f32::NEG_INFINITY, f32::max);
255 let exp_sum: f32 = scaled.iter().map(|&x| (x - max_val).exp()).sum();
256 let probs: Vec<f32> = scaled.iter()
257 .map(|&x| (x - max_val).exp() / exp_sum)
258 .collect();
259
260 use rand::Rng;
261 let mut rng = rand::thread_rng();
262 let random_val: f32 = rng.gen();
263 let mut cumulative = 0.0;
264 let mut sampled = 0i64;
265
266 for (idx, &prob) in probs.iter().enumerate() {
267 cumulative += prob;
268 if random_val <= cumulative {
269 sampled = idx as i64;
270 break;
271 }
272 }
273 sampled
274 } else {
275 let mut max_idx = 0;
277 let mut max_val = logits_vec[0];
278 for (idx, &val) in logits_vec.iter().enumerate() {
279 if val > max_val {
280 max_val = val;
281 max_idx = idx;
282 }
283 }
284 max_idx as i64
285 };
286
287 if next_token == self.config.eos_token_id {
289 break;
290 }
291
292 decoder_tokens.push(next_token);
293 }
294
295 let output_tokens: Vec<u32> = decoder_tokens.iter()
297 .skip(1) .map(|&t| t as u32)
299 .collect();
300
301 Ok(tokenizer.decode(&output_tokens))
302 }
303
304 fn config(&self) -> &Self::Config {
305 &self.config
306 }
307
308 fn memory_requirements(&self) -> MemoryRequirements {
309 let param_size = (self.config.vocab_size * self.config.d_model +
310 self.config.num_layers * 2 * self.config.d_model * self.config.d_model * 4 +
311 self.config.num_decoder_layers * 2 * self.config.d_model * self.config.d_model * 4) * 4;
312 MemoryRequirements {
313 gpu_memory: param_size,
314 cpu_memory: param_size / 4,
315 kv_cache_memory: 2048 * self.config.d_model * 4 * 4, peak_memory: param_size + param_size / 2,
317 }
318 }
319
320 fn to_device(&mut self, device: &Device) -> Result<()> {
321 self.device = device.clone();
322 self.shared = self.shared.to_device(device)?;
323 if let Some(ref mut lm_head) = self.lm_head {
324 *lm_head = lm_head.to_device(device)?;
325 }
326 self.encoder.to_device(device)?;
327 self.decoder.to_device(device)?;
328 Ok(())
329 }
330}
331
332impl T5Stack {
333 fn new(config: &T5Config, device: &Device, is_decoder: bool) -> Result<Self> {
334 let num_layers = if is_decoder {
335 config.num_decoder_layers
336 } else {
337 config.num_layers
338 };
339
340 let mut block = Vec::with_capacity(num_layers);
341 for i in 0..num_layers {
342 let has_relative_attention_bias = i == 0;
344 block.push(T5Block::new(config, device, is_decoder, has_relative_attention_bias)?);
345 }
346
347 let final_layer_norm = ops_fn::zeros(&[config.d_model], DataType::Float32, device)?;
348
349 Ok(Self {
350 block,
351 final_layer_norm,
352 config: config.clone(),
353 is_decoder,
354 })
355 }
356
357 fn forward(
358 &self,
359 input_ids: &Tensor,
360 shared_embedding: &Tensor,
361 encoder_hidden_states: Option<&Tensor>,
362 attention_mask: Option<&Tensor>,
363 ) -> Result<Tensor> {
364 let mut hidden_states = ops_fn::embedding(input_ids, shared_embedding)?;
366
367 let shape = hidden_states.shape();
369 let seq_len = if shape.len() == 3 { shape[1] } else { shape[0] };
370 let position_bias = self.compute_position_bias(seq_len, seq_len)?;
371
372 let cross_position_bias = if self.is_decoder {
374 if let Some(enc_hidden) = encoder_hidden_states {
375 let enc_shape = enc_hidden.shape();
376 let enc_seq_len = if enc_shape.len() == 3 { enc_shape[1] } else { enc_shape[0] };
377 Some(self.compute_position_bias(seq_len, enc_seq_len)?)
378 } else {
379 None
380 }
381 } else {
382 None
383 };
384
385 for layer in &self.block {
387 hidden_states = layer.forward(
388 &hidden_states,
389 encoder_hidden_states,
390 &position_bias,
391 cross_position_bias.as_ref(),
392 attention_mask,
393 )?;
394 }
395
396 ops_fn::layer_norm(&hidden_states, &self.final_layer_norm, None, self.config.layer_norm_epsilon)
398 }
399
400 fn compute_position_bias(&self, query_length: usize, key_length: usize) -> Result<Tensor> {
402 ops_fn::zeros(
405 &[1, self.config.num_heads, query_length, key_length],
406 DataType::Float32,
407 &Device::CPU,
408 )
409 }
410
411 fn load_weights(&mut self, weights: &ModelWeights, prefix: &str) -> Result<()> {
412 if let Some(w) = weights.get(&format!("{}.final_layer_norm.weight", prefix)) {
414 self.final_layer_norm = w.clone();
415 }
416
417 for (i, block) in self.block.iter_mut().enumerate() {
419 block.load_weights(weights, &format!("{}.block.{}", prefix, i))?;
420 }
421
422 Ok(())
423 }
424
425 fn to_device(&mut self, device: &Device) -> Result<()> {
426 self.final_layer_norm = self.final_layer_norm.to_device(device)?;
427 for block in &mut self.block {
428 block.to_device(device)?;
429 }
430 Ok(())
431 }
432}
433
434impl T5Block {
435 fn new(config: &T5Config, device: &Device, is_decoder: bool, has_relative_attention_bias: bool) -> Result<Self> {
436 let self_attention = T5LayerSelfAttention::new(config, device, is_decoder, has_relative_attention_bias)?;
438
439 let cross_attention = if is_decoder {
441 Some(T5LayerCrossAttention::new(config, device, has_relative_attention_bias)?)
442 } else {
443 None
444 };
445
446 let ff = T5LayerFF::new(config, device)?;
448
449 Ok(Self {
450 self_attention,
451 cross_attention,
452 ff,
453 config: config.clone(),
454 is_decoder,
455 })
456 }
457
458 fn forward(
459 &self,
460 hidden_states: &Tensor,
461 encoder_hidden_states: Option<&Tensor>,
462 position_bias: &Tensor,
463 cross_position_bias: Option<&Tensor>,
464 attention_mask: Option<&Tensor>,
465 ) -> Result<Tensor> {
466 let attn_output = self.self_attention.forward(
468 hidden_states,
469 hidden_states,
470 position_bias,
471 attention_mask,
472 )?;
473 let hidden_states = ops_fn::add(hidden_states, &attn_output)?;
474
475 let hidden_states = if let (Some(cross_attn), Some(enc_hidden)) = (&self.cross_attention, encoder_hidden_states) {
477 let cross_output = cross_attn.forward(
478 &hidden_states,
479 enc_hidden,
480 cross_position_bias.unwrap_or(position_bias),
481 None, )?;
483 ops_fn::add(&hidden_states, &cross_output)?
484 } else {
485 hidden_states
486 };
487
488 let ff_output = self.ff.forward(&hidden_states)?;
490 ops_fn::add(&hidden_states, &ff_output)
491 }
492
493 fn load_weights(&mut self, weights: &ModelWeights, prefix: &str) -> Result<()> {
494 self.self_attention.load_weights(weights, &format!("{}.layer.0", prefix))?;
496
497 if let Some(ref mut cross_attn) = self.cross_attention {
499 cross_attn.load_weights(weights, &format!("{}.layer.1", prefix))?;
500 }
501
502 let ff_layer_idx = if self.is_decoder { 2 } else { 1 };
504 self.ff.load_weights(weights, &format!("{}.layer.{}", prefix, ff_layer_idx))?;
505
506 Ok(())
507 }
508
509 fn to_device(&mut self, device: &Device) -> Result<()> {
510 self.self_attention.to_device(device)?;
511 if let Some(ref mut cross_attn) = self.cross_attention {
512 cross_attn.to_device(device)?;
513 }
514 self.ff.to_device(device)?;
515 Ok(())
516 }
517}
518
519impl T5LayerSelfAttention {
520 fn new(config: &T5Config, device: &Device, is_decoder: bool, has_relative_attention_bias: bool) -> Result<Self> {
521 let attention = T5Attention::new(config, device, is_decoder, has_relative_attention_bias)?;
522 let layer_norm = ops_fn::zeros(&[config.d_model], DataType::Float32, device)?;
523
524 Ok(Self { attention, layer_norm })
525 }
526
527 fn forward(
528 &self,
529 hidden_states: &Tensor,
530 key_value_states: &Tensor,
531 position_bias: &Tensor,
532 attention_mask: Option<&Tensor>,
533 ) -> Result<Tensor> {
534 let normed = ops_fn::layer_norm(hidden_states, &self.layer_norm, None, 1e-6)?;
536
537 self.attention.forward(&normed, key_value_states, position_bias, attention_mask)
539 }
540
541 fn load_weights(&mut self, weights: &ModelWeights, prefix: &str) -> Result<()> {
542 if let Some(w) = weights.get(&format!("{}.layer_norm.weight", prefix)) {
544 self.layer_norm = w.clone();
545 }
546
547 self.attention.load_weights(weights, &format!("{}.SelfAttention", prefix))?;
549
550 Ok(())
551 }
552
553 fn to_device(&mut self, device: &Device) -> Result<()> {
554 self.layer_norm = self.layer_norm.to_device(device)?;
555 self.attention.to_device(device)?;
556 Ok(())
557 }
558}
559
560impl T5LayerCrossAttention {
561 fn new(config: &T5Config, device: &Device, has_relative_attention_bias: bool) -> Result<Self> {
562 let attention = T5Attention::new(config, device, false, has_relative_attention_bias)?;
564 let layer_norm = ops_fn::zeros(&[config.d_model], DataType::Float32, device)?;
565
566 Ok(Self { attention, layer_norm })
567 }
568
569 fn forward(
570 &self,
571 hidden_states: &Tensor,
572 key_value_states: &Tensor,
573 position_bias: &Tensor,
574 attention_mask: Option<&Tensor>,
575 ) -> Result<Tensor> {
576 let normed = ops_fn::layer_norm(hidden_states, &self.layer_norm, None, 1e-6)?;
578
579 self.attention.forward(&normed, key_value_states, position_bias, attention_mask)
581 }
582
583 fn load_weights(&mut self, weights: &ModelWeights, prefix: &str) -> Result<()> {
584 if let Some(w) = weights.get(&format!("{}.layer_norm.weight", prefix)) {
586 self.layer_norm = w.clone();
587 }
588
589 self.attention.load_weights(weights, &format!("{}.EncDecAttention", prefix))?;
591
592 Ok(())
593 }
594
595 fn to_device(&mut self, device: &Device) -> Result<()> {
596 self.layer_norm = self.layer_norm.to_device(device)?;
597 self.attention.to_device(device)?;
598 Ok(())
599 }
600}
601
602impl T5LayerFF {
603 fn new(config: &T5Config, device: &Device) -> Result<Self> {
604 let is_gated = config.is_gated_act ||
605 config.feed_forward_proj.contains("gated");
606
607 let wi = ops_fn::zeros(&[config.d_model, config.d_ff], DataType::Float32, device)?;
608 let wo = ops_fn::zeros(&[config.d_ff, config.d_model], DataType::Float32, device)?;
609
610 let wi_1 = if is_gated {
611 Some(ops_fn::zeros(&[config.d_model, config.d_ff], DataType::Float32, device)?)
612 } else {
613 None
614 };
615
616 let layer_norm = ops_fn::zeros(&[config.d_model], DataType::Float32, device)?;
617
618 let activation = if config.feed_forward_proj.contains("gelu") {
620 "gelu".to_string()
621 } else {
622 "relu".to_string()
623 };
624
625 Ok(Self {
626 wi,
627 wo,
628 wi_1,
629 layer_norm,
630 is_gated,
631 activation,
632 })
633 }
634
635 fn forward(&self, hidden_states: &Tensor) -> Result<Tensor> {
636 let normed = ops_fn::layer_norm(hidden_states, &self.layer_norm, None, 1e-6)?;
638
639 let hidden = if self.is_gated {
641 let gate = ops_fn::matmul(&normed, &self.wi)?;
643 let up = ops_fn::matmul(&normed, self.wi_1.as_ref().unwrap())?;
644
645 let activated = match self.activation.as_str() {
646 "gelu" => ops_fn::gelu(&gate)?,
647 _ => relu(&gate)?,
648 };
649
650 ops_fn::mul(&activated, &up)?
651 } else {
652 let up = ops_fn::matmul(&normed, &self.wi)?;
654 match self.activation.as_str() {
655 "gelu" => ops_fn::gelu(&up)?,
656 _ => relu(&up)?,
657 }
658 };
659
660 ops_fn::matmul(&hidden, &self.wo)
662 }
663
664 fn load_weights(&mut self, weights: &ModelWeights, prefix: &str) -> Result<()> {
665 if let Some(w) = weights.get(&format!("{}.layer_norm.weight", prefix)) {
667 self.layer_norm = w.clone();
668 }
669
670 if let Some(w) = weights.get(&format!("{}.DenseReluDense.wi.weight", prefix)) {
672 self.wi = ops_fn::transpose(w)?;
673 }
674 if let Some(w) = weights.get(&format!("{}.DenseReluDense.wi_0.weight", prefix)) {
676 self.wi = ops_fn::transpose(w)?;
677 }
678 if let Some(w) = weights.get(&format!("{}.DenseReluDense.wi_1.weight", prefix)) {
679 self.wi_1 = Some(ops_fn::transpose(w)?);
680 }
681 if let Some(w) = weights.get(&format!("{}.DenseReluDense.wo.weight", prefix)) {
682 self.wo = ops_fn::transpose(w)?;
683 }
684
685 Ok(())
686 }
687
688 fn to_device(&mut self, device: &Device) -> Result<()> {
689 self.wi = self.wi.to_device(device)?;
690 self.wo = self.wo.to_device(device)?;
691 if let Some(ref mut wi_1) = self.wi_1 {
692 *wi_1 = wi_1.to_device(device)?;
693 }
694 self.layer_norm = self.layer_norm.to_device(device)?;
695 Ok(())
696 }
697}
698
699impl T5Attention {
700 fn new(config: &T5Config, device: &Device, is_decoder: bool, has_relative_attention_bias: bool) -> Result<Self> {
701 let inner_dim = config.num_heads * config.d_kv;
702
703 let q = ops_fn::zeros(&[config.d_model, inner_dim], DataType::Float32, device)?;
705 let k = ops_fn::zeros(&[config.d_model, inner_dim], DataType::Float32, device)?;
706 let v = ops_fn::zeros(&[config.d_model, inner_dim], DataType::Float32, device)?;
707 let o = ops_fn::zeros(&[inner_dim, config.d_model], DataType::Float32, device)?;
708
709 let relative_attention_bias = if has_relative_attention_bias {
711 Some(ops_fn::zeros(
712 &[config.relative_attention_num_buckets, config.num_heads],
713 DataType::Float32,
714 device,
715 )?)
716 } else {
717 None
718 };
719
720 Ok(Self {
721 q,
722 k,
723 v,
724 o,
725 relative_attention_bias,
726 num_heads: config.num_heads,
727 d_kv: config.d_kv,
728 is_decoder,
729 has_relative_attention_bias,
730 })
731 }
732
733 fn forward(
734 &self,
735 hidden_states: &Tensor,
736 key_value_states: &Tensor,
737 position_bias: &Tensor,
738 attention_mask: Option<&Tensor>,
739 ) -> Result<Tensor> {
740 let shape = hidden_states.shape();
741 let (batch_size, seq_len) = if shape.len() == 3 {
742 (shape[0], shape[1])
743 } else {
744 (1, shape[0])
745 };
746
747 let kv_shape = key_value_states.shape();
748 let kv_seq_len = if kv_shape.len() == 3 { kv_shape[1] } else { kv_shape[0] };
749
750 let query = ops_fn::matmul(hidden_states, &self.q)?;
752 let key = ops_fn::matmul(key_value_states, &self.k)?;
753 let value = ops_fn::matmul(key_value_states, &self.v)?;
754
755 let q_candle = query.to_candle()?;
757 let k_candle = key.to_candle()?;
758 let v_candle = value.to_candle()?;
759
760 let q_reshaped = q_candle
761 .reshape(&[batch_size, seq_len, self.num_heads, self.d_kv])?
762 .transpose(1, 2)?;
763
764 let k_reshaped = k_candle
765 .reshape(&[batch_size, kv_seq_len, self.num_heads, self.d_kv])?
766 .transpose(1, 2)?;
767
768 let v_reshaped = v_candle
769 .reshape(&[batch_size, kv_seq_len, self.num_heads, self.d_kv])?
770 .transpose(1, 2)?;
771
772 let k_t = k_reshaped.transpose(2, 3)?;
774 let scores = q_reshaped.contiguous()?.matmul(&k_t.contiguous()?)?;
775
776 let position_bias_candle = position_bias.to_candle()?;
778 let scores = scores.broadcast_add(&position_bias_candle)?;
779
780 let scores = if self.is_decoder && seq_len == kv_seq_len {
782 let device = scores.device();
784 let mut mask_data = vec![0.0f32; seq_len * kv_seq_len];
785 for i in 0..seq_len {
786 for j in 0..kv_seq_len {
787 if j > i {
788 mask_data[i * kv_seq_len + j] = f32::NEG_INFINITY;
789 }
790 }
791 }
792 let causal_mask = candle_core::Tensor::from_vec(mask_data, &[1, 1, seq_len, kv_seq_len], device)?;
793 scores.broadcast_add(&causal_mask)?
794 } else {
795 scores
796 };
797
798 let scores = if let Some(mask) = attention_mask {
800 let mask_candle = mask.to_candle()?;
801 scores.broadcast_add(&mask_candle)?
802 } else {
803 scores
804 };
805
806 let attention_weights = candle_nn::ops::softmax_last_dim(&scores)?;
808
809 let attn_output = attention_weights.matmul(&v_reshaped.contiguous()?)?;
811
812 let attn_output = attn_output
814 .transpose(1, 2)?
815 .reshape(&[batch_size, seq_len, self.num_heads * self.d_kv])?;
816
817 let attn_output = Tensor::from_candle(attn_output);
818
819 ops_fn::matmul(&attn_output, &self.o)
821 }
822
823 fn load_weights(&mut self, weights: &ModelWeights, prefix: &str) -> Result<()> {
824 if let Some(w) = weights.get(&format!("{}.q.weight", prefix)) {
826 self.q = ops_fn::transpose(w)?;
827 }
828 if let Some(w) = weights.get(&format!("{}.k.weight", prefix)) {
829 self.k = ops_fn::transpose(w)?;
830 }
831 if let Some(w) = weights.get(&format!("{}.v.weight", prefix)) {
832 self.v = ops_fn::transpose(w)?;
833 }
834 if let Some(w) = weights.get(&format!("{}.o.weight", prefix)) {
835 self.o = ops_fn::transpose(w)?;
836 }
837
838 if let Some(ref mut bias) = self.relative_attention_bias {
840 if let Some(w) = weights.get(&format!("{}.relative_attention_bias.weight", prefix)) {
841 *bias = w.clone();
842 }
843 }
844
845 Ok(())
846 }
847
848 fn to_device(&mut self, device: &Device) -> Result<()> {
849 self.q = self.q.to_device(device)?;
850 self.k = self.k.to_device(device)?;
851 self.v = self.v.to_device(device)?;
852 self.o = self.o.to_device(device)?;
853
854 if let Some(ref mut bias) = self.relative_attention_bias {
855 *bias = bias.to_device(device)?;
856 }
857
858 Ok(())
859 }
860}
861
862fn relu(input: &Tensor) -> Result<Tensor> {
864 let x = input.to_candle()?;
865 let result = x.relu()?;
866 Ok(Tensor::from_candle(result))
867}
868
869#[allow(dead_code)]
872fn relative_position_bucket(
873 relative_position: i32,
874 bidirectional: bool,
875 num_buckets: usize,
876 max_distance: usize,
877) -> usize {
878 let mut relative_buckets = 0usize;
879 let mut relative_position = relative_position;
880
881 if bidirectional {
882 let num_buckets = num_buckets / 2;
883 if relative_position > 0 {
884 relative_buckets = num_buckets;
885 } else {
886 relative_position = -relative_position;
887 }
888 } else {
889 relative_position = (-relative_position).max(0);
890 }
891
892 let relative_position = relative_position as usize;
893
894 let max_exact = num_buckets / 2;
896
897 if relative_position < max_exact {
898 relative_buckets + relative_position
899 } else {
900 let relative_position_if_large = max_exact +
902 ((relative_position as f32 / max_exact as f32).ln() /
903 (max_distance as f32 / max_exact as f32).ln() *
904 (num_buckets - max_exact) as f32) as usize;
905 relative_buckets + relative_position_if_large.min(num_buckets - 1)
906 }
907}
908
909#[cfg(test)]
910mod tests {
911 use super::*;
912
913 #[test]
914 fn test_t5_model_creation() {
915 let config = T5Config {
916 vocab_size: 1000,
917 d_model: 64,
918 hidden_size: 64,
919 d_kv: 16,
920 d_ff: 256,
921 num_layers: 2,
922 num_decoder_layers: 2,
923 num_hidden_layers: 2,
924 num_heads: 4,
925 ..Default::default()
926 };
927
928 let model = T5ModelV2::new(config).unwrap();
929 assert_eq!(model.config().vocab_size(), 1000);
930 assert_eq!(model.config().hidden_size(), 64);
931 assert_eq!(model.config().num_layers(), 2);
932 }
933
934 #[test]
935 fn test_t5_forward_pass() {
936 let config = T5Config {
937 vocab_size: 100,
938 d_model: 32,
939 hidden_size: 32,
940 d_kv: 8,
941 d_ff: 64,
942 num_layers: 1,
943 num_decoder_layers: 1,
944 num_hidden_layers: 1,
945 num_heads: 4,
946 ..Default::default()
947 };
948
949 let model = T5ModelV2::new(config).unwrap();
950 let input_ids = ops_fn::zeros(&[2, 8], DataType::Int64, &Device::CPU).unwrap();
951 let inputs = ModelInputs::text(input_ids);
952
953 let outputs = model.forward(&inputs).unwrap();
954 match outputs {
955 ModelOutputs::Sequence { logits, encoder_hidden_states, decoder_hidden_states } => {
956 assert_eq!(logits.shape(), &[2, 8, 100]); assert!(encoder_hidden_states.is_some());
958 assert!(decoder_hidden_states.is_some());
959 }
960 _ => panic!("Expected sequence output"),
961 }
962 }
963
964 #[test]
965 fn test_t5_generation() {
966 let config = T5Config {
967 vocab_size: 256,
968 d_model: 32,
969 hidden_size: 32,
970 d_kv: 8,
971 d_ff: 64,
972 num_layers: 1,
973 num_decoder_layers: 1,
974 num_hidden_layers: 1,
975 num_heads: 4,
976 ..Default::default()
977 };
978 let model = T5ModelV2::new(config).unwrap();
979 let gen_config = GenerationConfig {
980 max_new_tokens: 5,
981 ..Default::default()
982 };
983
984 let output = model.generate("Hello", &gen_config).unwrap();
985 let _ = output;
988 }
989
990 #[test]
991 fn test_relative_position_bucket() {
992 let bucket = relative_position_bucket(0, true, 32, 128);
994 assert_eq!(bucket, 0);
995
996 let bucket = relative_position_bucket(1, true, 32, 128);
997 assert!(bucket > 0);
998
999 let bucket = relative_position_bucket(-1, true, 32, 128);
1000 assert!(bucket < 16); let bucket = relative_position_bucket(0, false, 32, 128);
1004 assert_eq!(bucket, 0);
1005 }
1006}