[CI] Trim DSV4 trtllm B200 tests (#39013)
Co-authored-by: Mohammad Angkad <mohammad.angkad@radixark.ai>
This commit is contained in:
co-authored by
Mohammad Angkad
parent
17fa5ad327
commit
ad5af539cd
@@ -604,11 +604,7 @@ class TestDSV4BreakableCudaGraphMetadataContract(CustomTestCase):
|
|||||||
def test_trtllm_semaphore_capacity_covers_configured_query_rows(self):
|
def test_trtllm_semaphore_capacity_covers_configured_query_rows(self):
|
||||||
from sglang.srt.layers.attention import deepseek_v4_trtllm_backend as trtllm
|
from sglang.srt.layers.attention import deepseek_v4_trtllm_backend as trtllm
|
||||||
|
|
||||||
schedule = SimpleNamespace(
|
schedule = SimpleNamespace(max_prefill_tokens=16384, max_running_requests=256)
|
||||||
max_prefill_tokens=16384,
|
|
||||||
chunked_prefill_size=4096,
|
|
||||||
max_running_requests=256,
|
|
||||||
)
|
|
||||||
spec = SimpleNamespace(
|
spec = SimpleNamespace(
|
||||||
speculative_algorithm="EAGLE", speculative_num_draft_tokens=4
|
speculative_algorithm="EAGLE", speculative_num_draft_tokens=4
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -1,214 +0,0 @@
|
|||||||
"""SM100/SM103 coverage for DSV4's uniform-FP8 trtllm backend.
|
|
||||||
|
|
||||||
Covers decode correctness, GSM8K accuracy, varlen and cached-prefix prefill,
|
|
||||||
chunking, and decode CUDA-graph replay. Long outputs use sanity checks because
|
|
||||||
the FlashMLA and uniform-FP8 cache formats need not be bit-reproducible.
|
|
||||||
"""
|
|
||||||
|
|
||||||
import concurrent.futures
|
|
||||||
import unittest
|
|
||||||
from types import SimpleNamespace
|
|
||||||
|
|
||||||
import requests
|
|
||||||
import torch
|
|
||||||
|
|
||||||
from sglang.srt.utils import kill_process_tree
|
|
||||||
from sglang.test.ci.ci_register import register_cuda_ci
|
|
||||||
from sglang.test.kits.basic_decode_correctness_kit import BasicDecodeCorrectnessMixin
|
|
||||||
from sglang.test.run_eval import run_eval
|
|
||||||
from sglang.test.test_utils import (
|
|
||||||
DEFAULT_URL_FOR_TEST,
|
|
||||||
CustomTestCase,
|
|
||||||
popen_launch_server,
|
|
||||||
try_cached_model,
|
|
||||||
)
|
|
||||||
|
|
||||||
register_cuda_ci(est_time=900, stage="base-c", runner_config="4-gpu-b200")
|
|
||||||
|
|
||||||
DSV4_FLASH_MODEL_PATH = try_cached_model("deepseek-ai/DeepSeek-V4-Flash")
|
|
||||||
SERVER_LAUNCH_TIMEOUT = 3600
|
|
||||||
DSV4_BASE_ENV = {
|
|
||||||
"SGLANG_JIT_DEEPGEMM_FAST_WARMUP": "1",
|
|
||||||
}
|
|
||||||
|
|
||||||
SERVER_ARGS = [
|
|
||||||
"--trust-remote-code",
|
|
||||||
"--dsv4-attn-backend",
|
|
||||||
"trtllm",
|
|
||||||
"--tp",
|
|
||||||
"4",
|
|
||||||
"--max-running-requests",
|
|
||||||
"32",
|
|
||||||
"--mem-fraction-static",
|
|
||||||
"0.85",
|
|
||||||
"--chunked-prefill-size",
|
|
||||||
"4096",
|
|
||||||
# V4-Flash ships MXFP4 routed experts, and the auto-selected Triton MoE runner
|
|
||||||
# cannot consume the packed layout. Matches the B200 Flash cookbook recipe.
|
|
||||||
"--moe-runner-backend",
|
|
||||||
"flashinfer_mxfp4",
|
|
||||||
"--disable-flashinfer-autotune",
|
|
||||||
]
|
|
||||||
|
|
||||||
# Mixed lengths cover c4/c128 selection and VarSeq packing; the longest prompt
|
|
||||||
# exceeds the 4096-token prefill chunk.
|
|
||||||
_FILLER_SENTENCES = [
|
|
||||||
"The expedition recorded water temperature, salinity, and current speed "
|
|
||||||
"at every station along the transect. ",
|
|
||||||
"Archival records from the observatory describe decades of nightly "
|
|
||||||
"measurements taken with remarkable consistency. ",
|
|
||||||
"Each greenhouse module recycles condensate through a gravel bed before "
|
|
||||||
"returning it to the irrigation loop. ",
|
|
||||||
"The survey team catalogued the masonry of the aqueduct arch by arch, "
|
|
||||||
"noting repairs from three distinct centuries. ",
|
|
||||||
]
|
|
||||||
_LONG_PROMPT_QUESTION = (
|
|
||||||
"\n\nIn one short sentence, what kind of activity do the paragraphs above describe?"
|
|
||||||
)
|
|
||||||
|
|
||||||
|
|
||||||
def _make_long_prompt(idx: int, target_chars: int) -> str:
|
|
||||||
sentence = _FILLER_SENTENCES[idx % len(_FILLER_SENTENCES)]
|
|
||||||
body = ""
|
|
||||||
n = 0
|
|
||||||
while len(body) < target_chars:
|
|
||||||
body += f"[Entry {idx}-{n}] " + sentence
|
|
||||||
n += 1
|
|
||||||
return body + _LONG_PROMPT_QUESTION
|
|
||||||
|
|
||||||
|
|
||||||
# Roughly 2.5k, 4.5k, and 7k tokens.
|
|
||||||
LONG_PROMPTS = [
|
|
||||||
_make_long_prompt(0, 10_000),
|
|
||||||
_make_long_prompt(1, 18_000),
|
|
||||||
_make_long_prompt(2, 28_000),
|
|
||||||
]
|
|
||||||
LONG_MAX_NEW_TOKENS = 32
|
|
||||||
MIN_PRINTABLE_ASCII_RATIO = 0.85
|
|
||||||
|
|
||||||
GSM8K_NUM_EXAMPLES = 200
|
|
||||||
GSM8K_MIN_SCORE = 0.90
|
|
||||||
|
|
||||||
_REQUEST_TIMEOUT = 600
|
|
||||||
|
|
||||||
|
|
||||||
def _is_sm100() -> bool:
|
|
||||||
if not torch.cuda.is_available():
|
|
||||||
return False
|
|
||||||
return torch.cuda.get_device_capability() in ((10, 0), (10, 3))
|
|
||||||
|
|
||||||
|
|
||||||
def _greedy_generate(base_url: str, prompt: str, max_new_tokens: int) -> str:
|
|
||||||
resp = requests.post(
|
|
||||||
base_url + "/generate",
|
|
||||||
json={
|
|
||||||
"text": prompt,
|
|
||||||
"sampling_params": {
|
|
||||||
"temperature": 0.0,
|
|
||||||
"max_new_tokens": max_new_tokens,
|
|
||||||
},
|
|
||||||
},
|
|
||||||
timeout=_REQUEST_TIMEOUT,
|
|
||||||
)
|
|
||||||
resp.raise_for_status()
|
|
||||||
return resp.json()["text"]
|
|
||||||
|
|
||||||
|
|
||||||
def _printable_ascii_ratio(text: str) -> float:
|
|
||||||
if not text:
|
|
||||||
return 0.0
|
|
||||||
return sum(32 <= ord(c) < 127 or c in "\n\t" for c in text) / len(text)
|
|
||||||
|
|
||||||
|
|
||||||
class TestDSV4Fp8TrtllmBackend(BasicDecodeCorrectnessMixin, CustomTestCase):
|
|
||||||
"""TP4 DSv4-Flash-FP8 with --dsv4-attn-backend trtllm."""
|
|
||||||
|
|
||||||
@classmethod
|
|
||||||
def setUpClass(cls):
|
|
||||||
if not _is_sm100():
|
|
||||||
raise unittest.SkipTest(
|
|
||||||
"DSv4 trtllm uniform-FP8 attention requires SM100/SM103 (Blackwell)"
|
|
||||||
)
|
|
||||||
cls.base_url = DEFAULT_URL_FOR_TEST
|
|
||||||
cls.process = popen_launch_server(
|
|
||||||
DSV4_FLASH_MODEL_PATH,
|
|
||||||
cls.base_url,
|
|
||||||
timeout=SERVER_LAUNCH_TIMEOUT,
|
|
||||||
other_args=SERVER_ARGS,
|
|
||||||
env=dict(DSV4_BASE_ENV),
|
|
||||||
)
|
|
||||||
|
|
||||||
@classmethod
|
|
||||||
def tearDownClass(cls):
|
|
||||||
if hasattr(cls, "process") and cls.process is not None:
|
|
||||||
kill_process_tree(cls.process.pid)
|
|
||||||
|
|
||||||
def _assert_sane(self, out: str, what: str) -> None:
|
|
||||||
self.assertGreater(len(out.strip()), 0, f"{what}: empty output")
|
|
||||||
ratio = _printable_ascii_ratio(out)
|
|
||||||
self.assertGreater(
|
|
||||||
ratio,
|
|
||||||
MIN_PRINTABLE_ASCII_RATIO,
|
|
||||||
f"{what}: output looks like gibberish (ascii ratio={ratio:.2f}): {out!r}",
|
|
||||||
)
|
|
||||||
|
|
||||||
def test_long_prompt_varlen_prefill(self):
|
|
||||||
"""Exercise mixed-length VarSeq and cached-prefix chunked prefill.
|
|
||||||
|
|
||||||
Sanity checks avoid flaky exact matches from split-KV reduction order.
|
|
||||||
"""
|
|
||||||
|
|
||||||
with concurrent.futures.ThreadPoolExecutor(len(LONG_PROMPTS)) as pool:
|
|
||||||
outs = list(
|
|
||||||
pool.map(
|
|
||||||
lambda p: _greedy_generate(self.base_url, p, LONG_MAX_NEW_TOKENS),
|
|
||||||
LONG_PROMPTS,
|
|
||||||
)
|
|
||||||
)
|
|
||||||
for i, out in enumerate(outs):
|
|
||||||
print(f"[long-prefill] prompt_chars={len(LONG_PROMPTS[i])} out={out!r}")
|
|
||||||
self._assert_sane(out, f"concurrent long prompt {i}")
|
|
||||||
|
|
||||||
cached = _greedy_generate(self.base_url, LONG_PROMPTS[-1], LONG_MAX_NEW_TOKENS)
|
|
||||||
print(f"[long-prefill] cached-prefix rerun out={cached!r}")
|
|
||||||
self._assert_sane(cached, "cached-prefix extend")
|
|
||||||
|
|
||||||
def test_gsm8k_sanity(self):
|
|
||||||
args = SimpleNamespace(
|
|
||||||
base_url=self.base_url,
|
|
||||||
model=DSV4_FLASH_MODEL_PATH,
|
|
||||||
eval_name="gsm8k",
|
|
||||||
api="completion",
|
|
||||||
max_tokens=512,
|
|
||||||
num_examples=GSM8K_NUM_EXAMPLES,
|
|
||||||
num_threads=64,
|
|
||||||
)
|
|
||||||
metrics = run_eval(args)
|
|
||||||
print(f"GSM8K sanity on trtllm decode: {metrics=}")
|
|
||||||
self.assertGreater(metrics["score"], GSM8K_MIN_SCORE)
|
|
||||||
|
|
||||||
def test_cuda_graph_capture_replay_smoke(self):
|
|
||||||
"""Replay several decode graph buckets and recheck a greedy anchor."""
|
|
||||||
anchor_prompt = "Q: What is the capital of France?\nA:"
|
|
||||||
anchor_out = _greedy_generate(self.base_url, anchor_prompt, 32)
|
|
||||||
|
|
||||||
for concurrency in (2, 4, 8, 16):
|
|
||||||
prompts = [f"Count from {i} to {i + 5}: " for i in range(concurrency)]
|
|
||||||
with concurrent.futures.ThreadPoolExecutor(concurrency) as pool:
|
|
||||||
outs = list(
|
|
||||||
pool.map(lambda p: _greedy_generate(self.base_url, p, 32), prompts)
|
|
||||||
)
|
|
||||||
self.assertEqual(len(outs), concurrency)
|
|
||||||
for out in outs:
|
|
||||||
self.assertGreater(len(out), 0)
|
|
||||||
|
|
||||||
anchor_out_replayed = _greedy_generate(self.base_url, anchor_prompt, 32)
|
|
||||||
self.assertEqual(
|
|
||||||
anchor_out,
|
|
||||||
anchor_out_replayed,
|
|
||||||
"greedy output changed after batched decode-graph replays",
|
|
||||||
)
|
|
||||||
|
|
||||||
|
|
||||||
if __name__ == "__main__":
|
|
||||||
unittest.main(verbosity=3)
|
|
||||||
@@ -1,7 +1,9 @@
|
|||||||
"""B200 per-commit CI: DeepSeek-V4-Flash FP4 with the trtllm attention backend.
|
"""B200 per-commit CI: DeepSeek-V4-Flash FP4 with the trtllm attention backend.
|
||||||
|
|
||||||
Mirrors the four FlashMLA recipes with a uniform-FP8 KV pool and trtllm-gen
|
Mirrors two of the FlashMLA recipes with a uniform-FP8 KV pool and trtllm-gen
|
||||||
sparse MLA for decode and prefill.
|
sparse MLA for decode and prefill: the spec-decoding recipe (draft extend /
|
||||||
|
target verify / multi-step backend) and the breakable-CUDA-graph DP recipe
|
||||||
|
(DP padding, graph replay refresh, mixed chunk).
|
||||||
"""
|
"""
|
||||||
|
|
||||||
import unittest
|
import unittest
|
||||||
@@ -18,7 +20,7 @@ from sglang.test.test_utils import (
|
|||||||
try_cached_model,
|
try_cached_model,
|
||||||
)
|
)
|
||||||
|
|
||||||
register_cuda_ci(est_time=700, stage="base-c", runner_config="4-gpu-b200")
|
register_cuda_ci(est_time=500, stage="base-c", runner_config="4-gpu-b200")
|
||||||
|
|
||||||
MODEL = "deepseek-ai/DeepSeek-V4-Flash"
|
MODEL = "deepseek-ai/DeepSeek-V4-Flash"
|
||||||
SERVER_LAUNCH_TIMEOUT = 3600
|
SERVER_LAUNCH_TIMEOUT = 3600
|
||||||
@@ -40,6 +42,8 @@ class TestDSV4FlashFP4B200Trtllm(
|
|||||||
gsm8k_accuracy_thres = 0.93
|
gsm8k_accuracy_thres = 0.93
|
||||||
accept_length_thres = 2.8
|
accept_length_thres = 2.8
|
||||||
bs_1_speed_thres = 220
|
bs_1_speed_thres = 220
|
||||||
|
# Arbitrary distinctive digits; only needs to survive tokenization intact.
|
||||||
|
NEEDLE = "48173"
|
||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
def setUpClass(cls):
|
def setUpClass(cls):
|
||||||
@@ -76,100 +80,31 @@ class TestDSV4FlashFP4B200Trtllm(
|
|||||||
if hasattr(cls, "process") and cls.process:
|
if hasattr(cls, "process") and cls.process:
|
||||||
kill_process_tree(cls.process.pid)
|
kill_process_tree(cls.process.pid)
|
||||||
|
|
||||||
|
def test_long_prompt_chunked_prefill_recall(self):
|
||||||
class TestDSV4FlashFP4B200BalancedTrtllm(
|
# The needle sits in the first chunk and the question in the last, so
|
||||||
SpecDecodingMixin,
|
# only a correct multi-chunk _forward_trtllm_prefill can recall it.
|
||||||
BasicDecodeCorrectnessMixin,
|
filler = (
|
||||||
GSM8KMixin,
|
"The expedition recorded water temperature, salinity, and current "
|
||||||
CustomTestCase,
|
"speed at every station along the transect. "
|
||||||
):
|
|
||||||
"""Balanced recipe: TP=4, DP=4, DeepEP, EAGLE (1-step spec)."""
|
|
||||||
|
|
||||||
gsm8k_accuracy_thres = 0.93
|
|
||||||
accept_length_thres = 1.8
|
|
||||||
bs_1_speed_thres = 100
|
|
||||||
|
|
||||||
@classmethod
|
|
||||||
def setUpClass(cls):
|
|
||||||
cls.model = try_cached_model(MODEL)
|
|
||||||
cls.base_url = DEFAULT_URL_FOR_TEST
|
|
||||||
cls.process = popen_launch_server(
|
|
||||||
cls.model,
|
|
||||||
cls.base_url,
|
|
||||||
timeout=SERVER_LAUNCH_TIMEOUT,
|
|
||||||
other_args=[
|
|
||||||
"--trust-remote-code",
|
|
||||||
"--dsv4-attn-backend",
|
|
||||||
"trtllm",
|
|
||||||
"--tp",
|
|
||||||
"4",
|
|
||||||
"--dp",
|
|
||||||
"4",
|
|
||||||
"--enable-dp-attention",
|
|
||||||
"--moe-a2a-backend",
|
|
||||||
"deepep",
|
|
||||||
"--speculative-algorithm",
|
|
||||||
"EAGLE",
|
|
||||||
"--speculative-num-steps",
|
|
||||||
"1",
|
|
||||||
"--speculative-eagle-topk",
|
|
||||||
"1",
|
|
||||||
"--speculative-num-draft-tokens",
|
|
||||||
"2",
|
|
||||||
"--deepep-config",
|
|
||||||
DEEPEP_CONFIG,
|
|
||||||
],
|
|
||||||
env=_DEEPEP_ENV,
|
|
||||||
)
|
)
|
||||||
|
prompt = (
|
||||||
@classmethod
|
f"The station beacon identifier is {self.NEEDLE}.\n\n"
|
||||||
def tearDownClass(cls):
|
+ "".join(f"[Entry {i}] {filler}" for i in range(220))
|
||||||
if hasattr(cls, "process") and cls.process:
|
+ "\n\nQ: What is the station beacon identifier? Reply with just "
|
||||||
kill_process_tree(cls.process.pid)
|
"the number.\nA:"
|
||||||
|
|
||||||
|
|
||||||
class TestDSV4FlashFP4NonMTPB200Trtllm(
|
|
||||||
BasicDecodeCorrectnessMixin, GSM8KMixin, CustomTestCase
|
|
||||||
):
|
|
||||||
"""Non-MTP recipe: TP=4, DP=4, DeepEP, no speculative decoding."""
|
|
||||||
|
|
||||||
gsm8k_accuracy_thres = 0.93
|
|
||||||
|
|
||||||
@classmethod
|
|
||||||
def setUpClass(cls):
|
|
||||||
cls.model = try_cached_model(MODEL)
|
|
||||||
cls.base_url = DEFAULT_URL_FOR_TEST
|
|
||||||
cls.process = popen_launch_server(
|
|
||||||
cls.model,
|
|
||||||
cls.base_url,
|
|
||||||
timeout=SERVER_LAUNCH_TIMEOUT,
|
|
||||||
other_args=[
|
|
||||||
"--trust-remote-code",
|
|
||||||
"--dsv4-attn-backend",
|
|
||||||
"trtllm",
|
|
||||||
"--tp",
|
|
||||||
"4",
|
|
||||||
"--dp",
|
|
||||||
"4",
|
|
||||||
"--enable-dp-attention",
|
|
||||||
"--moe-a2a-backend",
|
|
||||||
"deepep",
|
|
||||||
"--deepep-config",
|
|
||||||
DEEPEP_CONFIG,
|
|
||||||
],
|
|
||||||
env=_DEEPEP_ENV,
|
|
||||||
)
|
)
|
||||||
|
# Second pass extends from the radix-cached prefix instead of prefilling it.
|
||||||
@classmethod
|
for label in ("cold", "cached-prefix"):
|
||||||
def tearDownClass(cls):
|
out = self._decode_generate(
|
||||||
if hasattr(cls, "process") and cls.process:
|
prompt=prompt, max_new_tokens=self.sanity_max_new_tokens_short
|
||||||
kill_process_tree(cls.process.pid)
|
)
|
||||||
|
self.assertIn(self.NEEDLE, out, f"{label}: {out!r}")
|
||||||
|
|
||||||
|
|
||||||
class TestDSV4FlashFP4BreakableCudaGraphB200Trtllm(
|
class TestDSV4FlashFP4BreakableCudaGraphB200Trtllm(
|
||||||
BasicDecodeCorrectnessMixin, GSM8KMixin, CustomTestCase
|
BasicDecodeCorrectnessMixin, GSM8KMixin, CustomTestCase
|
||||||
):
|
):
|
||||||
"""BCG recipe: TP=4, DP=4, DeepEP, DP attention, mixed chunk."""
|
"""BCG recipe: TP=4, DP=4, DeepEP, DP attention, mixed chunk, no spec."""
|
||||||
|
|
||||||
gsm8k_accuracy_thres = 0.93
|
gsm8k_accuracy_thres = 0.93
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user