[Memory] Size the CUDA graph pool from warmup measurements and fix graph-pool borrowing (#36911)

Co-authored-by: cctry <cctry@fb.com>
This commit is contained in:
cctry
2026-09-01 09:32:38 -07:00
committed by GitHub
co-authored by cctry
parent c34f378342
commit 9a05b470fa
17 changed files with 681 additions and 198 deletions
+3
View File
@@ -1293,6 +1293,9 @@ class Envs:
SGLANG_MEMORY_SAVER_CUDA_GRAPH = EnvBool(False)
# Reuse wholly-free graph-pool segments for step-local eager allocations.
SGLANG_ENABLE_GRAPH_POOL_BORROW = EnvBool(False)
# Mint capture's measured footprint as one span so the graph pool is carved
# out of a single contiguous region instead of grown segment by segment.
SGLANG_ENABLE_GRAPH_POOL_PRECARVE = EnvBool(False)
# Eager forward wraps the ForwardBatch's own tensors instead of copying them
# into the CUDA graph buffer registry (no per-iter device-to-device copy).
SGLANG_EAGER_INPUT_NO_COPY = EnvBool(False)
+7
View File
@@ -1060,6 +1060,13 @@ class Scheduler(
self.init_all_cuda_graphs()
model_runner = self.tp_worker.model_runner
with torch.get_device_module(model_runner.device).stream(
model_runner.forward_stream
):
if self.draft_worker is None:
model_runner.prewarm_sampling()
else:
self.draft_worker.prewarm_sampling()
if model_runner.token_to_kv_pool.post_capture_active:
tic = time.perf_counter()
model_runner.post_capture_resize_kv_pool()
@@ -253,6 +253,14 @@ elif current_platform.is_out_of_tree():
logger = logging.getLogger(__name__)
@dataclass(frozen=True)
class SamplingPrewarmResult:
"""Memory requirements observed while pre-warming a sampling path."""
sampling_input_bytes: int = 0
sampling_headroom_bytes: int = 0
def _prefill_cuda_graph_allows_context_parallel(
prefill_runner, forward_batch: ForwardBatch
) -> bool:
@@ -377,6 +385,7 @@ class ModelRunner:
self.draft_model_idx = draft_model_idx
self.enable_hisparse = get_memory().enable_hisparse
self._sampling_observer: Optional[SamplingObserver] = None
self.sampling_prewarm_result = SamplingPrewarmResult()
self.init_startup_observability()
@@ -1058,6 +1067,11 @@ class ModelRunner:
n_prepared,
)
def prewarm_sampling(self) -> SamplingPrewarmResult:
"""Warm the sampling path after graph initialization."""
self.sampling_prewarm_result = SamplingPrewarmResult()
return self.sampling_prewarm_result
def init_cuda_graphs(self, capture_decode_cuda_graph: bool = True):
capture = capture_cuda_graphs(
model_runner=self, capture_decode_cuda_graph=capture_decode_cuda_graph
@@ -11,6 +11,7 @@ from sglang.srt.configs.hybrid_arch import mambaish_config
from sglang.srt.distributed import get_world_group
from sglang.srt.mem_cache.kv_cache_configurator import mm_runtime_reservation_gb
from sglang.srt.model_executor.cuda_graph_config import Backend
from sglang.srt.model_executor.runner_utils.pool import graph_pool_borrow_enabled
from sglang.srt.platforms import current_platform
from sglang.srt.runtime_context import (
get_disagg,
@@ -85,6 +86,13 @@ def compute_post_capture_kv_resize(
)
/ 1024,
)
if not graph_pool_borrow_enabled():
# Borrowing serves the sampling temporaries out of idle graph storage;
# without it they need real headroom the KV pool must not claim.
headroom_gb = max(
headroom_gb,
model_runner.sampling_prewarm_result.sampling_headroom_bytes / (1 << 30),
)
mm_reservation_gb = mm_runtime_reservation_gb(
is_multimodal=model_runner.model_config.is_multimodal,
mm_feature_transport=get_mm().mm_feature_transport,
@@ -40,6 +40,7 @@ from sglang.srt.model_executor.runner_backend_utils.breakable_cuda_graph import
enable_breakable_cuda_graph,
)
from sglang.srt.model_executor.runner_utils.pool import (
GraphPoolPrecarve,
get_or_create_global_graph_memory_pool,
graph_pool_capture_scope,
graph_pool_replay_scope,
@@ -77,6 +78,7 @@ class BreakableCudaGraphBackend(DedupedCudaGraphMixin, BaseCudaGraphBackend):
self._capture_stream: Optional[torch.cuda.Stream] = None
self._debug_eager = debug_eager
self._shared_output_buffer: Optional[Any] = None
self._precarve = GraphPoolPrecarve()
self._memory_saver_adapter: Optional[Any] = TorchMemorySaverAdapter.create(
enable=enable_memory_saver
and get_bool_env_var("SGLANG_MEMORY_SAVER_CUDA_GRAPH")
@@ -117,7 +119,8 @@ class BreakableCudaGraphBackend(DedupedCudaGraphMixin, BaseCudaGraphBackend):
for _ in range(2):
self._device_module.synchronize()
self._tp_group.barrier()
warmup_out = forward_fn()
with self._precarve.measure():
warmup_out = forward_fn()
if post_warmup_hook is not None:
post_warmup_hook()
@@ -134,6 +137,7 @@ class BreakableCudaGraphBackend(DedupedCudaGraphMixin, BaseCudaGraphBackend):
stream=self._capture_stream,
barrier_fn=self._tp_group.barrier,
):
self._precarve.mint()
out = captured_fn()
out_rows = self._output_rows(out, size)
self._copy_output_to_buffer(out, self._shared_output_buffer, out_rows)
@@ -31,6 +31,7 @@ from sglang.srt.model_executor.runner_backend.base_cuda_graph_backend import (
BaseCudaGraphBackend,
)
from sglang.srt.model_executor.runner_utils.pool import (
GraphPoolPrecarve,
get_or_create_global_graph_memory_pool,
graph_pool_capture_scope,
graph_pool_replay_scope,
@@ -64,6 +65,7 @@ class FullCudaGraphBackend(BaseCudaGraphBackend):
self._device_module = cuda_graph_runner.device_module
self._tp_group = cuda_graph_runner.model_runner.tp_group
self._capture_stream: Optional[torch.cuda.Stream] = None
self._precarve = GraphPoolPrecarve()
self._memory_saver_adapter: Optional[Any] = TorchMemorySaverAdapter.create(
enable=enable_memory_saver
and get_bool_env_var("SGLANG_MEMORY_SAVER_CUDA_GRAPH")
@@ -107,7 +109,8 @@ class FullCudaGraphBackend(BaseCudaGraphBackend):
for _ in range(2):
self._device_module.synchronize()
self._tp_group.barrier()
forward_fn()
with self._precarve.measure():
forward_fn()
if profiler is not None:
profiler.step()
if post_warmup_hook is not None:
@@ -131,6 +134,7 @@ class FullCudaGraphBackend(BaseCudaGraphBackend):
graph_pool_capture_scope(),
graph_ctx(cuda_graph=graph, pool=self._pool, stream=self._capture_stream),
):
self._precarve.mint()
out = forward_fn()
if profiler is not None:
@@ -46,6 +46,7 @@ def disable_graph_pool_borrow(reason: str) -> None:
"""Disable borrowing when graph storage is managed outside the shared pool."""
global _borrow_disabled_reason
_borrow_disabled_reason = reason
_teardown_borrow_pool()
logger.info("Graph pool borrow disabled: %s", reason)
@@ -93,6 +94,43 @@ def get_or_create_global_graph_capture_stream() -> Any:
return get_stream(_CAPTURE_STREAM_NAME)
class GraphPoolPrecarve:
"""Pre-carve the memory pool to reduce fragmentation."""
def __init__(self) -> None:
self.nbytes = 0
self.minted = False
@contextmanager
def measure(self) -> Iterator[None]:
"""Wrap one eager warmup. the last one before ``mint`` sets the size."""
if self.minted or not envs.SGLANG_ENABLE_GRAPH_POOL_PRECARVE.get():
yield
return
torch.cuda.synchronize()
# Shrink the cache first so the warmup's reserved growth is its own
# footprint. Reserved (not allocated) is the stat to use: allocated
# peak is the live-byte sum and undershoots by exactly the packing
# holes the carved span has to absorb.
torch.cuda.empty_cache()
torch.cuda.reset_peak_memory_stats()
base = torch.cuda.memory_stats()["reserved_bytes.all.current"]
yield
torch.cuda.synchronize()
self.nbytes = torch.cuda.memory_stats()["reserved_bytes.all.peak"] - base
def mint(self) -> None:
"""Pre-allocate the space"""
if self.minted:
return
self.minted = True
if self.nbytes <= 0:
return
span = torch.empty(self.nbytes, dtype=torch.uint8, device="cuda")
del span
logger.info("Graph pool pre-carved: %.2f GB", self.nbytes / 2**30)
def graph_pool_borrow_enabled() -> bool:
if (
_borrow_disabled_reason is not None
@@ -143,11 +181,11 @@ def graph_pool_capture_scope() -> Iterator[None]:
def find_free_graph_pool_runs(pool_id: Any) -> list[tuple[int, int]]:
"""Return contiguous inactive runs inside wholly-free segments, largest first."""
"""Return contiguous inactive runs across all pool segments, largest first;
live and pending-free blocks break runs, so a run never overlaps live data.
"""
runs: list[tuple[int, int]] = []
for segment in torch.cuda.memory_snapshot(pool_id, include_traces=False):
if segment["allocated_size"] != 0:
continue
run_address = 0
run_bytes = 0
for block in segment["blocks"]:
@@ -195,16 +233,18 @@ def borrow_graph_pool(user: str) -> Iterator[None]:
# Return completed cross-stream frees to the cache. The allocator
# processes their events on a later allocation.
torch.empty(1, device="cuda")
if (
_borrow_stub.freed_bytes
or _borrow_stub.cursor_bytes > _borrow_extents_total // 2
):
# Rebuild if empty_cache released the arena's segments or if
# unresolved deferred frees consumed half of the extents.
if _borrow_stub.freed_bytes:
# Rebuild if the caching allocator released an arena segment.
# A high cursor alone means the cache owns reusable segments;
# rebuilding discards them and can turn the next large borrow
# into a fragmented cold-allocation OOM.
_teardown_borrow_pool()
if _borrow_mem_pool is None:
if _borrow_stub is None:
_borrow_stub = BumpArenaStub()
# Runs are sorted largest first, so first fit would carve every
# small allocation out of the run a probability matrix needs.
_borrow_stub.set_best_fit(True)
if _borrow_static_runs is not None:
runs = _borrow_static_runs
else:
@@ -19,7 +19,10 @@ if TYPE_CHECKING:
UpdateWeightsFromIPCReqInput,
)
from sglang.srt.managers.tp_worker import TpModelWorker
from sglang.srt.model_executor.model_runner import ModelRunner
from sglang.srt.model_executor.model_runner import (
ModelRunner,
SamplingPrewarmResult,
)
from sglang.srt.speculative.spec_info import SpeculativeAlgorithm
@@ -215,6 +218,9 @@ class BaseSpecWorker(ABC):
return 0.0
return self.draft_worker.weight_load_time
def prewarm_sampling(self) -> SamplingPrewarmResult:
return self.target_worker.model_runner.prewarm_sampling()
@property
def preloaded_weights_bytes(self) -> int:
if self.draft_worker is None:
+67 -54
View File
@@ -17,6 +17,7 @@ from sglang.srt.layers.sampler import (
top_p_normalize_probs_torch,
)
from sglang.srt.managers.schedule_batch import Req
from sglang.srt.model_executor.runner_utils.pool import borrow_graph_pool
from sglang.srt.runtime_context import get_spec
from sglang.srt.speculative.spec_utils import sample_simulated_acc_len
from sglang.srt.utils import is_cuda, is_hip, is_musa, is_npu
@@ -924,44 +925,22 @@ def compute_dflash_sampling_correct_drafts_and_bonus(
device = next_token_logits.device
if uniform_samples is None:
uniform_samples = torch.rand(
(bs, draft_token_num), dtype=torch.float32, device=device
if uniform_samples is not None and uniform_samples.shape != (bs, draft_token_num):
raise ValueError(
"uniform_samples shape mismatch. "
f"Expected {(bs, draft_token_num)}, got {tuple(uniform_samples.shape)}."
)
else:
if uniform_samples.shape != (bs, draft_token_num):
raise ValueError(
"uniform_samples shape mismatch. "
f"Expected {(bs, draft_token_num)}, got {tuple(uniform_samples.shape)}."
)
uniform_samples = uniform_samples.to(device=device, dtype=torch.float32)
if uniform_samples_for_final_sampling is None:
uniform_samples_for_final_sampling = torch.rand(
(bs,), dtype=torch.float32, device=device
)
else:
if uniform_samples_for_final_sampling.shape != (bs,):
raise ValueError(
"uniform_samples_for_final_sampling shape mismatch. "
f"Expected {(bs,)}, got {tuple(uniform_samples_for_final_sampling.shape)}."
)
uniform_samples_for_final_sampling = uniform_samples_for_final_sampling.to(
device=device,
dtype=torch.float32,
if (
uniform_samples_for_final_sampling is not None
and uniform_samples_for_final_sampling.shape != (bs,)
):
raise ValueError(
"uniform_samples_for_final_sampling shape mismatch. "
f"Expected {(bs,)}, got {tuple(uniform_samples_for_final_sampling.shape)}."
)
target_probs = build_dflash_verify_target_probs(
next_token_logits=next_token_logits,
sampling_info=sampling_info,
draft_token_num=draft_token_num,
bs=bs,
max_top_k=max_top_k,
uniform_top_k_value=uniform_top_k_value,
use_sparse_topk=use_sparse_topk,
)
draft_probs = torch.zeros_like(target_probs)
# Cached across steps, and `correct_len` below aliases `accept_token_num`,
# so these must predate the borrow scope the next replay reclaims.
(
retrieve_index,
retrieve_next_token,
@@ -974,25 +953,59 @@ def compute_dflash_sampling_correct_drafts_and_bonus(
draft_token_num=draft_token_num,
device=device,
)
candidates_i64 = (
candidates if candidates.dtype == torch.int64 else candidates.to(torch.int64)
)
tree_speculative_sampling_target_only(
predicts=predicts,
accept_index=accept_index,
accept_token_num=accept_token_num,
candidates=candidates_i64,
retrive_index=retrieve_index,
retrive_next_token=retrieve_next_token,
retrive_next_sibling=retrieve_next_sibling,
uniform_samples=uniform_samples,
uniform_samples_for_final_sampling=uniform_samples_for_final_sampling,
target_probs=target_probs,
draft_probs=draft_probs,
threshold_single=threshold_single,
threshold_acc=threshold_acc,
deterministic=True,
)
# The full-vocabulary matrices die with this step, so their bytes may come
# from the graph pool's idle storage. Anything outliving the scope is not.
with borrow_graph_pool(user="DFLASH verify probabilities"):
if uniform_samples is None:
coins = torch.rand(
(bs, draft_token_num), dtype=torch.float32, device=device
)
else:
coins = uniform_samples.to(device=device, dtype=torch.float32)
if uniform_samples_for_final_sampling is None:
coins_for_final_sampling = torch.rand(
(bs,), dtype=torch.float32, device=device
)
else:
coins_for_final_sampling = uniform_samples_for_final_sampling.to(
device=device,
dtype=torch.float32,
)
target_probs = build_dflash_verify_target_probs(
next_token_logits=next_token_logits,
sampling_info=sampling_info,
draft_token_num=draft_token_num,
bs=bs,
max_top_k=max_top_k,
uniform_top_k_value=uniform_top_k_value,
use_sparse_topk=use_sparse_topk,
)
draft_probs = torch.zeros_like(target_probs)
candidates_i64 = (
candidates
if candidates.dtype == torch.int64
else candidates.to(torch.int64)
)
tree_speculative_sampling_target_only(
predicts=predicts,
accept_index=accept_index,
accept_token_num=accept_token_num,
candidates=candidates_i64,
retrive_index=retrieve_index,
retrive_next_token=retrieve_next_token,
retrive_next_sibling=retrieve_next_sibling,
uniform_samples=coins,
uniform_samples_for_final_sampling=coins_for_final_sampling,
target_probs=target_probs,
draft_probs=draft_probs,
threshold_single=threshold_single,
threshold_acc=threshold_acc,
deterministic=True,
)
del target_probs, draft_probs, candidates_i64
del coins, coins_for_final_sampling
correct_len = accept_token_num
row_ids = torch.arange(bs, dtype=torch.long, device=device)
@@ -33,17 +33,24 @@ from sglang.srt.model_executor.forward_batch_info import (
ForwardMode,
compute_position,
)
from sglang.srt.model_executor.model_runner import SamplingPrewarmResult
from sglang.srt.model_executor.runner_utils.pool import (
disable_graph_pool_borrow,
graph_pool_borrow_enabled,
)
from sglang.srt.runtime_context import (
get_exec,
get_schedule,
get_spec,
mamba_track_grid,
)
from sglang.srt.sampling.sampling_batch_info import SamplingBatchInfo
from sglang.srt.server_args import ServerArgs
from sglang.srt.speculative.base_spec_worker import BaseSpecWorker
from sglang.srt.speculative.dflash_info import DFlashVerifyInput
from sglang.srt.speculative.dflash_info_v2 import DFlashDraftInputV2
from sglang.srt.speculative.dflash_utils import (
_get_or_create_chain_verify_buffers,
apply_dflash_simulated_acceptance,
apply_dflash_verify_logits_adjustments,
can_dflash_use_fused_qkv_proj,
@@ -491,6 +498,119 @@ class DFlashWorkerV2(BaseSpecWorker):
capture_decode_cuda_graph=capture_decode_cuda_graph
)
def _prewarm_batch_size(self, block_size: int) -> int:
"""Largest batch the non-greedy verify path can see in one step."""
bs = self.model_runner.max_decode_logits_rows() // block_size
max_running_requests = self.model_runner.max_running_requests
if max_running_requests is not None:
bs = min(bs, int(max_running_requests))
return max(1, bs)
def _measure_sampling_peak(
self,
*,
candidates: torch.Tensor,
target_logits: torch.Tensor,
sampling_info: SamplingBatchInfo,
) -> int:
torch.cuda.synchronize()
torch.cuda.reset_peak_memory_stats()
base_bytes = torch.cuda.memory_stats()["allocated_bytes.all.current"]
compute_dflash_sampling_correct_drafts_and_bonus(
candidates=candidates,
next_token_logits=target_logits,
sampling_info=sampling_info,
max_top_k=None,
uniform_top_k_value=None,
)
torch.cuda.synchronize()
return torch.cuda.memory_stats()["allocated_bytes.all.peak"] - base_bytes
def prewarm_sampling(self) -> SamplingPrewarmResult:
"""Rehearse the non-greedy verify path so an undersized graph-pool
borrow surfaces at startup rather than under traffic."""
if not is_cuda() or not is_dflash_sampling_verify_available():
return self.model_runner.prewarm_sampling()
block_size = int(self.block_size)
vocab_size = self.model_config.vocab_size
bs = self._prewarm_batch_size(block_size)
device = self.device
sampling_info = SamplingBatchInfo(
temperatures=torch.ones((bs, 1), dtype=torch.float32, device=device),
top_ps=torch.full((bs,), 0.95, dtype=torch.float32, device=device),
top_ks=torch.zeros(bs, dtype=torch.int32, device=device),
min_ps=torch.zeros(bs, dtype=torch.float32, device=device),
is_all_greedy=False,
is_any_greedy=False,
need_top_p_sampling=True,
# Top-k off: the sparse path is cheaper, so the dense full-vocab
# renormalization is the peak this rehearsal has to reproduce.
need_top_k_sampling=False,
need_min_p_sampling=False,
vocab_size=vocab_size,
)
candidates = torch.zeros((bs, block_size), dtype=torch.int64, device=device)
target_logits = torch.zeros(
(bs * block_size, vocab_size), dtype=torch.float32, device=device
)
# Size the process-lifetime chain buffers before anything borrows. The
# cache is keyed on device index, so pass the resolved device, not "cuda".
_get_or_create_chain_verify_buffers(
bs=bs, draft_token_num=block_size, device=target_logits.device
)
# One unborrowed pass, so any workspace a kernel caches for the process
# lifetime lands outside graph storage and off the measured peak.
with envs.SGLANG_ENABLE_GRAPH_POOL_BORROW.override(False):
self._measure_sampling_peak(
candidates=candidates,
target_logits=target_logits,
sampling_info=sampling_info,
)
borrowing = graph_pool_borrow_enabled()
try:
peak_bytes = self._measure_sampling_peak(
candidates=candidates,
target_logits=target_logits,
sampling_info=sampling_info,
)
except torch.OutOfMemoryError:
if not borrowing:
raise
reason = (
"DFLASH sampling rehearsal exhausted graph-pool memory for the "
f"{bs}x{block_size}x{vocab_size} verify probability matrices"
)
logger.warning(
"Graph pool %s; disabling borrowing and reserving the measured "
"headroom in the post-capture KV sizing instead",
reason,
)
disable_graph_pool_borrow(reason)
peak_bytes = self._measure_sampling_peak(
candidates=candidates,
target_logits=target_logits,
sampling_info=sampling_info,
)
sampling_input_bytes = target_logits.numel() * target_logits.element_size()
result = SamplingPrewarmResult(
sampling_input_bytes=sampling_input_bytes,
sampling_headroom_bytes=peak_bytes + sampling_input_bytes,
)
self.model_runner.sampling_prewarm_result = result
logger.info(
"DFLASH sampling pre-warm: borrowing=%s bs=%d peak=%d B headroom=%d B",
graph_pool_borrow_enabled(),
bs,
peak_bytes,
result.sampling_headroom_bytes,
)
return result
def _maybe_build_draft_sampler(self):
def _eager(reason):
if self.ps.tp_rank == 0:
+69 -78
View File
@@ -19,7 +19,6 @@ from sglang.srt.mem_cache.allocation_sizing import (
get_alloc_reserve_per_decode,
page_aligned_decode_alloc_lens,
)
from sglang.srt.model_executor.runner_utils.pool import borrow_graph_pool
from sglang.srt.runtime_context import get_parallel, get_spec
from sglang.srt.utils import (
is_cpu,
@@ -789,86 +788,78 @@ def eagle_sample(
else tree_speculative_sampling_target_only
)
# These full-vocabulary matrices are consumed by the sampling kernel
# within this step. Returned tensors were allocated before the scope,
# so the next CUDA graph replay may safely reclaim these borrowed bytes.
with borrow_graph_pool(user="EAGLE probability borrow"):
expanded_temperature = torch.repeat_interleave(
sampling_info.temperatures, verify_input.draft_token_num, dim=0
) # (bs * num_draft_tokens, 1)
expanded_temperature = torch.repeat_interleave(
sampling_info.temperatures, verify_input.draft_token_num, dim=0
) # (bs * num_draft_tokens, 1)
target_probs = F.softmax(
next_token_logits / expanded_temperature, dim=-1
) # (bs * num_draft_tokens, vocab_size)
maybe_detect_nan(target_probs, "v2 verify: target_probs after softmax")
if sampling_info.need_top_k_sampling:
target_probs = top_k_renorm_prob(
target_probs,
torch.repeat_interleave(
sampling_info.top_ks, verify_input.draft_token_num, dim=0
),
) # (bs * num_draft_tokens, vocab_size)
maybe_detect_nan(
target_probs, "v2 verify: target_probs after top_k_renorm"
)
if sampling_info.need_top_p_sampling:
target_probs = top_p_renorm_prob(
target_probs,
torch.repeat_interleave(
sampling_info.top_ps, verify_input.draft_token_num, dim=0
),
)
maybe_detect_nan(
target_probs, "v2 verify: target_probs after top_p_renorm"
)
target_probs = target_probs.reshape(bs, verify_input.draft_token_num, -1)
draft_probs = (
verify_input.draft_probs
if use_rejection_sampling
else torch.zeros_like(target_probs)
)
# Defense-in-depth behind the spec_hook startup allowlist: validate
# the actual kernel inputs before the Triton kernel.
if use_rejection_sampling and (
draft_probs is None or draft_probs.shape[-1] != target_probs.shape[-1]
):
raise ValueError(
"Rejection sampling requires a target-vocab draft proposal "
"distribution; the current speculative algorithm/draft worker "
"does not produce one (draft_probs missing or vocab-mismatched)."
)
coins, coins_for_final_sampling = _verify_coins(
sampling_info=sampling_info,
seq_lens=batch.seq_lens,
draft_token_num=verify_input.draft_token_num,
candidates=candidates,
device=device,
)
sampling_fn(
predicts=predict, # mutable
accept_index=accept_index, # mutable
accept_token_num=num_correct_drafts, # mutable
candidates=candidates,
# kwarg LHS retained as `retrive_*` to match sgl_kernel op schema.
retrive_index=verify_input.retrieve_index,
retrive_next_token=verify_input.retrieve_next_token,
retrive_next_sibling=verify_input.retrieve_next_sibling,
uniform_samples=coins,
uniform_samples_for_final_sampling=coins_for_final_sampling,
target_probs=target_probs,
draft_probs=draft_probs,
threshold_single=get_spec().speculative_accept_threshold_single,
threshold_acc=get_spec().speculative_accept_threshold_acc,
deterministic=True,
)
del (
expanded_temperature,
target_probs = F.softmax(
next_token_logits / expanded_temperature, dim=-1
) # (bs * num_draft_tokens, vocab_size)
maybe_detect_nan(target_probs, "v2 verify: target_probs after softmax")
if sampling_info.need_top_k_sampling:
target_probs = top_k_renorm_prob(
target_probs,
draft_probs,
coins,
coins_for_final_sampling,
torch.repeat_interleave(
sampling_info.top_ks, verify_input.draft_token_num, dim=0
),
) # (bs * num_draft_tokens, vocab_size)
maybe_detect_nan(target_probs, "v2 verify: target_probs after top_k_renorm")
if sampling_info.need_top_p_sampling:
target_probs = top_p_renorm_prob(
target_probs,
torch.repeat_interleave(
sampling_info.top_ps, verify_input.draft_token_num, dim=0
),
)
maybe_detect_nan(target_probs, "v2 verify: target_probs after top_p_renorm")
target_probs = target_probs.reshape(bs, verify_input.draft_token_num, -1)
draft_probs = (
verify_input.draft_probs
if use_rejection_sampling
else torch.zeros_like(target_probs)
)
# Defense-in-depth behind the spec_hook startup allowlist: validate
# the actual kernel inputs before the Triton kernel.
if use_rejection_sampling and (
draft_probs is None or draft_probs.shape[-1] != target_probs.shape[-1]
):
raise ValueError(
"Rejection sampling requires a target-vocab draft proposal "
"distribution; the current speculative algorithm/draft worker "
"does not produce one (draft_probs missing or vocab-mismatched)."
)
coins, coins_for_final_sampling = _verify_coins(
sampling_info=sampling_info,
seq_lens=batch.seq_lens,
draft_token_num=verify_input.draft_token_num,
candidates=candidates,
device=device,
)
sampling_fn(
predicts=predict, # mutable
accept_index=accept_index, # mutable
accept_token_num=num_correct_drafts, # mutable
candidates=candidates,
# kwarg LHS retained as `retrive_*` to match sgl_kernel op schema.
retrive_index=verify_input.retrieve_index,
retrive_next_token=verify_input.retrieve_next_token,
retrive_next_sibling=verify_input.retrieve_next_sibling,
uniform_samples=coins,
uniform_samples_for_final_sampling=coins_for_final_sampling,
target_probs=target_probs,
draft_probs=draft_probs,
threshold_single=get_spec().speculative_accept_threshold_single,
threshold_acc=get_spec().speculative_accept_threshold_acc,
deterministic=True,
)
del (
expanded_temperature,
target_probs,
draft_probs,
coins,
coins_for_final_sampling,
)
# Sync sampling results across TP ranks: different GPUs may
# produce slightly different target_probs due to floating-point
+23 -8
View File
@@ -358,6 +358,7 @@ static size_t g_cursors[BUMPARENA_MAX_EXTENTS];
static size_t g_num_extents = 0;
static size_t g_freed_bytes = 0;
static size_t g_align = 512;
static int g_best_fit = 0;
static std::mutex g_mu;
static size_t align_up(size_t v, size_t a){{ return (v + a - 1) / a * a; }}
void bumparena_set_extents_{sfx}(const uintptr_t* bases, const size_t* sizes, size_t n){{
@@ -372,6 +373,7 @@ void bumparena_set_extents_{sfx}(const uintptr_t* bases, const size_t* sizes, si
}}
}}
void bumparena_set_align_{sfx}(size_t a){{ std::lock_guard<std::mutex> lk(g_mu); if (a) g_align=a; }}
void bumparena_set_best_fit_{sfx}(int on){{ std::lock_guard<std::mutex> lk(g_mu); g_best_fit = on; }}
size_t bumparena_cursor_{sfx}(void){{
std::lock_guard<std::mutex> lk(g_mu);
size_t total = 0;
@@ -380,14 +382,21 @@ size_t bumparena_cursor_{sfx}(void){{
}}
void* bumparena_malloc_{sfx}(size_t size, int device, void* stream){{
std::lock_guard<std::mutex> lk(g_mu);
size_t need = align_up(size, g_align);
size_t pick = BUMPARENA_MAX_EXTENTS;
for (size_t i = 0; i < g_num_extents; ++i) {{
size_t need = g_cursors[i] + align_up(size, g_align);
if (need > g_reserved[i]) continue;
void* p = reinterpret_cast<void*>(g_bases[i] + g_cursors[i]);
g_cursors[i] = need;
return p;
size_t avail = g_reserved[i] - g_cursors[i];
if (avail < need) continue;
// First fit is for callers that map physical pages at the offsets they
// are handed, so extent order is theirs to choose.
if (!g_best_fit) {{ pick = i; break; }}
if (pick == BUMPARENA_MAX_EXTENTS ||
avail < g_reserved[pick] - g_cursors[pick]) pick = i;
}}
return 0; // no extent fits -- surfaces as an allocator OOM
if (pick == BUMPARENA_MAX_EXTENTS) return 0; // no extent fits -- surfaces as an allocator OOM
void* p = reinterpret_cast<void*>(g_bases[pick] + g_cursors[pick]);
g_cursors[pick] += need;
return p;
}}
size_t bumparena_freed_{sfx}(void){{ std::lock_guard<std::mutex> lk(g_mu); return g_freed_bytes; }}
void bumparena_free_{sfx}(void* ptr, size_t size, int device, void* stream){{
@@ -401,8 +410,8 @@ void bumparena_free_{sfx}(void* ptr, size_t size, int device, void* stream){{
class BumpArenaStub:
"""JIT-built pluggable bump allocator over caller-provided device VA extents.
``malloc`` first-fits an extent and hands out ``base + cursor``; ``free``
is a no-op. Plain ``torch.empty`` can thus be placed on externally managed
``malloc`` picks an extent and hands out ``base + cursor``; ``free`` is a
no-op. Plain ``torch.empty`` can thus be placed on externally managed
storage by wrapping ``allocator`` in a ``torch.cuda.MemPool``.
``set_extents`` re-points the arena and resets every cursor, letting one
stub serve successive region sets.
@@ -462,6 +471,9 @@ class BumpArenaStub:
self._fn_set_align = lib[f"bumparena_set_align_{self.sfx}"]
self._fn_set_align.argtypes = [ctypes.c_size_t]
self._fn_set_align.restype = None
self._fn_set_best_fit = lib[f"bumparena_set_best_fit_{self.sfx}"]
self._fn_set_best_fit.argtypes = [ctypes.c_int]
self._fn_set_best_fit.restype = None
self._fn_cursor = lib[f"bumparena_cursor_{self.sfx}"]
self._fn_cursor.argtypes = []
self._fn_cursor.restype = ctypes.c_size_t
@@ -486,6 +498,9 @@ class BumpArenaStub:
def set_align(self, nbytes: int) -> None:
self._fn_set_align(ctypes.c_size_t(nbytes))
def set_best_fit(self, on: bool) -> None:
self._fn_set_best_fit(ctypes.c_int(1 if on else 0))
@property
def cursor_bytes(self) -> int:
return int(self._fn_cursor())
@@ -194,6 +194,10 @@ class GSM8KMixin:
gsm8k_thinking: bool = False # sgl_eval backend
gsm8k_max_tokens: Optional[int] = None # sgl_eval backend
gsm8k_n_repeats: int = 1 # sgl_eval backend
# None keeps run_eval's greedy default; set both to route the run through
# the sampling path.
gsm8k_temperature: Optional[float] = None
gsm8k_top_p: Optional[float] = None
def test_gsm8k(self):
requests.get(self.base_url + "/flush_cache")
@@ -228,6 +232,8 @@ class GSM8KMixin:
api="completion",
max_tokens=512,
num_shots=self.gsm8k_num_shots,
temperature=self.gsm8k_temperature,
top_p=self.gsm8k_top_p,
)
+35 -7
View File
@@ -1,4 +1,5 @@
import unittest
from contextlib import ExitStack
import openai
@@ -22,7 +23,7 @@ from sglang.test.test_utils import (
popen_launch_server,
)
register_cuda_ci(est_time=420, stage="base-b", runner_config="1-gpu-small")
register_cuda_ci(est_time=500, stage="base-b", runner_config="1-gpu-small")
register_amd_ci(est_time=420, stage="stage-b", runner_config="1-gpu-small-amd")
@@ -45,6 +46,8 @@ class TestDFlashServerBase(
draft_model = DEFAULT_DRAFT_MODEL_DFLASH
gsm8k_accuracy_thres = 0.75
gsm8k_accept_length_thres = 2.8
# (env, value) pairs applied around the server launch.
extra_env_overrides: tuple = ()
@classmethod
def setUpClass(cls):
@@ -71,12 +74,15 @@ class TestDFlashServerBase(
if cls.disable_overlap:
launch_args.append("--disable-overlap-schedule")
launch_args.extend(cls.other_launch_args)
with (
envs.SGLANG_ENABLE_OVERLAP_PLAN_STREAM.override(cls.overlap_plan_stream),
envs.SGLANG_ENABLE_STRICT_MEM_CHECK_DURING_BUSY.override(1),
envs.SGLANG_ENABLE_ASYNC_ASSERT.override(True),
envs.SGLANG_ALLOW_OVERWRITE_LONGER_CONTEXT_LEN.override(True),
):
with ExitStack() as stack:
for env, value in (
(envs.SGLANG_ENABLE_OVERLAP_PLAN_STREAM, cls.overlap_plan_stream),
(envs.SGLANG_ENABLE_STRICT_MEM_CHECK_DURING_BUSY, 1),
(envs.SGLANG_ENABLE_ASYNC_ASSERT, True),
(envs.SGLANG_ALLOW_OVERWRITE_LONGER_CONTEXT_LEN, True),
*cls.extra_env_overrides,
):
stack.enter_context(env.override(value))
cls.process = popen_launch_server(
cls.model,
cls.base_url,
@@ -173,5 +179,27 @@ class TestDFlashServerOverlapPlanStream(TestDFlashServerOverlap):
overlap_plan_stream = True
@unittest.skipIf(
is_hip(),
"borrowing is CUDA-only and ROCm has no DFLASH sampling-verify kernel",
)
class TestDFlashServerGraphPoolBorrow(TestDFlashServerBase):
"""Verify probabilities served out of idle CUDA graph storage. Only the
non-greedy accept path borrows, hence the sampled GSM8K below."""
extra_env_overrides = (
(envs.SGLANG_ENABLE_GRAPH_POOL_BORROW, 1),
(envs.SGLANG_ENABLE_GRAPH_POOL_PRECARVE, 1),
)
disable_overlap = False
gsm8k_temperature = 0.6
gsm8k_top_p = 0.95
# Measured over 4x200 examples per arm: score 0.7525 (sd 0.021) and accept
# length 2.80, borrowing on or off. Corruption collapses accept length
# toward 1.0; the accuracy floor is loose to absorb other hardware.
gsm8k_accuracy_thres = 0.60
gsm8k_accept_length_thres = 2.0
if __name__ == "__main__":
unittest.main()
@@ -674,6 +674,9 @@ class _SchedulerWorker:
def __init__(self, trace, *, post_capture_active=False):
self._trace = trace
self.model_runner = SimpleNamespace(
device="cuda",
forward_stream=object(),
prewarm_sampling=lambda: trace.append("prewarm"),
token_to_kv_pool=SimpleNamespace(post_capture_active=post_capture_active),
post_capture_resize_kv_pool=lambda: trace.append("resize"),
)
@@ -687,7 +690,7 @@ class _SchedulerWorker:
class TestStartupWeightLoadSchedulerRouting(CustomTestCase):
@staticmethod
def _scheduler(worker, trace, *, mode):
def _scheduler(worker, trace, *, mode, draft_worker=None):
from sglang.srt.managers.scheduler import Scheduler
scheduler = Scheduler.__new__(Scheduler)
@@ -696,17 +699,38 @@ class TestStartupWeightLoadSchedulerRouting(CustomTestCase):
)
scheduler.init_tp_model_worker = lambda: setattr(scheduler, "tp_worker", worker)
scheduler.maybe_init_draft_worker = lambda: setattr(
scheduler, "draft_worker", None
scheduler, "draft_worker", draft_worker
)
scheduler.init_memory_pools = lambda: trace.append("memory_pool")
scheduler.init_all_attention_backends = lambda: trace.append("attention")
scheduler.init_all_cuda_graphs = lambda: trace.append("capture")
return scheduler
def _run_startup(self, mode):
def _run_startup(self, mode, *, use_draft_worker=False):
trace = []
worker = _SchedulerWorker(trace, post_capture_active=True)
scheduler = self._scheduler(worker, trace, mode=mode)
draft_worker = (
SimpleNamespace(prewarm_sampling=lambda: trace.append("draft_prewarm"))
if use_draft_worker
else None
)
scheduler = self._scheduler(
worker,
trace,
mode=mode,
draft_worker=draft_worker,
)
class StreamContext:
def __enter__(self):
trace.append("stream_enter")
def __exit__(self, *_args):
trace.append("stream_exit")
def stream_context(stream):
self.assertIs(stream, worker.model_runner.forward_stream)
return StreamContext()
def stop_after_startup():
raise RuntimeError("stop after startup")
@@ -723,6 +747,10 @@ class TestStartupWeightLoadSchedulerRouting(CustomTestCase):
)
),
),
patch(
"sglang.srt.managers.scheduler.torch.get_device_module",
return_value=SimpleNamespace(stream=stream_context),
),
self.assertRaisesRegex(RuntimeError, "stop after startup"),
):
scheduler.init_model_worker()
@@ -732,13 +760,45 @@ class TestStartupWeightLoadSchedulerRouting(CustomTestCase):
def test_serial_path_skips_overlap_hooks(self):
self.assertEqual(
self._run_startup("serial"),
["memory_pool", "attention", "capture", "resize"],
[
"memory_pool",
"attention",
"capture",
"stream_enter",
"prewarm",
"stream_exit",
"resize",
],
)
def test_overlap_starts_before_capture_and_finalizes_after(self):
self.assertEqual(
self._run_startup("overlap"),
["start", "memory_pool", "attention", "capture", "resize", "finalize"],
[
"start",
"memory_pool",
"attention",
"capture",
"stream_enter",
"prewarm",
"stream_exit",
"resize",
"finalize",
],
)
def test_draft_worker_prewarm_uses_target_forward_stream(self):
self.assertEqual(
self._run_startup("serial", use_draft_worker=True),
[
"memory_pool",
"attention",
"capture",
"stream_enter",
"draft_prewarm",
"stream_exit",
"resize",
],
)
@@ -56,6 +56,9 @@ def _make_backend(runner):
backend._outputs = {}
backend._pool = None
backend._capture_stream = None
backend._precarve = SimpleNamespace(
measure=contextlib.nullcontext, mint=mock.Mock()
)
backend._memory_saver_adapter = None
backend._cuda_graph_runner = runner
backend._device_module = runner.device_module
@@ -12,7 +12,8 @@ from sglang.srt.model_executor.runner_backend.full_cuda_graph_backend import (
FullCudaGraphBackend,
)
from sglang.srt.model_executor.runner_utils import pool
from sglang.srt.speculative import eagle_utils
from sglang.srt.speculative import dflash_utils, dflash_worker_v2, eagle_utils
from sglang.srt.speculative.dflash_worker_v2 import DFlashWorkerV2
from sglang.test.ci.ci_register import register_cuda_ci
from sglang.test.test_utils import CustomTestCase
@@ -43,6 +44,25 @@ class TestGraphPoolBorrow(CustomTestCase):
pool._borrow_extents_total = 0
pool._largest_logged_graph_pool_borrow = 0
def test_mixed_segment_runs_exclude_live_blocks(self):
"""A mixed segment's free runs are borrowable, but a returned run must
never overlap the live block overlap would silently corrupt
graph-owned data instead of raising an OOM."""
snapshot = [
{
"allocated_size": 4096,
"total_size": 3 * 4096,
"blocks": [
{"state": "inactive", "address": 0x1000, "size": 4096},
{"state": "active_allocated", "address": 0x2000, "size": 4096},
{"state": "inactive", "address": 0x3000, "size": 4096},
],
}
]
with patch.object(pool.torch.cuda, "memory_snapshot", return_value=snapshot):
runs = pool.find_free_graph_pool_runs((0, 1))
self.assertEqual(sorted(runs), [(0x1000, 4096), (0x3000, 4096)])
def test_graph_replay_fails_during_active_pool_borrow(self):
graph = Mock()
backend = object.__new__(FullCudaGraphBackend)
@@ -74,6 +94,24 @@ class TestGraphPoolBorrow(CustomTestCase):
graph.replay.assert_not_called()
def test_high_cursor_keeps_reusable_cached_segments(self):
stub = MagicMock(cursor_bytes=600, freed_bytes=0)
mem_pool = MagicMock()
with (
patch.object(pool, "graph_pool_borrow_enabled", return_value=True),
patch.object(pool, "_borrow_stub", stub),
patch.object(pool, "_borrow_mem_pool", mem_pool),
patch.object(pool, "_borrow_extents_total", 1000),
patch.object(pool, "_teardown_borrow_pool") as teardown,
patch.object(pool.torch, "empty"),
patch.object(pool.torch.cuda, "use_mem_pool"),
):
with pool.borrow_graph_pool(user="test"):
pass
teardown.assert_not_called()
def test_external_graph_storage_can_disable_borrowing(self):
with (
envs.SGLANG_ENABLE_GRAPH_POOL_BORROW.override(True),
@@ -84,29 +122,8 @@ class TestGraphPoolBorrow(CustomTestCase):
pool.disable_graph_pool_borrow("graph storage is externally managed")
self.assertFalse(pool.graph_pool_borrow_enabled())
def test_eagle_non_greedy_probabilities_use_borrow_scope(self):
state = {"active": False, "users": []}
@contextmanager
def tracking_borrow(*, user):
self.assertFalse(state["active"])
state["active"] = True
state["users"].append(user)
try:
yield
finally:
state["active"] = False
import torch.nn.functional as F
real_softmax = F.softmax
def checked_softmax(*args, **kwargs):
self.assertTrue(state["active"])
return real_softmax(*args, **kwargs)
def test_eagle_non_greedy_probabilities_do_not_borrow_graph_pool(self):
def fake_sampling(**kwargs):
self.assertTrue(state["active"])
kwargs["predicts"].fill_(3)
kwargs["accept_index"].fill_(0)
kwargs["accept_token_num"].fill_(1)
@@ -146,9 +163,8 @@ class TestGraphPoolBorrow(CustomTestCase):
tp_group = SimpleNamespace(world_size=1)
with (
patch.object(eagle_utils, "borrow_graph_pool", tracking_borrow),
patch.object(pool, "borrow_graph_pool") as borrow_graph_pool,
patch.object(eagle_utils, "get_spec", return_value=spec_config),
patch("torch.nn.functional.softmax", side_effect=checked_softmax),
patch(
"sglang.srt.layers.dp_attention.is_dp_attention_enabled",
return_value=False,
@@ -163,8 +179,7 @@ class TestGraphPoolBorrow(CustomTestCase):
verify_input, batch, logits_output
)
self.assertEqual(state["users"], ["EAGLE probability borrow"])
self.assertFalse(state["active"])
borrow_graph_pool.assert_not_called()
self.assertTrue(torch.equal(predict, torch.full_like(predict, 3)))
self.assertTrue(torch.equal(accept_lens, torch.full_like(accept_lens, 2)))
self.assertTrue(torch.equal(accept_index, torch.zeros_like(accept_index)))
@@ -176,8 +191,9 @@ class TestGraphPoolBorrow(CustomTestCase):
graph = torch.cuda.CUDAGraph()
x = torch.zeros(8, device="cuda")
stream = torch.cuda.Stream()
with torch.cuda.stream(stream), torch.cuda.graph(
graph, pool=handle, stream=stream
with (
torch.cuda.stream(stream),
torch.cuda.graph(graph, pool=handle, stream=stream),
):
# Two capture-only transients become disjoint free graph-pool runs.
transient_a = torch.empty(48 << 20, dtype=torch.uint8, device="cuda")
@@ -237,8 +253,9 @@ class TestGraphPoolBorrow(CustomTestCase):
graph = torch.cuda.CUDAGraph()
x = torch.zeros(8, device="cuda")
stream = torch.cuda.Stream()
with torch.cuda.stream(stream), torch.cuda.graph(
graph, pool=handle, stream=stream
with (
torch.cuda.stream(stream),
torch.cuda.graph(graph, pool=handle, stream=stream),
):
transient = torch.empty(64 << 20, dtype=torch.uint8, device="cuda")
y = x + 1
@@ -267,6 +284,40 @@ class TestGraphPoolBorrow(CustomTestCase):
del graph, y
@unittest.skipUnless(torch.cuda.is_available(), "requires CUDA")
def test_oversized_borrow_raises_oom_then_regular_allocation_succeeds(self):
handle = torch.cuda.graph_pool_handle()
graph = torch.cuda.CUDAGraph()
x = torch.zeros(8, device="cuda")
stream = torch.cuda.Stream()
with (
torch.cuda.stream(stream),
torch.cuda.graph(graph, pool=handle, stream=stream),
):
transient = torch.empty(64 << 20, dtype=torch.uint8, device="cuda")
y = x + 1
del transient
torch.cuda.synchronize()
runs = pool.find_free_graph_pool_runs(handle)
address, run_bytes = next(run for run in runs if run[1] >= 16 << 20)
with (
envs.SGLANG_ENABLE_GRAPH_POOL_BORROW.override(True),
patch.object(pool, "get_global_graph_memory_pool", return_value=None),
):
pool.set_graph_pool_borrow_runs([(address, 8 << 20)])
with self.assertRaises(torch.OutOfMemoryError):
with pool.borrow_graph_pool(user="undersized-test"):
torch.empty(16 << 20, dtype=torch.uint8, device="cuda")
pool.disable_graph_pool_borrow("undersized test pool")
regular = torch.empty(16 << 20, dtype=torch.uint8, device="cuda")
self.assertEqual(regular.nbytes, 16 << 20)
del regular
self.assertGreaterEqual(run_bytes, 16 << 20)
del graph, y
@unittest.skipUnless(torch.cuda.is_available(), "requires CUDA")
def test_cross_stream_borrow_frees_resolve_before_pointer_reuse(self):
"""Deferred record_stream frees must not collide on the next borrow."""
@@ -274,8 +325,9 @@ class TestGraphPoolBorrow(CustomTestCase):
graph = torch.cuda.CUDAGraph()
x = torch.zeros(8, device="cuda")
stream = torch.cuda.Stream()
with torch.cuda.stream(stream), torch.cuda.graph(
graph, pool=handle, stream=stream
with (
torch.cuda.stream(stream),
torch.cuda.graph(graph, pool=handle, stream=stream),
):
transient = torch.empty(128 << 20, dtype=torch.uint8, device="cuda")
y = x + 1
@@ -301,6 +353,115 @@ class TestGraphPoolBorrow(CustomTestCase):
del graph, y
def test_dflash_verify_output_buffers_predate_the_borrow_scope(self):
"""The chain verify buffers outlive the step, so creating them inside
the borrow scope would let the next replay overwrite the accept
length instead of raising."""
events = []
@contextmanager
def recording_borrow(user):
events.append(f"borrow:{user}")
yield
events.append("release")
real_buffers = dflash_utils._get_or_create_chain_verify_buffers
def recording_buffers(**kwargs):
events.append("buffers")
return real_buffers(**kwargs)
def fake_sampling(**kwargs):
kwargs["predicts"].fill_(3)
kwargs["accept_index"].fill_(0)
kwargs["accept_token_num"].fill_(1)
sampling_info = SimpleNamespace(
temperatures=torch.ones((1, 1)),
top_ks=torch.ones(1, dtype=torch.int32),
top_ps=torch.ones(1),
need_top_k_sampling=False,
need_top_p_sampling=False,
)
with (
patch.object(dflash_utils, "borrow_graph_pool", recording_borrow),
patch.object(
dflash_utils,
"_get_or_create_chain_verify_buffers",
recording_buffers,
),
patch.object(dflash_utils, "_DFLASH_SAMPLING_VERIFY_AVAILABLE", True),
patch.object(
dflash_utils,
"tree_speculative_sampling_target_only",
fake_sampling,
),
):
correct_len, bonus = (
dflash_utils.compute_dflash_sampling_correct_drafts_and_bonus(
candidates=torch.zeros((1, 2), dtype=torch.int64),
next_token_logits=torch.randn((2, 8)),
sampling_info=sampling_info,
threshold_single=1.0,
threshold_acc=1.0,
)
)
self.assertEqual(
events, ["buffers", "borrow:DFLASH verify probabilities", "release"]
)
self.assertTrue(torch.equal(correct_len, torch.ones_like(correct_len)))
self.assertTrue(torch.equal(bonus, torch.full_like(bonus, 3)))
@unittest.skipUnless(torch.cuda.is_available(), "requires CUDA")
def test_dflash_prewarm_falls_back_when_the_rehearsal_exhausts_the_pool(self):
"""A rehearsal too large for the pool must retire borrowing and
re-measure, rather than crash startup or leave KV sizing without the
headroom it now has to reserve."""
worker = object.__new__(DFlashWorkerV2)
worker.block_size = 4
worker.device = "cuda"
worker._target_worker = SimpleNamespace(
model_runner=SimpleNamespace(
max_running_requests=2,
max_decode_logits_rows=lambda: 8,
sampling_prewarm_result=None,
),
model_config=SimpleNamespace(vocab_size=32),
)
worker.model_runner = worker._target_worker.model_runner
calls = []
def rehearse(**kwargs):
calls.append(pool.graph_pool_borrow_enabled())
if len(calls) == 2:
raise torch.OutOfMemoryError("rehearsal too large")
return torch.zeros(2), torch.zeros(2)
with (
envs.SGLANG_ENABLE_GRAPH_POOL_BORROW.override(True),
patch.object(pool, "get_global_graph_memory_pool", return_value=(1, 2)),
patch.object(
dflash_worker_v2,
"compute_dflash_sampling_correct_drafts_and_bonus",
rehearse,
),
):
self.assertTrue(pool.graph_pool_borrow_enabled())
result = worker.prewarm_sampling()
self.assertFalse(pool.graph_pool_borrow_enabled())
# Warm pass outside the pool, borrowed pass that OOMs, retry after the
# fallback retires borrowing.
self.assertEqual(calls, [False, True, False])
# 2 rows x 4 draft tokens x 32 vocab x 4 bytes.
self.assertEqual(result.sampling_input_bytes, 2 * 4 * 32 * 4)
self.assertGreaterEqual(
result.sampling_headroom_bytes, result.sampling_input_bytes
)
self.assertIs(worker.model_runner.sampling_prewarm_result, result)
if __name__ == "__main__":
unittest.main()