fix: always capture default prefill CUDA graph (#33352)

This commit is contained in:
Mick
2026-08-08 19:24:49 +08:00
committed by GitHub
parent cf2d4fd679
commit db75dfe10f
3 changed files with 48 additions and 45 deletions
@@ -3,28 +3,15 @@ from types import SimpleNamespace
import pytest
from sglang.srt.model_executor.cuda_graph_config import Phase
from sglang.srt.model_executor.model_runner_components import cuda_graph_setup
from sglang.srt.model_executor.model_runner_components.cuda_graph_setup import (
capture_decode_graph,
should_skip_auto_prefill_cuda_graph_for_memory,
)
from sglang.test.ci.ci_register import register_cpu_ci
register_cpu_ci(est_time=5, suite="base-a-test-cpu")
def test_auto_prefill_cuda_graph_memory_gate():
assert should_skip_auto_prefill_cuda_graph_for_memory(3.99, set())
assert not should_skip_auto_prefill_cuda_graph_for_memory(4.0, set())
def test_explicit_prefill_backend_bypasses_memory_gate():
assert not should_skip_auto_prefill_cuda_graph_for_memory(
0.0, {(Phase.PREFILL, "backend")}
)
def test_model_runner_can_override_decode_graph_runner(monkeypatch):
class CustomGraphRunner:
def __init__(self, model_runner):