diff --git a/tests/unit/model_bridge/supported_architectures/test_ast_adapter.py b/tests/unit/model_bridge/supported_architectures/test_ast_adapter.py new file mode 100644 index 000000000..e35820ac2 --- /dev/null +++ b/tests/unit/model_bridge/supported_architectures/test_ast_adapter.py @@ -0,0 +1,56 @@ +import torch +from transformers import ASTConfig, ASTForAudioClassification + +from transformer_lens.model_bridge import TransformerBridge + + +def test_ast_parity(): + # 1. setup a tiny HF AST config for instantaneous local testing + hf_config = ASTConfig( + hidden_size=12, + num_hidden_layers=2, + num_attention_heads=2, + intermediate_size=24, + patch_size=16, + ) + + # explicitly define the architecture list for the local config + hf_config.architectures = ["ASTForAudioClassification"] + + hf_model = ASTForAudioClassification(hf_config) + hf_model.eval() + + # 2. boot the v3 transformerbridge natively + # bridge auotmatically detecty ASTForAudioClassification and use new ASTArchitectureAdapater + tl_model = TransformerBridge.boot_transformers("ast", hf_model=hf_model) + tl_model.eval() + + # 3. create dummy spectrogram input [batch, freq, time] + dummy_spectrogram = torch.randn(1, 1024, 128) + + # 4. compare forward passes + with torch.no_grad(): + hf_logits = hf_model(dummy_spectrogram).logits + + # bridge forward pass (extract sequence before HF's pooling) + _, cache = tl_model.run_with_cache(dummy_spectrogram) + resid = cache["ln_final.hook_normalized"] + + # apply AST pooling logic to the bridges residual stream + cls_token = resid[:, 0, :] + dist_token = resid[:, 1, :] + pooled_out = (cls_token + dist_token) / 2.0 + + # apply the final classifier (which holds 2nd layernorm + dense) + tl_logits = hf_model.classifier(pooled_out) + + diff = (hf_logits - tl_logits).abs().max().item() + print(f"Bridge vs HF Max Logit Diff: {diff:.6e}") + + # 5. assert parity + assert diff < 1e-4, "Parity failed: Bridge tensors do not match HuggingFace." + print("Parity dub. V3 Bridge connected excellent") + + +if __name__ == "__main__": + test_ast_parity() diff --git a/transformer_lens/factories/architecture_adapter_factory.py b/transformer_lens/factories/architecture_adapter_factory.py index 63d8c119e..fb72d0596 100644 --- a/transformer_lens/factories/architecture_adapter_factory.py +++ b/transformer_lens/factories/architecture_adapter_factory.py @@ -13,6 +13,7 @@ AfmoeArchitectureAdapter, ApertusArchitectureAdapter, ArceeArchitectureAdapter, + ASTArchitectureAdapter, AudioFlamingo3ArchitectureAdapter, BaichuanArchitectureAdapter, BambaArchitectureAdapter, @@ -158,6 +159,7 @@ "AfmoeForCausalLM": AfmoeArchitectureAdapter, "ApertusForCausalLM": ApertusArchitectureAdapter, "ArceeForCausalLM": ArceeArchitectureAdapter, + "ASTForAudioClassification": ASTArchitectureAdapter, "BaiChuanForCausalLM": BaichuanArchitectureAdapter, "BaichuanForCausalLM": BaichuanArchitectureAdapter, "BambaForCausalLM": BambaArchitectureAdapter, diff --git a/transformer_lens/model_bridge/supported_architectures/__init__.py b/transformer_lens/model_bridge/supported_architectures/__init__.py index ba60f5309..35f335f09 100644 --- a/transformer_lens/model_bridge/supported_architectures/__init__.py +++ b/transformer_lens/model_bridge/supported_architectures/__init__.py @@ -9,6 +9,9 @@ from transformer_lens.model_bridge.supported_architectures.audio_flamingo3 import ( AudioFlamingo3ArchitectureAdapter, ) +from transformer_lens.model_bridge.supported_architectures.ast import ( + ASTArchitectureAdapter, +) from transformer_lens.model_bridge.supported_architectures.baichuan import ( BaichuanArchitectureAdapter, ) @@ -271,6 +274,7 @@ "AfmoeArchitectureAdapter", "ApertusArchitectureAdapter", "ArceeArchitectureAdapter", + "ASTArchitectureAdapter", "AudioFlamingo3ArchitectureAdapter", "BD3LMArchitectureAdapter", "BaichuanArchitectureAdapter", diff --git a/transformer_lens/model_bridge/supported_architectures/ast.py b/transformer_lens/model_bridge/supported_architectures/ast.py new file mode 100644 index 000000000..d714ac30c --- /dev/null +++ b/transformer_lens/model_bridge/supported_architectures/ast.py @@ -0,0 +1,103 @@ +from typing import Any + +from transformer_lens.conversion_utils.conversion_steps import RearrangeTensorConversion +from transformer_lens.conversion_utils.param_processing_conversion import ( + ParamProcessingConversion, +) +from transformer_lens.model_bridge.architecture_adapter import ArchitectureAdapter +from transformer_lens.model_bridge.generalized_components import ( + AttentionBridge, + BlockBridge, + LinearBridge, + MLPBridge, + NormalizationBridge, + UnembeddingBridge, +) +from transformer_lens.model_bridge.generalized_components.base import ( + GeneralizedComponent, +) + + +class ASTArchitectureAdapter(ArchitectureAdapter): + def __init__(self, cfg: Any) -> None: + super().__init__(cfg) + + # essential flag for audio models in V3 + self.cfg.is_audio_model = True + self.cfg.normalization_type = "LN" + + n_heads = self.cfg.n_heads + + # Q/K/V/O rearrangement: splits hidden dims into (heads, head_dim) + self.weight_processing_conversions = { + "blocks.{i}.attn.q.weight": ParamProcessingConversion( + tensor_conversion=RearrangeTensorConversion( + "(h d_head) d_model -> h d_model d_head", h=n_heads + ), + ), + "blocks.{i}.attn.k.weight": ParamProcessingConversion( + tensor_conversion=RearrangeTensorConversion( + "(h d_head) d_model -> h d_model d_head", h=n_heads + ), + ), + "blocks.{i}.attn.v.weight": ParamProcessingConversion( + tensor_conversion=RearrangeTensorConversion( + "(h d_head) d_model -> h d_model d_head", h=n_heads + ), + ), + "blocks.{i}.attn.q.bias": ParamProcessingConversion( + tensor_conversion=RearrangeTensorConversion("(h d_head) -> h d_head", h=n_heads), + ), + "blocks.{i}.attn.k.bias": ParamProcessingConversion( + tensor_conversion=RearrangeTensorConversion("(h d_head) -> h d_head", h=n_heads), + ), + "blocks.{i}.attn.v.bias": ParamProcessingConversion( + tensor_conversion=RearrangeTensorConversion("(h d_head) -> h d_head", h=n_heads), + ), + "blocks.{i}.attn.o.weight": ParamProcessingConversion( + tensor_conversion=RearrangeTensorConversion( + "d_model (h d_head) -> h d_head d_model", h=n_heads + ), + ), + } + + # V3 bridge pattern: hierarchical mapping using bridge components + self.component_mapping = { + "embed": GeneralizedComponent(name="audio_spectrogram_transformer.embeddings"), + "ln_final": NormalizationBridge( + name="audio_spectrogram_transformer.layernorm", config=self.cfg + ), + "unembed": UnembeddingBridge(name="classifier.dense"), + "blocks": BlockBridge( + name="audio_spectrogram_transformer.layers", + submodules={ + "ln1": NormalizationBridge(name="layernorm_before", config=self.cfg), + "ln2": NormalizationBridge(name="layernorm_after", config=self.cfg), + "attn": AttentionBridge( + name="attention", + config=self.cfg, + submodules={ + "q": LinearBridge(name="q_proj"), + "k": LinearBridge(name="k_proj"), + "v": LinearBridge(name="v_proj"), + "o": LinearBridge(name="o_proj"), + }, + ), + "mlp": MLPBridge( + name="mlp", + config=self.cfg, + submodules={ + "in": LinearBridge(name="fc1"), + "out": LinearBridge(name="fc2"), + }, + ), + }, + ), + } + + def prepare_model(self, hf_model: Any) -> None: + # hook to access the live Huggingface model before boot + # calculate n_ctx dynamically from the instantiated position embeddings + self.cfg.n_ctx = ( + hf_model.audio_spectrogram_transformer.embeddings.position_embeddings.shape[1] + ) diff --git a/transformer_lens/tools/model_registry/__init__.py b/transformer_lens/tools/model_registry/__init__.py index aa31a3b0b..d8d7212e3 100644 --- a/transformer_lens/tools/model_registry/__init__.py +++ b/transformer_lens/tools/model_registry/__init__.py @@ -48,6 +48,7 @@ "AfmoeForCausalLM", "ApertusForCausalLM", "ArceeForCausalLM", + "ASTForAudioClassification", "BaiChuanForCausalLM", "BaichuanForCausalLM", "BambaForCausalLM", @@ -200,6 +201,7 @@ "AfmoeForCausalLM": ["arcee-ai"], "ApertusForCausalLM": ["swiss-ai"], "ArceeForCausalLM": ["arcee-ai"], + "ASTForAudioClassification": ["MIT"], "BaiChuanForCausalLM": ["baichuan-inc"], "BaichuanForCausalLM": ["baichuan-inc"], "BambaForCausalLM": ["ibm-ai-platform"],