Add deterministic inference for eagle parity test (#30026)
This commit is contained in:
@@ -161,6 +161,11 @@ class SpecParityKit:
|
|||||||
"--dtype",
|
"--dtype",
|
||||||
cls.dtype,
|
cls.dtype,
|
||||||
*(["--trust-remote-code"] if cls.trust_remote_code else []),
|
*(["--trust-remote-code"] if cls.trust_remote_code else []),
|
||||||
|
*(
|
||||||
|
["--enable-deterministic-inference"]
|
||||||
|
if cls.enable_deterministic_inference
|
||||||
|
else []
|
||||||
|
),
|
||||||
],
|
],
|
||||||
)
|
)
|
||||||
try:
|
try:
|
||||||
|
|||||||
@@ -69,6 +69,7 @@ class SpecEagleServerBase(CustomTestCase):
|
|||||||
# Launch with --enable-return-hidden-states so SpecHiddenStatesKit can probe
|
# Launch with --enable-return-hidden-states so SpecHiddenStatesKit can probe
|
||||||
# per-request hidden states; per-request gated, so other requests don't pay.
|
# per-request hidden states; per-request gated, so other requests don't pay.
|
||||||
enable_return_hidden_states = False
|
enable_return_hidden_states = False
|
||||||
|
enable_deterministic_inference = False
|
||||||
|
|
||||||
# -- extras --
|
# -- extras --
|
||||||
# env_overrides: (env_var_obj, value) pairs applied only around launch.
|
# env_overrides: (env_var_obj, value) pairs applied only around launch.
|
||||||
@@ -109,6 +110,8 @@ class SpecEagleServerBase(CustomTestCase):
|
|||||||
args.append("--trust-remote-code")
|
args.append("--trust-remote-code")
|
||||||
if cls.enable_return_hidden_states:
|
if cls.enable_return_hidden_states:
|
||||||
args.append("--enable-return-hidden-states")
|
args.append("--enable-return-hidden-states")
|
||||||
|
if cls.enable_deterministic_inference:
|
||||||
|
args.append("--enable-deterministic-inference")
|
||||||
if cls.cuda_graph_max_bs_decode is not None:
|
if cls.cuda_graph_max_bs_decode is not None:
|
||||||
args += ["--cuda-graph-max-bs-decode", str(cls.cuda_graph_max_bs_decode)]
|
args += ["--cuda-graph-max-bs-decode", str(cls.cuda_graph_max_bs_decode)]
|
||||||
args += [str(a) for a in cls.extra_args]
|
args += [str(a) for a in cls.extra_args]
|
||||||
|
|||||||
@@ -13,12 +13,7 @@ from sglang.test.kits.spec_server_kits import SpecParityKit
|
|||||||
from sglang.test.server_fixtures.spec_eagle_fixture import Eagle3Base
|
from sglang.test.server_fixtures.spec_eagle_fixture import Eagle3Base
|
||||||
|
|
||||||
register_cuda_ci(est_time=360, stage="base-b", runner_config="1-gpu-large")
|
register_cuda_ci(est_time=360, stage="base-b", runner_config="1-gpu-large")
|
||||||
register_xpu_ci(
|
register_xpu_ci(est_time=360, stage="stage-b", runner_config="1-gpu-xpu")
|
||||||
est_time=360,
|
|
||||||
stage="stage-b",
|
|
||||||
runner_config="1-gpu-xpu",
|
|
||||||
disabled="EAGLE3 numerical parity mismatches on XPU",
|
|
||||||
)
|
|
||||||
|
|
||||||
_is_xpu = is_xpu()
|
_is_xpu = is_xpu()
|
||||||
|
|
||||||
@@ -26,6 +21,7 @@ _is_xpu = is_xpu()
|
|||||||
class _Eagle3ParityBase(Eagle3Base):
|
class _Eagle3ParityBase(Eagle3Base):
|
||||||
"""Shared knobs for EAGLE3 parity variants; no test methods."""
|
"""Shared knobs for EAGLE3 parity variants; no test methods."""
|
||||||
|
|
||||||
|
enable_deterministic_inference = True
|
||||||
env_overrides = ((envs.SGLANG_ENABLE_STRICT_MEM_CHECK_DURING_BUSY, 1),)
|
env_overrides = ((envs.SGLANG_ENABLE_STRICT_MEM_CHECK_DURING_BUSY, 1),)
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user