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

1"""T5 architecture adapter.""" 

2 

3from typing import Any 

4 

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) 

16 

17 

18class T5ArchitectureAdapter(ArchitectureAdapter): 

19 """Architecture adapter for T5 models. 

20 

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 

26 

27 Supports both standard T5 (DenseReluDense with wi/wo) and gated variants 

28 like Flan-T5 (T5DenseGatedActDense with wi_0/wi_1/wo). 

29 """ 

30 

31 def __init__(self, cfg: Any) -> None: 

32 """Initialize the T5 architecture adapter. 

33 

34 Args: 

35 cfg: The configuration object. 

36 """ 

37 super().__init__(cfg) 

38 

39 # T5 RMSNorm: disable fold_ln to avoid corrupting weights. 

40 self.supports_fold_ln = False 

41 

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 

47 

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 

51 

52 self.weight_processing_conversions = {} 

53 

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") 

57 

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 } 

137 

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 )