[Spec] Support MegaMoE for DSpark under dp attention (#34844)
Co-authored-by: weireweire <20922698+weireweire@users.noreply.github.com>
This commit is contained in:
@@ -286,11 +286,24 @@ def _handle_dspark(server_args: ServerArgs) -> None:
|
|||||||
if server_args.enable_dp_attention and server_args.dp_size > 1:
|
if server_args.enable_dp_attention and server_args.dp_size > 1:
|
||||||
if not server_args.enable_dp_lm_head:
|
if not server_args.enable_dp_lm_head:
|
||||||
raise ValueError("DSpark with dp attention requires --enable-dp-lm-head.")
|
raise ValueError("DSpark with dp attention requires --enable-dp-lm-head.")
|
||||||
if not _is_npu and server_args.moe_a2a_backend != "none":
|
if not _is_npu and server_args.moe_a2a_backend not in ("none", "megamoe"):
|
||||||
raise ValueError(
|
raise ValueError(
|
||||||
"DSpark with dp attention only supports the built-in TP MoE "
|
"DSpark with dp attention supports moe_a2a_backend 'none' "
|
||||||
f"(moe_a2a_backend='none'), got {server_args.moe_a2a_backend!r}."
|
"(built-in TP MoE) or 'megamoe', got "
|
||||||
|
f"{server_args.moe_a2a_backend!r}."
|
||||||
)
|
)
|
||||||
|
if not _is_npu and server_args.moe_a2a_backend != "none":
|
||||||
|
from sglang.srt.speculative.ragged_verify import (
|
||||||
|
RaggedVerifyMode,
|
||||||
|
read_ragged_verify_mode,
|
||||||
|
)
|
||||||
|
|
||||||
|
if read_ragged_verify_mode() is not RaggedVerifyMode.STATIC:
|
||||||
|
raise ValueError(
|
||||||
|
"DSpark with dp attention + "
|
||||||
|
f"moe_a2a_backend={server_args.moe_a2a_backend!r} requires "
|
||||||
|
"SGLANG_RAGGED_VERIFY_MODE=static."
|
||||||
|
)
|
||||||
if server_args.attn_cp_size > 1:
|
if server_args.attn_cp_size > 1:
|
||||||
raise ValueError(
|
raise ValueError(
|
||||||
"DSpark with dp attention does not support context parallel "
|
"DSpark with dp attention does not support context parallel "
|
||||||
|
|||||||
@@ -5,6 +5,7 @@ from sglang.srt.arg_groups.speculative_hook import (
|
|||||||
_handle_dspark,
|
_handle_dspark,
|
||||||
_target_checkpoint_bundles_dspark_draft,
|
_target_checkpoint_bundles_dspark_draft,
|
||||||
)
|
)
|
||||||
|
from sglang.srt.environ import envs
|
||||||
from sglang.srt.server_args import ServerArgs
|
from sglang.srt.server_args import ServerArgs
|
||||||
from sglang.test.ci.ci_register import register_cpu_ci
|
from sglang.test.ci.ci_register import register_cpu_ci
|
||||||
from sglang.test.test_utils import CustomTestCase
|
from sglang.test.test_utils import CustomTestCase
|
||||||
@@ -84,5 +85,34 @@ class TestDsparkDraftPathDefaulting(CustomTestCase):
|
|||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
|
class TestDsparkDpAttentionMoeA2aGate(CustomTestCase):
|
||||||
|
"""Gate contract for DSpark + dp attention + MoE a2a backends."""
|
||||||
|
|
||||||
|
def _dp_server_args(self, *, moe_a2a_backend: str) -> ServerArgs:
|
||||||
|
server_args = _make_dspark_server_args(
|
||||||
|
model_path=_BUNDLED_MODEL_PATH, hf_config=_bundled_hf_config()
|
||||||
|
)
|
||||||
|
server_args.enable_dp_attention = True
|
||||||
|
server_args.enable_dp_lm_head = True
|
||||||
|
server_args.dp_size = 2
|
||||||
|
server_args.tp_size = 2
|
||||||
|
server_args.moe_a2a_backend = moe_a2a_backend
|
||||||
|
return server_args
|
||||||
|
|
||||||
|
def test_only_megamoe_is_admitted(self):
|
||||||
|
"""Both sides of the allowlist: megamoe passes, others raise by name."""
|
||||||
|
with envs.SGLANG_RAGGED_VERIFY_MODE.override("static"):
|
||||||
|
_handle_dspark(self._dp_server_args(moe_a2a_backend="megamoe"))
|
||||||
|
for backend in ("deepep", "pplx"):
|
||||||
|
with self.assertRaisesRegex(ValueError, backend):
|
||||||
|
_handle_dspark(self._dp_server_args(moe_a2a_backend=backend))
|
||||||
|
|
||||||
|
def test_a2a_backend_with_compact_verify_mode_raises(self):
|
||||||
|
server_args = self._dp_server_args(moe_a2a_backend="megamoe")
|
||||||
|
with envs.SGLANG_RAGGED_VERIFY_MODE.override("compact"):
|
||||||
|
with self.assertRaisesRegex(ValueError, "static"):
|
||||||
|
_handle_dspark(server_args)
|
||||||
|
|
||||||
|
|
||||||
if __name__ == "__main__":
|
if __name__ == "__main__":
|
||||||
unittest.main()
|
unittest.main()
|
||||||
|
|||||||
Reference in New Issue
Block a user