ci: combine H200 8-GPU warmup steps and surface server log on every path (#24253)

This commit is contained in:
Alison Shao
2026-05-13 20:09:52 -07:00
committed by GitHub
parent 8b9ff4a68c
commit b71d74673c
3 changed files with 268 additions and 49 deletions
+10 -13
View File
@@ -780,23 +780,20 @@ jobs:
CUSTOM_BUILD_SGL_KERNEL=${{needs.check-changes.outputs.sgl_kernel}} bash scripts/ci/cuda/ci_install_dependency.sh CUSTOM_BUILD_SGL_KERNEL=${{needs.check-changes.outputs.sgl_kernel}} bash scripts/ci/cuda/ci_install_dependency.sh
- name: Warmup DeepGEMM JIT Compilation - name: Warmup DeepGEMM JIT Compilation
timeout-minutes: 25 # Per-model TP must match the test's launch in test/registered/ — see
# FALLBACK_ARGS in scripts/ci/cuda/warmup_deep_gemm.py for extra dp/ep
# flags. Only models that actually invoke DeepGEMM kernels at runtime
# are listed. Cold-cache ~13 min; warm-cache ≤30 s via marker file.
timeout-minutes: 60
run: | run: |
# Activate venv if available (GITHUB_ENV may have failed to propagate)
[ -f "${SGLANG_CI_VENV_PATH:-/dev/null}/bin/activate" ] && source "${SGLANG_CI_VENV_PATH}/bin/activate" [ -f "${SGLANG_CI_VENV_PATH:-/dev/null}/bin/activate" ] && source "${SGLANG_CI_VENV_PATH}/bin/activate"
[ -f "${SGLANG_CI_VENV_PATH:-/dev/null}/env.sh" ] && source "${SGLANG_CI_VENV_PATH}/env.sh" [ -f "${SGLANG_CI_VENV_PATH:-/dev/null}/env.sh" ] && source "${SGLANG_CI_VENV_PATH}/env.sh"
python3 scripts/ci/cuda/warmup_deep_gemm.py \ python3 scripts/ci/cuda/warmup_deep_gemm.py \
deepseek-ai/DeepSeek-V3-0324:8 \ deepseek-ai/DeepSeek-V3-0324:8 \
deepseek-ai/DeepSeek-V3.2-Exp:8 deepseek-ai/DeepSeek-V3.2:8 \
zai-org/GLM-5-FP8:8 \
- name: Warmup Server CUDA Graphs XiaomiMiMo/MiMo-V2-Flash:4 \
timeout-minutes: 25 XiaomiMiMo/MiMo-V2.5:8
run: |
[ -f "${SGLANG_CI_VENV_PATH:-/dev/null}/bin/activate" ] && source "${SGLANG_CI_VENV_PATH}/bin/activate"
[ -f "${SGLANG_CI_VENV_PATH:-/dev/null}/env.sh" ] && source "${SGLANG_CI_VENV_PATH}/env.sh"
python3 scripts/ci/cuda/warmup_server.py \
deepseek-ai/DeepSeek-V3-0324:8 \
inclusionAI/Ring-2.5-1T:8
- name: Run test - name: Run test
timeout-minutes: 30 timeout-minutes: 30
@@ -998,7 +995,7 @@ jobs:
[ -f "${SGLANG_CI_VENV_PATH:-/dev/null}/env.sh" ] && source "${SGLANG_CI_VENV_PATH}/env.sh" [ -f "${SGLANG_CI_VENV_PATH:-/dev/null}/env.sh" ] && source "${SGLANG_CI_VENV_PATH}/env.sh"
python3 scripts/ci/cuda/warmup_deep_gemm.py \ python3 scripts/ci/cuda/warmup_deep_gemm.py \
deepseek-ai/DeepSeek-V3-0324:8 \ deepseek-ai/DeepSeek-V3-0324:8 \
deepseek-ai/DeepSeek-V3.2-Exp:8 deepseek-ai/DeepSeek-V3.2:8
- name: Warmup Server CUDA Graphs - name: Warmup Server CUDA Graphs
timeout-minutes: 25 timeout-minutes: 25
+214 -14
View File
@@ -11,16 +11,61 @@ for unsupported architectures.
Usage: Usage:
python3 scripts/ci/cuda/warmup_deep_gemm.py \ python3 scripts/ci/cuda/warmup_deep_gemm.py \
deepseek-ai/DeepSeek-V3-0324:8 \ deepseek-ai/DeepSeek-V3-0324:8 \
deepseek-ai/DeepSeek-V3.2-Exp:8 deepseek-ai/DeepSeek-V3.2:8
""" """
import hashlib
import json import json
import os import os
import signal
import subprocess import subprocess
import sys import sys
import threading
import time import time
from math import ceil from math import ceil
from pathlib import Path from pathlib import Path
from typing import Dict, List
# Shared with warmup_server.py. Wipe alongside /root/.cache/deep_gemm if you
# clear the DeepGEMM JIT cache — a stale marker → in-test JIT compile.
MARKER_DIR = os.path.join(os.path.expanduser("~"), ".cache", "sglang", "warmup_markers")
# Outer cap for stuck fallback subprocesses; CRASH_MARKERS abort sooner.
FALLBACK_TIMEOUT_SEC = 600
# Per-model launch flags forwarded to `sglang.compile_deep_gemm`. DeepGEMM
# cache key includes per-rank N/K (depends on tp/dp/ep) — must match each
# model's `other_args` in test/registered/ or warmed shapes won't be hit.
FALLBACK_ARGS: Dict[str, List[str]] = {
"deepseek-ai/DeepSeek-V3.2": ["--dp", "8", "--enable-dp-attention"],
"zai-org/GLM-5-FP8": ["--dp", "8", "--enable-dp-attention"],
"XiaomiMiMo/MiMo-V2-Flash": [
"--dp",
"2",
"--enable-dp-attention",
"--attention-backend",
"fa3",
],
# --mm-enable-dp-encoder is required: without it DP0 runs the vision
# encoder alone and DP1 deadlocks at the next collective.
"XiaomiMiMo/MiMo-V2.5": [
"--dp",
"2",
"--enable-dp-attention",
"--mm-enable-dp-encoder",
"--attention-backend",
"fa3",
"--mm-attention-backend",
"fa3",
],
}
# compile_deep_gemm polls /v1/models for the full timeout even after a TP rank
# dies; the watcher uses these to kill the group within seconds instead.
CRASH_MARKERS = (
"Scheduler hit an exception",
"Received sigquit from a child",
)
# Configure DeepGEMM cache before importing deep_gemm # Configure DeepGEMM cache before importing deep_gemm
os.environ["DG_JIT_CACHE_DIR"] = os.getenv( os.environ["DG_JIT_CACHE_DIR"] = os.getenv(
@@ -294,9 +339,89 @@ def compile_shapes_lightweight(shapes, m_list):
print(f" Done in {elapsed:.1f}s") print(f" Done in {elapsed:.1f}s")
def _kill_pg_and_wait(proc):
"""SIGTERM the subprocess's process group, escalate to SIGKILL if needed."""
try:
os.killpg(os.getpgid(proc.pid), signal.SIGTERM)
except (ProcessLookupError, OSError):
pass
try:
return proc.wait(timeout=10)
except subprocess.TimeoutExpired:
try:
os.killpg(os.getpgid(proc.pid), signal.SIGKILL)
except (ProcessLookupError, OSError):
pass
try:
return proc.wait(timeout=5)
except subprocess.TimeoutExpired:
return -1
def get_version_key():
"""Hash of Python + Triton + PyTorch versions; invalidates markers on upgrade."""
parts = [sys.version]
try:
import triton # noqa: WPS433
parts.append(f"triton={triton.__version__}")
except ImportError:
parts.append("triton=none")
try:
import torch # noqa: WPS433
parts.append(f"torch={torch.__version__}")
except ImportError:
parts.append("torch=none")
return hashlib.sha256("|".join(parts).encode()).hexdigest()[:12]
def get_fallback_marker_path(model, tp, extra_args):
"""Marker path for one (model, tp, extra_args) fallback invocation."""
args_blob = json.dumps(list(extra_args))
args_hash = hashlib.md5(args_blob.encode()).hexdigest()[:8]
safe_model = model.replace("/", "--")
return os.path.join(
MARKER_DIR,
f"deepgemm_fallback_{safe_model}_tp{tp}_{args_hash}_{get_version_key()}.done",
)
def check_fallback_marker(model, tp, extra_args):
return os.path.exists(get_fallback_marker_path(model, tp, extra_args))
def write_fallback_marker(model, tp, extra_args):
marker = get_fallback_marker_path(model, tp, extra_args)
os.makedirs(os.path.dirname(marker), exist_ok=True)
Path(marker).write_text(
json.dumps(
{
"model": model,
"tp": tp,
"extra_args": list(extra_args),
"version_key": get_version_key(),
"timestamp": time.time(),
}
)
)
print(f" Wrote marker: {marker}")
def fallback_compile_deep_gemm(model, tp): def fallback_compile_deep_gemm(model, tp):
"""Fall back to full sglang.compile_deep_gemm (loads model weights).""" """Fall back to full sglang.compile_deep_gemm (loads model weights).
print(f"Falling back to full compile_deep_gemm for {model} (tp={tp})...")
Runs in its own process group so a hung subprocess (e.g. one TP rank
crashes and the rest deadlock on NCCL collectives) can be killed
cleanly without leaking children. Watches subprocess output for crash
markers so a deterministic failure aborts in seconds rather than burning
the full FALLBACK_TIMEOUT_SEC.
"""
extra_args = FALLBACK_ARGS.get(model, [])
print(
f"Falling back to full compile_deep_gemm for {model} "
f"(tp={tp}, extra_args={extra_args})..."
)
cmd = [ cmd = [
sys.executable, sys.executable,
"-m", "-m",
@@ -308,11 +433,57 @@ def fallback_compile_deep_gemm(model, tp):
"--trust-remote-code", "--trust-remote-code",
"--model-loader-extra-config", "--model-loader-extra-config",
'{"enable_multithread_load": true, "num_threads": 64}', '{"enable_multithread_load": true, "num_threads": 64}',
# Cap compile_deep_gemm's own /v1/models polling loop so it gives up
# before our outer timeout has to SIGTERM it.
"--timeout",
str(FALLBACK_TIMEOUT_SEC),
*extra_args,
] ]
result = subprocess.run(cmd)
if result.returncode != 0: crashed = threading.Event()
print(f"Warning: fallback failed for {model} (exit code {result.returncode})") proc = subprocess.Popen(
return result.returncode == 0 cmd,
preexec_fn=os.setsid,
stdout=subprocess.PIPE,
stderr=subprocess.STDOUT,
bufsize=1,
text=True,
)
def _watch():
# Stream child output to our stdout while scanning for crash markers.
for line in proc.stdout:
sys.stdout.write(line)
sys.stdout.flush()
if not crashed.is_set() and any(m in line for m in CRASH_MARKERS):
crashed.set()
watcher = threading.Thread(target=_watch, daemon=True)
watcher.start()
deadline = time.monotonic() + FALLBACK_TIMEOUT_SEC
while True:
rc = proc.poll()
if rc is not None:
watcher.join(timeout=2)
if rc != 0:
print(f"Warning: fallback failed for {model} (exit code {rc})")
return rc == 0
if crashed.is_set():
print(
f"Warning: detected crash marker in {model} (tp={tp}) subprocess; "
"killing process group and continuing."
)
_kill_pg_and_wait(proc)
return False
if time.monotonic() >= deadline:
print(
f"Warning: fallback timed out after {FALLBACK_TIMEOUT_SEC}s for "
f"{model} (tp={tp}); killing process group and continuing."
)
_kill_pg_and_wait(proc)
return False
time.sleep(2)
def main(): def main():
@@ -353,22 +524,39 @@ def main():
print(f" SKIP {model} (tp={tp}): config.json not in HF cache") print(f" SKIP {model} (tp={tp}): config.json not in HF cache")
continue continue
# Models with FALLBACK_ARGS launch with extra dp/ep/dp-attention flags
# that change per-rank N/K. The lightweight path doesn't model those —
# it computes attention shapes assuming TP-only sharding — so dedup'ing
# such a model to a no-override DeepSeek V2/V3 lookalike (e.g. V3.2 →
# V3-0324) silently picks the wrong attention shapes for the test's
# actual launch config. Force these through fallback so the populated
# cache matches the real test.
has_fallback_override = model in FALLBACK_ARGS
key = get_architecture_key(config, tp) key = get_architecture_key(config, tp)
if key in seen_keys: if key in seen_keys and not has_fallback_override:
print(f" DEDUP {model} (tp={tp}): same shapes as {seen_keys[key]}") print(f" DEDUP {model} (tp={tp}): same shapes as {seen_keys[key]}")
continue continue
if is_deepseek_v2v3(config): if is_deepseek_v2v3(config) and not has_fallback_override:
shapes = compute_deepseek_v2v3_shapes(config, tp) shapes = compute_deepseek_v2v3_shapes(config, tp)
seen_keys[key] = model seen_keys[key] = model
to_process.append((model, tp, config, shapes)) to_process.append((model, tp, config, shapes))
print(f" FOUND {model} (tp={tp}): {len(shapes)} DeepGEMM shape(s)") print(f" FOUND {model} (tp={tp}): {len(shapes)} DeepGEMM shape(s)")
else: else:
# Unknown architecture: will use fallback
seen_keys[key] = model seen_keys[key] = model
to_process.append((model, tp, config, None)) to_process.append((model, tp, config, None))
arch = config.get("architectures", ["unknown"]) if has_fallback_override:
print(f" FOUND {model} (tp={tp}): unknown arch {arch}, will use fallback") print(
f" FOUND {model} (tp={tp}): forced fallback (extra args "
f"{FALLBACK_ARGS[model]})"
)
else:
arch = config.get("architectures", ["unknown"])
print(
f" FOUND {model} (tp={tp}): unknown arch {arch}, "
"will use fallback"
)
if not to_process: if not to_process:
print("\nNo models to process. Done.") print("\nNo models to process. Done.")
@@ -383,8 +571,20 @@ def main():
print(f"{'=' * 60}") print(f"{'=' * 60}")
if shapes is None: if shapes is None:
# Unknown architecture: fall back to full compile_deep_gemm # Fallback path: full sglang.compile_deep_gemm with the test's launch
fallback_compile_deep_gemm(model, tp) # flags. Loading model weights inside that subprocess is the dominant
# cost (45-170s/model) and dwarfs the actual DeepGEMM compile, so we
# skip the whole fallback when the marker says we already populated
# this cache. Cache is keyed on (model, tp, extra_args, version_key).
extra_args = FALLBACK_ARGS.get(model, [])
if check_fallback_marker(model, tp, extra_args):
print(
f" SKIP fallback (warm marker found): {model} (tp={tp}, "
f"extra_args={extra_args})"
)
continue
if fallback_compile_deep_gemm(model, tp):
write_fallback_marker(model, tp, extra_args)
continue continue
# Print shape summary # Print shape summary
+44 -22
View File
@@ -11,8 +11,7 @@ invalidated when Python, Triton, or PyTorch versions change.
Usage: Usage:
python3 scripts/ci/cuda/warmup_server.py \ python3 scripts/ci/cuda/warmup_server.py \
deepseek-ai/DeepSeek-V3-0324:8 \ deepseek-ai/DeepSeek-V3-0324:8
inclusionAI/Ring-2.5-1T:8
""" """
import hashlib import hashlib
@@ -87,23 +86,41 @@ def write_marker(model, tp):
def kill_server(proc): def kill_server(proc):
"""Kill server process tree.""" """Kill server process tree."""
if proc.poll() is not None: if proc.poll() is None:
return
try:
os.killpg(os.getpgid(proc.pid), signal.SIGTERM)
except (ProcessLookupError, OSError):
pass
try:
proc.wait(timeout=15)
except subprocess.TimeoutExpired:
try: try:
os.killpg(os.getpgid(proc.pid), signal.SIGKILL) os.killpg(os.getpgid(proc.pid), signal.SIGTERM)
except (ProcessLookupError, OSError): except (ProcessLookupError, OSError):
pass pass
try: try:
proc.wait(timeout=5) proc.wait(timeout=15)
except subprocess.TimeoutExpired: except subprocess.TimeoutExpired:
try:
os.killpg(os.getpgid(proc.pid), signal.SIGKILL)
except (ProcessLookupError, OSError):
pass
try:
proc.wait(timeout=5)
except subprocess.TimeoutExpired:
pass
# sglang's scheduler_TP* and detokenizer workers spawn through
# multiprocessing with their own session/process group, so they escape
# killpg on launch_server and stay alive holding GPU memory after a
# readiness-timeout or unclean exit. Kill any survivors by name so the
# next model (or the next CI step) starts with empty GPUs.
for pattern in ("sglang::scheduler", "sglang::detokenizer"):
try:
subprocess.run(
["pkill", "-9", "-f", pattern],
timeout=5,
check=False,
stdout=subprocess.DEVNULL,
stderr=subprocess.DEVNULL,
)
except (FileNotFoundError, subprocess.TimeoutExpired):
pass pass
# Let the driver release device memory before the caller measures it.
time.sleep(2)
def wait_for_server(base_url, proc, timeout): def wait_for_server(base_url, proc, timeout):
@@ -192,15 +209,6 @@ def warmup_one_model(model, tp, port):
ok, err = wait_for_server(base_url, proc, SERVER_STARTUP_TIMEOUT) ok, err = wait_for_server(base_url, proc, SERVER_STARTUP_TIMEOUT)
if not ok: if not ok:
print(f" Warning: server not ready: {err}") print(f" Warning: server not ready: {err}")
# Dump last lines of server log for debugging
try:
log_file.flush()
with open(log_path) as f:
lines = f.readlines()
for line in lines[-20:]:
print(f" | {line.rstrip()}")
except Exception:
pass
return False return False
print(" Server ready, sending generate request...") print(" Server ready, sending generate request...")
@@ -208,6 +216,20 @@ def warmup_one_model(model, tp, port):
return True return True
finally: finally:
# Surface the tail of the server log so CI captures validation
# messages, exceptions, and warmup progress (the launch_server
# subprocess writes stdout/stderr to the tempfile, not our stdout).
try:
log_file.flush()
with open(log_path) as f:
lines = f.readlines()
print(f" --- server log tail ({len(lines)} lines, last 30) ---")
for line in lines[-30:]:
print(f" | {line.rstrip()}")
print(" --- end server log ---")
except Exception:
pass
print(" Killing server...") print(" Killing server...")
kill_server(proc) kill_server(proc)
log_file.close() log_file.close()