[CI] Prune redundant CPU test overhead (#34309)

Co-authored-by: hnyls2002 <lsyincs@gmail.com>
This commit is contained in:
Xinyuan Tong
2026-08-13 19:51:17 -07:00
committed by GitHub
co-authored by hnyls2002
parent 6ad3f2d8fd
commit 85cdf1178d
78 changed files with 884 additions and 981 deletions
+3 -2
View File
@@ -26,8 +26,9 @@ A new unit test case must fall into one of these categories:
3. **Critical-path bookkeeping.** Defends conventions that are easy to break by
forgetting to sync -- registry completeness, field lifecycle, serialization
compatibility. Enumerating assertions are fine here; the guarded failure
mode is "someone extended X without updating Y". Example: the ratchet tests
(`test/registered/unit/test_module_state_ratchet.py`).
mode is "someone extended X without updating Y". Example: the namespace
coverage tests (`test/registered/unit/test_server_args_namespaces.py`). Static
source ratchets belong in `scripts/lint/` as checkers, not unit tests.
Not admissible:
+14 -14
View File
@@ -136,7 +136,7 @@ requires today). A process-global seed field-read of one of these
sizes (`get_server_args().tp_size`, or an alias of it) is a read-ratchet failure; the
sites that legitimately go around the live property are the `configured_*_size()`
readers, and those are what the ratchet registers, each with its reason
(`_CONFIGURED_SIZE_CALL_SITES` in `test_global_config_read_ratchet.py`). A
(`_CONFIGURED_SIZE_CALL_SITES` in `scripts/lint/check_global_config_read_ratchet.py`). A
`server_args` the object was *handed* is a different thing and not a ratchet
matter — see "Reads that legitimately stay on a ServerArgs instance".
Fail-loud is narrower: before dist init, a live size/group read raises — except
@@ -189,11 +189,11 @@ this).
topology (`1` / `False` when no group is installed). A site
that must know the *requested* DCP size before dist init needs its own
`configured_dcp_size()` (and an entry in `_CONFIGURED_SIZE_CALL_SITES`, which lives
in the ratchet test, not in this skill); note the live pair does not *need*
in the ratchet checker, not in this skill); note the live pair does not *need*
dist init — with no group it answers `1` / `False` — it just cannot answer
with the requested size. Every (file, accessor) pair is registered
with its reason in `test_global_config_read_ratchet.py`
(`_CONFIGURED_SIZE_CALL_SITES`), and that test fails if the code and the list
with its reason in `scripts/lint/check_global_config_read_ratchet.py`
(`_CONFIGURED_SIZE_CALL_SITES`), and that checker fails if the code and the list
disagree — a new file, or a new accessor in a listed file, has to be added — so a new site needs both an answer the live property cannot give and
an entry saying what it is.
- **this runner's resolved value** → the runner
@@ -378,31 +378,31 @@ ONE thread — do not design for TBO threads that don't exist.
resolved config with `get_context().override`; hand a per-runner value to its
runner as a constructor argument. Projected bags are sealed the same way (leaf
assignment raises).
2. **Mutation ratchet** (`test_server_args_mutation_ratchet.py`, exact pin 0 over the whole
2. **Mutation guard** (`scripts/lint/check_server_args_mutation_ratchet.py`, pinned at 0 over the whole
package minus the pipeline / multimodal_gen): textual scan for assignment forms. Never
raise the baseline.
3. **No-copy contract** (`test_server_args_no_instance_mutation_entry.py`): neither
`ServerArgs.override` nor `ServerArgs.derive` exists, and nothing in the package
calls either form. Rerouting a writer to the bags means flipping **all its readers
in the same commit** (no transitional dual-write).
4. **Legacy-accessor ratchet** (`test_legacy_global_ratchet.py`): `get_global_server_args`
4. **Legacy-accessor ratchet** (`scripts/lint/check_legacy_global_ratchet.py`): `get_global_server_args`
call sites must not grow. The replacement for a *decision* is a bag leaf, a named
accessor, or the owning runner's stamp — not `get_server_args().field`, which the
read ratchet below pins at zero. `runtime_context.get_server_args()` is only for the
whole-object shapes (dumps, provenance, a hand-off to a callee that takes a config).
5. **Global config read ratchet** (`test_global_config_read_ratchet.py`): baselines are
5. **Global config read ratchet** (`scripts/lint/check_global_config_read_ratchet.py`): baselines are
**0** for both the direct `get_server_args().field` and the alias form (function-local
— including local copies of an alias, `cfg = sa` — module-level, or parked on an
instance attribute, plus the `getattr(..., "field")` spelling of each; a name
computed at runtime or indirection deeper than a local name copy is census-tool
territory, per the test's docstring). The scanners match `get_server_args` and
territory, per the checker's module docstring). The scanners match `get_server_args` and
`configured_*_size` by their literal names, and the same file *bans*
`import ... as` renames of them so that matching stays sound. Exempt by owner
module only (`runtime_context.py`, `server_args.py`, `arg_groups/`). The same file
carries `_CONFIGURED_SIZE_CALL_SITES`, the (file, accessor) map of every
`configured_*_size()` reader with the reason the live property cannot serve it — a new
file or a new accessor in a listed file must be added there.
6. **Module-state ratchet** (`test_module_state_ratchet.py`): `global` statements in the
6. **Module-state ratchet** (`scripts/lint/check_module_state_ratchet.py`): `global` statements in the
flag-owning layers are pinned by name. A new module-level runtime global belongs on a
flags group / resources slot instead; migrating a pinned survivor must shrink the pin.
7. **Namespace coverage** (`test_server_args_namespaces.py`,
@@ -456,9 +456,9 @@ Key source files: `python/sglang/srt/runtime_context.py` (the container, every t
`publish`, `_ConfigBag`, `preserve_config`, `override_server_args`),
`python/sglang/srt/arg_groups/overrides.py` (override registry, passes,
`declare_late_resolution`), `python/sglang/srt/server_args.py` (`NS` metadata,
`Arg(..., resolvable=True)`, `__setattr__` strict guard), and the guardrail tests under
`test/registered/unit/` (`test_server_args_mutation_ratchet.py`,
`test_global_config_read_ratchet.py`, `test_legacy_global_ratchet.py`,
`test_module_state_ratchet.py`, `test_server_args_namespaces.py`,
`test_runtime_context.py` — the last one doubles
`Arg(..., resolvable=True)`, `__setattr__` strict guard), the static guardrails under
`scripts/lint/` (`check_server_args_mutation_ratchet.py`,
`check_global_config_read_ratchet.py`, `check_legacy_global_ratchet.py`,
`check_module_state_ratchet.py`), and the runtime guardrail tests under
`test/registered/unit/` (`test_server_args_namespaces.py`, `test_runtime_context.py` — the latter doubles
as executable documentation of every tier's semantics).
+6
View File
@@ -50,6 +50,12 @@ jobs:
- name: Run pre-commit checks
run: SKIP=no-commit-to-branch pre-commit run --all-files --show-diff-on-failure
# Not in the rust-ext build job: its cache key covers the built .so files,
# and a test script is not a build input. Tests take ~1s; the timeout is
# for a cold cache, which codegens the dependency graph first.
- name: Run rust/ workspace tests
run: cd rust && timeout 900 cargo test --workspace
- name: Run lychee docs checks (offline references)
uses: lycheeverse/lychee-action@8646ba30535128ac92d33dfc9133794bfdd9b411 # v2
with:
+21 -5
View File
@@ -99,26 +99,42 @@ repos:
pass_filenames: false
- id: check-workflow-job-names
name: check for duplicate workflow job names
entry: python3 scripts/ci/check_workflow_job_names.py
entry: python3 scripts/lint/check_workflow_job_names.py
language: system
files: ^\.github/workflows/.*\.yml$
pass_filenames: false
- id: check-rust-ext-cache-prefix
name: check rust-ext cache key prefix defaults match
entry: python3 scripts/ci/check_rust_ext_cache_prefix.py
entry: python3 scripts/lint/check_rust_ext_cache_prefix.py
language: system
files: ^(\.github/actions/download-rust-ext/action\.yml|\.github/workflows/_pr-test-rust-ext-build\.yml)$
pass_filenames: false
- id: check-no-bare-pytest-main
name: reject bare pytest.main calls in __main__ blocks
entry: python3 scripts/lint/check_no_bare_pytest_main.py
language: system
files: ^(python|test)/.*\.py$
- id: check-static-ratchets
name: validate static runtime ratchets
entry: python3 scripts/lint/check_static_ratchets.py
language: system
files: ^(python/sglang/.*\.py|scripts/lint/check_.*\.py)$
pass_filenames: false
- id: check-lint-script-tests
name: unit tests for lint checkers
entry: python3 -m unittest discover -s scripts/lint -p 'test_check_*.py'
language: system
files: ^scripts/lint/(check_|test_check_).*\.py$
pass_filenames: false
- id: check-registered-tests
name: validate registered test CI registries
entry: python3 scripts/ci/check_registered_tests.py
entry: python3 scripts/lint/check_registered_tests.py
language: system
files: ^test/registered/.*\.py$
exclude: ^test/registered/.*/utils\.py$
pass_filenames: false
- id: check-no-registered-tests-in-package
name: reject CI-registered tests inside the sglang package
entry: python3 scripts/ci/check_no_registered_tests_in_package.py
entry: python3 scripts/lint/check_no_registered_tests_in_package.py
language: system
files: ^python/sglang/.*\.py$
pass_filenames: false
+1 -1
View File
@@ -9603,7 +9603,7 @@ def m3_fp8_attn_gemm_enabled(args) -> bool:
# (sglang.srt.runtime_context). The two functions below are LEGACY shims kept
# for the existing call-sites; they publish/read the same live object by
# reference. Do not add new call-sites — the counts are ratcheted
# (decrease-only) by test/registered/unit/test_legacy_global_ratchet.py.
# (decrease-only) by scripts/lint/check_legacy_global_ratchet.py.
# Imports are in-function so the two modules stay cycle-free at import time.
def set_global_server_args_for_scheduler(server_args: ServerArgs):
"""Legacy publish shim (role=scheduler) — prefer
+22
View File
@@ -0,0 +1,22 @@
"""Deterministic batch-shape coverage for the input-logprob sweeps."""
import itertools
def coverage_cases(menu, max_seqs):
"""Every singleton, every ordered pair, and wider heterogeneous cases.
`menu` order is load-bearing: width >= 3 walks it cyclically, so reordering
it changes which wide combinations run.
"""
yield from ((item,) for item in menu)
yield from itertools.product(menu, repeat=2)
for width in range(3, max_seqs + 1):
for offset in range(len(menu)):
yield tuple(menu[(offset + step) % len(menu)] for step in range(width))
yield tuple(menu[(offset - step) % len(menu)] for step in range(width))
# Cyclic windows never repeat an item; adjacent duplicates need their own.
for index, item in enumerate(menu):
other = menu[(index + 1) % len(menu)]
yield (item,) * (width - 1) + (other,)
yield (other,) + (item,) * (width - 1)
-39
View File
@@ -1,39 +0,0 @@
#!/usr/bin/env python3
"""Check that the rust-ext cache_key_prefix defaults stay in sync.
The build workflow saves cache entries under its default; the download action
restores with its own. Neither file can reference the other, and a mismatch
makes every pool silently fall back to source builds at install time.
"""
import sys
import yaml
BUILD_WORKFLOW = ".github/workflows/_pr-test-rust-ext-build.yml"
DOWNLOAD_ACTION = ".github/actions/download-rust-ext/action.yml"
def main() -> int:
with open(BUILD_WORKFLOW, encoding="utf-8") as f:
workflow = yaml.safe_load(f)
with open(DOWNLOAD_ACTION, encoding="utf-8") as f:
action = yaml.safe_load(f)
# yaml 1.1 parses the `on:` key as boolean True
triggers = workflow.get("on", workflow.get(True))
save_prefix = triggers["workflow_call"]["inputs"]["cache_key_prefix"]["default"]
restore_prefix = action["inputs"]["cache_key_prefix"]["default"]
if save_prefix != restore_prefix:
print("ERROR: rust-ext cache_key_prefix defaults do not match.")
print(f" {BUILD_WORKFLOW} saves under: {save_prefix}")
print(f" {DOWNLOAD_ACTION} restores with: {restore_prefix}")
print("Bump both together, or every pool falls back to source builds.")
return 1
return 0
if __name__ == "__main__":
sys.exit(main())
+5 -9
View File
@@ -104,12 +104,6 @@ _FILE_EXTENSIONS = (
".onnx",
)
# Non-test helper files under test/registered/ (skipped by basename, matching
# scripts/ci/check_registered_tests.py). run_suite.py skips `cpu/utils.py` by
# path; excluding every `utils.py` by basename is a superset that drops no
# CUDA-registered test (the other `utils.py` registers CPU only).
_NON_TEST_BASENAMES = frozenset({"conftest.py", "__init__.py", "utils.py"})
def looks_like_model_id(value: str, deny: Optional[Set[str]] = None) -> bool:
"""Heuristic: does ``value`` look like a HuggingFace repo id?
@@ -273,11 +267,13 @@ def collect_suite_files(
ci_register = _load_ci_register(repo_root)
backend = getattr(ci_register.HWBackend, backend_name.upper())
pattern = os.path.join(repo_root, "test", "registered", "**", "*.py")
# Same exclusion as run_suite.py: pytest+package structure files.
files = sorted(
f
for f in glob.glob(pattern, recursive=True)
if os.path.basename(f) not in _NON_TEST_BASENAMES
for f in glob.glob(
os.path.join(repo_root, "test", "registered", "**", "*.py"), recursive=True
)
if os.path.basename(f) not in ("conftest.py", "__init__.py")
)
suite_files: Dict[str, List[str]] = {}
+2 -1
View File
@@ -88,7 +88,8 @@ def discover_files(repo_root: str) -> list[str]:
for f in glob.glob(
os.path.join(test_dir, "registered", "**", "*.py"), recursive=True
)
if not f.endswith("/conftest.py") and not f.endswith("/__init__.py")
# Same exclusion as run_suite.py: pytest+package structure files.
if os.path.basename(f) not in ("conftest.py", "__init__.py")
]
jit_kernel_dir = os.path.join(repo_root, "python", "sglang", "jit_kernel")
files += glob.glob(
@@ -12,17 +12,11 @@ leak checker, hence this AST-level guard.
"""
import ast
import unittest
import warnings
from collections import Counter
from pathlib import Path
from sglang.test.ci.ci_register import register_cpu_ci
from sglang.test.test_utils import CustomTestCase
register_cpu_ci(est_time=8, suite="base-a-test-cpu")
_REPO_ROOT = Path(__file__).resolve().parents[4]
_REPO_ROOT = Path(__file__).resolve().parents[2]
_SRT_DIR = _REPO_ROOT / "python" / "sglang" / "srt"
_SPECULATIVE_DIR = _SRT_DIR / "speculative"
assert _SRT_DIR.is_dir(), f"srt dir not found: {_SRT_DIR}"
@@ -193,40 +187,41 @@ def _scan_class_subtree(class_node):
return sites
class TestDecodeBookkeepingOwnership(CustomTestCase):
def test_bookkeeping_sites_match_owner_allowlist(self):
found = _scan_srt()
allow = Counter(_OWNER_SITES)
unexpected = found - allow
missing = allow - found
msg = []
if unexpected:
msg.append(
"New bookkeeping mutation(s) beyond the recorded counts:\n "
+ "\n ".join(f"{site} x{n}" for site, n in sorted(unexpected.items()))
+ "\nThese are owned by the sites in _OWNER_SITES -- do not "
"repeat them; a genuinely new owner must be recorded there."
)
if missing:
msg.append(
"Recorded site(s) no longer exist (update _OWNER_SITES):\n "
+ "\n ".join(f"{site} x{n}" for site, n in sorted(missing.items()))
)
self.assertFalse(msg, "\n\n".join(msg))
def check_bookkeeping_sites_match_owner_allowlist():
found = _scan_srt()
allow = Counter(_OWNER_SITES)
unexpected = found - allow
missing = allow - found
messages = []
if unexpected:
messages.append(
"New bookkeeping mutation(s) beyond the recorded counts:\n "
+ "\n ".join(f"{site} x{n}" for site, n in sorted(unexpected.items()))
+ "\nThese are owned by the sites in _OWNER_SITES -- do not "
"repeat them; a genuinely new owner must be recorded there."
)
if missing:
messages.append(
"Recorded site(s) no longer exist (update _OWNER_SITES):\n "
+ "\n ".join(f"{site} x{n}" for site, n in sorted(missing.items()))
)
if messages:
raise AssertionError("\n\n".join(messages))
def test_spec_v2_draft_workers_do_no_scheduler_bookkeeping(self):
classes = _draft_worker_classes()
names = {node.name for _, node in classes}
# Discovery sanity: fail loudly instead of silently guarding nothing.
self.assertIn("EagleDraftWorker", names)
self.assertIn("FrozenKVMTPDraftWorker", names)
violations = []
for rel, node in classes:
for scope, kind in _scan_class_subtree(node):
violations.append((rel, f"{node.name}.{scope}", kind))
self.assertFalse(
violations,
def check_spec_v2_draft_workers_do_no_scheduler_bookkeeping():
classes = _draft_worker_classes()
names = {node.name for _, node in classes}
for expected in ("EagleDraftWorker", "FrozenKVMTPDraftWorker"):
if expected not in names:
raise AssertionError(f"draft worker discovery missed {expected}")
violations = []
for rel, node in classes:
for scope, kind in _scan_class_subtree(node):
violations.append((rel, f"{node.name}.{scope}", kind))
if violations:
raise AssertionError(
"Spec-v2 draft worker(s) repeat scheduler-owned bookkeeping:\n "
+ "\n ".join(map(str, sorted(violations)))
+ "\nUnder spec v2 the iter-clock ticks, `maybe_evict_swa`, and "
@@ -236,4 +231,5 @@ class TestDecodeBookkeepingOwnership(CustomTestCase):
if __name__ == "__main__":
unittest.main(verbosity=3)
check_bookkeeping_sites_match_owner_allowlist()
check_spec_v2_draft_workers_do_no_scheduler_bookkeeping()
@@ -1,73 +1,24 @@
"""Ratchet guard: process-global config reads may only decrease.
"""Guard: business code never reads a config field off the process-global record.
``get_server_args()`` returns the published ``ServerArgs`` one process's
``get_server_args()`` returns the published ``ServerArgs`` -- one process's
startup record. Config decisions read the namespace accessors instead
(``get_exec()`` / ``get_memory()`` / ), which carry the resolved value
including post-publish overrides, and per-runner values come from the runner
that owns them.
(``get_exec()`` / ``get_memory()`` / ...); per-runner values come from the
runner that owns them. Both baselines are zero, over the whole package minus
the modules that own the slot.
Business code no longer reads the published record for a config value at all:
the baselines are zero for both shapes, over the whole package minus the two
modules that own the slot.
Where the remaining reads live (``runtime_context.py``, exempt by module):
- **Derived members.** ``@property`` / method members of ``ServerArgs``
(``mamba_cache_chunk_size``, ``max_speculative_num_draft_tokens``,
``use_mla_backend()``, ``get_attention_backends()``, ``get_model_config()``,
``cutedsl_moe_max_num_tokens()``) are computed from several fields plus the HF
config, so they are not namespace leaves and ``ServerArgs`` is their only
home. ``runtime_context`` exposes each one as a named accessor
(``mamba_cache_chunk_size()`` ) and is the only module that reads the slot
for them.
- **Config-intent reads of live-shadowed sizes.** ``get_parallel()`` shadows
``tp/pp/dcp/attn_cp/moe_dp_size`` with the live topology, and a few call sites
need what was *configured*: the ``configured_*_size()`` accessors. Their
reasons, per call site:
- ``dsa_indexer.pp_size`` gates ``pp_size > 1 and not get_pp_group()...``, and
the short circuit is the point: with PP off the group is never touched, which
is what lets the ``Indexer`` be constructed before distributed init. The live
property would demand the group either way.
- ``dp_attention.attn_cp_size`` / ``moe_dp_size``: the configuration the
predicate detects (``attn_cp_size > moe_dp_size``) is the one where
``initialize_model_parallel`` aliases ``_MOE_DP`` to ``_ATTN_CP``, so the live
sizes are equal there and a live comparison is always false.
- ``model_loader/loader.py`` reports both: the same dict carries the live
``moe_dp_size`` under ``"dp"``, so this entry is the configured intent.
What the ratchet sees, syntactically: ``get_server_args().field``,
``sa = get_server_args()`` followed by ``sa.field`` (function-local, module-level,
or parked on an instance attribute -- ``self._sa = get_server_args()`` read from
another method of the same class), function-local copies of an alias to a
fixpoint (``cfg = sa`` then ``cfg.field``), and the dynamic form of each --
``getattr(<either>, "field")`` -- since a string-named read reaches the same
slot. What it cannot see is a name computed at runtime (``getattr(sa, name)``)
or indirection deeper than a local name copy (through a container, an
attribute of another object, a cross-scope copy); the census tool in the
context repo is what audits those.
A whole-object pass (``def f(server_args)``) is not a global read and is not
counted -- there the caller decided which instance to hand over. An optional
parameter that falls back to the global (``f(server_args=None)``) hides one,
so those fallbacks were removed; the ratchet cannot see them and the census
tool in the context repo is what audits that shape.
The scanners match ``get_server_args`` and ``configured_*_size`` by their
literal names, which is why import-renaming them is banned below. A name
computed at runtime, or indirection deeper than a local name copy, is invisible
here -- the census tool in the context repo audits that shape.
"""
from sglang.test.ci.ci_register import register_cpu_ci
register_cpu_ci(est_time=5, suite="base-a-test-cpu")
import ast
import unittest
from functools import cache
from pathlib import Path
import sglang
from sglang.test.test_utils import CustomTestCase
# srt is the migrated surface; the rest of the package has no reads today and is
# scanned so a new one cannot appear there unnoticed.
_PACKAGE_ROOT = Path(next(iter(sglang.__path__)))
_PACKAGE_ROOT = Path(__file__).resolve().parents[2] / "python" / "sglang"
# The modules that own the slot: runtime_context publishes it and exposes the
# named accessors for the derived members, server_args/arg_groups ARE the
@@ -75,7 +26,7 @@ _PACKAGE_ROOT = Path(next(iter(sglang.__path__)))
_SLOT_OWNERS = ("srt/runtime_context.py", "srt/server_args.py", "srt/arg_groups/")
# Every call site of a ``configured_*_size()`` accessor, with the reason the
# live topology cannot answer there. The test below asserts this map is exactly
# live topology cannot answer there. The checker below asserts this map is exactly
# the set of call sites, so the reasons cannot drift away from the code.
_CONFIGURED_SIZE_CALL_SITES = {
("srt/layers/attention/dsa/dsa_indexer.py", "configured_pp_size"): (
@@ -392,16 +343,24 @@ def _collect(rel: str, tree: ast.AST, inert: frozenset = frozenset()):
return direct, alias
def _field_reads():
direct, alias = [], []
@cache
def _parsed_modules():
"""(rel, tree) per parseable module; the three scanners below share it."""
modules = []
for path in sorted(_PACKAGE_ROOT.rglob("*.py")):
rel = path.relative_to(_PACKAGE_ROOT).as_posix()
if rel.startswith(_SLOT_OWNERS):
continue
try:
tree = ast.parse(path.read_text())
except SyntaxError:
continue
modules.append((path.relative_to(_PACKAGE_ROOT).as_posix(), tree))
return modules
def _field_reads():
direct, alias = [], []
for rel, tree in _parsed_modules():
if rel.startswith(_SLOT_OWNERS):
continue
inert = frozenset(name for path_, name in _INERT_DYNAMIC_READS if path_ == rel)
module_direct, module_alias = _collect(rel, tree, inert)
direct += module_direct
@@ -409,29 +368,61 @@ def _field_reads():
return direct, alias
class TestGlobalConfigReadRatchet(CustomTestCase):
def _check(self, kind, reads, baseline):
if len(reads) > baseline:
self.fail(
f"{kind} process-global config field reads grew: {len(reads)} > "
f"baseline {baseline}. Read the namespace accessor for the "
"field's namespace, or the owning runner for a per-runner "
"field:\n" + "\n".join(reads)
def _configured_size_call_sites():
found = set()
for rel, tree in _parsed_modules():
if rel.startswith(_SLOT_OWNERS):
continue
for node in ast.walk(tree):
if not isinstance(node, ast.Call):
continue
func = node.func
name = (
func.id
if isinstance(func, ast.Name)
else (func.attr if isinstance(func, ast.Attribute) else None)
)
if len(reads) < baseline:
self.fail(
f"{kind} process-global config field reads shrank: {len(reads)} < "
f"baseline {baseline}. Lower the baseline in this file to lock "
"in the progress."
)
def test_global_field_reads_match_the_baseline(self):
direct, alias = _field_reads()
self._check("direct", direct, _DIRECT_BASELINE)
self._check("alias-form", alias, _ALIAS_BASELINE)
if name and name.startswith("configured_") and name.endswith("_size"):
found.add((rel, name))
return found
class TestConfiguredSizeCallSites(CustomTestCase):
def _renamed_accessor_imports():
offenders = []
for rel, tree in _parsed_modules():
for node in ast.walk(tree):
if not isinstance(node, (ast.ImportFrom, ast.Import)):
continue
for imported in node.names:
if imported.asname is None or imported.asname == imported.name:
continue
base = imported.name.rsplit(".", 1)[-1]
if base == "get_server_args" or (
base.startswith("configured_") and base.endswith("_size")
):
offenders.append(
f"{rel}:{node.lineno}: {imported.name} as {imported.asname}"
)
return offenders
def _check_count(kind, reads, baseline):
if len(reads) > baseline:
raise AssertionError(
f"{kind} process-global config field reads grew: {len(reads)} > "
f"baseline {baseline}. Read the namespace accessor for the "
"field's namespace, or the owning runner for a per-runner "
"field:\n" + "\n".join(reads)
)
def check_global_config_read_ratchet():
direct, alias = _field_reads()
_check_count("direct", direct, _DIRECT_BASELINE)
_check_count("alias-form", alias, _ALIAS_BASELINE)
def check_configured_size_call_sites():
"""The configured-vs-live exceptions are enumerated, with reasons.
``configured_*_size()`` answers what the user asked for where
@@ -446,38 +437,17 @@ class TestConfiguredSizeCallSites(CustomTestCase):
what this catches -- in either call form (bare or module-qualified).
"""
def test_the_call_sites_match_the_documented_set(self):
found = set()
for path in sorted(_PACKAGE_ROOT.rglob("*.py")):
rel = path.relative_to(_PACKAGE_ROOT).as_posix()
if rel.startswith(_SLOT_OWNERS):
continue
try:
tree = ast.parse(path.read_text())
except SyntaxError:
continue
for node in ast.walk(tree):
if not isinstance(node, ast.Call):
continue
func = node.func
name = (
func.id
if isinstance(func, ast.Name)
else (func.attr if isinstance(func, ast.Attribute) else None)
)
if name and name.startswith("configured_") and name.endswith("_size"):
found.add((rel, name))
documented = set(_CONFIGURED_SIZE_CALL_SITES)
self.assertEqual(
documented,
found,
found = _configured_size_call_sites()
documented = set(_CONFIGURED_SIZE_CALL_SITES)
if documented != found:
raise AssertionError(
"configured-size call sites drifted from their documented reasons.\n"
f" undocumented: {sorted(found - documented)}\n"
f" stale entries: {sorted(documented - found)}",
)
class TestNoRenamedAccessorImports(CustomTestCase):
def check_no_renamed_accessor_imports():
"""The scanners above match ``get_server_args`` and ``configured_*_size``
by their literal names, so an ``import ... as`` rename would walk a read
straight past both the zero baseline and the call-site registry. Renaming
@@ -485,30 +455,9 @@ class TestNoRenamedAccessorImports(CustomTestCase):
so it is banned outright which is exactly what makes literal-name
matching sound."""
def test_the_scanned_accessors_are_never_import_renamed(self):
offenders = []
for path in sorted(_PACKAGE_ROOT.rglob("*.py")):
rel = path.relative_to(_PACKAGE_ROOT).as_posix()
try:
tree = ast.parse(path.read_text())
except SyntaxError:
continue
for node in ast.walk(tree):
if not isinstance(node, (ast.ImportFrom, ast.Import)):
continue
for imported in node.names:
if imported.asname is None or imported.asname == imported.name:
continue
base = imported.name.rsplit(".", 1)[-1]
if base == "get_server_args" or (
base.startswith("configured_") and base.endswith("_size")
):
offenders.append(
f"{rel}:{node.lineno}: {imported.name} as "
f"{imported.asname}"
)
self.assertFalse(
offenders,
offenders = _renamed_accessor_imports()
if offenders:
raise AssertionError(
"get_server_args / configured_*_size imported under another name; "
"the read ratchet and the configured-size registry match these "
"accessors by their literal names, so a rename silently escapes "
@@ -517,4 +466,6 @@ class TestNoRenamedAccessorImports(CustomTestCase):
if __name__ == "__main__":
unittest.main()
check_global_config_read_ratchet()
check_configured_size_call_sites()
check_no_renamed_accessor_imports()
@@ -0,0 +1,52 @@
"""Ratchet guard: legacy global-accessor call-sites may only decrease.
The process-wide ``ServerArgs`` is owned by the runtime context; the legacy
``get_global_server_args`` / ``set_global_server_args_for_*`` names survive as
thin shims for the existing call-sites. New code should use the
``sglang.srt.runtime_context`` accessors (``get_server_args()`` /
``get_context().set_server_args()``), so the shim call-site counts below must
never grow. When your change removes call-sites, lower the matching baseline
to the new count.
"""
import re
from pathlib import Path
_SRT_ROOT = Path(__file__).resolve().parents[2] / "python" / "sglang" / "srt"
# Baselines counted over python/sglang/srt/**/*.py, including each function's
# own def line. Ratchet: decrease-only.
_RATCHETS = [
# Down to the shim definition itself; every call-site now goes through
# runtime_context.get_server_args().
("get_global_server_args", r"\bget_global_server_args\s*\(", 1),
(
"set_global_server_args_for_*",
r"\bset_global_server_args_for_(?:scheduler|tokenizer)\s*\(",
4,
),
]
def check_legacy_global_ratchet():
sources = [
path.read_text(encoding="utf-8", errors="replace")
for path in sorted(_SRT_ROOT.rglob("*.py"))
]
for name, pattern, baseline in _RATCHETS:
count = sum(len(re.findall(pattern, source)) for source in sources)
if count > baseline:
raise AssertionError(
f"{name} call-sites grew: {count} > baseline {baseline}. "
"New code must use the sglang.srt.runtime_context accessors "
"(get_server_args() / get_context().set_server_args())."
)
if count < baseline:
raise AssertionError(
f"{name} call-sites shrank: {count} < baseline {baseline}. "
"Lower the baseline in this file to lock in the progress."
)
if __name__ == "__main__":
check_legacy_global_ratchet()
@@ -0,0 +1,52 @@
"""Ratchet guard: module-level runtime state in the flag-owning layers may
only shrink.
Runtime flags belong on ``get_flags()`` groups, which have lifecycle reset and
a scoped test-override primitive; a module-level ``global`` has neither and
leaks across test teardowns. The pin below names the survivors -- migrating one
must shrink it.
"""
import ast
from pathlib import Path
_SRT_ROOT = Path(__file__).resolve().parents[2] / "python" / "sglang" / "srt"
_PINNED_GLOBALS = {
"layers/moe/utils.py": frozenset(),
"layers/dp_attention.py": frozenset(
{
# DP-attention topology (parallel vertical scope).
"_ATTN_DP_RANK",
"_ATTN_DP_SIZE",
}
),
}
def check_module_state_ratchet():
for rel, pinned in _PINNED_GLOBALS.items():
tree = ast.parse((_SRT_ROOT / rel).read_text())
declared = {
name
for node in ast.walk(tree)
if isinstance(node, ast.Global)
for name in node.names
}
grown = declared - pinned
if grown:
raise AssertionError(
f"{rel} declares new module-level runtime state {sorted(grown)}; "
"put runtime flags on a get_flags() group instead "
"(see runtime_context.MoeFlags / DpFlags).",
)
shrunk = pinned - declared
if shrunk:
raise AssertionError(
f"{rel} no longer declares {sorted(shrunk)}; "
"shrink the pin in this file to lock in the progress.",
)
if __name__ == "__main__":
check_module_state_ratchet()
+154
View File
@@ -0,0 +1,154 @@
#!/usr/bin/env python3
import ast
import pathlib
import re
import sys
_PYTEST_MAIN = re.compile(r"pytest\s*\.\s*main")
def is_main_guard(node: ast.expr) -> bool:
if not isinstance(node, ast.Compare) or len(node.ops) != 1:
return False
if not isinstance(node.ops[0], ast.Eq):
return False
sides = [node.left, *node.comparators]
has_name = any(
isinstance(side, ast.Name) and side.id == "__name__" for side in sides
)
has_main = any(
isinstance(side, ast.Constant) and side.value == "__main__" for side in sides
)
return has_name and has_main
def is_pytest_main_call(node: ast.AST) -> bool:
if not isinstance(node, ast.Call):
return False
func = node.func
return (
isinstance(func, ast.Attribute)
and func.attr == "main"
and isinstance(func.value, ast.Name)
and func.value.id == "pytest"
)
def is_exit_call(node: ast.AST, parents: dict[int, ast.AST]) -> bool:
"""``sys.exit(...)``, or a ``SystemExit(...)`` that is actually raised."""
if not isinstance(node, ast.Call):
return False
func = node.func
if (
isinstance(func, ast.Attribute)
and func.attr == "exit"
and isinstance(func.value, ast.Name)
and func.value.id == "sys"
):
return True
parent = parents.get(id(node))
return (
isinstance(func, ast.Name)
and func.id == "SystemExit"
and isinstance(parent, ast.Raise)
and parent.exc is node
)
def assigned_names(node: ast.AST) -> list[str]:
if isinstance(node, ast.Assign):
return [t.id for t in node.targets if isinstance(t, ast.Name)]
if isinstance(node, (ast.AnnAssign, ast.NamedExpr)):
return [node.target.id] if isinstance(node.target, ast.Name) else []
return []
def exited_names(nodes: list[ast.AST], parents: dict[int, ast.AST]) -> set[str]:
"""Names handed to an exit call, so the two-step form still propagates."""
return {
arg.id
for node in nodes
if is_exit_call(node, parents)
for arg in node.args
if isinstance(arg, ast.Name)
}
def propagates_exit_code(
node: ast.Call, parents: dict[int, ast.AST], exited: set[str]
) -> bool:
parent = parents.get(id(node))
if (
isinstance(parent, ast.Call)
and node in parent.args
and is_exit_call(parent, parents)
):
return True
return any(name in exited for name in assigned_names(parent))
def runtime_nodes(node: ast.AST):
yield node
if isinstance(
node, (ast.FunctionDef, ast.AsyncFunctionDef, ast.ClassDef, ast.Lambda)
):
return
for child in ast.iter_child_nodes(node):
yield from runtime_nodes(child)
def find_bare_pytest_main(path: pathlib.Path) -> int | None:
try:
source = path.read_text(encoding="utf-8")
except (OSError, UnicodeDecodeError):
return None
if "__main__" not in source or _PYTEST_MAIN.search(source) is None:
return None
try:
tree = ast.parse(source, filename=str(path))
except SyntaxError:
return None
for node in ast.walk(tree):
if not isinstance(node, ast.If) or not is_main_guard(node.test):
continue
# Whole body at once: the call and the sys.exit() that propagates it
# are separate statements.
nodes = [n for statement in node.body for n in runtime_nodes(statement)]
parents = {
id(child): parent
for parent in nodes
for child in ast.iter_child_nodes(parent)
}
exited = exited_names(nodes, parents)
for candidate in nodes:
if is_pytest_main_call(candidate) and not propagates_exit_code(
candidate, parents, exited
):
return candidate.lineno
return None
def main(paths: list[str]) -> int:
offenders = []
for path_string in paths:
path = pathlib.Path(path_string)
line = find_bare_pytest_main(path)
if line is not None:
offenders.append(f"{path}:{line}")
if not offenders:
return 0
print(
"ERROR: pytest.main(...) in an __main__ block must propagate its exit "
"code with sys.exit(...) or raise SystemExit(...):"
)
for offender in offenders:
print(f" {offender}")
return 1
if __name__ == "__main__":
sys.exit(main(sys.argv[1:]))
@@ -1,30 +1,16 @@
"""Ratchet guard: legacy parallel-getter calls in swept directories may only
shrink.
"""Guard: no legacy parallel-getter calls in the swept directories.
``models/`` and ``layers/`` read parallel topology through
``get_parallel().<dim>`` (the read-through wrapper in ``runtime_context``),
which gives one import, one naming scheme, and the scoped ``override()``
test primitive. Direct calls to the ``parallel_state`` size/rank getters in
these directories are regressions against that sweep.
Exemptions, pinned by path: ``runtime_context.py`` and
``layers/dp_attention.py`` are delegation substrate, while
``layers/dcp/comm.py`` retains deprecated DCP compatibility shims for
out-of-tree callers. Sweeping an exempt path must remove it from the pin.
which gives one import, one naming scheme, and the scoped ``override()`` test
primitive. Exemptions are pinned in ``_EXEMPT``, each with its reason; sweeping
one must remove it from there.
"""
from sglang.test.ci.ci_register import register_cpu_ci
register_cpu_ci(est_time=5, suite="base-a-test-cpu")
import re
import unittest
from pathlib import Path
import sglang.srt
from sglang.test.test_utils import CustomTestCase
_SRT_ROOT = Path(next(iter(sglang.srt.__path__)))
_SRT_ROOT = Path(__file__).resolve().parents[2] / "python" / "sglang" / "srt"
_BANNED_CALLS = re.compile(
r"\b(?:dcp_enabled|get_(?:"
@@ -56,23 +42,22 @@ _EXEMPT = (
)
class TestParallelAdoptionRatchet(CustomTestCase):
def test_no_legacy_parallel_getters_in_swept_dirs(self):
offenders = []
for top in _SWEPT_DIRS:
for path in sorted((_SRT_ROOT / top).rglob("*.py")):
rel = path.relative_to(_SRT_ROOT).as_posix()
if rel.startswith(_EXEMPT):
continue
for i, line in enumerate(path.read_text().split("\n"), 1):
if _BANNED_CALLS.search(line):
offenders.append(f"{rel}:{i}")
self.assertFalse(
offenders,
def check_parallel_adoption_ratchet():
offenders = []
for top in _SWEPT_DIRS:
for path in sorted((_SRT_ROOT / top).rglob("*.py")):
rel = path.relative_to(_SRT_ROOT).as_posix()
if rel.startswith(_EXEMPT):
continue
for line_number, line in enumerate(path.read_text().split("\n"), 1):
if _BANNED_CALLS.search(line):
offenders.append(f"{rel}:{line_number}")
if offenders:
raise AssertionError(
"legacy parallel-getter calls in swept directories (use "
f"get_parallel().<dim> instead): {offenders}",
)
if __name__ == "__main__":
unittest.main()
check_parallel_adoption_ratchet()
@@ -80,11 +80,11 @@ def main() -> int:
spec.loader.exec_module(ci_register)
cuda = ci_register.HWBackend.CUDA
# Same filter as run_suite.py: skip conftest.py, __init__.py, and utils.py
# Same exclusion as run_suite.py: pytest+package structure files.
files = sorted(
f
for f in glob.glob("test/registered/**/*.py", recursive=True)
if os.path.basename(f) not in ("conftest.py", "__init__.py", "utils.py")
if os.path.basename(f) not in ("conftest.py", "__init__.py")
)
if not files:
return 0
+66
View File
@@ -0,0 +1,66 @@
#!/usr/bin/env python3
"""Check that the rust-ext cache key stays in sync across its sites.
The build workflow looks up and saves cache entries under its own prefix and
hashed inputs; the download action restores with its own. Neither file can
reference the other, and a mismatch in EITHER half makes every pool silently
fall back to source builds at install time.
"""
import re
import sys
import yaml
BUILD_WORKFLOW = ".github/workflows/_pr-test-rust-ext-build.yml"
DOWNLOAD_ACTION = ".github/actions/download-rust-ext/action.yml"
_HASH_FILES = re.compile(r"hashFiles\(([^)]*)\)")
_QUOTED = re.compile(r"'([^']*)'")
def hashed_inputs(path: str) -> list[tuple[str, ...]]:
"""The argument tuple of every ``hashFiles(...)`` cache key in a file."""
with open(path, encoding="utf-8") as f:
text = f.read()
return [tuple(_QUOTED.findall(args)) for args in _HASH_FILES.findall(text)]
def main() -> int:
with open(BUILD_WORKFLOW, encoding="utf-8") as f:
workflow = yaml.safe_load(f)
with open(DOWNLOAD_ACTION, encoding="utf-8") as f:
action = yaml.safe_load(f)
# yaml 1.1 parses the `on:` key as boolean True
triggers = workflow.get("on", workflow.get(True))
save_prefix = triggers["workflow_call"]["inputs"]["cache_key_prefix"]["default"]
restore_prefix = action["inputs"]["cache_key_prefix"]["default"]
if save_prefix != restore_prefix:
print("ERROR: rust-ext cache_key_prefix defaults do not match.")
print(f" {BUILD_WORKFLOW} saves under: {save_prefix}")
print(f" {DOWNLOAD_ACTION} restores with: {restore_prefix}")
print("Bump both together, or every pool falls back to source builds.")
return 1
# Adding a file to one key alone permanently misses the other's entries.
sites = [(BUILD_WORKFLOW, inputs) for inputs in hashed_inputs(BUILD_WORKFLOW)]
sites += [(DOWNLOAD_ACTION, inputs) for inputs in hashed_inputs(DOWNLOAD_ACTION)]
if not sites:
print("ERROR: no hashFiles(...) cache key found; this check is dead.")
return 1
if len({inputs for _, inputs in sites}) > 1:
print("ERROR: rust-ext cache key inputs do not match.")
for path, inputs in sites:
print(f" {path}: {list(inputs)}")
print("Every lookup/save/restore site must hash the same inputs.")
return 1
return 0
if __name__ == "__main__":
sys.exit(main())
@@ -0,0 +1,57 @@
"""Guard: no server_args mutation outside the resolution pipeline, pinned at 0.
``ServerArgs.__setattr__`` already raises on a bare assignment after
resolution; this static scan is what reaches the sites tests never execute.
"""
import re
from pathlib import Path
_SGLANG_ROOT = Path(__file__).resolve().parents[2] / "python" / "sglang"
# Assignments to a server_args attribute (``server_args.x = ...``,
# ``self.server_args.x = ...``, and the ``sa`` alias used by a few helpers).
# ``==`` comparisons are excluded by the negative lookahead.
_MUTATION_PATTERNS = [
# (?![=}]) skips ``==`` comparisons and f-string ``{x=}`` debug specs.
re.compile(r"\bserver_args\.[a-z0-9_]+\s*=(?![=}])"),
re.compile(r"\bsa\.[a-z0-9_]+\s*=(?![=}])"),
re.compile(r"get_(?:global_)?server_args\(\)\.[a-z0-9_]+\s*=(?![=}])"),
# setattr is the same write with the attribute name behind a variable.
re.compile(
r"setattr\(\s*(?:[\w.]+\.)?(?:server_args|sa|get_(?:global_)?server_args\(\))\s*,"
),
]
# The resolution pipeline itself (mutation is its job) and multimodal_gen,
# whose ServerArgs is a different class outside this contract.
_EXCLUDED = (
"srt/server_args.py",
"srt/arg_groups",
"multimodal_gen",
)
_BASELINE = 0
def check_server_args_mutation_ratchet():
count = 0
for path in sorted(_SGLANG_ROOT.rglob("*.py")):
rel = path.relative_to(_SGLANG_ROOT).as_posix()
if rel.startswith(_EXCLUDED):
continue
source = path.read_text()
count += sum(len(pattern.findall(source)) for pattern in _MUTATION_PATTERNS)
if count > _BASELINE:
raise AssertionError(
f"server_args mutations outside the resolution pipeline grew: "
f"{count} > baseline {_BASELINE}. Configuration is resolved in "
"ServerArgs.__post_init__; declare through the pipeline "
"(passes / declare_late_resolution), change resolved config "
"with get_context().override(source, ...), or hand the value "
"to its runner as a constructor argument — do not assign fields."
)
if __name__ == "__main__":
check_server_args_mutation_ratchet()
+48
View File
@@ -0,0 +1,48 @@
#!/usr/bin/env python3
"""Runs every static ratchet in one process, so the package is parsed once."""
import sys
from check_decode_bookkeeping_ownership import (
check_bookkeeping_sites_match_owner_allowlist,
check_spec_v2_draft_workers_do_no_scheduler_bookkeeping,
)
from check_global_config_read_ratchet import (
check_configured_size_call_sites,
check_global_config_read_ratchet,
check_no_renamed_accessor_imports,
)
from check_legacy_global_ratchet import check_legacy_global_ratchet
from check_module_state_ratchet import check_module_state_ratchet
from check_parallel_adoption_ratchet import check_parallel_adoption_ratchet
from check_server_args_mutation_ratchet import check_server_args_mutation_ratchet
def main():
checks = (
check_bookkeeping_sites_match_owner_allowlist,
check_spec_v2_draft_workers_do_no_scheduler_bookkeeping,
check_global_config_read_ratchet,
check_configured_size_call_sites,
check_no_renamed_accessor_imports,
check_legacy_global_ratchet,
check_module_state_ratchet,
check_parallel_adoption_ratchet,
check_server_args_mutation_ratchet,
)
# They guard independent invariants, so one failure must not hide the rest.
failures = []
for check in checks:
try:
check()
except AssertionError as exc:
failures.append(f"[{check.__name__}] {exc}")
for failure in failures:
print(failure, file=sys.stderr)
print(file=sys.stderr)
return 1 if failures else 0
if __name__ == "__main__":
sys.exit(main())
@@ -0,0 +1,86 @@
import pathlib
import tempfile
import unittest
from check_no_bare_pytest_main import find_bare_pytest_main
class TestFindBarePytestMain(unittest.TestCase):
def check_source(self, source: str) -> int | None:
with tempfile.TemporaryDirectory() as directory:
path = pathlib.Path(directory) / "example.py"
path.write_text(source, encoding="utf-8")
return find_bare_pytest_main(path)
def test_rejects_discarded_result(self):
source = """
if __name__ == "__main__":
pytest.main([__file__])
"""
self.assertEqual(self.check_source(source), 3)
def test_rejects_discarded_result_with_whitespace(self):
source = """
if __name__ == "__main__":
pytest . main([__file__])
"""
self.assertEqual(self.check_source(source), 3)
def test_accepts_propagated_result(self):
source = """
if __name__ == "__main__":
sys.exit(pytest.main([__file__]))
"""
self.assertIsNone(self.check_source(source))
def test_rejects_assigned_result(self):
source = """
if "__main__" == __name__:
exit_code = pytest.main([__file__])
"""
self.assertEqual(self.check_source(source), 3)
def test_accepts_assigned_result_that_is_later_exited(self):
source = """
if __name__ == "__main__":
exit_code = pytest.main([__file__])
sys.exit(exit_code)
"""
self.assertIsNone(self.check_source(source))
def test_accepts_assigned_result_that_is_raised(self):
source = """
if __name__ == "__main__":
exit_code = pytest.main([__file__])
raise SystemExit(exit_code)
"""
self.assertIsNone(self.check_source(source))
def test_rejects_nested_discarded_result(self):
source = """
if __name__ == "__main__":
if enabled:
pytest.main([__file__])
"""
self.assertEqual(self.check_source(source), 4)
def test_accepts_raised_system_exit(self):
source = """
if __name__ == "__main__":
raise SystemExit(pytest.main([__file__]))
"""
self.assertIsNone(self.check_source(source))
def test_rejects_unraised_system_exit(self):
source = """
if __name__ == "__main__":
error = SystemExit(pytest.main([__file__]))
"""
self.assertEqual(self.check_source(source), 3)
def test_ignores_call_outside_main_guard(self):
self.assertIsNone(self.check_source("pytest.main([__file__])\n"))
if __name__ == "__main__":
unittest.main()
@@ -1,10 +1,9 @@
import sys
from pathlib import Path
# Add this directory to sys.path so that test files can do
# `sys.path.insert(0, str(Path(__file__).resolve().parents[1]))` equivalently,
# and so pytest can import subpackages (dense/, mla/, etc.) without
# confusing this directory with the Python stdlib `unittest` module.
# Put this directory on sys.path so pytest can import the subpackages
# (dense/, mla/, ...) without confusing this directory with the stdlib
# `unittest` module.
_here = str(Path(__file__).resolve().parent)
if _here not in sys.path:
sys.path.insert(0, _here)
@@ -23,23 +23,18 @@ specific regression #26735 introduced and then fixed
``breakable_cuda_graph_runner.py`` capture sites).
"""
import sys
import unittest
from pathlib import Path
import torch
from sglang.srt.model_executor.forward_batch_info import ForwardMode
from sglang.srt.utils import get_device_sm
from sglang.test.test_utils import CustomTestCase
sys.path.insert(0, str(Path(__file__).resolve().parents[1]))
from sglang.test.ci.ci_register import register_cuda_ci
from sglang.test.kits.attention_unittest.attention_methods.dense_attention import (
DenseAttentionCase,
build_dense_attention_fixture,
)
from sglang.test.test_utils import CustomTestCase
register_cuda_ci(est_time=10, stage="base-a", runner_config="1-gpu-small")
@@ -1,15 +1,9 @@
import sys
import unittest
from pathlib import Path
import torch
from sglang.srt.model_executor.forward_batch_info import ForwardMode
from sglang.srt.utils import get_device_sm
from sglang.test.test_utils import CustomTestCase
sys.path.insert(0, str(Path(__file__).resolve().parents[1]))
from sglang.test.ci.ci_register import register_cuda_ci
from sglang.test.kits.attention_unittest.attention_methods.dense_attention import (
DenseAttentionCase,
@@ -34,6 +28,7 @@ from sglang.test.kits.attention_unittest.runner_modes.speculative_target_verify_
from sglang.test.kits.attention_unittest.runner_modes.split_op_runner import (
run_dense_split_op_extend_case,
)
from sglang.test.test_utils import CustomTestCase
register_cuda_ci(est_time=25, stage="base-b", runner_config="1-gpu-large")
@@ -1,14 +1,8 @@
import sys
import unittest
from pathlib import Path
import torch
from sglang.srt.model_executor.forward_batch_info import ForwardMode
from sglang.test.test_utils import CustomTestCase
sys.path.insert(0, str(Path(__file__).resolve().parents[1]))
from sglang.test.ci.ci_register import register_cuda_ci
from sglang.test.kits.attention_unittest.attention_methods.dense_attention import (
DenseAttentionCase,
@@ -33,6 +27,7 @@ from sglang.test.kits.attention_unittest.runner_modes.speculative_target_verify_
from sglang.test.kits.attention_unittest.runner_modes.split_op_runner import (
run_dense_split_op_extend_case,
)
from sglang.test.test_utils import CustomTestCase
register_cuda_ci(est_time=45, stage="base-b", runner_config="4-gpu-b200")
register_cuda_ci(est_time=45, stage="base-b", runner_config="1-gpu-large")
@@ -1,15 +1,9 @@
import sys
import unittest
from pathlib import Path
import torch
from sglang.srt.model_executor.forward_batch_info import ForwardMode
from sglang.srt.utils import is_flashinfer_available
from sglang.test.test_utils import CustomTestCase
sys.path.insert(0, str(Path(__file__).resolve().parents[1]))
from sglang.test.ci.ci_register import register_cuda_ci
from sglang.test.kits.attention_unittest.attention_methods.dense_attention import (
DenseAttentionCase,
@@ -30,6 +24,7 @@ from sglang.test.kits.attention_unittest.runner_modes.speculative_target_verify_
from sglang.test.kits.attention_unittest.runner_modes.split_op_runner import (
run_dense_split_op_extend_case,
)
from sglang.test.test_utils import CustomTestCase
register_cuda_ci(est_time=25, stage="base-b", runner_config="4-gpu-b200")
register_cuda_ci(est_time=25, stage="base-b", runner_config="1-gpu-large")
@@ -1,14 +1,8 @@
import sys
import unittest
from pathlib import Path
import torch
from sglang.srt.model_executor.forward_batch_info import ForwardMode
from sglang.test.test_utils import CustomTestCase
sys.path.insert(0, str(Path(__file__).resolve().parents[1]))
from sglang.test.ci.ci_register import register_cuda_ci
from sglang.test.kits.attention_unittest.attention_methods.dense_attention import (
DenseAttentionCase,
@@ -18,6 +12,7 @@ from sglang.test.kits.attention_unittest.attention_methods.dense_attention impor
from sglang.test.kits.attention_unittest.runner_modes.split_op_runner import (
run_dense_split_op_extend_case,
)
from sglang.test.test_utils import CustomTestCase
register_cuda_ci(est_time=25, stage="base-b", runner_config="4-gpu-b200")
register_cuda_ci(est_time=25, stage="base-b", runner_config="1-gpu-large")
@@ -1,6 +1,4 @@
import sys
import unittest
from pathlib import Path
import torch
@@ -8,10 +6,6 @@ from sglang.srt.layers.attention.attention_registry import ATTENTION_BACKENDS
from sglang.srt.layers.attention.hybrid_attn_backend import HybridAttnBackend
from sglang.srt.model_executor.forward_batch_info import ForwardMode
from sglang.srt.utils import is_flashinfer_available
from sglang.test.test_utils import CustomTestCase
sys.path.insert(0, str(Path(__file__).resolve().parents[1]))
from sglang.test.ci.ci_register import register_cuda_ci
from sglang.test.kits.attention_unittest.attention_methods.dense_attention import (
DENSE_ATOL,
@@ -22,6 +16,7 @@ from sglang.test.kits.attention_unittest.attention_methods.dense_attention impor
replace_backend,
run_dense_fixture_eager,
)
from sglang.test.test_utils import CustomTestCase
register_cuda_ci(est_time=20, stage="base-b", runner_config="4-gpu-b200")
register_cuda_ci(est_time=20, stage="base-b", runner_config="1-gpu-large")
@@ -1,6 +1,4 @@
import sys
import unittest
from pathlib import Path
import torch
@@ -8,10 +6,6 @@ from sglang.srt.layers.attention.attention_registry import ATTENTION_BACKENDS
from sglang.srt.layers.attention.tbo_backend import TboAttnBackend
from sglang.srt.model_executor.forward_batch_info import ForwardMode
from sglang.srt.utils import get_device_sm
from sglang.test.test_utils import CustomTestCase
sys.path.insert(0, str(Path(__file__).resolve().parents[1]))
from sglang.test.ci.ci_register import register_cuda_ci
from sglang.test.kits.attention_unittest.attention_methods.dense_attention import (
DENSE_ATOL,
@@ -25,6 +19,7 @@ from sglang.test.kits.attention_unittest.attention_methods.dense_attention impor
from sglang.test.kits.attention_unittest.runner_modes.speculative_target_verify_runner import (
_prepare_spec_verify_batch,
)
from sglang.test.test_utils import CustomTestCase
register_cuda_ci(est_time=20, stage="base-b", runner_config="4-gpu-b200")
register_cuda_ci(est_time=20, stage="base-b", runner_config="1-gpu-large")
@@ -1,20 +1,15 @@
import sys
import unittest
from pathlib import Path
import torch
from sglang.srt.model_executor.forward_batch_info import ForwardMode
from sglang.test.test_utils import CustomTestCase
sys.path.insert(0, str(Path(__file__).resolve().parents[1]))
from sglang.test.ci.ci_register import register_amd_ci, register_cuda_ci
from sglang.test.kits.attention_unittest.attention_methods.dense_attention import (
DenseAttentionCase,
make_dense_cases,
run_dense_attention_case,
)
from sglang.test.test_utils import CustomTestCase
register_cuda_ci(est_time=20, stage="base-b", runner_config="4-gpu-b200")
register_cuda_ci(est_time=20, stage="base-b", runner_config="1-gpu-large")
@@ -1,15 +1,9 @@
import sys
import unittest
from pathlib import Path
import torch
from sglang.srt.model_executor.forward_batch_info import ForwardMode
from sglang.srt.utils import is_hip
from sglang.test.test_utils import CustomTestCase
sys.path.insert(0, str(Path(__file__).resolve().parents[1]))
from sglang.test.ci.ci_register import register_amd_ci, register_cuda_ci
from sglang.test.kits.attention_unittest.attention_methods.dense_attention import (
DenseAttentionCase,
@@ -34,6 +28,7 @@ from sglang.test.kits.attention_unittest.runner_modes.speculative_target_verify_
from sglang.test.kits.attention_unittest.runner_modes.split_op_runner import (
run_dense_split_op_extend_case,
)
from sglang.test.test_utils import CustomTestCase
register_cuda_ci(est_time=25, stage="base-b", runner_config="4-gpu-b200")
register_cuda_ci(est_time=25, stage="base-b", runner_config="1-gpu-large")
@@ -1,6 +1,4 @@
import sys
import unittest
from pathlib import Path
import torch
@@ -11,10 +9,6 @@ from sglang.srt.utils.common import (
is_sm100_supported,
is_sm120_supported,
)
from sglang.test.test_utils import CustomTestCase
sys.path.insert(0, str(Path(__file__).resolve().parents[1]))
from sglang.test.ci.ci_register import register_cuda_ci
from sglang.test.kits.attention_unittest.attention_methods.dense_attention import (
DenseAttentionCase,
@@ -31,6 +25,7 @@ from sglang.test.kits.attention_unittest.runner_modes.speculative_draft_runner i
from sglang.test.kits.attention_unittest.runner_modes.split_op_runner import (
run_dense_split_op_extend_case,
)
from sglang.test.test_utils import CustomTestCase
register_cuda_ci(est_time=20, stage="base-b", runner_config="4-gpu-b200")
register_cuda_ci(est_time=20, stage="base-b", runner_config="1-gpu-large")
@@ -1,14 +1,8 @@
import sys
import unittest
from pathlib import Path
import torch
from sglang.srt.model_executor.forward_batch_info import ForwardMode
from sglang.test.test_utils import CustomTestCase
sys.path.insert(0, str(Path(__file__).resolve().parents[1]))
from sglang.test.ci.ci_register import register_cuda_ci
from sglang.test.kits.attention_unittest.attention_methods.dsa_attention import (
DSA_DECODE_IMPL_VARIANTS,
@@ -34,6 +28,7 @@ from sglang.test.kits.attention_unittest.runner_modes.cuda_graph_decode_runner i
from sglang.test.kits.attention_unittest.runner_modes.speculative_draft_runner import (
run_dsa_eagle_draft_cuda_graph_runner_case,
)
from sglang.test.test_utils import CustomTestCase
register_cuda_ci(est_time=25, stage="base-b", runner_config="4-gpu-b200")
register_cuda_ci(est_time=25, stage="base-b", runner_config="1-gpu-large")
@@ -13,9 +13,7 @@ gate+norm+rotate compression itself) is a deferred follow-up.
"""
import importlib.util
import sys
import unittest
from pathlib import Path
from types import SimpleNamespace
from unittest import mock
@@ -24,8 +22,6 @@ import torch
from sglang.srt.model_executor.forward_batch_info import ForwardMode
from sglang.test.test_utils import CustomTestCase
sys.path.insert(0, str(Path(__file__).resolve().parents[1]))
_FLASH_MLA_AVAILABLE = (
importlib.util.find_spec("sgl_kernel") is not None
and importlib.util.find_spec("sgl_kernel.flash_mla") is not None
@@ -1,14 +1,8 @@
import sys
import unittest
from pathlib import Path
import torch
from sglang.srt.model_executor.forward_batch_info import ForwardMode
from sglang.test.test_utils import CustomTestCase
sys.path.insert(0, str(Path(__file__).resolve().parents[1]))
from sglang.test.kits.attention_unittest.attention_methods.dual_chunk_attention import (
DualChunkAttentionCase,
make_dual_chunk_cases,
@@ -23,6 +17,7 @@ from sglang.test.kits.attention_unittest.attention_methods.dual_chunk_attention
from sglang.test.kits.attention_unittest.runner_modes.cuda_graph_decode_runner import (
run_dual_chunk_cuda_graph_decode_case,
)
from sglang.test.test_utils import CustomTestCase
# Container gate (KNOWN_FAILURES.md §1): `DualChunkFlashAttentionBackend` calls
@@ -1,16 +1,10 @@
import sys
import unittest
from pathlib import Path
import torch
from sglang.srt.layers.attention.linear.kernels.gdn_triton import TritonGDNKernel
from sglang.srt.model_executor.forward_batch_info import ForwardMode
from sglang.srt.utils import is_flashinfer_available
from sglang.test.test_utils import CustomTestCase
sys.path.insert(0, str(Path(__file__).resolve().parents[1]))
from sglang.srt.layers.attention.linear.kernels.gdn_triton import TritonGDNKernel
from sglang.test.ci.ci_register import register_cuda_ci
from sglang.test.kits.attention_unittest.attention_methods.gdn_attention import (
GDNAttentionCase,
@@ -29,6 +23,7 @@ from sglang.test.kits.attention_unittest.runner_modes.speculative_target_verify_
from sglang.test.kits.attention_unittest.runner_modes.split_op_runner import (
run_gdn_split_op_extend_case,
)
from sglang.test.test_utils import CustomTestCase
register_cuda_ci(est_time=20, stage="base-b", runner_config="4-gpu-b200")
register_cuda_ci(est_time=20, stage="base-b", runner_config="1-gpu-large")
@@ -33,21 +33,17 @@ L sweep:
Runnable as ``pytest`` and as ``__main__``.
"""
import sys
import unittest
from pathlib import Path
import torch
from sglang.test.test_utils import CustomTestCase
# Mirror sibling GDN unittests: register for CUDA/AMD CI. This is a kernel-math
# unit test; it lives with the other linear-attention kernel correctness tests.
# The registry calls MUST be module-level (the CI collector / check-registered-
# tests hook parses them statically via AST and only scans top-level statements
# -- a try/except wrapper hides them and fails the hook).
sys.path.insert(0, str(Path(__file__).resolve().parents[1]))
from sglang.test.ci.ci_register import register_amd_ci, register_cuda_ci
from sglang.test.test_utils import CustomTestCase
register_cuda_ci(est_time=20, stage="base-b", runner_config="1-gpu-large")
register_amd_ci(est_time=20, suite="stage-b-test-1-gpu-large-amd")
@@ -1,15 +1,9 @@
import sys
import unittest
from pathlib import Path
import torch
from sglang.srt.model_executor.forward_batch_info import ForwardMode
from sglang.srt.utils import is_hip
from sglang.test.test_utils import CustomTestCase
sys.path.insert(0, str(Path(__file__).resolve().parents[1]))
from sglang.test.ci.ci_register import register_amd_ci, register_cuda_ci
from sglang.test.kits.attention_unittest.attention_methods.gdn_attention import (
GDNAttentionCase,
@@ -19,6 +13,7 @@ from sglang.test.kits.attention_unittest.attention_methods.gdn_attention import
from sglang.test.kits.attention_unittest.runner_modes.split_op_runner import (
run_gdn_split_op_extend_case,
)
from sglang.test.test_utils import CustomTestCase
register_cuda_ci(est_time=20, stage="base-b", runner_config="4-gpu-b200")
register_cuda_ci(est_time=20, stage="base-b", runner_config="1-gpu-large")
@@ -1,6 +1,4 @@
import sys
import unittest
from pathlib import Path
from types import SimpleNamespace
from unittest.mock import MagicMock
@@ -12,10 +10,6 @@ from sglang.srt.layers.attention.hybrid_linear_attn_backend import (
)
from sglang.srt.model_executor.forward_batch_info import ForwardMode
from sglang.srt.utils import is_hip
from sglang.test.test_utils import CustomTestCase
sys.path.insert(0, str(Path(__file__).resolve().parents[1]))
from sglang.test.ci.ci_register import register_amd_ci, register_cuda_ci
from sglang.test.kits.attention_unittest.attention_methods.gdn_attention import (
GDNAttentionCase,
@@ -32,6 +26,7 @@ from sglang.test.kits.attention_unittest.runner_modes.speculative_target_verify_
from sglang.test.kits.attention_unittest.runner_modes.split_op_runner import (
run_gdn_split_op_extend_case,
)
from sglang.test.test_utils import CustomTestCase
register_cuda_ci(est_time=20, stage="base-b", runner_config="4-gpu-b200")
register_cuda_ci(est_time=20, stage="base-b", runner_config="1-gpu-large")
@@ -13,7 +13,6 @@ so the chunked-MHA path never runs.
"""
import unittest
from pathlib import Path
from types import SimpleNamespace
import torch
@@ -143,8 +142,4 @@ class TestHybridLinearChunkMetadataDelegation(CustomTestCase):
if __name__ == "__main__":
sys_path_parent = str(Path(__file__).resolve().parents[1])
import sys
sys.path.insert(0, sys_path_parent)
unittest.main()
@@ -1,15 +1,9 @@
import sys
import unittest
from pathlib import Path
import torch
from sglang.srt.model_executor.forward_batch_info import ForwardMode
from sglang.srt.utils import is_hip
from sglang.test.test_utils import CustomTestCase
sys.path.insert(0, str(Path(__file__).resolve().parents[1]))
from sglang.test.ci.ci_register import register_amd_ci, register_cuda_ci
from sglang.test.kits.attention_unittest.attention_methods.kda_attention import (
KDAAttentionCase,
@@ -26,6 +20,7 @@ from sglang.test.kits.attention_unittest.runner_modes.speculative_target_verify_
from sglang.test.kits.attention_unittest.runner_modes.split_op_runner import (
run_kda_split_op_extend_case,
)
from sglang.test.test_utils import CustomTestCase
register_cuda_ci(est_time=20, stage="base-b", runner_config="4-gpu-b200")
register_cuda_ci(est_time=20, stage="base-b", runner_config="1-gpu-large")
@@ -1,15 +1,9 @@
import sys
import unittest
from pathlib import Path
import torch
from sglang.kernels.ops.attention.linear.seg_la import SegLaMeta, seg_la_fwd
from sglang.srt.model_executor.forward_batch_info import ForwardMode
from sglang.test.test_utils import CustomTestCase
sys.path.insert(0, str(Path(__file__).resolve().parents[1]))
from sglang.test.ci.ci_register import register_amd_ci, register_cuda_ci
from sglang.test.kits.attention_unittest.attention_methods.lightning_attention import (
LightningAttentionCase,
@@ -23,6 +17,7 @@ from sglang.test.kits.attention_unittest.runner_modes.speculative_target_verify_
run_lightning_eagle_verify_case,
run_lightning_eagle_verify_cuda_graph_case,
)
from sglang.test.test_utils import CustomTestCase
register_cuda_ci(est_time=20, stage="base-b", runner_config="4-gpu-b200")
register_cuda_ci(est_time=20, stage="base-b", runner_config="1-gpu-large")
@@ -1,6 +1,4 @@
import sys
import unittest
from pathlib import Path
from types import SimpleNamespace
from unittest.mock import MagicMock
@@ -11,10 +9,6 @@ from sglang.srt.layers.attention.hybrid_linear_attn_backend import (
MambaAttnBackendBase,
)
from sglang.srt.model_executor.forward_batch_info import ForwardMode
from sglang.test.test_utils import CustomTestCase
sys.path.insert(0, str(Path(__file__).resolve().parents[1]))
from sglang.test.ci.ci_register import register_cuda_ci
from sglang.test.kits.attention_unittest.attention_methods.mamba2_attention import (
DEFAULT_CONV_KERNEL,
@@ -36,6 +30,7 @@ from sglang.test.kits.attention_unittest.runner_modes.speculative_target_verify_
run_mamba2_eagle_verify_case,
run_mamba2_eagle_verify_cuda_graph_case,
)
from sglang.test.test_utils import CustomTestCase
register_cuda_ci(est_time=20, stage="base-b", runner_config="4-gpu-b200")
register_cuda_ci(est_time=20, stage="base-b", runner_config="1-gpu-large")
@@ -1,18 +1,13 @@
import sys
import unittest
from pathlib import Path
import torch
from sglang.srt.model_executor.forward_batch_info import ForwardMode
from sglang.test.test_utils import CustomTestCase
sys.path.insert(0, str(Path(__file__).resolve().parents[1]))
from sglang.test.kits.attention_unittest.attention_methods.mla_attention import (
MLAAttentionCase,
run_mla_attention_case,
)
from sglang.test.test_utils import CustomTestCase
# Cutlass MLA requires exactly Blackwell SM 10.0. The sgl-kernel
# `cutlass_mla_decode` checks `sm_version == 100` (major*10+minor), so
@@ -1,14 +1,8 @@
import sys
import unittest
from pathlib import Path
import torch
from sglang.srt.model_executor.forward_batch_info import ForwardMode
from sglang.test.test_utils import CustomTestCase
sys.path.insert(0, str(Path(__file__).resolve().parents[1]))
from sglang.test.kits.attention_unittest.attention_methods.mla_attention import (
MLAAttentionCase,
make_mla_cases,
@@ -27,6 +21,7 @@ from sglang.test.kits.attention_unittest.runner_modes.speculative_target_verify_
from sglang.test.kits.attention_unittest.runner_modes.split_op_runner import (
run_mla_split_op_extend_case,
)
from sglang.test.test_utils import CustomTestCase
MLA_SHAPE_KWARGS = dict(
kv_lora_rank=512,
@@ -1,16 +1,10 @@
import sys
import unittest
from pathlib import Path
import torch
import triton
from sglang.srt.model_executor.forward_batch_info import ForwardMode
from sglang.srt.model_executor.forward_context import ForwardContext, forward_context
from sglang.test.test_utils import CustomTestCase
sys.path.insert(0, str(Path(__file__).resolve().parents[1]))
from sglang.test.kits.attention_unittest.attention_methods.mla_attention import (
MLAAttentionCase,
build_mla_attention_fixture,
@@ -33,6 +27,7 @@ from sglang.test.kits.attention_unittest.runner_modes.speculative_target_verify_
from sglang.test.kits.attention_unittest.runner_modes.split_op_runner import (
run_mla_split_op_extend_case,
)
from sglang.test.test_utils import CustomTestCase
MLA_SHAPE_KWARGS = dict(
kv_lora_rank=512,
@@ -1,19 +1,14 @@
import importlib.util
import sys
import unittest
from pathlib import Path
import torch
from sglang.srt.model_executor.forward_batch_info import ForwardMode
from sglang.test.test_utils import CustomTestCase
sys.path.insert(0, str(Path(__file__).resolve().parents[1]))
from sglang.test.kits.attention_unittest.attention_methods.mla_attention import (
MLAAttentionCase,
run_mla_attention_case,
)
from sglang.test.test_utils import CustomTestCase
# tokenspeed_mla is a CuTe DSL backend for Blackwell (SM100). It additionally
# enforces:
@@ -1,14 +1,8 @@
import sys
import unittest
from pathlib import Path
import torch
from sglang.srt.model_executor.forward_batch_info import ForwardMode
from sglang.test.test_utils import CustomTestCase
sys.path.insert(0, str(Path(__file__).resolve().parents[1]))
from sglang.test.ci.ci_register import register_cuda_ci
from sglang.test.kits.attention_unittest.attention_methods.mla_attention import (
MLAAttentionCase,
@@ -32,6 +26,7 @@ from sglang.test.kits.attention_unittest.runner_modes.speculative_target_verify_
from sglang.test.kits.attention_unittest.runner_modes.split_op_runner import (
run_mla_split_op_extend_case,
)
from sglang.test.test_utils import CustomTestCase
register_cuda_ci(est_time=25, stage="base-b", runner_config="4-gpu-b200")
register_cuda_ci(est_time=25, stage="base-b", runner_config="1-gpu-large")
@@ -1,18 +1,13 @@
import sys
import unittest
from pathlib import Path
import torch
from sglang.srt.model_executor.forward_batch_info import ForwardMode
from sglang.test.test_utils import CustomTestCase
sys.path.insert(0, str(Path(__file__).resolve().parents[1]))
from sglang.test.kits.attention_unittest.attention_methods.mla_attention import (
MLAAttentionCase,
run_mla_attention_case,
)
from sglang.test.test_utils import CustomTestCase
# trtllm_mla goes through FlashInfer's XQA MLA path. Per PLAN.md and the
# project's is_sm120_supported helper (device_capability_majors=[12]), the
@@ -1,16 +1,10 @@
import sys
import unittest
from pathlib import Path
import torch
from sglang.srt.environ import envs
from sglang.srt.model_executor.forward_batch_info import ForwardMode
from sglang.srt.utils import is_flashinfer_available
from sglang.test.test_utils import CustomTestCase
sys.path.insert(0, str(Path(__file__).resolve().parents[1]))
from sglang.test.ci.ci_register import register_cuda_ci
from sglang.test.kits.attention_unittest.attention_methods.dense_attention import (
DenseAttentionCase,
@@ -28,6 +22,7 @@ from sglang.test.kits.attention_unittest.runner_modes.speculative_target_verify_
from sglang.test.kits.attention_unittest.runner_modes.split_op_runner import (
run_dense_split_op_extend_case,
)
from sglang.test.test_utils import CustomTestCase
register_cuda_ci(est_time=20, stage="base-b", runner_config="4-gpu-b200")
register_cuda_ci(est_time=20, stage="base-b", runner_config="1-gpu-large")
@@ -6,20 +6,15 @@ uses swa_loc directly for SWA layers and asserts it is provided. The per-backend
cuda-graph buffer plumbing is covered by the backend SWA integration tests.
"""
import sys
import unittest
from pathlib import Path
from types import SimpleNamespace
import torch
from sglang.srt.mem_cache.memory_pool import KVWriteLoc
from sglang.srt.mem_cache.swa_memory_pool import SWAKVPool
from sglang.test.test_utils import CustomTestCase
sys.path.insert(0, str(Path(__file__).resolve().parents[1]))
from sglang.test.ci.ci_register import register_cpu_ci
from sglang.test.test_utils import CustomTestCase
register_cpu_ci(est_time=5, suite="base-a-test-cpu")
@@ -1,14 +1,8 @@
import sys
import unittest
from pathlib import Path
import torch
from sglang.srt.model_executor.forward_batch_info import ForwardMode
from sglang.test.test_utils import CustomTestCase
sys.path.insert(0, str(Path(__file__).resolve().parents[1]))
from sglang.test.ci.ci_register import register_amd_ci, register_cuda_ci
from sglang.test.kits.attention_unittest.attention_methods.dense_attention import (
DenseAttentionCase,
@@ -16,6 +10,7 @@ from sglang.test.kits.attention_unittest.attention_methods.dense_attention impor
make_swa_prefix_input_config_cases,
run_dense_attention_case,
)
from sglang.test.test_utils import CustomTestCase
register_cuda_ci(est_time=20, stage="base-b", runner_config="4-gpu-b200")
register_cuda_ci(est_time=20, stage="base-b", runner_config="1-gpu-large")
@@ -1,15 +1,9 @@
import sys
import unittest
from pathlib import Path
import torch
from sglang.srt.model_executor.forward_batch_info import ForwardMode
from sglang.srt.utils import is_hip
from sglang.test.test_utils import CustomTestCase
sys.path.insert(0, str(Path(__file__).resolve().parents[1]))
from sglang.test.ci.ci_register import register_amd_ci, register_cuda_ci
from sglang.test.kits.attention_unittest.attention_methods.dense_attention import (
DenseAttentionCase,
@@ -27,6 +21,7 @@ from sglang.test.kits.attention_unittest.runner_modes.speculative_target_verify_
from sglang.test.kits.attention_unittest.runner_modes.split_op_runner import (
run_dense_split_op_extend_case,
)
from sglang.test.test_utils import CustomTestCase
register_cuda_ci(est_time=20, stage="base-b", runner_config="4-gpu-b200")
register_cuda_ci(est_time=20, stage="base-b", runner_config="1-gpu-large")
+1 -10
View File
@@ -11,17 +11,13 @@ register_cpu_ci(est_time=10, suite="base-b-test-cpu-arm64")
import itertools
import math
import os
import platform
import sys
import unittest
import torch
# Add parent dir (test/srt/cpu/) to path for utils import
sys.path.insert(0, os.path.join(os.path.dirname(__file__), ".."))
from sglang.srt.layers.amx_utils import CPUQuantMethod
from sglang.test.cpu_test_utils import precision, torch_w8a8_per_column_fused_moe
from sglang.test.test_utils import CustomTestCase
kernel = torch.ops.sgl_kernel
@@ -29,11 +25,6 @@ IS_ARM64 = platform.machine().lower() in ("aarch64", "arm64")
torch.manual_seed(128)
from utils import (
precision,
torch_w8a8_per_column_fused_moe,
)
class TestFusedExpertsInt8(CustomTestCase):
M = [1, 6, 32, 64]
+1 -1
View File
@@ -2,10 +2,10 @@ import sys
import pytest
import torch
from utils import GeluAndMul, SiluAndMul, precision
from sglang.srt.server_args import ServerArgs, set_global_server_args_for_scheduler
from sglang.test.ci.ci_register import register_cpu_ci
from sglang.test.cpu_test_utils import GeluAndMul, SiluAndMul, precision
register_cpu_ci(est_time=10, suite="base-b-test-cpu")
register_cpu_ci(est_time=10, suite="base-b-test-cpu-arm64")
+1 -1
View File
@@ -4,10 +4,10 @@ import unittest
# TODO: use interface in cpu.py
import torch
import torch.nn as nn
from utils import precision
from sglang.srt.layers.quantization.fp8_utils import input_to_float8
from sglang.test.ci.ci_register import register_cpu_ci
from sglang.test.cpu_test_utils import precision
from sglang.test.test_utils import CustomTestCase
register_cpu_ci(est_time=10, suite="base-b-test-cpu")
+1 -1
View File
@@ -4,9 +4,9 @@ from typing import Optional
import sgl_kernel # noqa: F401
import torch
import torch.nn.functional as F
from utils import parametrize, precision
from sglang.test.ci.ci_register import register_cpu_ci
from sglang.test.cpu_test_utils import parametrize, precision
from sglang.test.test_utils import CustomTestCase
register_cpu_ci(est_time=10, suite="base-b-test-cpu")
+1 -1
View File
@@ -3,9 +3,9 @@ import unittest
import sgl_kernel # noqa: F401
import torch
import torch.nn.functional as F
from utils import parametrize, precision
from sglang.test.ci.ci_register import register_cpu_ci
from sglang.test.cpu_test_utils import parametrize, precision
from sglang.test.test_utils import CustomTestCase
register_cpu_ci(est_time=10, suite="base-b-test-cpu")
+3 -3
View File
@@ -3,7 +3,9 @@ import unittest
# TODO: use interface in cpu.py
import torch
import torch.nn as nn
from utils import (
from sglang.test.ci.ci_register import register_cpu_ci
from sglang.test.cpu_test_utils import (
MXFP4QuantizeUtil,
convert_weight,
native_w8a8_per_token_matmul,
@@ -13,8 +15,6 @@ from utils import (
unpack_and_dequant_awq,
unpack_and_dequant_gptq,
)
from sglang.test.ci.ci_register import register_cpu_ci
from sglang.test.test_utils import CustomTestCase
register_cpu_ci(est_time=10, suite="base-b-test-cpu")
+1 -1
View File
@@ -4,9 +4,9 @@ import pytest
import torch
import torch.nn.functional as F
from torch.nn.functional import softplus
from utils import precision
from sglang.test.ci.ci_register import register_cpu_ci
from sglang.test.cpu_test_utils import precision
register_cpu_ci(est_time=10, suite="base-b-test-cpu")
+1 -1
View File
@@ -2,9 +2,9 @@ import unittest
import torch
from torch.nn.functional import scaled_dot_product_attention
from utils import precision
from sglang.test.ci.ci_register import register_cpu_ci
from sglang.test.cpu_test_utils import precision
from sglang.test.test_utils import CustomTestCase
register_cpu_ci(est_time=10, suite="base-b-test-cpu")
+2 -3
View File
@@ -16,7 +16,8 @@ prepack = True
alpha = 1.702
limit = 7.0
from utils import (
from sglang.test.ci.ci_register import register_cpu_ci
from sglang.test.cpu_test_utils import (
BLOCK_K,
BLOCK_N,
MXFP4QuantizeUtil,
@@ -32,8 +33,6 @@ from utils import (
unpack_and_dequant_awq,
)
from sglang.test.ci.ci_register import register_cpu_ci
register_cpu_ci(est_time=10, suite="base-b-test-cpu")
+1 -1
View File
@@ -3,9 +3,9 @@ from typing import Optional, Tuple, Union
import pytest
import torch
from utils import make_non_contiguous, precision
from sglang.test.ci.ci_register import register_cpu_ci
from sglang.test.cpu_test_utils import make_non_contiguous, precision
register_cpu_ci(est_time=10, suite="base-b-test-cpu")
register_cpu_ci(est_time=10, suite="base-b-test-cpu-arm64")
@@ -1,16 +1,16 @@
import unittest
import torch
from utils import (
from sglang.srt.layers.quantization.fp8_utils import input_to_float8
from sglang.srt.layers.rotary_embedding.utils import apply_rotary_emb
from sglang.test.ci.ci_register import register_cpu_ci
from sglang.test.cpu_test_utils import (
convert_weight,
native_w8a8_per_token_matmul,
per_token_quant_int8,
precision,
)
from sglang.srt.layers.quantization.fp8_utils import input_to_float8
from sglang.srt.layers.rotary_embedding.utils import apply_rotary_emb
from sglang.test.ci.ci_register import register_cpu_ci
from sglang.test.test_utils import CustomTestCase
register_cpu_ci(est_time=10, suite="base-b-test-cpu")
+1 -1
View File
@@ -2,10 +2,10 @@ import sys
import pytest
import torch
from utils import precision
from sglang.srt.utils import is_host_cpu_arm64
from sglang.test.ci.ci_register import register_cpu_ci
from sglang.test.cpu_test_utils import precision
register_cpu_ci(est_time=10, suite="base-b-test-cpu")
register_cpu_ci(est_time=10, suite="base-b-test-cpu-arm64")
+1 -1
View File
@@ -1,7 +1,6 @@
import unittest
import torch
from utils import precision
from sglang.srt.layers.rotary_embedding import (
MRotaryEmbedding,
@@ -14,6 +13,7 @@ from sglang.srt.layers.rotary_embedding.rope_variant import (
from sglang.srt.layers.rotary_embedding.utils import apply_rotary_pos_emb_native_eager
from sglang.srt.server_args import ServerArgs, set_global_server_args_for_scheduler
from sglang.test.ci.ci_register import register_cpu_ci
from sglang.test.cpu_test_utils import precision
from sglang.test.test_utils import CustomTestCase
register_cpu_ci(est_time=10, suite="base-b-test-cpu")
+3 -3
View File
@@ -3,7 +3,9 @@ import math
import unittest
import torch
from utils import (
from sglang.test.ci.ci_register import register_cpu_ci
from sglang.test.cpu_test_utils import (
BLOCK_K,
BLOCK_N,
factor_for_scale,
@@ -15,8 +17,6 @@ from utils import (
torch_naive_moe,
torch_w8a8_per_column_moe,
)
from sglang.test.ci.ci_register import register_cpu_ci
from sglang.test.test_utils import CustomTestCase
register_cpu_ci(est_time=10, suite="base-b-test-cpu")
+1 -1
View File
@@ -3,10 +3,10 @@ import unittest
import sgl_kernel # noqa: F401
import torch
import torch.nn.functional as F
from utils import precision
from sglang.srt.speculative.eagle_utils import TreeMaskMode, organize_draft_results
from sglang.test.ci.ci_register import register_cpu_ci
from sglang.test.cpu_test_utils import precision
from sglang.test.test_utils import CustomTestCase
register_cpu_ci(est_time=20, suite="base-b-test-cpu")
@@ -1,76 +0,0 @@
"""Run the `rust/` Cargo workspace's unit tests from the CPU CI suite.
The `rust/` workspace (sglang-grpc, sglang-mm, sglang-server) is compiled into
the wheel by setuptools-rust, but until now nothing ran `cargo test` in CI --
`.github/workflows/pr-test-rust.yml` and `pr-benchmark-rust.yml` are both
path-scoped to `sgl-model-gateway/**`, a different workspace. `lint.yml` covers
rustfmt/clippy via the pre-commit hooks, so this file only adds the test run.
The debug profile is deliberate: these are pure-logic tests (no timing or
codegen assertions), and the release profile costs a full LTO build for the
same coverage.
"""
import shutil
import subprocess
import unittest
from pathlib import Path
from sglang.test.ci.ci_register import register_cpu_ci
from sglang.test.test_utils import CustomTestCase
# base-c-test-cpu is where this was asked for, and it matches the repo's
# base-a + base-c dual-registration convention -- but base-c-test-cpu currently
# has no runner job in any workflow (it was carved out of base-b in #28623 to
# *reduce* CPU CI scope), so base-a-test-cpu is what actually executes.
register_cpu_ci(est_time=300, suite="base-a-test-cpu")
# repo root: test/registered/rust/<this file>
RUST_WORKSPACE = Path(__file__).resolve().parents[3] / "rust"
# Not `est_time`: that is a scheduling hint for partition balancing (a rough
# average), this is a hard ceiling for the worst case. The 136 tests run in ~1s;
# what varies is the build. Cache-warm the workspace crates recompile in ~15s,
# but a Swatinem/rust-cache miss rebuilds all ~370 dependencies -- measured at
# 48s on 4 fast cores, so several minutes on a hosted runner.
#
# Capped below the 600s `timeout-minutes` on the suite's "Run test" step so a
# hang fails here, with output, instead of being killed as an opaque job
# timeout. The harness `--timeout-per-file` (1200s) is looser still.
BUILD_AND_RUN_TIMEOUT_S = 300
class TestCargoWorkspace(CustomTestCase):
def test_cargo_test_workspace(self):
# Not skipUnless: cargo is a hard dependency of the editable install
# (setuptools-rust builds sglang-grpc), so a missing toolchain is a
# broken environment, and a silently-skipped CI test is worthless.
self.assertIsNotNone(
shutil.which("cargo"),
"cargo not found on PATH; install a Rust toolchain "
"(scripts/ci/utils/install_rust_protoc.sh)",
)
self.assertTrue(
(RUST_WORKSPACE / "Cargo.toml").is_file(),
f"rust workspace manifest not found at {RUST_WORKSPACE}",
)
proc = subprocess.run(
["cargo", "test", "--workspace"],
cwd=RUST_WORKSPACE,
capture_output=True,
text=True,
timeout=BUILD_AND_RUN_TIMEOUT_S,
)
# Print unconditionally so a green run still shows which tests ran.
print(proc.stdout)
self.assertEqual(
proc.returncode,
0,
f"`cargo test --workspace` failed in {RUST_WORKSPACE}\n"
f"--- stdout ---\n{proc.stdout}\n--- stderr ---\n{proc.stderr}",
)
if __name__ == "__main__":
unittest.main()
@@ -6,7 +6,6 @@ were skipped or double-emitted, drifting the per-request entry counts that
the scheduler asserts on.
"""
import itertools
import unittest
from types import SimpleNamespace
@@ -14,6 +13,7 @@ import torch
from sglang.srt.layers.logprob_processor import InputLogprobProcessor
from sglang.test.ci.ci_register import register_cpu_ci
from sglang.test.logprob_test_utils import coverage_cases
from sglang.test.test_utils import CustomTestCase
register_cpu_ci(est_time=30, suite="base-a-test-cpu")
@@ -23,6 +23,11 @@ VOCAB = 11
TOPK_CYCLE = [2, 0, 3]
# [] is a valid probe set distinct from None (opt-out).
TOKEN_IDS_CYCLE = [[0, 3], None, [1], []]
# start == extend_len is the zero-logprob-row shape. Order determines the cyclic
# width-3/4 heterogeneous coverage cases.
SEQ_SPEC_MENU = ((1, 1), (2, 2), (3, 0), (4, 1), (5, 5), (2, 0), (6, 2))
# 7 singletons + 7*7 ordered pairs + 4*7 wide cases each at width 3 and 4.
EXPECTED_CASES = 112
def _build_batch(seq_specs, with_token_ids):
@@ -92,40 +97,38 @@ class TestLogprobChunkStitching(CustomTestCase):
def _sweep(self, with_token_ids):
torch.manual_seed(0)
proc = InputLogprobProcessor()
# (extend_len, start); start == extend_len is the degenerate
# zero-logprob-row shape.
menu = [(1, 1), (2, 2), (3, 0), (4, 1), (5, 5), (2, 0), (6, 2)]
combos = list(coverage_cases(SEQ_SPEC_MENU, max_seqs=4))
self.assertEqual(len(combos), EXPECTED_CASES)
tried = 0
for n_seqs in (1, 2, 3, 4):
for combo in itertools.product(menu, repeat=n_seqs):
batch = _build_batch(list(combo), with_token_ids)
# Same unit as the production gate: grid rows, not logprob rows.
total_rows = batch[0].shape[0]
for chunk_size in (1, 2, 3, 5):
if total_rows <= chunk_size:
continue
tried += 1
ref, ref_sampled = _run(proc, batch, False, 10**9)
got, got_sampled = _run(proc, batch, True, chunk_size)
label = f"specs={list(combo)} chunk={chunk_size}"
self.assertEqual(ref.top_logprobs_val, got.top_logprobs_val, label)
self.assertEqual(ref.top_logprobs_idx, got.top_logprobs_idx, label)
if with_token_ids:
self.assertEqual(
ref.token_ids_logprobs_val,
got.token_ids_logprobs_val,
label,
)
self.assertEqual(
ref.token_ids_logprobs_idx,
got.token_ids_logprobs_idx,
label,
)
torch.testing.assert_close(
ref.token_logprobs, got.token_logprobs, msg=label
for combo in combos:
batch = _build_batch(list(combo), with_token_ids)
# Same unit as the production gate: grid rows, not logprob rows.
total_rows = batch[0].shape[0]
for chunk_size in (1, 2, 3, 5):
if total_rows <= chunk_size:
continue
tried += 1
ref, ref_sampled = _run(proc, batch, False, 10**9)
got, got_sampled = _run(proc, batch, True, chunk_size)
label = f"specs={list(combo)} chunk={chunk_size}"
self.assertEqual(ref.top_logprobs_val, got.top_logprobs_val, label)
self.assertEqual(ref.top_logprobs_idx, got.top_logprobs_idx, label)
if with_token_ids:
self.assertEqual(
ref.token_ids_logprobs_val,
got.token_ids_logprobs_val,
label,
)
torch.testing.assert_close(ref_sampled, got_sampled, msg=label)
self.assertGreater(tried, 1000)
self.assertEqual(
ref.token_ids_logprobs_idx,
got.token_ids_logprobs_idx,
label,
)
torch.testing.assert_close(
ref.token_logprobs, got.token_logprobs, msg=label
)
torch.testing.assert_close(ref_sampled, got_sampled, msg=label)
self.assertGreater(tried, 100)
def test_top_logprobs_stitching(self):
self._sweep(with_token_ids=False)
@@ -7,7 +7,6 @@ agree with the reference path to floating-point tolerance, with identical
top-k indices, across chunk splits and heterogeneous per-sequence params.
"""
import itertools
import unittest
from types import SimpleNamespace
@@ -18,6 +17,7 @@ from sglang.srt.layers.logprob_processor import (
compute_row_log_normalizer,
)
from sglang.test.ci.ci_register import register_cpu_ci
from sglang.test.logprob_test_utils import coverage_cases
from sglang.test.test_utils import CustomTestCase
register_cpu_ci(est_time=30, suite="base-a-test-cpu")
@@ -27,6 +27,11 @@ VOCAB = 11
TOPK_CYCLE = [2, 0, 3]
# [] is a valid probe set distinct from None (opt-out).
TOKEN_IDS_CYCLE = [[0, 3], None, [1], []]
# start == extend_len is the zero-logprob-row shape. Order determines the cyclic
# width-3 heterogeneous coverage cases.
SEQ_SPEC_MENU = ((1, 1), (3, 0), (4, 1), (5, 5), (6, 2))
# 5 singletons + 5*5 ordered pairs + 4*5 wide cases at width 3.
EXPECTED_CASES = 50
def _build_batch(seq_specs, dtype, vocab=VOCAB):
@@ -121,48 +126,46 @@ class TestFastInputLogprobs(CustomTestCase):
def _sweep(self, dtype, rtol, atol):
torch.manual_seed(0)
proc = InputLogprobProcessor()
# (extend_len, start); start == extend_len is the degenerate
# zero-logprob-row shape.
menu = [(1, 1), (3, 0), (4, 1), (5, 5), (6, 2)]
combos = list(coverage_cases(SEQ_SPEC_MENU, max_seqs=3))
self.assertEqual(len(combos), EXPECTED_CASES)
tried = 0
for n_seqs in (1, 2, 3):
for combo in itertools.product(menu, repeat=n_seqs):
batch = _build_batch(list(combo), dtype)
for chunk_size in (None, 1, 2, 3, 5):
tried += 1
ref, ref_sampled = _run(proc, batch, False, chunk_size)
got, got_sampled = _run(proc, batch, True, chunk_size)
label = f"specs={list(combo)} chunk={chunk_size} dtype={dtype}"
# Top-k order comes from the same values shifted by a
# per-row constant, so indices must match exactly.
self.assertEqual(ref.top_logprobs_idx, got.top_logprobs_idx, label)
self.assertEqual(
ref.token_ids_logprobs_idx, got.token_ids_logprobs_idx, label
)
_assert_nested_close(
self,
ref.top_logprobs_val,
got.top_logprobs_val,
label,
rtol,
atol,
)
_assert_nested_close(
self,
ref.token_ids_logprobs_val,
got.token_ids_logprobs_val,
label,
rtol,
atol,
)
torch.testing.assert_close(
ref.token_logprobs.float(),
got.token_logprobs.float(),
rtol=rtol,
atol=atol,
msg=label,
)
torch.testing.assert_close(ref_sampled, got_sampled, msg=label)
for combo in combos:
batch = _build_batch(list(combo), dtype)
for chunk_size in (None, 1, 2, 3, 5):
tried += 1
ref, ref_sampled = _run(proc, batch, False, chunk_size)
got, got_sampled = _run(proc, batch, True, chunk_size)
label = f"specs={list(combo)} chunk={chunk_size} dtype={dtype}"
# Top-k order comes from the same values shifted by a
# per-row constant, so indices must match exactly.
self.assertEqual(ref.top_logprobs_idx, got.top_logprobs_idx, label)
self.assertEqual(
ref.token_ids_logprobs_idx, got.token_ids_logprobs_idx, label
)
_assert_nested_close(
self,
ref.top_logprobs_val,
got.top_logprobs_val,
label,
rtol,
atol,
)
_assert_nested_close(
self,
ref.token_ids_logprobs_val,
got.token_ids_logprobs_val,
label,
rtol,
atol,
)
torch.testing.assert_close(
ref.token_logprobs.float(),
got.token_logprobs.float(),
rtol=rtol,
atol=atol,
msg=label,
)
torch.testing.assert_close(ref_sampled, got_sampled, msg=label)
self.assertGreater(tried, 100)
def test_fast_matches_reference_fp32(self):
@@ -176,20 +179,16 @@ class TestFastInputLogprobs(CustomTestCase):
# sits much closer to the truth than bf16 resolution.
torch.manual_seed(0)
proc = InputLogprobProcessor()
menu = [(1, 1), (3, 0), (4, 1), (5, 5), (6, 2)]
for n_seqs in (1, 2, 3):
for combo in itertools.product(menu, repeat=n_seqs):
batch = _build_batch(list(combo), torch.bfloat16)
pruned_states, _, input_logprob_indices, _, metadata = batch
truth = torch.log_softmax(pruned_states.double(), dim=-1)[
input_logprob_indices
]
for chunk_size in (None, 2, 5):
got, _ = _run(proc, batch, True, chunk_size)
label = f"specs={list(combo)} chunk={chunk_size}"
self._assert_rows_match_truth(
got, truth, metadata, label, atol=1e-4
)
for combo in coverage_cases(SEQ_SPEC_MENU, max_seqs=3):
batch = _build_batch(list(combo), torch.bfloat16)
pruned_states, _, input_logprob_indices, _, metadata = batch
truth = torch.log_softmax(pruned_states.double(), dim=-1)[
input_logprob_indices
]
for chunk_size in (None, 2, 5):
got, _ = _run(proc, batch, True, chunk_size)
label = f"specs={list(combo)} chunk={chunk_size}"
self._assert_rows_match_truth(got, truth, metadata, label, atol=1e-4)
def _assert_rows_match_truth(self, got, truth, metadata, label, atol):
pt = 0
@@ -1,65 +0,0 @@
"""Ratchet guard: legacy global-accessor call-sites may only decrease.
The process-wide ``ServerArgs`` is owned by the runtime context; the legacy
``get_global_server_args`` / ``set_global_server_args_for_*`` names survive as
thin shims for the existing call-sites. New code should use the
``sglang.srt.runtime_context`` accessors (``get_server_args()`` /
``get_context().set_server_args()``), so the shim call-site counts below must
never grow. When your change removes call-sites, lower the matching baseline
to the new count.
"""
from sglang.test.ci.ci_register import register_cpu_ci
register_cpu_ci(est_time=5, suite="base-a-test-cpu")
import re
import unittest
from pathlib import Path
import sglang.srt
from sglang.test.test_utils import CustomTestCase
_SRT_ROOT = Path(next(iter(sglang.srt.__path__)))
# Baselines counted over python/sglang/srt/**/*.py, including each function's
# own def line. Ratchet: decrease-only.
_RATCHETS = [
# Down to the shim definition itself; every call-site now goes through
# runtime_context.get_server_args().
("get_global_server_args", r"\bget_global_server_args\s*\(", 1),
(
"set_global_server_args_for_*",
r"\bset_global_server_args_for_(?:scheduler|tokenizer)\s*\(",
4,
),
]
class TestLegacyGlobalRatchet(CustomTestCase):
def test_legacy_accessor_call_sites_match_the_baselines(self):
# Exact pin, failing in BOTH directions: a grown count means new code
# bypassed the runtime_context accessors; a shrunk count means a
# removal forgot to lower the baseline, which would let later changes
# silently re-add call-sites up to the stale ceiling.
sources = [
path.read_text(encoding="utf-8", errors="replace")
for path in sorted(_SRT_ROOT.rglob("*.py"))
]
for name, pattern, baseline in _RATCHETS:
count = sum(len(re.findall(pattern, source)) for source in sources)
if count > baseline:
self.fail(
f"{name} call-sites grew: {count} > baseline {baseline}. "
"New code must use the sglang.srt.runtime_context accessors "
"(get_server_args() / get_context().set_server_args())."
)
if count < baseline:
self.fail(
f"{name} call-sites shrank: {count} < baseline {baseline}. "
"Lower the baseline in this file to lock in the progress."
)
if __name__ == "__main__":
unittest.main()
@@ -1,68 +0,0 @@
"""Ratchet guard: module-level runtime state in the flag-owning layers may
only shrink.
Runtime flags live on ``get_flags()`` groups (``moe`` / ``dp`` / ``capture``),
where they get lifecycle reset, typo-safe writes, and the transactional
test-override primitive. A new module-level global written through a
``global`` statement in these modules recreates the pattern this replaced:
state with ad-hoc lifecycle that leaks across unit-test teardowns and cannot
be overridden scoped.
The pin lists the survivors by name: the DP-attention topology values (owned
by the parallel vertical) and the TBO comm stream (a resource, owned by the
resources vertical). Migrating one of them must shrink its pin; adding a name
fails the ratchet.
"""
from sglang.test.ci.ci_register import register_cpu_ci
register_cpu_ci(est_time=5, suite="base-a-test-cpu")
import ast
import unittest
from pathlib import Path
import sglang.srt
from sglang.test.test_utils import CustomTestCase
_SRT_ROOT = Path(next(iter(sglang.srt.__path__)))
_PINNED_GLOBALS = {
"layers/moe/utils.py": frozenset(),
"layers/dp_attention.py": frozenset(
{
# DP-attention topology (parallel vertical scope).
"_ATTN_DP_RANK",
"_ATTN_DP_SIZE",
}
),
}
class TestModuleStateRatchet(CustomTestCase):
def test_global_statements_match_the_pins(self):
for rel, pinned in _PINNED_GLOBALS.items():
tree = ast.parse((_SRT_ROOT / rel).read_text())
declared = {
name
for node in ast.walk(tree)
if isinstance(node, ast.Global)
for name in node.names
}
grown = declared - pinned
self.assertFalse(
grown,
f"{rel} declares new module-level runtime state {sorted(grown)}; "
"put runtime flags on a get_flags() group instead "
"(see runtime_context.MoeFlags / DpFlags).",
)
shrunk = pinned - declared
self.assertFalse(
shrunk,
f"{rel} no longer declares {sorted(shrunk)}; "
"shrink the pin in this file to lock in the progress.",
)
if __name__ == "__main__":
unittest.main()
@@ -1,90 +0,0 @@
import ast
import pathlib
import unittest
from sglang.test.ci.ci_register import register_cpu_ci
from sglang.test.test_utils import CustomTestCase
register_cpu_ci(est_time=5, suite="base-a-test-cpu")
_REPO_ROOT = pathlib.Path(__file__).resolve().parents[3]
_SCAN_ROOTS = [_REPO_ROOT / "python", _REPO_ROOT / "test"]
class TestNoBarePytestMain(CustomTestCase):
def test_no_bare_pytest_main_in_repo(self):
offenders = []
for root in _SCAN_ROOTS:
if not root.exists():
continue
for path in root.rglob("*.py"):
violation = _find_bare_pytest_main(path)
if violation is not None:
offenders.append(violation)
self.assertFalse(
offenders,
msg=(
"Found bare `pytest.main(...)` in __main__ blocks (must be "
"wrapped in sys.exit(...) so failing tests propagate the exit "
"code to the CI runner):\n " + "\n ".join(offenders)
),
)
def _find_bare_pytest_main(path: pathlib.Path):
"""Return `<rel_path>:<lineno>` if `path` has a bare pytest.main(...) call
inside `if __name__ == "__main__":`, else None."""
try:
source = path.read_text(encoding="utf-8")
except (OSError, UnicodeDecodeError):
return None
try:
tree = ast.parse(source, filename=str(path))
except SyntaxError:
return None
for node in ast.walk(tree):
if not isinstance(node, ast.If):
continue
if not _is_main_guard(node.test):
continue
for stmt in node.body:
if _is_bare_pytest_main_call(stmt):
rel = path.relative_to(_REPO_ROOT)
return f"{rel}:{stmt.lineno}"
return None
def _is_main_guard(test: ast.expr) -> bool:
"""Match `__name__ == "__main__"` (either side)."""
if not isinstance(test, ast.Compare) or len(test.ops) != 1:
return False
if not isinstance(test.ops[0], ast.Eq):
return False
sides = [test.left, *test.comparators]
has_name = any(isinstance(s, ast.Name) and s.id == "__name__" for s in sides)
has_main = any(isinstance(s, ast.Constant) and s.value == "__main__" for s in sides)
return has_name and has_main
def _is_bare_pytest_main_call(stmt: ast.stmt) -> bool:
"""Match `pytest.main(...)` whose return value is discarded.
`sys.exit(pytest.main(...))` and `code = pytest.main(...)` are fine."""
if not isinstance(stmt, ast.Expr):
return False
call = stmt.value
if not isinstance(call, ast.Call):
return False
func = call.func
return (
isinstance(func, ast.Attribute)
and func.attr == "main"
and isinstance(func.value, ast.Name)
and func.value.id == "pytest"
)
if __name__ == "__main__":
unittest.main()
@@ -1,85 +0,0 @@
"""Ratchet guard: server_args mutations outside the resolution pipeline may
only decrease.
After ``ServerArgs.__post_init__`` returns, the instance carries the resolved
configuration; the resolution pipeline (``server_args.py`` and
``arg_groups/``) is the only place that computes it. Every assignment to a
``server_args`` field elsewhere weakens that contract, so the count below is
an exact pin: new mutations must not appear, and removals must lower the
baseline to lock in the progress.
There is no post-resolution mutation entry point on the instance any more:
resolved config changes go to the context bags via
``get_context().override(source, **fields)``, and a value that differs for one
runner or worker travels as a constructor argument to it. The baseline is
therefore zero. ``ServerArgs.__setattr__`` raises
on a bare assignment after resolution; this ratchet catches the sites the tests
never execute.
"""
from sglang.test.ci.ci_register import register_cpu_ci
register_cpu_ci(est_time=5, suite="base-a-test-cpu")
import re
import unittest
from pathlib import Path
import sglang
from sglang.test.test_utils import CustomTestCase
_SGLANG_ROOT = Path(next(iter(sglang.__path__)))
# Assignments to a server_args attribute (``server_args.x = ...``,
# ``self.server_args.x = ...``, and the ``sa`` alias used by a few helpers).
# ``==`` comparisons are excluded by the negative lookahead.
_MUTATION_PATTERNS = [
# (?![=}]) skips ``==`` comparisons and f-string ``{x=}`` debug specs.
re.compile(r"\bserver_args\.[a-z0-9_]+\s*=(?![=}])"),
re.compile(r"\bsa\.[a-z0-9_]+\s*=(?![=}])"),
re.compile(r"get_(?:global_)?server_args\(\)\.[a-z0-9_]+\s*=(?![=}])"),
# setattr is the same write with the attribute name behind a variable.
re.compile(
r"setattr\(\s*(?:[\w.]+\.)?(?:server_args|sa|get_(?:global_)?server_args\(\))\s*,"
),
]
# The resolution pipeline itself (mutation is its job) and multimodal_gen,
# whose ServerArgs is a different class outside this contract.
_EXCLUDED = (
"srt/server_args.py",
"srt/arg_groups",
"multimodal_gen",
)
_BASELINE = 0
class TestServerArgsMutationRatchet(CustomTestCase):
def test_out_of_pipeline_mutations_match_the_baseline(self):
count = 0
for path in sorted(_SGLANG_ROOT.rglob("*.py")):
rel = path.relative_to(_SGLANG_ROOT).as_posix()
if rel.startswith(_EXCLUDED):
continue
source = path.read_text()
count += sum(len(p.findall(source)) for p in _MUTATION_PATTERNS)
if count > _BASELINE:
self.fail(
f"server_args mutations outside the resolution pipeline grew: "
f"{count} > baseline {_BASELINE}. Configuration is resolved in "
"ServerArgs.__post_init__; declare through the pipeline "
"(passes / declare_late_resolution), change resolved config "
"with get_context().override(source, ...), or hand the value "
"to its runner as a constructor argument — do not assign fields."
)
if count < _BASELINE:
self.fail(
f"server_args mutations outside the resolution pipeline "
f"shrank: {count} < baseline {_BASELINE}. Lower the baseline "
"in this file to lock in the progress."
)
if __name__ == "__main__":
unittest.main()
+3 -4
View File
@@ -319,10 +319,9 @@ def run_a_suite(args):
for f in glob.glob(
os.path.join(script_dir, "registered", "**", "*.py"), recursive=True
)
if not f.endswith("/conftest.py")
and not f.endswith("/__init__.py")
and not f.endswith("/cpu/utils.py")
and not f.endswith("/run_tests.py")
# conftest.py / __init__.py are pytest+package structure, never
# registered tests, and must not be executed as one.
if os.path.basename(f) not in ("conftest.py", "__init__.py")
]
# Strict: all discovered files must have proper registration