[CI] Trim DSV4 trtllm B200 tests (#39013)

Co-authored-by: Mohammad Angkad <mohammad.angkad@radixark.ai>
This commit is contained in:
Mohammad Miadh Angkad
2026-09-11 01:04:16 -07:00
committed by GitHub
co-authored by Mohammad Angkad
parent 17fa5ad327
commit ad5af539cd
3 changed files with 26 additions and 309 deletions
@@ -604,11 +604,7 @@ class TestDSV4BreakableCudaGraphMetadataContract(CustomTestCase):
def test_trtllm_semaphore_capacity_covers_configured_query_rows(self):
from sglang.srt.layers.attention import deepseek_v4_trtllm_backend as trtllm
schedule = SimpleNamespace(
max_prefill_tokens=16384,
chunked_prefill_size=4096,
max_running_requests=256,
)
schedule = SimpleNamespace(max_prefill_tokens=16384, max_running_requests=256)
spec = SimpleNamespace(
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.
Mirrors the four FlashMLA recipes with a uniform-FP8 KV pool and trtllm-gen
sparse MLA for decode and prefill.
Mirrors two of the FlashMLA recipes with a uniform-FP8 KV pool and trtllm-gen
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
@@ -18,7 +20,7 @@ from sglang.test.test_utils import (
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"
SERVER_LAUNCH_TIMEOUT = 3600
@@ -40,6 +42,8 @@ class TestDSV4FlashFP4B200Trtllm(
gsm8k_accuracy_thres = 0.93
accept_length_thres = 2.8
bs_1_speed_thres = 220
# Arbitrary distinctive digits; only needs to survive tokenization intact.
NEEDLE = "48173"
@classmethod
def setUpClass(cls):
@@ -76,100 +80,31 @@ class TestDSV4FlashFP4B200Trtllm(
if hasattr(cls, "process") and cls.process:
kill_process_tree(cls.process.pid)
class TestDSV4FlashFP4B200BalancedTrtllm(
SpecDecodingMixin,
BasicDecodeCorrectnessMixin,
GSM8KMixin,
CustomTestCase,
):
"""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,
def test_long_prompt_chunked_prefill_recall(self):
# The needle sits in the first chunk and the question in the last, so
# only a correct multi-chunk _forward_trtllm_prefill can recall it.
filler = (
"The expedition recorded water temperature, salinity, and current "
"speed at every station along the transect. "
)
@classmethod
def tearDownClass(cls):
if hasattr(cls, "process") and cls.process:
kill_process_tree(cls.process.pid)
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,
prompt = (
f"The station beacon identifier is {self.NEEDLE}.\n\n"
+ "".join(f"[Entry {i}] {filler}" for i in range(220))
+ "\n\nQ: What is the station beacon identifier? Reply with just "
"the number.\nA:"
)
@classmethod
def tearDownClass(cls):
if hasattr(cls, "process") and cls.process:
kill_process_tree(cls.process.pid)
# Second pass extends from the radix-cached prefix instead of prefilling it.
for label in ("cold", "cached-prefix"):
out = self._decode_generate(
prompt=prompt, max_new_tokens=self.sanity_max_new_tokens_short
)
self.assertIn(self.NEEDLE, out, f"{label}: {out!r}")
class TestDSV4FlashFP4BreakableCudaGraphB200Trtllm(
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