fix: always capture default prefill CUDA graph (#33352)
This commit is contained in:
@@ -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):
|
||||
|
||||
@@ -61,6 +61,54 @@ class _FakeKVIndexKernel:
|
||||
|
||||
|
||||
class TestPrefillCudaGraphRunnerChunkedPrefix(CustomTestCase):
|
||||
def test_low_free_memory_still_captures_prefill_graph(self):
|
||||
eager_runner = object()
|
||||
prefill_runner = object()
|
||||
model_runner = SimpleNamespace(
|
||||
device="cuda",
|
||||
gpu_id=0,
|
||||
is_draft_worker=False,
|
||||
spec_algorithm=SimpleNamespace(is_eagle=lambda: False),
|
||||
server_args=SimpleNamespace(
|
||||
enable_lora=False,
|
||||
cuda_graph_config=SimpleNamespace(
|
||||
prefill=SimpleNamespace(bs=[1], backend=Backend.BREAKABLE)
|
||||
),
|
||||
),
|
||||
model=SimpleNamespace(),
|
||||
model_config=SimpleNamespace(context_len=8192, num_hidden_layers=1),
|
||||
req_to_token_pool=SimpleNamespace(size=1),
|
||||
)
|
||||
language_model = SimpleNamespace(layers=[object()])
|
||||
|
||||
with (
|
||||
patch.object(graph_setup, "check_cuda_graph_backend", return_value=False),
|
||||
patch.object(
|
||||
graph_setup, "resolve_language_model", return_value=language_model
|
||||
),
|
||||
patch.object(
|
||||
graph_setup,
|
||||
"compute_attention_and_moe_layers",
|
||||
return_value=([object()], [], [], [], []),
|
||||
),
|
||||
patch.object(
|
||||
graph_setup,
|
||||
"get_available_gpu_memory",
|
||||
side_effect=[3.99, 3.5],
|
||||
),
|
||||
patch.object(
|
||||
graph_setup,
|
||||
"PrefillCudaGraphRunner",
|
||||
return_value=prefill_runner,
|
||||
),
|
||||
):
|
||||
capture = capture_prefill_graph(
|
||||
model_runner=model_runner,
|
||||
eager_runner=eager_runner,
|
||||
)
|
||||
|
||||
self.assertIs(capture.runner, prefill_runner)
|
||||
|
||||
def test_eagle_target_tc_piecewise_skips_last_mode_capture(self):
|
||||
eager_runner = object()
|
||||
model_runner = SimpleNamespace(
|
||||
|
||||
Reference in New Issue
Block a user