Coverage for transformer_lens/model_bridge/supported_architectures/ast.py: 89%

30 statements  

« prev     ^ index     » next       coverage.py v7.10.1, created at 2026-08-11 18:50 +0000

1"""AST (Audio Spectrogram Transformer) architecture adapter. 

2 

3Supports ASTForAudioClassification 

4""" 

5 

6from typing import Any 

7 

8from transformer_lens.conversion_utils.conversion_steps import RearrangeTensorConversion 

9from transformer_lens.conversion_utils.param_processing_conversion import ( 

10 ParamProcessingConversion, 

11) 

12from transformer_lens.model_bridge.architecture_adapter import ArchitectureAdapter 

13from transformer_lens.model_bridge.generalized_components import ( 

14 AttentionBridge, 

15 BlockBridge, 

16 LinearBridge, 

17 MLPBridge, 

18 NormalizationBridge, 

19 UnembeddingBridge, 

20) 

21from transformer_lens.model_bridge.generalized_components.base import ( 

22 GeneralizedComponent, 

23) 

24 

25 

26class ASTArchitectureAdapter(ArchitectureAdapter): 

27 """Architecture adapter for AST (Audio Spectrogram Transformer) audio classifiers. 

28 

29 Input is a [batch, time=1024, n_mels=128] spectrogram; positions 0/1 are the 

30 CLS and distillation tokens. HF pools (cls+dist)/2 after ln_final and applies 

31 an extra LayerNorm inside the classifier head, see the unembed mapping note. 

32 """ 

33 

34 supports_generation: bool = False 

35 

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

37 super().__init__(cfg) 

38 

39 # essential flag for audio models in V3 

40 self.cfg.is_audio_model = True 

41 self.cfg.normalization_type = "LN" 

42 

43 n_heads = self.cfg.n_heads 

44 

45 # Q/K/V/O rearrangement: splits hidden dims into (heads, head_dim) 

46 self.weight_processing_conversions = { 

47 "blocks.{i}.attn.q.weight": ParamProcessingConversion( 

48 tensor_conversion=RearrangeTensorConversion( 

49 "(h d_head) d_model -> h d_model d_head", h=n_heads 

50 ), 

51 ), 

52 "blocks.{i}.attn.k.weight": ParamProcessingConversion( 

53 tensor_conversion=RearrangeTensorConversion( 

54 "(h d_head) d_model -> h d_model d_head", h=n_heads 

55 ), 

56 ), 

57 "blocks.{i}.attn.v.weight": ParamProcessingConversion( 

58 tensor_conversion=RearrangeTensorConversion( 

59 "(h d_head) d_model -> h d_model d_head", h=n_heads 

60 ), 

61 ), 

62 "blocks.{i}.attn.q.bias": ParamProcessingConversion( 

63 tensor_conversion=RearrangeTensorConversion("(h d_head) -> h d_head", h=n_heads), 

64 ), 

65 "blocks.{i}.attn.k.bias": ParamProcessingConversion( 

66 tensor_conversion=RearrangeTensorConversion("(h d_head) -> h d_head", h=n_heads), 

67 ), 

68 "blocks.{i}.attn.v.bias": ParamProcessingConversion( 

69 tensor_conversion=RearrangeTensorConversion("(h d_head) -> h d_head", h=n_heads), 

70 ), 

71 "blocks.{i}.attn.o.weight": ParamProcessingConversion( 

72 tensor_conversion=RearrangeTensorConversion( 

73 "d_model (h d_head) -> h d_head d_model", h=n_heads 

74 ), 

75 ), 

76 } 

77 

78 # default mapping for bare ASTModel (prefix="") 

79 self.component_mapping = self._build_component_mapping(prefix="") 

80 

81 def _build_component_mapping(self, prefix: str) -> dict: 

82 """Build component mapping. prefix="" for ASTModel, "audio_spectrogram_transformer." for classification.""" 

83 p = prefix 

84 

85 return { 

86 "embed": GeneralizedComponent(name=f"{p}embeddings"), 

87 "ln_final": NormalizationBridge(name=f"{p}layernorm", config=self.cfg), 

88 "blocks": BlockBridge( 

89 name=f"{p}layers", 

90 submodules={ 

91 "ln1": NormalizationBridge(name="layernorm_before", config=self.cfg), 

92 "ln2": NormalizationBridge(name="layernorm_after", config=self.cfg), 

93 "attn": AttentionBridge( 

94 name="attention", 

95 config=self.cfg, 

96 submodules={ 

97 "q": LinearBridge(name="q_proj"), 

98 "k": LinearBridge(name="k_proj"), 

99 "v": LinearBridge(name="v_proj"), 

100 "o": LinearBridge(name="o_proj"), 

101 }, 

102 ), 

103 "mlp": MLPBridge( 

104 name="mlp", 

105 config=self.cfg, 

106 submodules={ 

107 "in": LinearBridge(name="fc1"), 

108 "out": LinearBridge(name="fc2"), 

109 }, 

110 ), 

111 }, 

112 ), 

113 } 

114 

115 def prepare_model(self, hf_model: Any) -> None: 

116 """Detect classification head, rebind prefixes, guard unembed, and set n_ctx.""" 

117 # 1. handle classification model vs bare encoder 

118 if hasattr(hf_model, "audio_spectrogram_transformer"): 118 ↛ 135line 118 didn't jump to line 135 because the condition on line 118 was always true

119 self.component_mapping = self._build_component_mapping( 

120 prefix="audio_spectrogram_transformer." 

121 ) 

122 base_model = hf_model.audio_spectrogram_transformer 

123 

124 # guard unembed: only map if num_labels > 0 and dense head is real 

125 num_labels = getattr(hf_model.config, "num_labels", 0) 

126 if ( 126 ↛ 138line 126 didn't jump to line 138 because the condition on line 126 was always true

127 num_labels > 0 

128 and hasattr(hf_model, "classifier") 

129 and hasattr(hf_model.classifier, "dense") 

130 ): 

131 self.component_mapping["unembed"] = UnembeddingBridge(name="classifier.dense") 

132 self.cfg.d_vocab = num_labels 

133 self.cfg.d_vocab_out = num_labels 

134 else: 

135 base_model = hf_model 

136 

137 # 2. dynamically grab n_ctx from whatever base model we resolved 

138 if hasattr(base_model, "embeddings") and hasattr( 138 ↛ exitline 138 didn't return from function 'prepare_model' because the condition on line 138 was always true

139 base_model.embeddings, "position_embeddings" 

140 ): 

141 self.cfg.n_ctx = base_model.embeddings.position_embeddings.shape[1]