Coverage for transformer_lens/model_bridge/supported_architectures/modernbert_decoder.py: 100%
19 statements
« prev ^ index » next coverage.py v7.10.1, created at 2026-09-21 19:27 +0000
« prev ^ index » next coverage.py v7.10.1, created at 2026-09-21 19:27 +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
30 # Delegated attention computes rotary inside HF; nothing to wire.
31 _testing_eager = None
32 _testing_wire_rotary = False
34 def __init__(self, cfg: Any) -> None:
35 """Initialize the ModernBERT Decoder architecture adapter."""
36 super().__init__(cfg)
38 self.cfg.normalization_type = "LN"
39 self.cfg.uses_rms_norm = False
40 self.cfg.positional_embedding_type = "rotary"
41 self.cfg.gated_mlp = False # fused-GLU: Wi carries both halves
42 self.cfg.attn_only = False
43 self.cfg.final_rms = False
45 self.weight_processing_conversions = {}
47 self.component_mapping = {
48 "embed": EmbeddingBridge(name="model.embeddings.tok_embeddings"),
49 "embed_ln": NormalizationBridge(
50 name="model.embeddings.norm",
51 config=self.cfg,
52 use_native_layernorm_autograd=True,
53 ),
54 "rotary_emb": RotaryEmbeddingBridge(name="model.rotary_emb"),
55 "blocks": BlockBridge(
56 name="model.layers",
57 config=self.cfg,
58 submodules={
59 # Identity on layer 0 (pre-normed by the embedding norm).
60 "ln1": NormalizationBridge(
61 name="attn_norm",
62 config=self.cfg,
63 use_native_layernorm_autograd=True,
64 ),
65 # Sliding-window/global mix per layer_types: delegate.
66 "attn": AttentionBridge(
67 name="attn",
68 config=self.cfg,
69 submodules={
70 "q": LinearBridge(name="q_proj"),
71 "k": LinearBridge(name="k_proj"),
72 "v": LinearBridge(name="v_proj"),
73 "o": LinearBridge(name="Wo"),
74 },
75 maintain_native_attention=True,
76 ),
77 "ln2": NormalizationBridge(
78 name="mlp_norm",
79 config=self.cfg,
80 use_native_layernorm_autograd=True,
81 ),
82 "mlp": self._ungated_mlp(up="Wi", down="Wo"),
83 },
84 ),
85 "ln_final": NormalizationBridge(
86 name="model.final_norm",
87 config=self.cfg,
88 use_native_layernorm_autograd=True,
89 ),
90 # BERT-style head (dense + act + norm) ahead of the vocab projection.
91 "prediction_head": GeneralizedComponent(name="lm_head"),
92 "unembed": UnembeddingBridge(name="decoder"),
93 }