fix: skip unsafe automatic prefill graph capture (#31204)
This commit is contained in:
@@ -40,6 +40,24 @@ if TYPE_CHECKING:
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
# Measured immediately before prefill graph construction, after model weights,
|
||||
# KV cache, and the eager runner's static buffers have been allocated. Below
|
||||
# this budget, compiling/capturing a multi-bucket prefill graph is likely to
|
||||
# OOM or make no forward progress for large models. Keep an explicitly chosen
|
||||
# backend untouched: an operator may intentionally trade KV capacity for it.
|
||||
_MIN_AUTO_PREFILL_CUDA_GRAPH_FREE_MEMORY_GB = 4.0
|
||||
|
||||
|
||||
def should_skip_auto_prefill_cuda_graph_for_memory(
|
||||
available_memory_gb: float,
|
||||
cuda_graph_config_locked: set[tuple[str, str]],
|
||||
) -> bool:
|
||||
"""Return whether an auto-selected prefill graph lacks capture headroom."""
|
||||
return (
|
||||
(Phase.PREFILL, "backend") not in cuda_graph_config_locked
|
||||
and available_memory_gb < _MIN_AUTO_PREFILL_CUDA_GRAPH_FREE_MEMORY_GB
|
||||
)
|
||||
|
||||
|
||||
class DecodeGraphCapture(msgspec.Struct, frozen=True, kw_only=True):
|
||||
runner: Optional[BaseRunner]
|
||||
@@ -217,6 +235,20 @@ def capture_prefill_graph(
|
||||
tic = time.perf_counter()
|
||||
before_mem = get_available_gpu_memory(model_runner.device, model_runner.gpu_id)
|
||||
prefill_backend = model_runner.server_args.cuda_graph_config.prefill.backend
|
||||
if should_skip_auto_prefill_cuda_graph_for_memory(
|
||||
before_mem,
|
||||
getattr(model_runner.server_args, "_cuda_graph_config_locked", set()),
|
||||
):
|
||||
logger.warning(
|
||||
"Disabling auto-selected prefill CUDA graph: only %.2f GiB is free "
|
||||
"after model/KV/eager-buffer allocation; at least %.2f GiB is "
|
||||
"required for capture. Set an explicit prefill CUDA graph backend "
|
||||
"to override this safety gate.",
|
||||
before_mem,
|
||||
_MIN_AUTO_PREFILL_CUDA_GRAPH_FREE_MEMORY_GB,
|
||||
)
|
||||
return eager_runner
|
||||
|
||||
role = "draft" if model_runner.is_draft_worker else "target"
|
||||
capture_name = f"{role} prefill"
|
||||
capture_num_tokens = sorted(model_runner.server_args.cuda_graph_config.prefill.bs)
|
||||
|
||||
@@ -0,0 +1,26 @@
|
||||
import sys
|
||||
|
||||
import pytest
|
||||
|
||||
from sglang.srt.model_executor.cuda_graph_config import Phase
|
||||
from sglang.srt.model_executor.model_runner_components.cuda_graph_setup import (
|
||||
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")}
|
||||
)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
sys.exit(pytest.main([__file__, "-v"]))
|
||||
Reference in New Issue
Block a user