[Feature] Add DeepEPv2 (ElasticBuffer) MoE A2A backend (#35634)

Co-authored-by: menyu <menyu@nvidia.com>
Co-authored-by: Jinyan Chen <93358689+liz-badada@users.noreply.github.com>
Co-authored-by: Han Yu <helloyu0903@gmail.com>
Co-authored-by: Cheng Wan <cheng.wan@radixark.ai>
This commit is contained in:
MengYu
2026-08-26 19:54:33 -07:00
committed by GitHub
co-authored by menyu Jinyan Chen Han Yu Cheng Wan
parent cbfe54fba8
commit a3ae667d67
20 changed files with 2414 additions and 35 deletions
@@ -2106,6 +2106,328 @@ class TestSamplingBackendTokenOracleEnvGate(CustomTestCase):
self.assertEqual(parsed.sampling_backend, "token_oracle")
class TestDeepEPv2Args(CustomTestCase):
"""DeepEP v2 server-argument resolution and validation."""
def _args(self, **overrides):
server_args = ServerArgs(model_path="dummy", moe_a2a_backend="deepep_v2")
server_args.model_config = SimpleNamespace(
hf_config=SimpleNamespace(architectures=["DeepseekV4ForCausalLM"])
)
# The dummy path does not initialize phase configs.
server_args.cuda_graph_config = CudaGraphConfig(
decode=PhaseConfig(backend=Backend.FULL, max_bs=512),
prefill=PhaseConfig(backend=Backend.FULL, max_bs=512),
)
server_args._resolved_overrides = []
valid = {f.name for f in dataclasses.fields(ServerArgs)}
for key, value in overrides.items():
# Reject stale field names before setattr silently accepts them.
assert key in valid, f"{key} is not a ServerArgs field"
setattr(server_args, key, value)
return server_args
def test_validated_architectures_allowed(self):
for architecture in (
"DeepseekV3ForCausalLM",
"DeepseekV4ForCausalLM",
"Qwen3MoeForCausalLM",
):
args = self._args(moe_runner_backend="deep_gemm")
args.model_config.hf_config.architectures = [architecture]
args._handle_a2a_moe()
def test_unvalidated_and_missing_architectures_rejected(self):
for architectures in (
["Qwen2MoeForCausalLM"],
["Qwen3_5MoeForCausalLM"],
[],
None,
):
args = self._args(moe_runner_backend="deep_gemm")
args.model_config.hf_config.architectures = architectures
with self.assertRaisesRegex(ValueError, "not validated"):
args._handle_a2a_moe()
def test_instance_connector_rejected(self):
args = self._args(
model_path="instance://worker/model",
moe_runner_backend="deep_gemm",
)
with self.assertRaisesRegex(ValueError, "instance connector"):
args._handle_a2a_moe()
def test_deterministic_inference_rejected(self):
args = self._args(
moe_runner_backend="deep_gemm",
enable_deterministic_inference=True,
)
with self.assertRaisesRegex(ValueError, "deterministic sorting"):
args._handle_a2a_moe()
def test_rl_on_policy_deterministic_inference_rejected(self):
args = self._args(
moe_runner_backend="deep_gemm",
rl_on_policy_target="fsdp",
)
args.model_config.hf_config.architectures = ["Qwen3MoeForCausalLM"]
with (
envs.SGLANG_VLM_CACHE_SIZE_MB.override(envs.SGLANG_VLM_CACHE_SIZE_MB.get()),
envs.SGLANG_ENABLE_DETERMINISTIC_INFERENCE.override(
envs.SGLANG_ENABLE_DETERMINISTIC_INFERENCE.get()
),
):
args._handle_deterministic_inference()
with self.assertRaisesRegex(ValueError, "deterministic sorting"):
args._handle_a2a_moe()
def test_deterministic_inference_does_not_affect_legacy_deepep(self):
args = self._args(
moe_a2a_backend="deepep",
moe_runner_backend="deep_gemm",
enable_deterministic_inference=True,
)
args._handle_a2a_moe()
def test_runner_restored_by_declaration_fails_fast(self):
# Validate the declaration-resolved runner rather than the raw field.
args = self._args(moe_runner_backend="auto")
args._resolved_overrides = [
("test_mxfp8", {"moe_runner_backend": "flashinfer_trtllm"})
]
with self.assertRaises(ValueError):
args._handle_a2a_moe()
def test_declarations_resolve_ep_size_and_fusion(self):
from sglang.srt.arg_groups.overrides import resolved_view
args = self._args(moe_runner_backend="auto", tp_size=2)
args._handle_a2a_moe()
self.assertEqual(resolved_view(args).ep_size, args.tp_size)
self.assertTrue(resolved_view(args).disable_shared_experts_fusion)
def test_auto_runner_defaults_to_deep_gemm(self):
from sglang.srt.arg_groups.overrides import resolved_view
args = self._args(moe_runner_backend="auto")
args._handle_a2a_moe()
self.assertEqual(resolved_view(args).moe_runner_backend, "deep_gemm")
def test_unsupported_runner_rejected(self):
args = self._args(moe_runner_backend="flashinfer_trtllm")
with self.assertRaises(ValueError):
args._handle_a2a_moe()
def test_triton_runner_rejected(self):
args = self._args(moe_runner_backend="triton")
with self.assertRaises(ValueError):
args._handle_a2a_moe()
def test_decode_graph_stays_enabled_in_both_comm_modes(self):
for mode in ("direct", "hybrid"):
args = self._args(moe_runner_backend="deep_gemm", deepep_v2_mode=mode)
args._handle_a2a_moe()
self.assertEqual(args.cuda_graph_config.decode.backend, Backend.FULL)
self.assertEqual(args.cuda_graph_config.prefill.backend, Backend.DISABLED)
def test_two_batch_overlap_rejected(self):
args = self._args(moe_runner_backend="deep_gemm", enable_two_batch_overlap=True)
with self.assertRaises(ValueError):
args._handle_a2a_moe()
def test_speculative_draft_backend_rejected(self):
for main_backend in ("none", "deepep", "deepep_v2"):
args = self._args(
moe_a2a_backend=main_backend,
moe_runner_backend="deep_gemm",
speculative_moe_a2a_backend="deepep_v2",
)
with self.assertRaisesRegex(ValueError, "speculative draft backend"):
args._validate_deepep_v2_speculative_draft()
def test_inherited_speculative_draft_backend_rejected(self):
args = self._args(
moe_runner_backend="deep_gemm",
speculative_algorithm="EAGLE",
)
with self.assertRaisesRegex(ValueError, "speculative draft backend"):
args._validate_deepep_v2_speculative_draft()
def test_ngram_does_not_inherit_a_draft_backend(self):
args = self._args(
moe_runner_backend="deep_gemm",
speculative_algorithm="NGRAM",
)
args._validate_deepep_v2_speculative_draft()
def test_explicit_legacy_speculative_backend_allowed(self):
args = self._args(
moe_runner_backend="deep_gemm",
speculative_algorithm="EAGLE",
speculative_moe_a2a_backend="deepep",
)
args._validate_deepep_v2_speculative_draft()
def test_resolved_legacy_speculative_backend_allowed(self):
args = self._args(
moe_runner_backend="deep_gemm",
speculative_algorithm="EAGLE",
)
args._resolved_overrides = [
(
"test_speculative_backend",
{"speculative_moe_a2a_backend": "deepep"},
)
]
args._validate_deepep_v2_speculative_draft()
def test_prefill_chunk_exceeding_cap_rejected(self):
args = self._args(moe_runner_backend="deep_gemm", chunked_prefill_size=2048)
with envs.SGLANG_DEEPEP_V2_NUM_MAX_DISPATCH_TOKENS_PER_RANK.override(1024):
with self.assertRaisesRegex(ValueError, "NUM_MAX_DISPATCH_TOKENS_PER_RANK"):
args._validate_deepep_v2_dispatch_token_budget()
def test_prefill_chunk_at_cap_boundary_accepted(self):
args = self._args(moe_runner_backend="deep_gemm", chunked_prefill_size=1024)
with envs.SGLANG_DEEPEP_V2_NUM_MAX_DISPATCH_TOKENS_PER_RANK.override(1024):
args._validate_deepep_v2_dispatch_token_budget()
def test_dynamic_chunking_probe_is_included(self):
args = self._args(
chunked_prefill_size=1024,
max_prefill_tokens=1024,
enable_dynamic_chunking=True,
pp_size=2,
disaggregation_mode="prefill",
)
with envs.SGLANG_DEEPEP_V2_NUM_MAX_DISPATCH_TOKENS_PER_RANK.override(1024):
with self.assertRaisesRegex(ValueError, "required=1280"):
args._validate_deepep_v2_dispatch_token_budget()
def test_disabled_chunking_uses_max_prefill_tokens(self):
for disabled in (None, 0, -1):
args = self._args(
chunked_prefill_size=disabled,
max_prefill_tokens=1024,
disaggregation_mode="prefill",
)
with envs.SGLANG_DEEPEP_V2_NUM_MAX_DISPATCH_TOKENS_PER_RANK.override(128):
with self.assertRaisesRegex(ValueError, "required=1024"):
args._validate_deepep_v2_dispatch_token_budget()
def test_decode_role_skips_prefill_capacity(self):
args = self._args(
chunked_prefill_size=4096,
disaggregation_mode="decode",
max_running_requests=32,
dp_size=1,
)
with envs.SGLANG_DEEPEP_V2_NUM_MAX_DISPATCH_TOKENS_PER_RANK.override(128):
args._validate_deepep_v2_dispatch_token_budget()
def test_decode_graph_capacity_boundaries(self):
for max_bs, raises in ((128, False), (129, True)):
args = self._args(
disaggregation_mode="decode",
max_running_requests=None,
)
args.cuda_graph_config.decode.max_bs = max_bs
with envs.SGLANG_DEEPEP_V2_NUM_MAX_DISPATCH_TOKENS_PER_RANK.override(128):
if raises:
with self.assertRaisesRegex(ValueError, "decode CUDA graph"):
args._validate_deepep_v2_dispatch_token_budget()
else:
args._validate_deepep_v2_dispatch_token_budget()
def test_dp_attention_divides_max_running_requests_per_rank(self):
args = self._args(
disaggregation_mode="decode",
max_running_requests=256,
tp_size=8,
dp_size=8,
enable_dp_attention=True,
)
with envs.SGLANG_DEEPEP_V2_NUM_MAX_DISPATCH_TOKENS_PER_RANK.override(128):
args._validate_deepep_v2_dispatch_token_budget()
def test_tp_only_max_running_requests_is_not_divided(self):
args = self._args(
disaggregation_mode="decode",
max_running_requests=256,
tp_size=8,
dp_size=1,
enable_dp_attention=False,
)
with envs.SGLANG_DEEPEP_V2_NUM_MAX_DISPATCH_TOKENS_PER_RANK.override(128):
with self.assertRaisesRegex(ValueError, "decode CUDA graph"):
args._validate_deepep_v2_dispatch_token_budget()
def test_memory_derived_eager_pool_remains_runtime_validated(self):
args = self._args(
disaggregation_mode="decode",
max_running_requests=None,
)
args.cuda_graph_config.decode.backend = Backend.DISABLED
with envs.SGLANG_DEEPEP_V2_NUM_MAX_DISPATCH_TOKENS_PER_RANK.override(1):
args._validate_deepep_v2_dispatch_token_budget()
def test_speculative_decode_width_is_included(self):
args = self._args(
disaggregation_mode="decode",
speculative_algorithm="EAGLE",
speculative_num_draft_tokens=8,
max_running_requests=256,
dp_size=8,
enable_dp_attention=True,
)
with envs.SGLANG_DEEPEP_V2_NUM_MAX_DISPATCH_TOKENS_PER_RANK.override(128):
with self.assertRaisesRegex(ValueError, "tokens/request=8"):
args._validate_deepep_v2_dispatch_token_budget()
def test_adaptive_speculative_uses_widest_candidate(self):
args = self._args(
disaggregation_mode="decode",
speculative_algorithm="EAGLE",
speculative_num_draft_tokens=4,
speculative_adaptive=True,
max_running_requests=128,
dp_size=8,
enable_dp_attention=True,
)
with patch.object(
ServerArgs,
"max_speculative_num_draft_tokens",
new=property(lambda _self: 16),
):
with envs.SGLANG_DEEPEP_V2_NUM_MAX_DISPATCH_TOKENS_PER_RANK.override(128):
with self.assertRaisesRegex(ValueError, "tokens/request=16"):
args._validate_deepep_v2_dispatch_token_budget()
def test_prefill_role_skips_decode_capacity(self):
args = self._args(
disaggregation_mode="prefill",
chunked_prefill_size=64,
max_running_requests=8192,
)
with envs.SGLANG_DEEPEP_V2_NUM_MAX_DISPATCH_TOKENS_PER_RANK.override(128):
args._validate_deepep_v2_dispatch_token_budget()
def test_other_backend_skips_capacity_validation(self):
args = self._args(
moe_a2a_backend="deepep",
chunked_prefill_size=4096,
max_running_requests=4096,
)
with envs.SGLANG_DEEPEP_V2_NUM_MAX_DISPATCH_TOKENS_PER_RANK.override(1):
args._validate_deepep_v2_dispatch_token_budget()
def test_capacity_validation_uses_resolved_backend(self):
args = self._args(chunked_prefill_size=4096)
args._resolved_overrides = [("test", {"moe_a2a_backend": "deepep"})]
with envs.SGLANG_DEEPEP_V2_NUM_MAX_DISPATCH_TOKENS_PER_RANK.override(1):
args._validate_deepep_v2_dispatch_token_budget()
class TestHandleCrashDumpEnv(CustomTestCase):
_COREDUMP_ENV_KEYS = (
"CUDA_ENABLE_COREDUMP_ON_EXCEPTION",