Coverage for transformer_lens/model_bridge/supported_architectures/modernbert_decoder.py: 100%
18 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"""ModernBERT Decoder (``ModernBertDecoderForCausalLM``) adapter: ModernBERT run
2causally. Attention delegates (per-layer sliding windows) and LN folding is disabled
3(layer-0 Identity attention norm)."""
5from typing import Any
7from transformer_lens.model_bridge.architecture_adapter import ArchitectureAdapter
8from transformer_lens.model_bridge.generalized_components import (
9 AttentionBridge,
10 BlockBridge,
11 EmbeddingBridge,
12 LinearBridge,
13 NormalizationBridge,
14 RotaryEmbeddingBridge,
15 UnembeddingBridge,
16)
17from transformer_lens.model_bridge.generalized_components.base import (
18 GeneralizedComponent,
19)
22class ModernBertDecoderArchitectureAdapter(ArchitectureAdapter):
23 """Architecture adapter for ModernBertDecoderForCausalLM models."""
25 # Layer 0's attention norm is Identity: folding assumes a real norm on
26 # every sublayer, and centering writing weights assumes every residual
27 # reader is mean-invariant (layer 0's attention reads the residual raw).
28 supports_fold_ln = False
29 supports_center_writing_weights = False
31 def __init__(self, cfg: Any) -> None:
32 """Initialize the ModernBERT Decoder architecture adapter."""
33 super().__init__(cfg)
35 self.cfg.normalization_type = "LN"
36 self.cfg.uses_rms_norm = False
37 self.cfg.positional_embedding_type = "rotary"
38 self.cfg.gated_mlp = False # fused-GLU: Wi carries both halves
39 self.cfg.attn_only = False
40 self.cfg.final_rms = False
42 self.weight_processing_conversions = {}
44 self.component_mapping = {
45 "embed": EmbeddingBridge(name="model.embeddings.tok_embeddings"),
46 "embed_ln": NormalizationBridge(
47 name="model.embeddings.norm",
48 config=self.cfg,
49 use_native_layernorm_autograd=True,
50 ),
51 "rotary_emb": RotaryEmbeddingBridge(name="model.rotary_emb"),
52 "blocks": BlockBridge(
53 name="model.layers",
54 config=self.cfg,
55 submodules={
56 # Identity on layer 0 (pre-normed by the embedding norm).
57 "ln1": NormalizationBridge(
58 name="attn_norm",
59 config=self.cfg,
60 use_native_layernorm_autograd=True,
61 ),
62 # Sliding-window/global mix per layer_types: delegate.
63 "attn": AttentionBridge(
64 name="attn",
65 config=self.cfg,
66 submodules={
67 "q": LinearBridge(name="q_proj"),
68 "k": LinearBridge(name="k_proj"),
69 "v": LinearBridge(name="v_proj"),
70 "o": LinearBridge(name="Wo"),
71 },
72 maintain_native_attention=True,
73 ),
74 "ln2": NormalizationBridge(
75 name="mlp_norm",
76 config=self.cfg,
77 use_native_layernorm_autograd=True,
78 ),
79 "mlp": self._ungated_mlp(up="Wi", down="Wo"),
80 },
81 ),
82 "ln_final": NormalizationBridge(
83 name="model.final_norm",
84 config=self.cfg,
85 use_native_layernorm_autograd=True,
86 ),
87 # BERT-style head (dense + act + norm) ahead of the vocab projection.
88 "prediction_head": GeneralizedComponent(name="lm_head"),
89 "unembed": UnembeddingBridge(name="decoder"),
90 }
92 def setup_component_testing(self, hf_model: Any, bridge_model: Any = None) -> None:
93 """Delegated attention computes rotary inside HF; nothing to wire."""