[CI] Collapse the EAGLE launch matrix and the scoring engine boots on the per-commit runners (#33756)
This commit is contained in:
@@ -16,7 +16,7 @@ from sglang.multimodal_gen.runtime.models.parameter import (
|
|||||||
PerTensorScaleParameter,
|
PerTensorScaleParameter,
|
||||||
)
|
)
|
||||||
from sglang.srt.layers.quantization.utils import is_layer_skipped
|
from sglang.srt.layers.quantization.utils import is_layer_skipped
|
||||||
from sglang.srt.utils import is_hip, mxfp_supported
|
from sglang.srt.utils import is_gfx95_supported, is_hip
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
_is_hip = is_hip()
|
_is_hip = is_hip()
|
||||||
@@ -70,7 +70,7 @@ class Mxfp4Config(QuantizationConfig):
|
|||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
def get_min_capability(cls) -> int:
|
def get_min_capability(cls) -> int:
|
||||||
return 95 # gfx95x, Note: mxfp_supported() is a better check
|
return 95 # gfx95x, Note: is_gfx95_supported() is a better check
|
||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
def get_config_filenames(cls) -> list[str]:
|
def get_config_filenames(cls) -> list[str]:
|
||||||
@@ -165,7 +165,7 @@ class Mxfp4LinearMethod(LinearMethodBase):
|
|||||||
- Packed uint8 (2 FP4 values per byte)
|
- Packed uint8 (2 FP4 values per byte)
|
||||||
- E8M0 scales (one per 32-element block)
|
- E8M0 scales (one per 32-element block)
|
||||||
"""
|
"""
|
||||||
if not mxfp_supported():
|
if not is_gfx95_supported():
|
||||||
platform = "unknown"
|
platform = "unknown"
|
||||||
if _is_hip:
|
if _is_hip:
|
||||||
try:
|
try:
|
||||||
@@ -217,7 +217,7 @@ class Mxfp4LinearMethod(LinearMethodBase):
|
|||||||
bias: Optional[torch.Tensor] = None,
|
bias: Optional[torch.Tensor] = None,
|
||||||
) -> torch.Tensor:
|
) -> torch.Tensor:
|
||||||
|
|
||||||
if not mxfp_supported():
|
if not is_gfx95_supported():
|
||||||
raise RuntimeError(
|
raise RuntimeError(
|
||||||
"MXFP4 inference requires ROCm and MI350+ (gfx95x). "
|
"MXFP4 inference requires ROCm and MI350+ (gfx95x). "
|
||||||
"Current platform not supported."
|
"Current platform not supported."
|
||||||
|
|||||||
@@ -58,13 +58,12 @@ from sglang.srt.utils import (
|
|||||||
cpu_has_amx_support,
|
cpu_has_amx_support,
|
||||||
is_cpu,
|
is_cpu,
|
||||||
is_cuda,
|
is_cuda,
|
||||||
is_hip,
|
is_gfx95_supported,
|
||||||
is_mps,
|
is_mps,
|
||||||
is_npu,
|
is_npu,
|
||||||
mxfp_supported,
|
|
||||||
)
|
)
|
||||||
|
|
||||||
_is_mxfp_supported = mxfp_supported()
|
_is_gfx95_supported = is_gfx95_supported()
|
||||||
|
|
||||||
if TYPE_CHECKING:
|
if TYPE_CHECKING:
|
||||||
from sglang.srt.layers.moe.topk import TopKOutput
|
from sglang.srt.layers.moe.topk import TopKOutput
|
||||||
@@ -102,7 +101,7 @@ BASE_QUANTIZATION_METHODS: Dict[str, Type[QuantizationConfig]] = {
|
|||||||
}
|
}
|
||||||
|
|
||||||
|
|
||||||
if is_cpu() or is_cuda() or (_is_mxfp_supported and is_hip()):
|
if is_cpu() or is_cuda() or _is_gfx95_supported:
|
||||||
BASE_QUANTIZATION_METHODS.update(
|
BASE_QUANTIZATION_METHODS.update(
|
||||||
{
|
{
|
||||||
"mxfp4": Mxfp4Config,
|
"mxfp4": Mxfp4Config,
|
||||||
|
|||||||
@@ -59,7 +59,6 @@ from sglang.srt.utils import (
|
|||||||
is_sm100_supported,
|
is_sm100_supported,
|
||||||
is_sm120_supported,
|
is_sm120_supported,
|
||||||
is_triton_kernels_available,
|
is_triton_kernels_available,
|
||||||
mxfp_supported,
|
|
||||||
next_power_of_2,
|
next_power_of_2,
|
||||||
round_up,
|
round_up,
|
||||||
set_weight_attrs,
|
set_weight_attrs,
|
||||||
@@ -256,7 +255,7 @@ class Mxfp4Config(QuantizationConfig):
|
|||||||
is_checkpoint_mxfp4_serialized = "mxfp4" in quant_method
|
is_checkpoint_mxfp4_serialized = "mxfp4" in quant_method
|
||||||
|
|
||||||
if _is_hip:
|
if _is_hip:
|
||||||
if mxfp_supported():
|
if is_gfx95_supported():
|
||||||
return cls(
|
return cls(
|
||||||
is_checkpoint_mxfp4_serialized=is_checkpoint_mxfp4_serialized
|
is_checkpoint_mxfp4_serialized=is_checkpoint_mxfp4_serialized
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -16,7 +16,7 @@ from sglang.srt.layers.quantization.fp8 import Fp8Config, Fp8LinearMethod
|
|||||||
from sglang.srt.layers.quantization.online_quantization import CopyNumelCounter
|
from sglang.srt.layers.quantization.online_quantization import CopyNumelCounter
|
||||||
from sglang.srt.layers.quantization.quark.schemes import QuarkLinearScheme
|
from sglang.srt.layers.quantization.quark.schemes import QuarkLinearScheme
|
||||||
from sglang.srt.utils import is_hip
|
from sglang.srt.utils import is_hip
|
||||||
from sglang.srt.utils.common import direct_register_custom_op, mxfp_supported
|
from sglang.srt.utils.common import direct_register_custom_op, is_gfx95_supported
|
||||||
|
|
||||||
_is_hip = is_hip()
|
_is_hip = is_hip()
|
||||||
if _is_hip:
|
if _is_hip:
|
||||||
@@ -180,7 +180,7 @@ class QuarkW4A4MXFP4(QuarkLinearScheme):
|
|||||||
self.dequantization_config = dequantization_config
|
self.dequantization_config = dequantization_config
|
||||||
|
|
||||||
if not self.is_checkpoint_mxfp4_serialized:
|
if not self.is_checkpoint_mxfp4_serialized:
|
||||||
if not mxfp_supported():
|
if not is_gfx95_supported():
|
||||||
raise NotImplementedError(
|
raise NotImplementedError(
|
||||||
"Online MXFP4 quantization requires an AMD ROCm device with "
|
"Online MXFP4 quantization requires an AMD ROCm device with "
|
||||||
"FP4 hardware support (gfx95x, e.g. MI355x)."
|
"FP4 hardware support (gfx95x, e.g. MI355x)."
|
||||||
|
|||||||
@@ -24,7 +24,7 @@ from sglang.srt.utils import (
|
|||||||
is_hip,
|
is_hip,
|
||||||
set_weight_attrs,
|
set_weight_attrs,
|
||||||
)
|
)
|
||||||
from sglang.srt.utils.common import mxfp_supported
|
from sglang.srt.utils.common import is_gfx95_supported
|
||||||
|
|
||||||
if TYPE_CHECKING:
|
if TYPE_CHECKING:
|
||||||
from sglang.srt.layers.moe.token_dispatcher import (
|
from sglang.srt.layers.moe.token_dispatcher import (
|
||||||
@@ -79,7 +79,7 @@ class QuarkW4A4MXFp4MoE(QuarkMoEScheme):
|
|||||||
self.with_bias = False
|
self.with_bias = False
|
||||||
|
|
||||||
if not self.is_checkpoint_mxfp4_serialized:
|
if not self.is_checkpoint_mxfp4_serialized:
|
||||||
if not mxfp_supported():
|
if not is_gfx95_supported():
|
||||||
raise NotImplementedError(
|
raise NotImplementedError(
|
||||||
"Online MXFP4 quantization for MoE layers requires an AMD ROCm "
|
"Online MXFP4 quantization for MoE layers requires an AMD ROCm "
|
||||||
"device with FP4 hardware support (gfx95x, e.g. MI355x)."
|
"device with FP4 hardware support (gfx95x, e.g. MI355x)."
|
||||||
|
|||||||
@@ -1023,21 +1023,11 @@ def set_cuda_arch():
|
|||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
def mxfp_supported():
|
|
||||||
"""
|
|
||||||
Returns whether the current platform supports MX types.
|
|
||||||
"""
|
|
||||||
if torch.version.hip:
|
|
||||||
gcn_arch = torch.cuda.get_device_properties(0).gcnArchName
|
|
||||||
return any(gfx in gcn_arch for gfx in ["gfx95"])
|
|
||||||
else:
|
|
||||||
return False
|
|
||||||
|
|
||||||
|
|
||||||
@lru_cache(maxsize=1)
|
@lru_cache(maxsize=1)
|
||||||
def is_gfx95_supported():
|
def is_gfx95_supported():
|
||||||
"""
|
"""Whether the device is an AMD gfx95 GPU (the MX-capable ROCm arch).
|
||||||
Returns whether the current platform supports MX types.
|
|
||||||
|
False on every non-HIP build, so callers do not need their own is_hip().
|
||||||
"""
|
"""
|
||||||
if torch.version.hip:
|
if torch.version.hip:
|
||||||
gcn_arch = torch.cuda.get_device_properties(0).gcnArchName
|
gcn_arch = torch.cuda.get_device_properties(0).gcnArchName
|
||||||
|
|||||||
@@ -386,8 +386,8 @@ class SpecLogprobKit:
|
|||||||
with ThreadPoolExecutor(8) as executor:
|
with ThreadPoolExecutor(8) as executor:
|
||||||
list(executor.map(func, args))
|
list(executor.map(func, args))
|
||||||
|
|
||||||
def test_logprob_spec_v2_match(self):
|
def test_logprob_decode_match_prefill(self):
|
||||||
"""Verify spec v2 decode logprobs match prefill scoring logprobs."""
|
"""Decode logprobs from the spec path must match prefill scoring."""
|
||||||
top_k = 5
|
top_k = 5
|
||||||
probe_token_ids = [1, 2, 10, 100, 1000]
|
probe_token_ids = [1, 2, 10, 100, 1000]
|
||||||
prompts = [
|
prompts = [
|
||||||
|
|||||||
@@ -50,13 +50,12 @@ class TestStep3p5FlashChainMTP(GSM8KMixin, DefaultServerBase):
|
|||||||
gsm8k_accuracy_thres = 0.83
|
gsm8k_accuracy_thres = 0.83
|
||||||
gsm8k_accept_length_thres = 2.6
|
gsm8k_accept_length_thres = 2.6
|
||||||
|
|
||||||
def test_logprob_spec_v2_match(self):
|
def test_logprob_decode_match_prefill(self):
|
||||||
"""Verify spec v2 decode logprobs match prefill scoring logprobs.
|
"""Decode logprobs from the spec path must match prefill scoring.
|
||||||
|
|
||||||
Generate tokens with chain MTP spec v2, then score the same sequence
|
Generate tokens with chain MTP, then score the same sequence via
|
||||||
via prefill-only (no speculation). The two sets of logprobs should be
|
prefill-only (no speculation). The two sets of logprobs should be
|
||||||
close, validating that spec v2 + multi-layer EAGLE computes logprobs
|
close, validating that multi-layer EAGLE computes logprobs correctly.
|
||||||
correctly.
|
|
||||||
"""
|
"""
|
||||||
requests.get(self.base_url + "/flush_cache")
|
requests.get(self.base_url + "/flush_cache")
|
||||||
|
|
||||||
|
|||||||
@@ -16,9 +16,9 @@ from sglang.test.test_utils import (
|
|||||||
popen_launch_server,
|
popen_launch_server,
|
||||||
)
|
)
|
||||||
|
|
||||||
register_cuda_ci(est_time=222, stage="base-b", runner_config="1-gpu-small")
|
register_cuda_ci(est_time=165, stage="base-b", runner_config="1-gpu-small")
|
||||||
register_amd_ci(
|
register_amd_ci(
|
||||||
est_time=186,
|
est_time=140,
|
||||||
suite="stage-b-test-1-gpu-small-amd",
|
suite="stage-b-test-1-gpu-small-amd",
|
||||||
disabled="see https://github.com/sgl-project/sglang/issues/11127",
|
disabled="see https://github.com/sgl-project/sglang/issues/11127",
|
||||||
)
|
)
|
||||||
@@ -307,52 +307,5 @@ class TestOpenAIServerWithEAGLEAndHiddenStatesEnabled(
|
|||||||
kill_process_tree(cls.process.pid)
|
kill_process_tree(cls.process.pid)
|
||||||
|
|
||||||
|
|
||||||
class TestOpenAIServerWithEAGLE3AndHiddenStatesEnabled(
|
|
||||||
CustomTestCase, BaseTestOpenAIServerWithHiddenStates
|
|
||||||
):
|
|
||||||
@classmethod
|
|
||||||
def setUpClass(cls):
|
|
||||||
cls.model = "meta-llama/Llama-3.1-8B-Instruct"
|
|
||||||
cls.base_url = DEFAULT_URL_FOR_TEST
|
|
||||||
cls.api_key = "sk-123456"
|
|
||||||
cls.speculative_algorithm = "EAGLE3"
|
|
||||||
cls.speculative_draft_model = "jamesliu1/sglang-EAGLE3-Llama-3.1-Instruct-8B"
|
|
||||||
cls.process = popen_launch_server(
|
|
||||||
cls.model,
|
|
||||||
cls.base_url,
|
|
||||||
timeout=DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH,
|
|
||||||
other_args=[
|
|
||||||
"--speculative-algorithm",
|
|
||||||
cls.speculative_algorithm,
|
|
||||||
"--speculative-draft-model-path",
|
|
||||||
cls.speculative_draft_model,
|
|
||||||
"--speculative-num-steps",
|
|
||||||
5,
|
|
||||||
"--speculative-eagle-topk",
|
|
||||||
16,
|
|
||||||
"--speculative-num-draft-tokens",
|
|
||||||
64,
|
|
||||||
"--mem-fraction-static",
|
|
||||||
0.7,
|
|
||||||
"--chunked-prefill-size",
|
|
||||||
128,
|
|
||||||
"--max-running-requests",
|
|
||||||
8,
|
|
||||||
"--dtype",
|
|
||||||
"float16",
|
|
||||||
"--enable-return-hidden-states",
|
|
||||||
],
|
|
||||||
)
|
|
||||||
cls.base_url += "/v1"
|
|
||||||
cls.tokenizer = get_tokenizer(cls.model)
|
|
||||||
cls.return_hidden_states = [False, True]
|
|
||||||
cls.use_list_input = [True, False]
|
|
||||||
cls.parallel_sample_nums = [1]
|
|
||||||
|
|
||||||
@classmethod
|
|
||||||
def tearDownClass(cls):
|
|
||||||
kill_process_tree(cls.process.pid)
|
|
||||||
|
|
||||||
|
|
||||||
if __name__ == "__main__":
|
if __name__ == "__main__":
|
||||||
unittest.main()
|
unittest.main()
|
||||||
|
|||||||
@@ -27,7 +27,7 @@ from sglang.test.test_utils import (
|
|||||||
CustomTestCase,
|
CustomTestCase,
|
||||||
)
|
)
|
||||||
|
|
||||||
register_cuda_ci(est_time=211, stage="base-b", runner_config="1-gpu-small")
|
register_cuda_ci(est_time=120, stage="base-b", runner_config="1-gpu-small")
|
||||||
|
|
||||||
TEST_MODEL_NAME = os.environ.get("TEST_MODEL_NAME", DEFAULT_SMALL_MODEL_NAME_FOR_TEST)
|
TEST_MODEL_NAME = os.environ.get("TEST_MODEL_NAME", DEFAULT_SMALL_MODEL_NAME_FOR_TEST)
|
||||||
TEST_CLASSIFICATION_BASE_MODEL = os.environ.get(
|
TEST_CLASSIFICATION_BASE_MODEL = os.environ.get(
|
||||||
@@ -145,16 +145,21 @@ class TestMultiItemScoringOptimization(CustomTestCase):
|
|||||||
|
|
||||||
|
|
||||||
class TestMultiItemScoringClassification(CustomTestCase):
|
class TestMultiItemScoringClassification(CustomTestCase):
|
||||||
"""Test MIS with classification models.
|
"""MIS on a classification model: basics, MIS-vs-single-item parity, score
|
||||||
|
distinctness / determinism / concurrency.
|
||||||
|
|
||||||
Uses a pre-trained Qwen3ForSequenceClassification model so that the
|
Pre-trained Qwen3ForSequenceClassification, so the head weights are
|
||||||
classification head weights are deterministic across Engine instances.
|
deterministic. One class rather than four because the CI harness demands an
|
||||||
|
idle GPU at every setUpClass -- splitting these means re-booting the same
|
||||||
|
two engines instead of sharing them. score() is stateless and the radix
|
||||||
|
cache is off, so sharing is safe.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
NUM_LABELS = _CLS_NUM_LABELS
|
NUM_LABELS = _CLS_NUM_LABELS
|
||||||
|
|
||||||
def setUp(self):
|
@classmethod
|
||||||
self.engine = Engine(
|
def setUpClass(cls):
|
||||||
|
cls.engine = Engine(
|
||||||
model_path=TEST_CLASSIFICATION_BASE_MODEL,
|
model_path=TEST_CLASSIFICATION_BASE_MODEL,
|
||||||
disable_radix_cache=True,
|
disable_radix_cache=True,
|
||||||
chunked_prefill_size=-1,
|
chunked_prefill_size=-1,
|
||||||
@@ -162,10 +167,17 @@ class TestMultiItemScoringClassification(CustomTestCase):
|
|||||||
attention_backend="flashinfer",
|
attention_backend="flashinfer",
|
||||||
mem_fraction_static=0.15,
|
mem_fraction_static=0.15,
|
||||||
)
|
)
|
||||||
|
cls.non_mis_engine = Engine(
|
||||||
|
model_path=TEST_CLASSIFICATION_BASE_MODEL,
|
||||||
|
disable_radix_cache=True,
|
||||||
|
mem_fraction_static=0.15,
|
||||||
|
)
|
||||||
|
|
||||||
def tearDown(self):
|
@classmethod
|
||||||
if self.engine is not None:
|
def tearDownClass(cls):
|
||||||
self.engine.shutdown()
|
for engine in (cls.engine, cls.non_mis_engine):
|
||||||
|
if engine is not None:
|
||||||
|
engine.shutdown()
|
||||||
torch.cuda.empty_cache()
|
torch.cuda.empty_cache()
|
||||||
|
|
||||||
def test_classification_mis_basic(self):
|
def test_classification_mis_basic(self):
|
||||||
@@ -203,13 +215,7 @@ class TestMultiItemScoringClassification(CustomTestCase):
|
|||||||
|
|
||||||
def test_classification_non_mis_fallback(self):
|
def test_classification_non_mis_fallback(self):
|
||||||
"""Classification model works correctly without --enable-mis."""
|
"""Classification model works correctly without --enable-mis."""
|
||||||
non_mis_engine = Engine(
|
scores = self.non_mis_engine.score(
|
||||||
model_path=TEST_CLASSIFICATION_BASE_MODEL,
|
|
||||||
disable_radix_cache=True,
|
|
||||||
mem_fraction_static=0.15,
|
|
||||||
)
|
|
||||||
try:
|
|
||||||
scores = non_mis_engine.score(
|
|
||||||
query="Test:", items=["A", "B"], apply_softmax=True
|
query="Test:", items=["A", "B"], apply_softmax=True
|
||||||
).scores
|
).scores
|
||||||
|
|
||||||
@@ -217,135 +223,6 @@ class TestMultiItemScoringClassification(CustomTestCase):
|
|||||||
for score_list in scores:
|
for score_list in scores:
|
||||||
self.assertEqual(len(score_list), self.NUM_LABELS)
|
self.assertEqual(len(score_list), self.NUM_LABELS)
|
||||||
self.assertAlmostEqual(sum(score_list), 1.0, places=5)
|
self.assertAlmostEqual(sum(score_list), 1.0, places=5)
|
||||||
finally:
|
|
||||||
non_mis_engine.shutdown()
|
|
||||||
torch.cuda.empty_cache()
|
|
||||||
|
|
||||||
|
|
||||||
class TestMultiItemScoringParity(CustomTestCase):
|
|
||||||
"""Test that MIS produces the same results as single-item scoring."""
|
|
||||||
|
|
||||||
@classmethod
|
|
||||||
def setUpClass(cls):
|
|
||||||
cls.engine_single = Engine(
|
|
||||||
model_path=TEST_MODEL_NAME,
|
|
||||||
disable_radix_cache=True,
|
|
||||||
log_level="error",
|
|
||||||
mem_fraction_static=0.15,
|
|
||||||
)
|
|
||||||
cls.engine_mis = Engine(
|
|
||||||
model_path=TEST_MODEL_NAME,
|
|
||||||
disable_radix_cache=True,
|
|
||||||
chunked_prefill_size=-1,
|
|
||||||
log_level="error",
|
|
||||||
enable_mis=True,
|
|
||||||
attention_backend="flashinfer",
|
|
||||||
mem_fraction_static=0.15,
|
|
||||||
)
|
|
||||||
|
|
||||||
@classmethod
|
|
||||||
def tearDownClass(cls):
|
|
||||||
if cls.engine_single is not None:
|
|
||||||
cls.engine_single.shutdown()
|
|
||||||
if cls.engine_mis is not None:
|
|
||||||
cls.engine_mis.shutdown()
|
|
||||||
torch.cuda.empty_cache()
|
|
||||||
|
|
||||||
def _compare_scores(
|
|
||||||
self, query, items, label_token_ids=None, apply_softmax=True, test_name=""
|
|
||||||
):
|
|
||||||
"""Compare MIS vs single-item scoring results."""
|
|
||||||
single_scores = self.engine_single.score(
|
|
||||||
query=query,
|
|
||||||
items=items,
|
|
||||||
label_token_ids=label_token_ids,
|
|
||||||
apply_softmax=apply_softmax,
|
|
||||||
).scores
|
|
||||||
|
|
||||||
mis_scores = self.engine_mis.score(
|
|
||||||
query=query,
|
|
||||||
items=items,
|
|
||||||
label_token_ids=label_token_ids,
|
|
||||||
apply_softmax=apply_softmax,
|
|
||||||
).scores
|
|
||||||
|
|
||||||
self.assertEqual(
|
|
||||||
len(mis_scores), len(single_scores), f"{test_name}: count mismatch"
|
|
||||||
)
|
|
||||||
for i, (ms, ss) in enumerate(zip(mis_scores, single_scores)):
|
|
||||||
self.assertEqual(len(ms), len(ss), f"{test_name}: item {i} length mismatch")
|
|
||||||
for j, (m, s) in enumerate(zip(ms, ss)):
|
|
||||||
self.assertAlmostEqual(
|
|
||||||
m,
|
|
||||||
s,
|
|
||||||
places=1,
|
|
||||||
msg=f"{test_name}: item {i} label {j}: MIS={m} vs single={s}",
|
|
||||||
)
|
|
||||||
|
|
||||||
def test_parity_basic(self):
|
|
||||||
tokenizer = AutoTokenizer.from_pretrained(TEST_MODEL_NAME)
|
|
||||||
query = "Rate this option:"
|
|
||||||
items = [" Option A", " Option B", " Option C"]
|
|
||||||
labels = [" good", " bad"]
|
|
||||||
label_ids = [tokenizer.encode(lb, add_special_tokens=False)[0] for lb in labels]
|
|
||||||
self._compare_scores(query, items, label_ids, test_name="basic")
|
|
||||||
|
|
||||||
def test_parity_tokenized_inputs(self):
|
|
||||||
tokenizer = AutoTokenizer.from_pretrained(TEST_MODEL_NAME)
|
|
||||||
query = "Rate this option:"
|
|
||||||
items = [" Option X", " Option Y"]
|
|
||||||
labels = [" good", " bad"]
|
|
||||||
query_ids = tokenizer.encode(query, add_special_tokens=False)
|
|
||||||
items_ids = [tokenizer.encode(i, add_special_tokens=False) for i in items]
|
|
||||||
label_ids = [tokenizer.encode(lb, add_special_tokens=False)[0] for lb in labels]
|
|
||||||
self._compare_scores(query_ids, items_ids, label_ids, test_name="tokenized")
|
|
||||||
|
|
||||||
def test_parity_without_softmax(self):
|
|
||||||
tokenizer = AutoTokenizer.from_pretrained(TEST_MODEL_NAME)
|
|
||||||
query = "The weather today is"
|
|
||||||
items = [" sunny", " cloudy", " rainy"]
|
|
||||||
labels = [" nice", " bad"]
|
|
||||||
label_ids = [tokenizer.encode(lb, add_special_tokens=False)[0] for lb in labels]
|
|
||||||
self._compare_scores(
|
|
||||||
query, items, label_ids, apply_softmax=False, test_name="no_softmax"
|
|
||||||
)
|
|
||||||
|
|
||||||
def test_parity_many_items(self):
|
|
||||||
tokenizer = AutoTokenizer.from_pretrained(TEST_MODEL_NAME)
|
|
||||||
query = "Rate this option from 1 to 5:"
|
|
||||||
items = [f" Option {i}" for i in range(10)]
|
|
||||||
labels = [" 1", " 2", " 3", " 4", " 5"]
|
|
||||||
label_ids = [tokenizer.encode(lb, add_special_tokens=False)[0] for lb in labels]
|
|
||||||
self._compare_scores(query, items, label_ids, test_name="many_items")
|
|
||||||
|
|
||||||
|
|
||||||
class TestMultiItemScoringClassificationParity(CustomTestCase):
|
|
||||||
"""Test that MIS multi-item batching matches single-item MIS scoring.
|
|
||||||
|
|
||||||
Both paths use the MIS engine (with delimiter tokens in the attention
|
|
||||||
context). The reference scores each item individually so each gets its
|
|
||||||
own forward pass; the batched path packs all items into one pass.
|
|
||||||
This isolates the MIS batching logic from the delimiter-presence effect.
|
|
||||||
"""
|
|
||||||
|
|
||||||
NUM_LABELS = _CLS_NUM_LABELS
|
|
||||||
|
|
||||||
@classmethod
|
|
||||||
def setUpClass(cls):
|
|
||||||
cls.engine = Engine(
|
|
||||||
model_path=TEST_CLASSIFICATION_BASE_MODEL,
|
|
||||||
disable_radix_cache=True,
|
|
||||||
chunked_prefill_size=-1,
|
|
||||||
enable_mis=True,
|
|
||||||
attention_backend="flashinfer",
|
|
||||||
mem_fraction_static=0.15,
|
|
||||||
)
|
|
||||||
|
|
||||||
@classmethod
|
|
||||||
def tearDownClass(cls):
|
|
||||||
if cls.engine is not None:
|
|
||||||
cls.engine.shutdown()
|
|
||||||
torch.cuda.empty_cache()
|
|
||||||
|
|
||||||
def _compare_scores(self, query, items, apply_softmax=True, test_name=""):
|
def _compare_scores(self, query, items, apply_softmax=True, test_name=""):
|
||||||
"""Compare MIS batched vs MIS single-item scoring results."""
|
"""Compare MIS batched vs MIS single-item scoring results."""
|
||||||
@@ -403,89 +280,6 @@ class TestMultiItemScoringClassificationParity(CustomTestCase):
|
|||||||
items = [f" Option {i}" for i in range(10)]
|
items = [f" Option {i}" for i in range(10)]
|
||||||
self._compare_scores(query, items, test_name="cls_many_items")
|
self._compare_scores(query, items, test_name="cls_many_items")
|
||||||
|
|
||||||
|
|
||||||
class TestMultiItemScoringClassificationMISvsNonMIS(CustomTestCase):
|
|
||||||
"""Test that MIS single-item approximates non-MIS single-item.
|
|
||||||
|
|
||||||
The MIS path inserts delimiter tokens into the attention context,
|
|
||||||
which slightly perturbs hidden states. After softmax the scores
|
|
||||||
should still be close. Uses places=1 (±0.05) tolerance.
|
|
||||||
|
|
||||||
Runs as a separate class so each engine is created and destroyed
|
|
||||||
independently to avoid GPU OOM.
|
|
||||||
"""
|
|
||||||
|
|
||||||
def test_mis_single_vs_non_mis(self):
|
|
||||||
non_mis_engine = Engine(
|
|
||||||
model_path=TEST_CLASSIFICATION_BASE_MODEL,
|
|
||||||
disable_radix_cache=True,
|
|
||||||
mem_fraction_static=0.15,
|
|
||||||
)
|
|
||||||
try:
|
|
||||||
query = "Rate this option:"
|
|
||||||
items = [" Option A", " Option B", " Option C"]
|
|
||||||
non_mis_scores = non_mis_engine.score(
|
|
||||||
query=query,
|
|
||||||
items=items,
|
|
||||||
apply_softmax=True,
|
|
||||||
).scores
|
|
||||||
finally:
|
|
||||||
non_mis_engine.shutdown()
|
|
||||||
torch.cuda.empty_cache()
|
|
||||||
|
|
||||||
mis_engine = Engine(
|
|
||||||
model_path=TEST_CLASSIFICATION_BASE_MODEL,
|
|
||||||
disable_radix_cache=True,
|
|
||||||
chunked_prefill_size=-1,
|
|
||||||
enable_mis=True,
|
|
||||||
attention_backend="flashinfer",
|
|
||||||
mem_fraction_static=0.15,
|
|
||||||
)
|
|
||||||
try:
|
|
||||||
mis_scores = mis_engine.score(
|
|
||||||
query=query,
|
|
||||||
items=items,
|
|
||||||
apply_softmax=True,
|
|
||||||
).scores
|
|
||||||
finally:
|
|
||||||
mis_engine.shutdown()
|
|
||||||
torch.cuda.empty_cache()
|
|
||||||
|
|
||||||
self.assertEqual(len(mis_scores), len(non_mis_scores))
|
|
||||||
for i, (ms, ns) in enumerate(zip(mis_scores, non_mis_scores)):
|
|
||||||
self.assertEqual(len(ms), len(ns))
|
|
||||||
for j, (m, n) in enumerate(zip(ms, ns)):
|
|
||||||
self.assertAlmostEqual(
|
|
||||||
m,
|
|
||||||
n,
|
|
||||||
places=1,
|
|
||||||
msg=f"item {i} label {j}: MIS={m} vs non-MIS={n}",
|
|
||||||
)
|
|
||||||
|
|
||||||
|
|
||||||
class TestMultiItemScoringClassificationAdvanced(CustomTestCase):
|
|
||||||
"""Advanced MIS tests for classification models: score distinctness,
|
|
||||||
determinism, and concurrent request handling."""
|
|
||||||
|
|
||||||
NUM_LABELS = _CLS_NUM_LABELS
|
|
||||||
|
|
||||||
@classmethod
|
|
||||||
def setUpClass(cls):
|
|
||||||
cls.engine = Engine(
|
|
||||||
model_path=TEST_CLASSIFICATION_BASE_MODEL,
|
|
||||||
disable_radix_cache=True,
|
|
||||||
chunked_prefill_size=-1,
|
|
||||||
enable_mis=True,
|
|
||||||
attention_backend="flashinfer",
|
|
||||||
mem_fraction_static=0.15,
|
|
||||||
)
|
|
||||||
|
|
||||||
@classmethod
|
|
||||||
def tearDownClass(cls):
|
|
||||||
if cls.engine is not None:
|
|
||||||
cls.engine.shutdown()
|
|
||||||
torch.cuda.empty_cache()
|
|
||||||
|
|
||||||
def test_items_produce_distinct_scores(self):
|
def test_items_produce_distinct_scores(self):
|
||||||
"""Different items must produce different score vectors.
|
"""Different items must produce different score vectors.
|
||||||
|
|
||||||
@@ -594,6 +388,130 @@ class TestMultiItemScoringClassificationAdvanced(CustomTestCase):
|
|||||||
f"concurrent={c} vs sequential={s}",
|
f"concurrent={c} vs sequential={s}",
|
||||||
)
|
)
|
||||||
|
|
||||||
|
def test_mis_single_vs_non_mis(self):
|
||||||
|
"""MIS single-item must approximate non-MIS single-item.
|
||||||
|
|
||||||
|
MIS inserts delimiter tokens into the attention context, which
|
||||||
|
perturbs hidden states; after softmax the scores should still land
|
||||||
|
within places=1 (+-0.05).
|
||||||
|
"""
|
||||||
|
query = "Rate this option:"
|
||||||
|
items = [" Option A", " Option B", " Option C"]
|
||||||
|
non_mis_scores = self.non_mis_engine.score(
|
||||||
|
query=query, items=items, apply_softmax=True
|
||||||
|
).scores
|
||||||
|
mis_scores = self.engine.score(
|
||||||
|
query=query, items=items, apply_softmax=True
|
||||||
|
).scores
|
||||||
|
|
||||||
|
self.assertEqual(len(mis_scores), len(non_mis_scores))
|
||||||
|
for i, (ms, ns) in enumerate(zip(mis_scores, non_mis_scores)):
|
||||||
|
self.assertEqual(len(ms), len(ns))
|
||||||
|
for j, (m, n) in enumerate(zip(ms, ns)):
|
||||||
|
self.assertAlmostEqual(
|
||||||
|
m,
|
||||||
|
n,
|
||||||
|
places=1,
|
||||||
|
msg=f"item {i} label {j}: MIS={m} vs non-MIS={n}",
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
class TestMultiItemScoringParity(CustomTestCase):
|
||||||
|
"""Test that MIS produces the same results as single-item scoring."""
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def setUpClass(cls):
|
||||||
|
cls.engine_single = Engine(
|
||||||
|
model_path=TEST_MODEL_NAME,
|
||||||
|
disable_radix_cache=True,
|
||||||
|
log_level="error",
|
||||||
|
mem_fraction_static=0.15,
|
||||||
|
)
|
||||||
|
cls.engine_mis = Engine(
|
||||||
|
model_path=TEST_MODEL_NAME,
|
||||||
|
disable_radix_cache=True,
|
||||||
|
chunked_prefill_size=-1,
|
||||||
|
log_level="error",
|
||||||
|
enable_mis=True,
|
||||||
|
attention_backend="flashinfer",
|
||||||
|
mem_fraction_static=0.15,
|
||||||
|
)
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def tearDownClass(cls):
|
||||||
|
if cls.engine_single is not None:
|
||||||
|
cls.engine_single.shutdown()
|
||||||
|
if cls.engine_mis is not None:
|
||||||
|
cls.engine_mis.shutdown()
|
||||||
|
torch.cuda.empty_cache()
|
||||||
|
|
||||||
|
def _compare_scores(
|
||||||
|
self, query, items, label_token_ids=None, apply_softmax=True, test_name=""
|
||||||
|
):
|
||||||
|
"""Compare MIS vs single-item scoring results."""
|
||||||
|
single_scores = self.engine_single.score(
|
||||||
|
query=query,
|
||||||
|
items=items,
|
||||||
|
label_token_ids=label_token_ids,
|
||||||
|
apply_softmax=apply_softmax,
|
||||||
|
).scores
|
||||||
|
|
||||||
|
mis_scores = self.engine_mis.score(
|
||||||
|
query=query,
|
||||||
|
items=items,
|
||||||
|
label_token_ids=label_token_ids,
|
||||||
|
apply_softmax=apply_softmax,
|
||||||
|
).scores
|
||||||
|
|
||||||
|
self.assertEqual(
|
||||||
|
len(mis_scores), len(single_scores), f"{test_name}: count mismatch"
|
||||||
|
)
|
||||||
|
for i, (ms, ss) in enumerate(zip(mis_scores, single_scores)):
|
||||||
|
self.assertEqual(len(ms), len(ss), f"{test_name}: item {i} length mismatch")
|
||||||
|
for j, (m, s) in enumerate(zip(ms, ss)):
|
||||||
|
self.assertAlmostEqual(
|
||||||
|
m,
|
||||||
|
s,
|
||||||
|
places=1,
|
||||||
|
msg=f"{test_name}: item {i} label {j}: MIS={m} vs single={s}",
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_parity_basic(self):
|
||||||
|
tokenizer = AutoTokenizer.from_pretrained(TEST_MODEL_NAME)
|
||||||
|
query = "Rate this option:"
|
||||||
|
items = [" Option A", " Option B", " Option C"]
|
||||||
|
labels = [" good", " bad"]
|
||||||
|
label_ids = [tokenizer.encode(lb, add_special_tokens=False)[0] for lb in labels]
|
||||||
|
self._compare_scores(query, items, label_ids, test_name="basic")
|
||||||
|
|
||||||
|
def test_parity_tokenized_inputs(self):
|
||||||
|
tokenizer = AutoTokenizer.from_pretrained(TEST_MODEL_NAME)
|
||||||
|
query = "Rate this option:"
|
||||||
|
items = [" Option X", " Option Y"]
|
||||||
|
labels = [" good", " bad"]
|
||||||
|
query_ids = tokenizer.encode(query, add_special_tokens=False)
|
||||||
|
items_ids = [tokenizer.encode(i, add_special_tokens=False) for i in items]
|
||||||
|
label_ids = [tokenizer.encode(lb, add_special_tokens=False)[0] for lb in labels]
|
||||||
|
self._compare_scores(query_ids, items_ids, label_ids, test_name="tokenized")
|
||||||
|
|
||||||
|
def test_parity_without_softmax(self):
|
||||||
|
tokenizer = AutoTokenizer.from_pretrained(TEST_MODEL_NAME)
|
||||||
|
query = "The weather today is"
|
||||||
|
items = [" sunny", " cloudy", " rainy"]
|
||||||
|
labels = [" nice", " bad"]
|
||||||
|
label_ids = [tokenizer.encode(lb, add_special_tokens=False)[0] for lb in labels]
|
||||||
|
self._compare_scores(
|
||||||
|
query, items, label_ids, apply_softmax=False, test_name="no_softmax"
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_parity_many_items(self):
|
||||||
|
tokenizer = AutoTokenizer.from_pretrained(TEST_MODEL_NAME)
|
||||||
|
query = "Rate this option from 1 to 5:"
|
||||||
|
items = [f" Option {i}" for i in range(10)]
|
||||||
|
labels = [" 1", " 2", " 3", " 4", " 5"]
|
||||||
|
label_ids = [tokenizer.encode(lb, add_special_tokens=False)[0] for lb in labels]
|
||||||
|
self._compare_scores(query, items, label_ids, test_name="many_items")
|
||||||
|
|
||||||
|
|
||||||
if __name__ == "__main__":
|
if __name__ == "__main__":
|
||||||
unittest.main()
|
unittest.main()
|
||||||
|
|||||||
@@ -3,9 +3,8 @@ import os
|
|||||||
import re
|
import re
|
||||||
import unittest
|
import unittest
|
||||||
|
|
||||||
from sglang.test.ci.ci_register import register_amd_ci, register_cuda_ci
|
from sglang.test.ci.ci_register import register_amd_ci
|
||||||
|
|
||||||
register_cuda_ci(est_time=103, stage="base-b", runner_config="1-gpu-small")
|
|
||||||
register_amd_ci(est_time=106, suite="stage-b-test-1-gpu-small-amd-mi35x")
|
register_amd_ci(est_time=106, suite="stage-b-test-1-gpu-small-amd-mi35x")
|
||||||
import time
|
import time
|
||||||
from types import SimpleNamespace
|
from types import SimpleNamespace
|
||||||
@@ -14,7 +13,7 @@ import requests
|
|||||||
import torch
|
import torch
|
||||||
|
|
||||||
from sglang.srt.utils import kill_process_tree
|
from sglang.srt.utils import kill_process_tree
|
||||||
from sglang.srt.utils.common import is_cuda_alike, mxfp_supported
|
from sglang.srt.utils.common import is_cuda_alike, is_gfx95_supported
|
||||||
from sglang.test.few_shot_gsm8k import run_eval
|
from sglang.test.few_shot_gsm8k import run_eval
|
||||||
from sglang.test.test_utils import (
|
from sglang.test.test_utils import (
|
||||||
DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH,
|
DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH,
|
||||||
@@ -36,7 +35,7 @@ class TestOnlineQuantizationMemoryLoad(CustomTestCase):
|
|||||||
f"test requires {cls.tp} devices, only {torch.cuda.device_count()} are available."
|
f"test requires {cls.tp} devices, only {torch.cuda.device_count()} are available."
|
||||||
)
|
)
|
||||||
|
|
||||||
if not mxfp_supported():
|
if not is_gfx95_supported():
|
||||||
raise unittest.SkipTest(
|
raise unittest.SkipTest(
|
||||||
"online MXFP4 quantization requires an AMD ROCm device with "
|
"online MXFP4 quantization requires an AMD ROCm device with "
|
||||||
"FP4 hardware support (gfx95x, e.g. MI355x)"
|
"FP4 hardware support (gfx95x, e.g. MI355x)"
|
||||||
|
|||||||
@@ -161,7 +161,7 @@ class TestDFlashServerNoCudaGraph(TestDFlashServerBase):
|
|||||||
other_launch_args = ["--disable-cuda-graph"]
|
other_launch_args = ["--disable-cuda-graph"]
|
||||||
|
|
||||||
|
|
||||||
class TestDFlashServerSpecV2(TestDFlashServerBase):
|
class TestDFlashServerOverlap(TestDFlashServerBase):
|
||||||
disable_overlap = False
|
disable_overlap = False
|
||||||
|
|
||||||
def test_radix_attention(self):
|
def test_radix_attention(self):
|
||||||
@@ -169,7 +169,7 @@ class TestDFlashServerSpecV2(TestDFlashServerBase):
|
|||||||
assert self.process.poll() is None
|
assert self.process.poll() is None
|
||||||
|
|
||||||
|
|
||||||
class TestDFlashServerSpecV2PlanStream(TestDFlashServerSpecV2):
|
class TestDFlashServerOverlapPlanStream(TestDFlashServerOverlap):
|
||||||
overlap_plan_stream = True
|
overlap_plan_stream = True
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -1,5 +1,7 @@
|
|||||||
"""EAGLE3 spec-decoding core: overlap (spec v2) x no-overlap (spec v1) matrix,
|
"""EAGLE3 spec-decoding core: overlap x no-overlap matrix at the standard
|
||||||
same standard config (topk=1, page_size=1), only ``disable_overlap`` differs.
|
config (topk=1, page_size=1); only ``disable_overlap`` differs. Both run the
|
||||||
|
same EAGLEWorkerV2 -- the scheduler just drives it synchronously when overlap
|
||||||
|
is off.
|
||||||
flashinfer is pinned (the 5090 default) so a default-selection change can't
|
flashinfer is pinned (the 5090 default) so a default-selection change can't
|
||||||
silently alter what this exercises.
|
silently alter what this exercises.
|
||||||
"""
|
"""
|
||||||
@@ -35,13 +37,13 @@ class _Core(Eagle3Base):
|
|||||||
|
|
||||||
|
|
||||||
class TestEagle3Overlap(_Core, *_KITS):
|
class TestEagle3Overlap(_Core, *_KITS):
|
||||||
"""Spec v2 (overlap scheduler on)."""
|
"""Overlap scheduler on."""
|
||||||
|
|
||||||
disable_overlap = False
|
disable_overlap = False
|
||||||
|
|
||||||
|
|
||||||
class TestEagle3NoOverlap(_Core, *_KITS):
|
class TestEagle3NoOverlap(_Core, *_KITS):
|
||||||
"""Spec v1 (overlap scheduler off)."""
|
"""Overlap scheduler off (synchronous)."""
|
||||||
|
|
||||||
disable_overlap = True
|
disable_overlap = True
|
||||||
|
|
||||||
|
|||||||
@@ -10,7 +10,6 @@ from sglang.srt.environ import envs
|
|||||||
from sglang.test.ci.ci_register import register_cuda_ci
|
from sglang.test.ci.ci_register import register_cuda_ci
|
||||||
from sglang.test.kits.spec_server_kits import (
|
from sglang.test.kits.spec_server_kits import (
|
||||||
SpecAccuracyKit,
|
SpecAccuracyKit,
|
||||||
SpecCorrectnessKit,
|
|
||||||
SpecFeatureKit,
|
SpecFeatureKit,
|
||||||
SpecLogprobKit,
|
SpecLogprobKit,
|
||||||
SpecPenaltyKit,
|
SpecPenaltyKit,
|
||||||
@@ -18,11 +17,16 @@ from sglang.test.kits.spec_server_kits import (
|
|||||||
)
|
)
|
||||||
from sglang.test.server_fixtures.spec_eagle_fixture import Eagle3Base, EagleLlama2Base
|
from sglang.test.server_fixtures.spec_eagle_fixture import Eagle3Base, EagleLlama2Base
|
||||||
|
|
||||||
register_cuda_ci(est_time=600, stage="base-b", runner_config="1-gpu-large")
|
register_cuda_ci(est_time=250, stage="base-b", runner_config="1-gpu-large")
|
||||||
|
|
||||||
|
|
||||||
class TestEagle3Fa3(Eagle3Base, SpecCorrectnessKit, SpecAccuracyKit, SpecLogprobKit):
|
class TestEagle3Fa3(Eagle3Base, SpecAccuracyKit, SpecLogprobKit):
|
||||||
"""EAGLE3 spec v2 topk=1 on fa3 (the H200 default backend)."""
|
"""EAGLE3 topk=1 on fa3 (the H200 default backend), overlap on.
|
||||||
|
|
||||||
|
No SpecCorrectnessKit: those checks are scheduler/sampling behaviour, which
|
||||||
|
the 5090 runs already cover. Logprob losslessness stays -- it reads through
|
||||||
|
the verify output, which the attention unit cases do not reach.
|
||||||
|
"""
|
||||||
|
|
||||||
attention_backend = "fa3"
|
attention_backend = "fa3"
|
||||||
disable_overlap = False
|
disable_overlap = False
|
||||||
@@ -37,7 +41,7 @@ class TestEagleLlama2Fa3Page256(
|
|||||||
SpecPerfKit,
|
SpecPerfKit,
|
||||||
SpecFeatureKit,
|
SpecFeatureKit,
|
||||||
):
|
):
|
||||||
"""EAGLE/Llama-2 topk=5 tree on fa3 + page_size=256 (spec v1)."""
|
"""EAGLE/Llama-2 topk=5 tree on fa3 + page_size=256, overlap off."""
|
||||||
|
|
||||||
spec_topk = 5
|
spec_topk = 5
|
||||||
spec_steps = 8
|
spec_steps = 8
|
||||||
|
|||||||
@@ -1,7 +1,7 @@
|
|||||||
"""page_size > 1 variants at topk=1 (flashinfer).
|
"""EAGLE3 chain drafting (topk=1) at page_size > 1, flashinfer.
|
||||||
|
|
||||||
EAGLE3 page64 (spec v2) + EAGLE/Llama-2 page4 (spec v1). topk>1 page variants
|
topk=1 takes its own fast path in the draft worker, so this cell is not
|
||||||
live in test_spec_eagle_topk.py. Runs on the cheap (5090) runner.
|
covered by the tree variants in test_spec_eagle_topk_page.py.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
import unittest
|
import unittest
|
||||||
@@ -10,30 +10,28 @@ from sglang.srt.environ import envs
|
|||||||
from sglang.test.ci.ci_register import register_cuda_ci
|
from sglang.test.ci.ci_register import register_cuda_ci
|
||||||
from sglang.test.kits.spec_server_kits import (
|
from sglang.test.kits.spec_server_kits import (
|
||||||
SpecAccuracyKit,
|
SpecAccuracyKit,
|
||||||
|
SpecCorrectnessKit,
|
||||||
SpecFeatureKit,
|
SpecFeatureKit,
|
||||||
SpecLogprobKit,
|
SpecLogprobKit,
|
||||||
)
|
)
|
||||||
from sglang.test.server_fixtures.spec_eagle_fixture import Eagle3Base, EagleLlama2Base
|
from sglang.test.server_fixtures.spec_eagle_fixture import Eagle3Base
|
||||||
|
|
||||||
register_cuda_ci(est_time=360, stage="base-b", runner_config="1-gpu-small")
|
register_cuda_ci(est_time=230, stage="base-b", runner_config="1-gpu-small")
|
||||||
|
|
||||||
|
|
||||||
class TestEagle3Page64(Eagle3Base, SpecAccuracyKit, SpecLogprobKit, SpecFeatureKit):
|
class TestEagle3Page64(
|
||||||
"""EAGLE3 spec v2, page_size=64 (flashinfer): + logprob losslessness."""
|
Eagle3Base,
|
||||||
|
SpecCorrectnessKit,
|
||||||
|
SpecAccuracyKit,
|
||||||
|
SpecLogprobKit,
|
||||||
|
SpecFeatureKit,
|
||||||
|
):
|
||||||
|
"""Overlap scheduler, page_size=64: + logprob losslessness."""
|
||||||
|
|
||||||
page_size = 64
|
page_size = 64
|
||||||
disable_overlap = False
|
disable_overlap = False
|
||||||
env_overrides = ((envs.SGLANG_ENABLE_STRICT_MEM_CHECK_DURING_BUSY, 1),)
|
env_overrides = ((envs.SGLANG_ENABLE_STRICT_MEM_CHECK_DURING_BUSY, 1),)
|
||||||
|
|
||||||
|
|
||||||
class TestEagleLlama2Page4Topk1(EagleLlama2Base, SpecAccuracyKit, SpecFeatureKit):
|
|
||||||
"""Llama-2 topk=1 + page_size=4."""
|
|
||||||
|
|
||||||
spec_topk = 1
|
|
||||||
spec_tokens = 6
|
|
||||||
page_size = 4
|
|
||||||
env_overrides = ((envs.SGLANG_ENABLE_STRICT_MEM_CHECK_DURING_BUSY, 1),)
|
|
||||||
|
|
||||||
|
|
||||||
if __name__ == "__main__":
|
if __name__ == "__main__":
|
||||||
unittest.main()
|
unittest.main()
|
||||||
|
|||||||
@@ -27,7 +27,7 @@ class _Eagle3ParityBase(Eagle3Base):
|
|||||||
|
|
||||||
@unittest.skipIf(_is_xpu, "CUDA runner only")
|
@unittest.skipIf(_is_xpu, "CUDA runner only")
|
||||||
class TestEagle3ParityCUDA(SpecParityKit, _Eagle3ParityBase):
|
class TestEagle3ParityCUDA(SpecParityKit, _Eagle3ParityBase):
|
||||||
"""EAGLE3 spec v2 (flashinfer, overlap) greedy output == non-spec reference.
|
"""EAGLE3 (flashinfer, overlap) greedy output == non-spec reference.
|
||||||
|
|
||||||
SpecParityKit is first so its setUpClass runs the reference server (and tears
|
SpecParityKit is first so its setUpClass runs the reference server (and tears
|
||||||
it down) before the fixture launches the spec server -- sequential, one model
|
it down) before the fixture launches the spec server -- sequential, one model
|
||||||
|
|||||||
@@ -1,18 +1,15 @@
|
|||||||
"""Perf + stress: throughput, retract-under-pressure, abort storms, timeouts.
|
"""Perf + stress: throughput and retract-under-pressure.
|
||||||
|
|
||||||
These need memory headroom / measure load behavior, so they run on the large
|
These need memory headroom / measure load behaviour, so they run on the large
|
||||||
(Hopper) runner.
|
(Hopper) runner. The scheduler timeout paths carry no spec-specific state, so
|
||||||
|
they live in unit/managers/test_scheduler_timeouts.py plus the cheap e2e in
|
||||||
|
scheduler/test_scheduler_control.py.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
import unittest
|
import unittest
|
||||||
|
|
||||||
from sglang.srt.environ import envs
|
from sglang.srt.environ import envs
|
||||||
from sglang.test.ci.ci_register import register_cuda_ci
|
from sglang.test.ci.ci_register import register_cuda_ci
|
||||||
from sglang.test.kits.abort_timeout_kit import (
|
|
||||||
AbortAllMixin,
|
|
||||||
RunningTimeoutTwoWaveMixin,
|
|
||||||
WaitingTimeoutMixin,
|
|
||||||
)
|
|
||||||
from sglang.test.kits.spec_server_kits import (
|
from sglang.test.kits.spec_server_kits import (
|
||||||
SpecAccuracyKit,
|
SpecAccuracyKit,
|
||||||
SpecFeatureKit,
|
SpecFeatureKit,
|
||||||
@@ -20,11 +17,11 @@ from sglang.test.kits.spec_server_kits import (
|
|||||||
)
|
)
|
||||||
from sglang.test.server_fixtures.spec_eagle_fixture import Eagle3Base, EagleLlama2Base
|
from sglang.test.server_fixtures.spec_eagle_fixture import Eagle3Base, EagleLlama2Base
|
||||||
|
|
||||||
register_cuda_ci(est_time=780, stage="base-b", runner_config="1-gpu-large")
|
register_cuda_ci(est_time=440, stage="base-b", runner_config="1-gpu-large")
|
||||||
|
|
||||||
|
|
||||||
class TestEagle3Perf(Eagle3Base, SpecPerfKit):
|
class TestEagle3Perf(Eagle3Base, SpecPerfKit):
|
||||||
"""Decode throughput (max_new_tokens=1) on EAGLE3 spec v2."""
|
"""Decode throughput (max_new_tokens=1) on EAGLE3."""
|
||||||
|
|
||||||
disable_overlap = False
|
disable_overlap = False
|
||||||
env_overrides = ((envs.SGLANG_ENABLE_STRICT_MEM_CHECK_DURING_BUSY, 1),)
|
env_overrides = ((envs.SGLANG_ENABLE_STRICT_MEM_CHECK_DURING_BUSY, 1),)
|
||||||
@@ -42,7 +39,7 @@ class TestEagleLlama2Retract(EagleLlama2Base, SpecAccuracyKit, SpecFeatureKit):
|
|||||||
|
|
||||||
|
|
||||||
class TestEagle3Topk16V2Retract(Eagle3Base, SpecAccuracyKit, SpecFeatureKit):
|
class TestEagle3Topk16V2Retract(Eagle3Base, SpecAccuracyKit, SpecFeatureKit):
|
||||||
"""EAGLE3 topk=16 tree on spec v2 under retract; must not leak KV. Stresses
|
"""EAGLE3 topk=16 tree under retract; must not leak KV. Stresses
|
||||||
the accepted-path KV move (move_accept_tokens_to_target_kvcache)."""
|
the accepted-path KV move (move_accept_tokens_to_target_kvcache)."""
|
||||||
|
|
||||||
spec_topk = 16
|
spec_topk = 16
|
||||||
@@ -58,27 +55,5 @@ class TestEagle3Topk16V2Retract(Eagle3Base, SpecAccuracyKit, SpecFeatureKit):
|
|||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
class TestEagleLlama2AbortAll(EagleLlama2Base, AbortAllMixin):
|
|
||||||
abort_all_max_new_tokens = 4000
|
|
||||||
env_overrides = ((envs.SGLANG_ENABLE_STRICT_MEM_CHECK_DURING_BUSY, 1),)
|
|
||||||
|
|
||||||
|
|
||||||
class TestEagleLlama2WaitingTimeout(EagleLlama2Base, WaitingTimeoutMixin):
|
|
||||||
max_running_requests = 1
|
|
||||||
env_overrides = (
|
|
||||||
(envs.SGLANG_REQ_WAITING_TIMEOUT, 0.001),
|
|
||||||
(envs.SGLANG_ENABLE_STRICT_MEM_CHECK_DURING_BUSY, 1),
|
|
||||||
)
|
|
||||||
|
|
||||||
|
|
||||||
class TestEagleLlama2RunningTimeout(EagleLlama2Base, RunningTimeoutTwoWaveMixin):
|
|
||||||
# Regression: https://github.com/sgl-project/sglang/pull/18760
|
|
||||||
max_running_requests = 16
|
|
||||||
env_overrides = (
|
|
||||||
(envs.SGLANG_REQ_RUNNING_TIMEOUT, 3),
|
|
||||||
(envs.SGLANG_ENABLE_STRICT_MEM_CHECK_DURING_BUSY, 1),
|
|
||||||
)
|
|
||||||
|
|
||||||
|
|
||||||
if __name__ == "__main__":
|
if __name__ == "__main__":
|
||||||
unittest.main()
|
unittest.main()
|
||||||
|
|||||||
@@ -1,8 +1,7 @@
|
|||||||
"""topk > 1 tree drafting (EAGLE3 topk16 + EAGLE/Llama-2 topk8).
|
"""topk > 1 tree drafting at page_size=1 (EAGLE3 topk16 + EAGLE/Llama-2 topk8).
|
||||||
|
|
||||||
topk > 1 routes to spec v1, except page_size==1 which can also stay on spec v2
|
flashinfer is pinned because this runs on the cheap (5090) runner, where fa3
|
||||||
(overlap). flashinfer is pinned because this runs on the cheap (5090) runner,
|
(Hopper-only) isn't available -- functional sanity only, no perf/stress.
|
||||||
where fa3 (Hopper-only) isn't available -- functional sanity only, no perf/stress.
|
|
||||||
(topk > 1 on fa3 is covered on the Hopper runner in test_spec_eagle_fa3.py.)
|
(topk > 1 on fa3 is covered on the Hopper runner in test_spec_eagle_fa3.py.)
|
||||||
"""
|
"""
|
||||||
|
|
||||||
@@ -10,24 +9,37 @@ import unittest
|
|||||||
|
|
||||||
from sglang.srt.environ import envs
|
from sglang.srt.environ import envs
|
||||||
from sglang.test.ci.ci_register import register_cuda_ci
|
from sglang.test.ci.ci_register import register_cuda_ci
|
||||||
|
from sglang.test.kits.abort_timeout_kit import AbortAllMixin
|
||||||
from sglang.test.kits.spec_server_kits import (
|
from sglang.test.kits.spec_server_kits import (
|
||||||
SpecAccuracyKit,
|
SpecAccuracyKit,
|
||||||
SpecCorrectnessKit,
|
SpecCorrectnessKit,
|
||||||
SpecFeatureKit,
|
SpecFeatureKit,
|
||||||
|
SpecHiddenStatesKit,
|
||||||
SpecLogprobKit,
|
SpecLogprobKit,
|
||||||
SpecPenaltyKit,
|
SpecPenaltyKit,
|
||||||
)
|
)
|
||||||
from sglang.test.server_fixtures.spec_eagle_fixture import Eagle3Base, EagleLlama2Base
|
from sglang.test.server_fixtures.spec_eagle_fixture import Eagle3Base, EagleLlama2Base
|
||||||
|
|
||||||
register_cuda_ci(est_time=1180, stage="base-b", runner_config="1-gpu-small")
|
register_cuda_ci(est_time=870, stage="base-b", runner_config="1-gpu-small")
|
||||||
|
|
||||||
|
|
||||||
class TestEagle3Topk16(Eagle3Base, SpecCorrectnessKit, SpecAccuracyKit, SpecLogprobKit):
|
class TestEagle3Topk16(
|
||||||
"""EAGLE3 topk=16 tree (spec v1): correctness + gsm8k + logprob losslessness."""
|
Eagle3Base,
|
||||||
|
SpecCorrectnessKit,
|
||||||
|
SpecAccuracyKit,
|
||||||
|
SpecLogprobKit,
|
||||||
|
SpecFeatureKit,
|
||||||
|
SpecHiddenStatesKit,
|
||||||
|
):
|
||||||
|
"""EAGLE3 topk=16 tree, overlap scheduler: guards the accepted-path
|
||||||
|
compaction (via logprob_decode_match_prefill) and the per-request
|
||||||
|
hidden-state stride slicing that the same compaction feeds.
|
||||||
|
"""
|
||||||
|
|
||||||
spec_topk = 16
|
spec_topk = 16
|
||||||
spec_tokens = 64
|
spec_tokens = 64
|
||||||
disable_overlap = True # synchronous baseline; SpecV2 subclass flips overlap on
|
disable_overlap = False
|
||||||
|
enable_return_hidden_states = True
|
||||||
cuda_graph_max_bs_decode = 5
|
cuda_graph_max_bs_decode = 5
|
||||||
acc_length_thres = 3.1
|
acc_length_thres = 3.1
|
||||||
batch_accept_len_thres = 1.75
|
batch_accept_len_thres = 1.75
|
||||||
@@ -35,13 +47,6 @@ class TestEagle3Topk16(Eagle3Base, SpecCorrectnessKit, SpecAccuracyKit, SpecLogp
|
|||||||
env_overrides = ((envs.SGLANG_ENABLE_STRICT_MEM_CHECK_DURING_BUSY, 1),)
|
env_overrides = ((envs.SGLANG_ENABLE_STRICT_MEM_CHECK_DURING_BUSY, 1),)
|
||||||
|
|
||||||
|
|
||||||
class TestEagle3Topk16SpecV2(TestEagle3Topk16, SpecFeatureKit):
|
|
||||||
"""EAGLE3 topk=16 tree on spec v2 (overlap, page1): guards the v2 tree path's
|
|
||||||
accepted-path compaction, validated by logprob_spec_v2_match."""
|
|
||||||
|
|
||||||
disable_overlap = False
|
|
||||||
|
|
||||||
|
|
||||||
class TestEagleLlama2Suite(
|
class TestEagleLlama2Suite(
|
||||||
EagleLlama2Base,
|
EagleLlama2Base,
|
||||||
SpecCorrectnessKit,
|
SpecCorrectnessKit,
|
||||||
@@ -49,9 +54,16 @@ class TestEagleLlama2Suite(
|
|||||||
SpecLogprobKit,
|
SpecLogprobKit,
|
||||||
SpecPenaltyKit,
|
SpecPenaltyKit,
|
||||||
SpecFeatureKit,
|
SpecFeatureKit,
|
||||||
|
AbortAllMixin,
|
||||||
):
|
):
|
||||||
"""EAGLE/Llama-2 topk=8 full coverage (kits listed in bases)."""
|
"""EAGLE/Llama-2 topk=8 full coverage (kits listed in bases).
|
||||||
|
|
||||||
|
Hosts AbortAllMixin: aborting mid-decode has to release the tree draft
|
||||||
|
state, and the strict mem check below turns a leak into a failure. It needs
|
||||||
|
no server flags of its own, so it rides this launch.
|
||||||
|
"""
|
||||||
|
|
||||||
|
abort_all_max_new_tokens = 4000
|
||||||
env_overrides = ((envs.SGLANG_ENABLE_STRICT_MEM_CHECK_DURING_BUSY, 1),)
|
env_overrides = ((envs.SGLANG_ENABLE_STRICT_MEM_CHECK_DURING_BUSY, 1),)
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -1,9 +1,10 @@
|
|||||||
"""topk > 1 tree drafting at page_size > 1 (EAGLE3 topk8 + EAGLE/Llama-2 topk8).
|
"""EAGLE3 tree drafting (topk > 1) at page_size > 1, flashinfer (fa3 is
|
||||||
|
Hopper-only).
|
||||||
|
|
||||||
page64 stays on spec v2 (overlap), page4 runs on spec v1 (no overlap). flashinfer is
|
page_size=4 with 32 draft tokens spreads the draft window over several pages --
|
||||||
pinned because this runs on the cheap (5090) runner, where fa3 (Hopper-only) isn't
|
the layout the unit fixture refuses to build (tree draft is pinned to
|
||||||
available -- functional sanity only, no perf/stress. (page>1 topk>1 on fa3 is covered
|
page_size=1 there, see speculative_draft_runner.py). The window-inside-one-page
|
||||||
on the Hopper runner in test_spec_eagle_fa3.py.)
|
regime is covered by test_spec_eagle_fa3.py page256 on the Hopper runner.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
import unittest
|
import unittest
|
||||||
@@ -13,29 +14,26 @@ from sglang.test.ci.ci_register import register_cuda_ci
|
|||||||
from sglang.test.kits.spec_server_kits import (
|
from sglang.test.kits.spec_server_kits import (
|
||||||
SpecAccuracyKit,
|
SpecAccuracyKit,
|
||||||
SpecFeatureKit,
|
SpecFeatureKit,
|
||||||
|
SpecLogprobKit,
|
||||||
)
|
)
|
||||||
from sglang.test.server_fixtures.spec_eagle_fixture import Eagle3Base, EagleLlama2Base
|
from sglang.test.server_fixtures.spec_eagle_fixture import Eagle3Base
|
||||||
|
|
||||||
register_cuda_ci(est_time=720, stage="base-b", runner_config="1-gpu-small")
|
register_cuda_ci(est_time=345, stage="base-b", runner_config="1-gpu-small")
|
||||||
|
|
||||||
|
|
||||||
class TestEagle3Page64Topk8(Eagle3Base, SpecAccuracyKit, SpecFeatureKit):
|
class TestEagle3Page4Topk8(Eagle3Base, SpecAccuracyKit, SpecLogprobKit, SpecFeatureKit):
|
||||||
"""EAGLE3 topk=8 tree + page_size=64 (spec v2)."""
|
"""Overlap scheduler, topk=8 tree, page_size=4."""
|
||||||
|
|
||||||
page_size = 64
|
page_size = 4
|
||||||
spec_topk = 8
|
spec_topk = 8
|
||||||
spec_tokens = 32
|
spec_tokens = 32
|
||||||
disable_overlap = False
|
disable_overlap = False
|
||||||
|
# Preset accept-length values are topk=1 numbers -- loose for a topk=8
|
||||||
|
# tree; tighten once CI reports the actuals.
|
||||||
|
gsm8k_accept_len_thres = 2.0
|
||||||
cuda_graph_max_bs_decode = 5
|
cuda_graph_max_bs_decode = 5
|
||||||
env_overrides = ((envs.SGLANG_ENABLE_STRICT_MEM_CHECK_DURING_BUSY, 1),)
|
env_overrides = ((envs.SGLANG_ENABLE_STRICT_MEM_CHECK_DURING_BUSY, 1),)
|
||||||
|
|
||||||
|
|
||||||
class TestEagleLlama2Page4Topk8(EagleLlama2Base, SpecAccuracyKit, SpecFeatureKit):
|
|
||||||
"""Llama-2 topk>1 tree + page_size=4 (spec v1)."""
|
|
||||||
|
|
||||||
page_size = 4
|
|
||||||
env_overrides = ((envs.SGLANG_ENABLE_STRICT_MEM_CHECK_DURING_BUSY, 1),)
|
|
||||||
|
|
||||||
|
|
||||||
if __name__ == "__main__":
|
if __name__ == "__main__":
|
||||||
unittest.main()
|
unittest.main()
|
||||||
|
|||||||
@@ -1,6 +1,8 @@
|
|||||||
"""triton attention backend (EAGLE3 topk=1 chain + EAGLE/Llama-2 topk=8 tree).
|
"""triton attention backend, EAGLE3 chain drafting.
|
||||||
|
|
||||||
triton runs everywhere, so this stays on the cheap (5090) runner.
|
triton runs everywhere, so this stays on the cheap (5090) runner. triton tree
|
||||||
|
verify is covered by attention/unittests/dense/test_triton.py, and the tree
|
||||||
|
accept-path compaction e2e lives in test_spec_eagle_topk.py.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
import unittest
|
import unittest
|
||||||
@@ -11,13 +13,12 @@ from sglang.test.kits.matched_stop_kit import MatchedStopMixin
|
|||||||
from sglang.test.kits.spec_server_kits import (
|
from sglang.test.kits.spec_server_kits import (
|
||||||
SpecAccuracyKit,
|
SpecAccuracyKit,
|
||||||
SpecFeatureKit,
|
SpecFeatureKit,
|
||||||
SpecHiddenStatesKit,
|
|
||||||
SpecLogprobKit,
|
SpecLogprobKit,
|
||||||
SpecPenaltyKit,
|
SpecPenaltyKit,
|
||||||
)
|
)
|
||||||
from sglang.test.server_fixtures.spec_eagle_fixture import Eagle3Base, EagleLlama2Base
|
from sglang.test.server_fixtures.spec_eagle_fixture import Eagle3Base
|
||||||
|
|
||||||
register_cuda_ci(est_time=350, stage="base-b", runner_config="1-gpu-small")
|
register_cuda_ci(est_time=230, stage="base-b", runner_config="1-gpu-small")
|
||||||
|
|
||||||
|
|
||||||
class TestEagle3Triton(
|
class TestEagle3Triton(
|
||||||
@@ -28,7 +29,7 @@ class TestEagle3Triton(
|
|||||||
SpecPenaltyKit,
|
SpecPenaltyKit,
|
||||||
SpecFeatureKit,
|
SpecFeatureKit,
|
||||||
):
|
):
|
||||||
"""EAGLE3 spec v2 on triton (kits listed in bases)."""
|
"""Overlap scheduler on triton (kits listed in bases)."""
|
||||||
|
|
||||||
attention_backend = "triton"
|
attention_backend = "triton"
|
||||||
max_running_requests = 64
|
max_running_requests = 64
|
||||||
@@ -38,19 +39,5 @@ class TestEagle3Triton(
|
|||||||
env_overrides = ((envs.SGLANG_ENABLE_STRICT_MEM_CHECK_DURING_BUSY, 1),)
|
env_overrides = ((envs.SGLANG_ENABLE_STRICT_MEM_CHECK_DURING_BUSY, 1),)
|
||||||
|
|
||||||
|
|
||||||
class TestEagleLlama2Triton(
|
|
||||||
EagleLlama2Base, SpecAccuracyKit, SpecFeatureKit, SpecHiddenStatesKit
|
|
||||||
):
|
|
||||||
"""EAGLE/Llama-2 topk=8 tree on triton.
|
|
||||||
|
|
||||||
Hosts SpecHiddenStatesKit: topk>1 exercises the tree accept-path
|
|
||||||
compaction that the per-req hidden-state stride slicing depends on.
|
|
||||||
"""
|
|
||||||
|
|
||||||
attention_backend = "triton"
|
|
||||||
enable_return_hidden_states = True
|
|
||||||
env_overrides = ((envs.SGLANG_ENABLE_STRICT_MEM_CHECK_DURING_BUSY, 1),)
|
|
||||||
|
|
||||||
|
|
||||||
if __name__ == "__main__":
|
if __name__ == "__main__":
|
||||||
unittest.main()
|
unittest.main()
|
||||||
|
|||||||
@@ -4,24 +4,29 @@ from sglang.test.ci.ci_register import register_cuda_ci
|
|||||||
from sglang.test.kits.json_constrained_kit import JSONConstrainedMixin
|
from sglang.test.kits.json_constrained_kit import JSONConstrainedMixin
|
||||||
from sglang.test.kits.regex_constrained_kit import RegexConstrainedMixin
|
from sglang.test.kits.regex_constrained_kit import RegexConstrainedMixin
|
||||||
from sglang.test.server_fixtures.standalone_fixture import StandaloneServerBase
|
from sglang.test.server_fixtures.standalone_fixture import StandaloneServerBase
|
||||||
from sglang.test.test_utils import CustomTestCase
|
from sglang.test.test_utils import CustomTestCase, is_in_ci
|
||||||
|
|
||||||
# V2 standalone speculative decoding tests (FA3, Triton, FlashInfer backends).
|
# V2 standalone speculative decoding. CI runs only fa3 (the backend this is
|
||||||
|
# deployed on); triton / flashinfer stay runnable locally, and their spec verify
|
||||||
|
# numerics live in attention/unittests/dense/test_{triton,flashinfer}.py.
|
||||||
# Non-V2 backends moved to test_spec_standalone_extra.py.
|
# Non-V2 backends moved to test_spec_standalone_extra.py.
|
||||||
register_cuda_ci(est_time=450, stage="base-b", runner_config="1-gpu-large")
|
register_cuda_ci(est_time=170, stage="base-b", runner_config="1-gpu-large")
|
||||||
|
|
||||||
|
|
||||||
class TestStandaloneV2SpeculativeDecodingBase(StandaloneServerBase, CustomTestCase):
|
class TestStandaloneV2SpeculativeDecodingBase(
|
||||||
|
StandaloneServerBase, CustomTestCase, RegexConstrainedMixin, JSONConstrainedMixin
|
||||||
|
):
|
||||||
|
# Hosts the constrained mixins: overlap is on, so they exercise the
|
||||||
|
# grammar barrier path.
|
||||||
attention_backend = "fa3"
|
attention_backend = "fa3"
|
||||||
|
|
||||||
|
|
||||||
class TestStandaloneV2SpeculativeDecodingTriton(
|
@unittest.skipIf(is_in_ci(), "CI covers fa3 only; run locally for triton.")
|
||||||
StandaloneServerBase, CustomTestCase, RegexConstrainedMixin, JSONConstrainedMixin
|
class TestStandaloneV2SpeculativeDecodingTriton(StandaloneServerBase, CustomTestCase):
|
||||||
):
|
|
||||||
# Constrained mixins reuse this server; overlap on -> grammar barrier path.
|
|
||||||
attention_backend = "triton"
|
attention_backend = "triton"
|
||||||
|
|
||||||
|
|
||||||
|
@unittest.skipIf(is_in_ci(), "CI covers fa3 only; run locally for flashinfer.")
|
||||||
class TestStandaloneV2SpeculativeDecodingFlashinfer(
|
class TestStandaloneV2SpeculativeDecodingFlashinfer(
|
||||||
StandaloneServerBase, CustomTestCase
|
StandaloneServerBase, CustomTestCase
|
||||||
):
|
):
|
||||||
|
|||||||
@@ -0,0 +1,121 @@
|
|||||||
|
"""Boundary tests for the scheduler's waiting / running request timeouts.
|
||||||
|
|
||||||
|
Both paths are pure bookkeeping over timestamps -- no model, no GPU, no draft
|
||||||
|
worker -- so they are driven here directly instead of through a server. The
|
||||||
|
e2e side (503 reaching the client, server stays up) is covered by
|
||||||
|
scheduler/test_scheduler_control.py.
|
||||||
|
"""
|
||||||
|
|
||||||
|
import time
|
||||||
|
import unittest
|
||||||
|
from types import SimpleNamespace
|
||||||
|
from unittest.mock import MagicMock
|
||||||
|
|
||||||
|
from sglang.srt.environ import envs
|
||||||
|
from sglang.test.ci.ci_register import register_cpu_ci
|
||||||
|
from sglang.test.test_utils import CustomTestCase, maybe_stub_sgl_kernel
|
||||||
|
|
||||||
|
maybe_stub_sgl_kernel()
|
||||||
|
|
||||||
|
from sglang.srt.managers.scheduler import Scheduler
|
||||||
|
|
||||||
|
register_cpu_ci(est_time=6, suite="base-a-test-cpu")
|
||||||
|
|
||||||
|
|
||||||
|
class _FakeReq:
|
||||||
|
"""Must stay hashable: the waiting-timeout path collects drops in a set."""
|
||||||
|
|
||||||
|
def __init__(self, rid, wait_entry=0.0, forward_entry=0.0, is_finished=False):
|
||||||
|
self.rid = rid
|
||||||
|
self.to_finish = None
|
||||||
|
self._finished = is_finished
|
||||||
|
self.time_stats = SimpleNamespace(
|
||||||
|
wait_queue_entry_time=wait_entry,
|
||||||
|
forward_entry_time=forward_entry,
|
||||||
|
trace_ctx=MagicMock(),
|
||||||
|
)
|
||||||
|
|
||||||
|
def finished(self):
|
||||||
|
return self._finished
|
||||||
|
|
||||||
|
|
||||||
|
def _req(
|
||||||
|
rid: str, *, wait_entry: float = 0.0, forward_entry: float = 0.0, finished=False
|
||||||
|
):
|
||||||
|
return _FakeReq(rid, wait_entry, forward_entry, finished)
|
||||||
|
|
||||||
|
|
||||||
|
def _scheduler(waiting_queue):
|
||||||
|
s = Scheduler.__new__(Scheduler)
|
||||||
|
s.waiting_queue = waiting_queue
|
||||||
|
s.enable_hicache_storage = False
|
||||||
|
s.ipc_channels = SimpleNamespace(send_to_tokenizer=MagicMock())
|
||||||
|
return s
|
||||||
|
|
||||||
|
|
||||||
|
class TestWaitingTimeout(CustomTestCase):
|
||||||
|
def test_drops_only_reqs_past_the_deadline(self):
|
||||||
|
now = time.perf_counter()
|
||||||
|
stale = _req("stale", wait_entry=now - 10)
|
||||||
|
fresh = _req("fresh", wait_entry=now)
|
||||||
|
s = _scheduler([stale, fresh])
|
||||||
|
|
||||||
|
with envs.SGLANG_REQ_WAITING_TIMEOUT.override(1.0):
|
||||||
|
s._abort_on_waiting_timeout()
|
||||||
|
|
||||||
|
self.assertEqual([r.rid for r in s.waiting_queue], ["fresh"])
|
||||||
|
self.assertEqual(s.ipc_channels.send_to_tokenizer.send_output.call_count, 1)
|
||||||
|
|
||||||
|
def test_unset_entry_time_is_never_dropped(self):
|
||||||
|
# 0 is the "not yet stamped" sentinel; the guard is `0 < entry_time`.
|
||||||
|
s = _scheduler([_req("unstamped", wait_entry=0.0)])
|
||||||
|
with envs.SGLANG_REQ_WAITING_TIMEOUT.override(1e-9):
|
||||||
|
s._abort_on_waiting_timeout()
|
||||||
|
self.assertEqual(len(s.waiting_queue), 1)
|
||||||
|
s.ipc_channels.send_to_tokenizer.send_output.assert_not_called()
|
||||||
|
|
||||||
|
def test_disabled_timeout_is_a_no_op(self):
|
||||||
|
s = _scheduler([_req("stale", wait_entry=time.perf_counter() - 100)])
|
||||||
|
with envs.SGLANG_REQ_WAITING_TIMEOUT.override(0):
|
||||||
|
s._abort_on_waiting_timeout()
|
||||||
|
self.assertEqual(len(s.waiting_queue), 1)
|
||||||
|
|
||||||
|
|
||||||
|
class TestRunningTimeout(CustomTestCase):
|
||||||
|
@staticmethod
|
||||||
|
def _batch(reqs):
|
||||||
|
return SimpleNamespace(reqs=reqs, is_empty=lambda: not reqs)
|
||||||
|
|
||||||
|
def test_marks_only_stale_unfinished_reqs(self):
|
||||||
|
now = time.perf_counter()
|
||||||
|
stale = _req("stale", forward_entry=now - 10)
|
||||||
|
fresh = _req("fresh", forward_entry=now)
|
||||||
|
done = _req("done", forward_entry=now - 10, finished=True)
|
||||||
|
s = _scheduler([])
|
||||||
|
|
||||||
|
with envs.SGLANG_REQ_RUNNING_TIMEOUT.override(1.0):
|
||||||
|
s._abort_on_running_timeout(self._batch([stale, fresh, done]))
|
||||||
|
|
||||||
|
self.assertIsNotNone(stale.to_finish)
|
||||||
|
self.assertIsNone(fresh.to_finish)
|
||||||
|
self.assertIsNone(done.to_finish, "a finished req must not be aborted")
|
||||||
|
|
||||||
|
def test_unset_forward_entry_time_is_never_marked(self):
|
||||||
|
s = _scheduler([])
|
||||||
|
req = _req("unstamped", forward_entry=0.0)
|
||||||
|
with envs.SGLANG_REQ_RUNNING_TIMEOUT.override(1e-9):
|
||||||
|
s._abort_on_running_timeout(self._batch([req]))
|
||||||
|
self.assertIsNone(req.to_finish)
|
||||||
|
|
||||||
|
def test_empty_batch_and_disabled_timeout_are_no_ops(self):
|
||||||
|
s = _scheduler([])
|
||||||
|
with envs.SGLANG_REQ_RUNNING_TIMEOUT.override(1.0):
|
||||||
|
s._abort_on_running_timeout(self._batch([]))
|
||||||
|
req = _req("stale", forward_entry=time.perf_counter() - 100)
|
||||||
|
with envs.SGLANG_REQ_RUNNING_TIMEOUT.override(0):
|
||||||
|
s._abort_on_running_timeout(self._batch([req]))
|
||||||
|
self.assertIsNone(req.to_finish)
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
unittest.main()
|
||||||
Reference in New Issue
Block a user