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

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 

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

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

33 super().__init__(cfg) 

34 

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 

41 

42 self.weight_processing_conversions = {} 

43 

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 } 

91 

92 def setup_component_testing(self, hf_model: Any, bridge_model: Any = None) -> None: 

93 """Delegated attention computes rotary inside HF; nothing to wire."""