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
« prev ^ index » next coverage.py v7.10.1, created at 2026-08-11 18:50 +0000
1"""AST (Audio Spectrogram Transformer) architecture adapter.
3Supports ASTForAudioClassification
4"""
6from typing import Any
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)
26class ASTArchitectureAdapter(ArchitectureAdapter):
27 """Architecture adapter for AST (Audio Spectrogram Transformer) audio classifiers.
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 """
34 supports_generation: bool = False
36 def __init__(self, cfg: Any) -> None:
37 super().__init__(cfg)
39 # essential flag for audio models in V3
40 self.cfg.is_audio_model = True
41 self.cfg.normalization_type = "LN"
43 n_heads = self.cfg.n_heads
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 }
78 # default mapping for bare ASTModel (prefix="")
79 self.component_mapping = self._build_component_mapping(prefix="")
81 def _build_component_mapping(self, prefix: str) -> dict:
82 """Build component mapping. prefix="" for ASTModel, "audio_spectrogram_transformer." for classification."""
83 p = prefix
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 }
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
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
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]