Coverage for transformer_lens/model_bridge/supported_architectures/deepseek_v3.py: 100%

16 statements  

« prev     ^ index     » next       coverage.py v7.10.1, created at 2026-09-01 16:23 +0000

1"""DeepSeek V3 architecture adapter. 

2 

3Supports DeepSeek V3 and DeepSeek-R1 models (both use DeepseekV3ForCausalLM). 

4Key features: 

5- Multi-Head Latent Attention (MLA): Q and KV compressed via LoRA-style projections 

6- Mixture of Experts (MoE) with shared experts on most layers 

7- Dense MLP on first `first_k_dense_replace` layers 

8""" 

9 

10from typing import Any 

11 

12from transformer_lens.model_bridge.architecture_adapter import ArchitectureAdapter 

13from transformer_lens.model_bridge.generalized_components import ( 

14 EmbeddingBridge, 

15 LinearBridge, 

16 MLAAttentionBridge, 

17 MLABlockBridge, 

18 MoEBridge, 

19 RMSNormalizationBridge, 

20 RotaryEmbeddingBridge, 

21 UnembeddingBridge, 

22) 

23from transformer_lens.model_bridge.generalized_components.base import ( 

24 GeneralizedComponent, 

25) 

26 

27 

28class DeepSeekV3ArchitectureAdapter(ArchitectureAdapter): 

29 """Architecture adapter for DeepSeek V3 / R1 models. 

30 

31 Uses RMSNorm, MLA with compressed Q/KV projections, partial RoPE, 

32 MoE on most layers (dense MLP on first few), and no biases. 

33 """ 

34 

35 _testing_eager = None 

36 

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

38 super().__init__(cfg) 

39 

40 self.cfg.normalization_type = "RMS" 

41 self.cfg.positional_embedding_type = "rotary" 

42 self.cfg.gated_mlp = True 

43 self.cfg.final_rms = True 

44 self.cfg.uses_rms_norm = True 

45 # HF defaults to SDPA which handles MLA correctly. 

46 # HF's eager attention crashes on MLA's asymmetric Q/K dimensions. 

47 

48 # MLA has no per-head q/k/v to fold into; skip LN folding. 

49 self.supports_fold_ln = False 

50 

51 self.weight_processing_conversions = {} 

52 

53 self.component_mapping = { 

54 "embed": EmbeddingBridge(name="model.embed_tokens"), 

55 "rotary_emb": RotaryEmbeddingBridge(name="model.rotary_emb", config=self.cfg), 

56 "blocks": MLABlockBridge( 

57 name="model.layers", 

58 submodules={ 

59 "ln1": RMSNormalizationBridge(name="input_layernorm", config=self.cfg), 

60 "ln2": RMSNormalizationBridge(name="post_attention_layernorm", config=self.cfg), 

61 "attn": MLAAttentionBridge( 

62 name="self_attn", 

63 config=self.cfg, 

64 submodules={ 

65 # Two-stage LoRA Q compression, built only when 

66 # q_lora_rank is set; a config that leaves it null 

67 # gets a single q_proj instead, and MLAAttentionBridge 

68 # already forwards down whichever path exists. 

69 "q_a_proj": LinearBridge(name="q_a_proj", optional=True), 

70 "q_a_layernorm": RMSNormalizationBridge( 

71 name="q_a_layernorm", config=self.cfg, optional=True 

72 ), 

73 "q_b_proj": LinearBridge(name="q_b_proj", optional=True), 

74 "q_proj": LinearBridge(name="q_proj", optional=True), 

75 "kv_a_proj_with_mqa": LinearBridge(name="kv_a_proj_with_mqa"), 

76 "kv_a_layernorm": RMSNormalizationBridge( 

77 name="kv_a_layernorm", config=self.cfg 

78 ), 

79 "kv_b_proj": LinearBridge(name="kv_b_proj"), 

80 "o": LinearBridge(name="o_proj"), 

81 }, 

82 ), 

83 # Dense-prefix layers (idx < first_k_dense_replace) bind as 

84 # gated MLPs with neuron-basis hooks; sparse layers keep the 

85 # MoE mapping with its optional router/shared experts (#1645). 

86 "mlp": MoEBridge( 

87 name="mlp", 

88 config=self.cfg, 

89 sparse_required=("gate",), 

90 submodules={ 

91 # Router is a custom Module, not nn.Linear 

92 "gate": GeneralizedComponent(name="gate", optional=True), 

93 "shared_experts": self._gated_mlp(name="shared_experts", optional=True), 

94 # Dense-layer projections (present only on the 

95 # dense layers of this interleaved stack); their 

96 # presence is what makes MoEBridge bind gated-MLP 

97 # neuron hooks there (#1645). 

98 "dense_gate": LinearBridge(name="gate_proj", optional=True), 

99 "dense_in": LinearBridge(name="up_proj", optional=True), 

100 "dense_out": LinearBridge(name="down_proj", optional=True), 

101 }, 

102 ), 

103 }, 

104 ), 

105 "ln_final": RMSNormalizationBridge(name="model.norm", config=self.cfg), 

106 "unembed": UnembeddingBridge(name="lm_head"), 

107 }