Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
21 changes: 14 additions & 7 deletions moe_infinity/models/deepseek.py
Original file line number Diff line number Diff line change
Expand Up @@ -48,25 +48,32 @@ def __init__(self, config):
self.num_expert = config.n_routed_experts

if self.config.model_type == "deepseek_v2":
import transformers.models.deepseek_v2.modeling_deepseek_v2 as _dsv2
from transformers.models.deepseek_v2.modeling_deepseek_v2 import (
DeepseekV2MLP,
DeepseekV2MoEGate,
)

self.mlp_cls = DeepseekV2MLP
self.gate_cls = DeepseekV2MoEGate
# DeepseekV2MoEGate was removed in newer Transformers; fall back to
# the local DeepseekMoEGate which has the same raw-logits interface.
_gate_cls = getattr(_dsv2, "DeepseekV2MoEGate", None)
self.gate_cls = (
_gate_cls if _gate_cls is not None else DeepseekMoEGate
)
if self.config.model_type == "deepseek_v3":
from transformers.models.deepseek_v3.modeling_deepseek_v3 import (
DeepseekV3MLP,
)

self.mlp_cls = DeepseekV3MLP
# V3 upstream has no standalone gate; use V2 gate (same interface)
from transformers.models.deepseek_v2.modeling_deepseek_v2 import (
DeepseekV2MoEGate,
)
# V3 upstream has no standalone gate; use V2 gate when available,
# otherwise fall back to the local DeepseekMoEGate.
import transformers.models.deepseek_v2.modeling_deepseek_v2 as _dsv2

self.gate_cls = DeepseekV2MoEGate
_gate_cls = getattr(_dsv2, "DeepseekV2MoEGate", None)
self.gate_cls = (
_gate_cls if _gate_cls is not None else DeepseekMoEGate
)

self.experts = nn.ModuleList(
[
Expand Down
21 changes: 17 additions & 4 deletions tests/python/ops/test_deepseek_v2_gate_consistency.py
Original file line number Diff line number Diff line change
Expand Up @@ -41,12 +41,23 @@ def decorator(fn):

_ensure_nvtx_stub_has_annotate()

import transformers.models.deepseek_v2.modeling_deepseek_v2 as _dsv2_mod
from transformers import DeepseekV2Config
from transformers.models.deepseek_v2.modeling_deepseek_v2 import (
DeepseekV2MoE,

DeepseekV2MoE = getattr(_dsv2_mod, "DeepseekV2Moe", None) or getattr(
_dsv2_mod, "DeepseekV2MoE", None
)
from transformers.models.deepseek_v2.modeling_deepseek_v2 import (
DeepseekV2MoEGate as MoEGate,
if DeepseekV2MoE is None:
raise ImportError(
"Neither 'DeepseekV2Moe' nor 'DeepseekV2MoE' found in "
"transformers.models.deepseek_v2.modeling_deepseek_v2"
)

MoEGate = getattr(_dsv2_mod, "DeepseekV2MoEGate", None)

requires_moe_gate = pytest.mark.skipif(
MoEGate is None,
reason="DeepseekV2MoEGate removed in this Transformers version",
)

from moe_infinity.models.deepseek import DeepseekMoEBlock, DeepseekMoEGate
Expand Down Expand Up @@ -132,6 +143,7 @@ def wait_dispatch_local(self):


@requires_cuda
@requires_moe_gate
def test_v2_gate_equivalence(seed_everything):
"""V2-Lite config: native MoEGate and simplified DeepseekMoEGate must
produce identical routing because topk_method='greedy' with
Expand Down Expand Up @@ -185,6 +197,7 @@ def test_v2_gate_equivalence(seed_everything):


@requires_cuda
@requires_moe_gate
def test_v2_gate_group_limited_greedy(seed_everything):
"""V2-full config: simplified gate MUST differ from native gate because
group_limited_greedy constrains which expert groups can be selected."""
Expand Down
Loading