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

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).""" 

4 

5from typing import Any 

6 

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) 

20 

21 

22class ModernBertDecoderArchitectureAdapter(ArchitectureAdapter): 

23 """Architecture adapter for ModernBertDecoderForCausalLM models.""" 

24 

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 

33 

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

35 """Initialize the ModernBERT Decoder architecture adapter.""" 

36 super().__init__(cfg) 

37 

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 

44 

45 self.weight_processing_conversions = {} 

46 

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 }