Skip to main content

runtime/models_v2/
t5.rs

1//! T5 Model V2 - Clean implementation using solid abstractions
2//!
3//! This implements the T5 (Text-to-Text Transfer Transformer) architecture including:
4//! - T5-small, T5-base, T5-large, T5-3B, T5-11B
5//! - Encoder-decoder architecture with relative position embeddings
6//! - Bidirectional attention in encoder, causal + cross-attention in decoder
7
8use 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    // Required by model_config macro for ModelConfig trait
35    hidden_size: usize = 512,
36    num_hidden_layers: usize = 6,
37});
38
39impl T5Config {
40    /// Create T5Config from GGUF model configuration
41    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, // Shared embedding table
62    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,      // Input projection
94    wo: Tensor,      // Output projection
95    wi_1: Option<Tensor>, // For gated activations (T5 v1.1)
96    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 // Will use shared embeddings
125        } 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        // Load shared embeddings
136        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        // Load lm_head if not tied
143        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        // Load encoder and decoder weights
150        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        // Encoder forward pass
163        let encoder_hidden_states = self.encoder.forward(input_ids, &self.shared, None, None)?;
164
165        // For simple forward, use same input for decoder (in practice, decoder input would be shifted)
166        let decoder_input = input_ids;
167
168        // Decoder forward pass with cross-attention to encoder outputs
169        let decoder_hidden_states = self.decoder.forward(
170            decoder_input,
171            &self.shared,
172            Some(&encoder_hidden_states),
173            None,
174        )?;
175
176        // Language modeling head: project to vocabulary
177        let logits = if let Some(ref lm_head) = self.lm_head {
178            ops_fn::matmul(&decoder_hidden_states, lm_head)?
179        } else {
180            // Tied embeddings: use transposed shared embeddings
181            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        // 1. Tokenize input prompt
196        let tokenizer = Tokenizer::new();
197        let input_tokens: Vec<u32> = tokenizer.encode(prompt);
198
199        // 2. Create encoder input tensor
200        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        // 3. Run encoder once to get hidden states
204        let encoder_hidden_states = self.encoder.forward(&input_tensor, &self.shared, None, None)?;
205
206        // 4. Initialize decoder input with decoder_start_token_id
207        let mut decoder_tokens: Vec<i64> = vec![self.config.decoder_start_token_id];
208
209        // 5. Autoregressive generation loop
210        for _ in 0..config.max_new_tokens {
211            // Create decoder input tensor
212            let decoder_tensor = Tensor::from_i64_slice(
213                &decoder_tokens,
214                &[1, decoder_tokens.len()],
215                &self.device,
216            )?;
217
218            // Run decoder with cached encoder hidden states
219            let decoder_hidden_states = self.decoder.forward(
220                &decoder_tensor,
221                &self.shared,
222                Some(&encoder_hidden_states),
223                None,
224            )?;
225
226            // Get logits for last position
227            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            // Extract last token logits and sample
235            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            // Greedy sampling (can be extended with temperature, top-k, top-p)
248            let next_token = if config.do_sample && config.temperature > 0.0 {
249                // Temperature sampling
250                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                // Greedy
276                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            // Check for EOS
288            if next_token == self.config.eos_token_id {
289                break;
290            }
291
292            decoder_tokens.push(next_token);
293        }
294
295        // 6. Decode output tokens (skip decoder_start_token)
296        let output_tokens: Vec<u32> = decoder_tokens.iter()
297            .skip(1) // Skip decoder_start_token
298            .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, // encoder + decoder KV cache
316            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            // Only first layer has relative position bias
343            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        // Token embeddings
365        let mut hidden_states = ops_fn::embedding(input_ids, shared_embedding)?;
366
367        // Compute position bias from first layer (shared across all layers)
368        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        // Cross-attention position bias (for decoder)
373        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        // Apply transformer blocks
386        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        // Final layer norm
397        ops_fn::layer_norm(&hidden_states, &self.final_layer_norm, None, self.config.layer_norm_epsilon)
398    }
399
400    /// Compute relative position bias for T5 attention
401    fn compute_position_bias(&self, query_length: usize, key_length: usize) -> Result<Tensor> {
402        // For now, return zeros - the actual bias is computed in attention
403        // This is a placeholder for the relative position encoding
404        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        // Load final layer norm
413        if let Some(w) = weights.get(&format!("{}.final_layer_norm.weight", prefix)) {
414            self.final_layer_norm = w.clone();
415        }
416
417        // Load block weights
418        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        // Self-attention layer
437        let self_attention = T5LayerSelfAttention::new(config, device, is_decoder, has_relative_attention_bias)?;
438
439        // Cross-attention layer (decoder only)
440        let cross_attention = if is_decoder {
441            Some(T5LayerCrossAttention::new(config, device, has_relative_attention_bias)?)
442        } else {
443            None
444        };
445
446        // Feed-forward layer
447        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        // 1. Self-attention with residual
467        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        // 2. Cross-attention (decoder only) with residual
476        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, // No mask for cross-attention typically
482            )?;
483            ops_fn::add(&hidden_states, &cross_output)?
484        } else {
485            hidden_states
486        };
487
488        // 3. Feed-forward with residual
489        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        // Load self-attention weights
495        self.self_attention.load_weights(weights, &format!("{}.layer.0", prefix))?;
496
497        // Load cross-attention weights (decoder only)
498        if let Some(ref mut cross_attn) = self.cross_attention {
499            cross_attn.load_weights(weights, &format!("{}.layer.1", prefix))?;
500        }
501
502        // Load feed-forward weights
503        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        // Pre-LayerNorm
535        let normed = ops_fn::layer_norm(hidden_states, &self.layer_norm, None, 1e-6)?;
536
537        // Self-attention
538        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        // Layer norm
543        if let Some(w) = weights.get(&format!("{}.layer_norm.weight", prefix)) {
544            self.layer_norm = w.clone();
545        }
546
547        // Attention weights
548        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        // Cross-attention is never causal
563        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        // Pre-LayerNorm
577        let normed = ops_fn::layer_norm(hidden_states, &self.layer_norm, None, 1e-6)?;
578
579        // Cross-attention (query from decoder, key/value from encoder)
580        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        // Layer norm
585        if let Some(w) = weights.get(&format!("{}.layer_norm.weight", prefix)) {
586            self.layer_norm = w.clone();
587        }
588
589        // Attention weights
590        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        // Determine activation function
619        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        // Pre-LayerNorm
637        let normed = ops_fn::layer_norm(hidden_states, &self.layer_norm, None, 1e-6)?;
638
639        // Feed-forward computation
640        let hidden = if self.is_gated {
641            // Gated activation: gate * activation(up)
642            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            // Standard: activation(up)
653            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        // Down projection
661        ops_fn::matmul(&hidden, &self.wo)
662    }
663
664    fn load_weights(&mut self, weights: &ModelWeights, prefix: &str) -> Result<()> {
665        // Layer norm
666        if let Some(w) = weights.get(&format!("{}.layer_norm.weight", prefix)) {
667            self.layer_norm = w.clone();
668        }
669
670        // Dense relu dense weights (transpose for matmul)
671        if let Some(w) = weights.get(&format!("{}.DenseReluDense.wi.weight", prefix)) {
672            self.wi = ops_fn::transpose(w)?;
673        }
674        // Alternative naming for gated variants
675        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        // Q, K, V, O projections
704        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        // Relative position bias (only first layer of each stack)
710        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        // Project to Q, K, V
751        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        // Reshape for multi-head attention: [batch, seq, heads * d_kv] -> [batch, heads, seq, d_kv]
756        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        // Compute attention scores: Q @ K^T
773        let k_t = k_reshaped.transpose(2, 3)?;
774        let scores = q_reshaped.contiguous()?.matmul(&k_t.contiguous()?)?;
775
776        // Add relative position bias
777        let position_bias_candle = position_bias.to_candle()?;
778        let scores = scores.broadcast_add(&position_bias_candle)?;
779
780        // Apply causal mask for decoder self-attention
781        let scores = if self.is_decoder && seq_len == kv_seq_len {
782            // Create causal mask
783            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        // Apply attention mask if provided
799        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        // Softmax
807        let attention_weights = candle_nn::ops::softmax_last_dim(&scores)?;
808
809        // Apply attention to values
810        let attn_output = attention_weights.matmul(&v_reshaped.contiguous()?)?;
811
812        // Reshape back: [batch, heads, seq, d_kv] -> [batch, seq, heads * d_kv]
813        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        // Output projection
820        ops_fn::matmul(&attn_output, &self.o)
821    }
822
823    fn load_weights(&mut self, weights: &ModelWeights, prefix: &str) -> Result<()> {
824        // Load and transpose projection weights
825        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        // Load relative attention bias
839        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
862/// ReLU activation function
863fn 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/// Compute relative position bucket for T5 attention
870/// This converts relative positions to bucket indices for the position bias lookup
871#[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    // Half of buckets are for exact positions
895    let max_exact = num_buckets / 2;
896
897    if relative_position < max_exact {
898        relative_buckets + relative_position
899    } else {
900        // The other half are for logarithmically bigger bins
901        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]); // batch, seq, vocab
957                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        // Should produce some output (even if random with uninitialized weights)
986        // Generation may produce empty output if EOS token is sampled early
987        let _ = output;
988    }
989
990    #[test]
991    fn test_relative_position_bucket() {
992        // Test bidirectional bucketing (encoder)
993        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); // Should be in first half
1001
1002        // Test unidirectional bucketing (decoder)
1003        let bucket = relative_position_bucket(0, false, 32, 128);
1004        assert_eq!(bucket, 0);
1005    }
1006}