Coverage for transformer_lens/model_bridge/supported_architectures/t5.py: 100%
21 statements
« prev ^ index » next coverage.py v7.10.1, created at 2026-09-01 16:23 +0000
« prev ^ index » next coverage.py v7.10.1, created at 2026-09-01 16:23 +0000
1"""T5 architecture adapter."""
3from typing import Any
5from transformer_lens.model_bridge.architecture_adapter import ArchitectureAdapter
6from transformer_lens.model_bridge.generalized_components import (
7 AttentionBridge,
8 EmbeddingBridge,
9 LinearBridge,
10 MLPBridge,
11 PosEmbedBridge,
12 RMSNormalizationBridge,
13 T5BlockBridge,
14 UnembeddingBridge,
15)
18class T5ArchitectureAdapter(ArchitectureAdapter):
19 """Architecture adapter for T5 models.
21 T5 is an encoder-decoder model with:
22 - Shared embeddings
23 - Encoder stack (self-attention + FFN)
24 - Decoder stack (self-attention + cross-attention + FFN)
25 - Language modeling head
27 Supports both standard T5 (DenseReluDense with wi/wo) and gated variants
28 like Flan-T5 (T5DenseGatedActDense with wi_0/wi_1/wo).
29 """
31 def __init__(self, cfg: Any) -> None:
32 """Initialize the T5 architecture adapter.
34 Args:
35 cfg: The configuration object.
36 """
37 super().__init__(cfg)
39 # T5 RMSNorm: disable fold_ln to avoid corrupting weights.
40 self.supports_fold_ln = False
42 # Set config variables for weight processing
43 self.cfg.normalization_type = "RMS"
44 self.cfg.positional_embedding_type = "relative_positional_bias"
45 self.cfg.final_rms = False
46 self.cfg.attn_only = False
48 # Detect gated MLP variant (Flan-T5 uses T5DenseGatedActDense)
49 is_gated = getattr(cfg, "is_gated_act", False)
50 self.cfg.gated_mlp = is_gated
52 self.weight_processing_conversions = {}
54 # Build MLP bridges via the seam (Switch swaps in sparse MoE FFs).
55 encoder_mlp = self._build_ff_bridge("layer.1")
56 decoder_mlp = self._build_ff_bridge("layer.2")
58 self.component_mapping = {
59 # Shared embeddings
60 "embed": EmbeddingBridge(name="shared"),
61 # Encoder positional embeddings (relative attention bias)
62 "pos_embed": PosEmbedBridge(
63 name="encoder.block.0.layer.0.SelfAttention.relative_attention_bias"
64 ),
65 # Encoder blocks (2 layers: self-attn, FFN)
66 "encoder_blocks": T5BlockBridge(
67 name="encoder.block",
68 config=self.cfg,
69 is_decoder=False,
70 submodules={
71 "ln1": RMSNormalizationBridge(name="layer.0.layer_norm", config=self.cfg),
72 "attn": AttentionBridge(
73 name="layer.0.SelfAttention",
74 config=self.cfg,
75 submodules={
76 "q": LinearBridge(name="q"),
77 "k": LinearBridge(name="k"),
78 "v": LinearBridge(name="v"),
79 "o": LinearBridge(name="o"),
80 },
81 requires_relative_position_bias=True,
82 ),
83 "ln2": RMSNormalizationBridge(name="layer.1.layer_norm", config=self.cfg),
84 "mlp": encoder_mlp,
85 },
86 ),
87 # Encoder final layer norm
88 "encoder_ln_final": RMSNormalizationBridge(
89 name="encoder.final_layer_norm", config=self.cfg
90 ),
91 # Decoder positional embeddings (relative attention bias)
92 "decoder_pos_embed": PosEmbedBridge(
93 name="decoder.block.0.layer.0.SelfAttention.relative_attention_bias"
94 ),
95 # Decoder blocks (3 layers: self-attn, cross-attn, FFN)
96 "decoder_blocks": T5BlockBridge(
97 name="decoder.block",
98 config=self.cfg,
99 is_decoder=True,
100 submodules={
101 "ln1": RMSNormalizationBridge(name="layer.0.layer_norm", config=self.cfg),
102 "self_attn": AttentionBridge(
103 name="layer.0.SelfAttention",
104 config=self.cfg,
105 submodules={
106 "q": LinearBridge(name="q"),
107 "k": LinearBridge(name="k"),
108 "v": LinearBridge(name="v"),
109 "o": LinearBridge(name="o"),
110 },
111 requires_relative_position_bias=True,
112 ),
113 "ln2": RMSNormalizationBridge(name="layer.1.layer_norm", config=self.cfg),
114 "cross_attn": AttentionBridge(
115 name="layer.1.EncDecAttention",
116 config=self.cfg,
117 submodules={
118 "q": LinearBridge(name="q"),
119 "k": LinearBridge(name="k"),
120 "v": LinearBridge(name="v"),
121 "o": LinearBridge(name="o"),
122 },
123 requires_relative_position_bias=True,
124 is_cross_attention=True,
125 ),
126 "ln3": RMSNormalizationBridge(name="layer.2.layer_norm", config=self.cfg),
127 "mlp": decoder_mlp,
128 },
129 ),
130 # Decoder final layer norm
131 "decoder_ln_final": RMSNormalizationBridge(
132 name="decoder.final_layer_norm", config=self.cfg
133 ),
134 # Language modeling head
135 "unembed": UnembeddingBridge(name="lm_head"),
136 }
138 def _build_ff_bridge(self, layer_prefix: str):
139 """Feed-forward bridge for one stack (gated or plain wi/wo); Switch
140 overrides with a sparse MoE variant."""
141 if self.cfg.gated_mlp:
142 return self._gated_mlp(
143 name=f"{layer_prefix}.DenseReluDense", gate="wi_0", up="wi_1", down="wo"
144 )
145 return MLPBridge(
146 name=f"{layer_prefix}.DenseReluDense",
147 submodules={
148 "in": LinearBridge(name="wi"),
149 "out": LinearBridge(name="wo"),
150 },
151 )