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
« prev ^ index » next coverage.py v7.10.1, created at 2026-09-01 16:23 +0000
1"""DeepSeek V3 architecture adapter.
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"""
10from typing import Any
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)
28class DeepSeekV3ArchitectureAdapter(ArchitectureAdapter):
29 """Architecture adapter for DeepSeek V3 / R1 models.
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 """
35 _testing_eager = None
37 def __init__(self, cfg: Any) -> None:
38 super().__init__(cfg)
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.
48 # MLA has no per-head q/k/v to fold into; skip LN folding.
49 self.supports_fold_ln = False
51 self.weight_processing_conversions = {}
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 }