[Logprob] Borrow graph-pool memory for input logprob logits construction (#40007)
Co-authored-by: cctry <csycfl@gmail.com>
This commit is contained in:
@@ -26,7 +26,7 @@ from sglang.kernels.ops.activation.softcap import (
|
|||||||
softcap_inplace_logits as fused_softcap,
|
softcap_inplace_logits as fused_softcap,
|
||||||
)
|
)
|
||||||
from sglang.srt.beam_search.logits_capture import BeamLogitsCapture
|
from sglang.srt.beam_search.logits_capture import BeamLogitsCapture
|
||||||
from sglang.srt.distributed import get_tp_group
|
from sglang.srt.distributed import get_attn_tp_group, get_tp_group
|
||||||
from sglang.srt.distributed.device_communicators import triton_symm_mem_ag
|
from sglang.srt.distributed.device_communicators import triton_symm_mem_ag
|
||||||
from sglang.srt.environ import envs
|
from sglang.srt.environ import envs
|
||||||
from sglang.srt.layers import layernorm_sp
|
from sglang.srt.layers import layernorm_sp
|
||||||
@@ -458,7 +458,16 @@ class LogitsProcessor(nn.Module):
|
|||||||
skip_entry_sync=True,
|
skip_entry_sync=True,
|
||||||
)
|
)
|
||||||
|
|
||||||
self.input_logprob_processor = InputLogprobProcessor()
|
chunking_group = None
|
||||||
|
if (
|
||||||
|
self.do_tensor_parallel_all_gather
|
||||||
|
and not self.do_tensor_parallel_all_gather_dp_attn
|
||||||
|
):
|
||||||
|
group = get_attn_tp_group() if self.use_attn_tp_group else get_tp_group()
|
||||||
|
chunking_group = group.cpu_group
|
||||||
|
self.input_logprob_processor = InputLogprobProcessor(
|
||||||
|
self.vocab_size, chunking_group=chunking_group
|
||||||
|
)
|
||||||
|
|
||||||
def forward(
|
def forward(
|
||||||
self,
|
self,
|
||||||
|
|||||||
@@ -2,12 +2,20 @@ from __future__ import annotations
|
|||||||
|
|
||||||
import dataclasses
|
import dataclasses
|
||||||
import logging
|
import logging
|
||||||
|
from contextlib import nullcontext
|
||||||
from enum import Enum, auto
|
from enum import Enum, auto
|
||||||
from typing import TYPE_CHECKING, Callable, List, Optional, Tuple
|
from typing import TYPE_CHECKING, Callable, List, Optional, Tuple
|
||||||
|
|
||||||
import torch
|
import torch
|
||||||
|
|
||||||
|
from sglang.srt.distributed.device_communicators.pynccl_allocator import (
|
||||||
|
is_symmetric_memory_enabled,
|
||||||
|
)
|
||||||
from sglang.srt.environ import envs
|
from sglang.srt.environ import envs
|
||||||
|
from sglang.srt.model_executor.runner_utils.pool import (
|
||||||
|
borrow_graph_pool,
|
||||||
|
graph_pool_borrow_largest_run,
|
||||||
|
)
|
||||||
from sglang.srt.runtime_context import get_exec
|
from sglang.srt.runtime_context import get_exec
|
||||||
|
|
||||||
if TYPE_CHECKING:
|
if TYPE_CHECKING:
|
||||||
@@ -434,6 +442,10 @@ def _deterministic_inference_enabled() -> bool:
|
|||||||
return False
|
return False
|
||||||
|
|
||||||
|
|
||||||
|
# Scratch allocations and caching-allocator segment rounding in the borrow scope.
|
||||||
|
_GRAPH_POOL_BORROW_SLACK_BYTES = 64 << 20
|
||||||
|
|
||||||
|
|
||||||
class InputLogprobProcessor:
|
class InputLogprobProcessor:
|
||||||
"""Input (prefill) logprob processing: single-pass or chunked.
|
"""Input (prefill) logprob processing: single-pass or chunked.
|
||||||
|
|
||||||
@@ -442,7 +454,13 @@ class InputLogprobProcessor:
|
|||||||
the lm_head / TP-gather machinery in LogitsProcessor.
|
the lm_head / TP-gather machinery in LogitsProcessor.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
def __init__(self):
|
def __init__(
|
||||||
|
self,
|
||||||
|
vocab_size: int,
|
||||||
|
chunking_group: Optional[torch.distributed.ProcessGroup] = None,
|
||||||
|
):
|
||||||
|
self.vocab_size = vocab_size
|
||||||
|
self.chunking_group = chunking_group
|
||||||
# enable chunked logprobs processing
|
# enable chunked logprobs processing
|
||||||
self.enable_logprobs_chunk = envs.SGLANG_ENABLE_LOGPROB_CHUNK.get()
|
self.enable_logprobs_chunk = envs.SGLANG_ENABLE_LOGPROB_CHUNK.get()
|
||||||
# chunk size for logprobs processing
|
# chunk size for logprobs processing
|
||||||
@@ -479,6 +497,10 @@ class InputLogprobProcessor:
|
|||||||
else:
|
else:
|
||||||
chunk_size = self.logprobs_chunk_size
|
chunk_size = self.logprobs_chunk_size
|
||||||
|
|
||||||
|
borrow_logits_memory = False
|
||||||
|
if pruned_states.is_cuda and not skip_chunking_for_dp_attn:
|
||||||
|
borrow_logits_memory = self._can_borrow_logits_memory(chunk_size)
|
||||||
|
|
||||||
return self._forward_by_chunk(
|
return self._forward_by_chunk(
|
||||||
pruned_states,
|
pruned_states,
|
||||||
sample_indices,
|
sample_indices,
|
||||||
@@ -488,8 +510,37 @@ class InputLogprobProcessor:
|
|||||||
get_logits_fn,
|
get_logits_fn,
|
||||||
logits_metadata,
|
logits_metadata,
|
||||||
chunk_size,
|
chunk_size,
|
||||||
|
borrow_logits_memory=borrow_logits_memory,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
def _can_borrow_logits_memory(self, chunk_size: int) -> bool:
|
||||||
|
"""Borrow only when the planned chunk fits on every TP rank.
|
||||||
|
|
||||||
|
Resizing chunks to graph-pool capacity changes LM-head GEMM shapes
|
||||||
|
and their rounding. Keep chunk boundaries independent of the graph
|
||||||
|
memory layout, including when borrowing is disabled on another runner.
|
||||||
|
"""
|
||||||
|
# TP gathering can hold the local projection, gathered tensor, and
|
||||||
|
# contiguous reshape together; FP32 bounds their possible dtypes.
|
||||||
|
bytes_per_row = 3 * self.vocab_size * 4
|
||||||
|
# NCCL's symmetric allocator owns its collective buffers, so those
|
||||||
|
# allocations cannot be counted as borrowed storage.
|
||||||
|
free_run = (
|
||||||
|
0 if is_symmetric_memory_enabled() else graph_pool_borrow_largest_run()
|
||||||
|
)
|
||||||
|
fit_rows = max(0, free_run - _GRAPH_POOL_BORROW_SLACK_BYTES) // bytes_per_row
|
||||||
|
if self.chunking_group is not None:
|
||||||
|
capacity = torch.tensor(fit_rows, dtype=torch.int64, device="cpu")
|
||||||
|
torch.distributed.all_reduce(
|
||||||
|
capacity,
|
||||||
|
op=torch.distributed.ReduceOp.MIN,
|
||||||
|
group=self.chunking_group,
|
||||||
|
)
|
||||||
|
fit_rows = int(capacity.item())
|
||||||
|
# Fall back to the existing reserved workspace if borrowing cannot
|
||||||
|
# hold the whole chunk; do not change LoRA or logprob chunk boundaries.
|
||||||
|
return fit_rows >= chunk_size
|
||||||
|
|
||||||
def _forward_by_chunk(
|
def _forward_by_chunk(
|
||||||
self,
|
self,
|
||||||
pruned_states: torch.Tensor,
|
pruned_states: torch.Tensor,
|
||||||
@@ -500,6 +551,7 @@ class InputLogprobProcessor:
|
|||||||
get_logits_fn: Callable,
|
get_logits_fn: Callable,
|
||||||
logits_metadata: LogitsMetadata,
|
logits_metadata: LogitsMetadata,
|
||||||
chunk_size: int,
|
chunk_size: int,
|
||||||
|
borrow_logits_memory: bool = False,
|
||||||
) -> Tuple[LogprobResult, torch.Tensor]:
|
) -> Tuple[LogprobResult, torch.Tensor]:
|
||||||
"""Compute input logprobs chunk by chunk to cap peak memory."""
|
"""Compute input logprobs chunk by chunk to cap peak memory."""
|
||||||
total_size = pruned_states.shape[0]
|
total_size = pruned_states.shape[0]
|
||||||
@@ -557,14 +609,20 @@ class InputLogprobProcessor:
|
|||||||
# writing through the shared graph logits buffer would alias
|
# writing through the shared graph logits buffer would alias
|
||||||
# chunks whose shape happens to match the buffer.
|
# chunks whose shape happens to match the buffer.
|
||||||
chunk_states = pruned_states[start_idx:end_idx]
|
chunk_states = pruned_states[start_idx:end_idx]
|
||||||
chunk_logits = get_logits_fn(
|
with (
|
||||||
chunk_states,
|
borrow_graph_pool(user="input logits")
|
||||||
lm_head,
|
if borrow_logits_memory
|
||||||
logits_metadata,
|
else nullcontext()
|
||||||
use_logits_buffer=num_chunks == 1,
|
):
|
||||||
)
|
chunk_logits = get_logits_fn(
|
||||||
|
chunk_states,
|
||||||
|
lm_head,
|
||||||
|
logits_metadata,
|
||||||
|
use_logits_buffer=num_chunks == 1,
|
||||||
|
)
|
||||||
|
|
||||||
# Initialize sampled_logits on first chunk
|
# Sampled outputs must survive graph replay, so they are allocated
|
||||||
|
# outside borrowing. The transient logits are released below.
|
||||||
if i == 0:
|
if i == 0:
|
||||||
sampled_logits = torch.empty(
|
sampled_logits = torch.empty(
|
||||||
(sample_indices.shape[0], chunk_logits.shape[1]),
|
(sample_indices.shape[0], chunk_logits.shape[1]),
|
||||||
|
|||||||
@@ -292,6 +292,7 @@ from sglang.srt.mem_cache.common import (
|
|||||||
retraction_discard,
|
retraction_discard,
|
||||||
)
|
)
|
||||||
from sglang.srt.model_executor.forward_batch_info import PPProxyTensors
|
from sglang.srt.model_executor.forward_batch_info import PPProxyTensors
|
||||||
|
from sglang.srt.model_executor.runner_utils.pool import prewarm_graph_pool_borrow
|
||||||
from sglang.srt.model_loader.utils import get_resolved_model_impl
|
from sglang.srt.model_loader.utils import get_resolved_model_impl
|
||||||
from sglang.srt.multiplex.multiplexing_mixin import SchedulerMultiplexMixin
|
from sglang.srt.multiplex.multiplexing_mixin import SchedulerMultiplexMixin
|
||||||
from sglang.srt.observability.metrics_collector import SchedulerMetricsCollector
|
from sglang.srt.observability.metrics_collector import SchedulerMetricsCollector
|
||||||
@@ -1134,6 +1135,7 @@ class Scheduler(
|
|||||||
else self.schedule_stream
|
else self.schedule_stream
|
||||||
)
|
)
|
||||||
with device_module.stream(forward_stream):
|
with device_module.stream(forward_stream):
|
||||||
|
prewarm_graph_pool_borrow()
|
||||||
if self.draft_worker is None:
|
if self.draft_worker is None:
|
||||||
model_runner.prewarm_sampling()
|
model_runner.prewarm_sampling()
|
||||||
else:
|
else:
|
||||||
|
|||||||
@@ -174,6 +174,16 @@ def graph_pool_borrow_enabled() -> bool:
|
|||||||
return get_global_graph_memory_pool() is not None
|
return get_global_graph_memory_pool() is not None
|
||||||
|
|
||||||
|
|
||||||
|
def prewarm_graph_pool_borrow() -> None:
|
||||||
|
"""Initialize cuBLAS's persistent workspace outside borrowed storage.
|
||||||
|
|
||||||
|
Run on the forward stream before final KV sizing. Otherwise the first
|
||||||
|
borrowed GEMM leaves a cached workspace alive across graph replay.
|
||||||
|
"""
|
||||||
|
if graph_pool_borrow_enabled():
|
||||||
|
torch.cuda.current_blas_handle()
|
||||||
|
|
||||||
|
|
||||||
@contextmanager
|
@contextmanager
|
||||||
def graph_pool_user_scope(user: str) -> Iterator[None]:
|
def graph_pool_user_scope(user: str) -> Iterator[None]:
|
||||||
state = _get_graph_pool_borrow_state()
|
state = _get_graph_pool_borrow_state()
|
||||||
|
|||||||
@@ -96,7 +96,7 @@ def _run(proc, batch, chunked, chunk_size):
|
|||||||
class TestLogprobChunkStitching(CustomTestCase):
|
class TestLogprobChunkStitching(CustomTestCase):
|
||||||
def _sweep(self, with_token_ids):
|
def _sweep(self, with_token_ids):
|
||||||
torch.manual_seed(0)
|
torch.manual_seed(0)
|
||||||
proc = InputLogprobProcessor()
|
proc = InputLogprobProcessor(vocab_size=VOCAB)
|
||||||
combos = list(coverage_cases(SEQ_SPEC_MENU, max_seqs=4))
|
combos = list(coverage_cases(SEQ_SPEC_MENU, max_seqs=4))
|
||||||
self.assertEqual(len(combos), EXPECTED_CASES)
|
self.assertEqual(len(combos), EXPECTED_CASES)
|
||||||
tried = 0
|
tried = 0
|
||||||
|
|||||||
@@ -125,7 +125,7 @@ def _shape_of(nested):
|
|||||||
class TestFastInputLogprobs(CustomTestCase):
|
class TestFastInputLogprobs(CustomTestCase):
|
||||||
def _sweep(self, dtype, rtol, atol):
|
def _sweep(self, dtype, rtol, atol):
|
||||||
torch.manual_seed(0)
|
torch.manual_seed(0)
|
||||||
proc = InputLogprobProcessor()
|
proc = InputLogprobProcessor(vocab_size=VOCAB)
|
||||||
combos = list(coverage_cases(SEQ_SPEC_MENU, max_seqs=3))
|
combos = list(coverage_cases(SEQ_SPEC_MENU, max_seqs=3))
|
||||||
self.assertEqual(len(combos), EXPECTED_CASES)
|
self.assertEqual(len(combos), EXPECTED_CASES)
|
||||||
tried = 0
|
tried = 0
|
||||||
@@ -178,7 +178,7 @@ class TestFastInputLogprobs(CustomTestCase):
|
|||||||
# rounds at the bf16 logits themselves (normalizer is fp32), so it
|
# rounds at the bf16 logits themselves (normalizer is fp32), so it
|
||||||
# sits much closer to the truth than bf16 resolution.
|
# sits much closer to the truth than bf16 resolution.
|
||||||
torch.manual_seed(0)
|
torch.manual_seed(0)
|
||||||
proc = InputLogprobProcessor()
|
proc = InputLogprobProcessor(vocab_size=VOCAB)
|
||||||
for combo in coverage_cases(SEQ_SPEC_MENU, max_seqs=3):
|
for combo in coverage_cases(SEQ_SPEC_MENU, max_seqs=3):
|
||||||
batch = _build_batch(list(combo), torch.bfloat16)
|
batch = _build_batch(list(combo), torch.bfloat16)
|
||||||
pruned_states, _, input_logprob_indices, _, metadata = batch
|
pruned_states, _, input_logprob_indices, _, metadata = batch
|
||||||
@@ -233,7 +233,7 @@ class TestFastInputLogprobs(CustomTestCase):
|
|||||||
# true precision of the result), while the log_softmax path keeps
|
# true precision of the result), while the log_softmax path keeps
|
||||||
# the logits dtype. Runs on CPU CI so the policy is pinned even
|
# the logits dtype. Runs on CPU CI so the policy is pinned even
|
||||||
# where the CUDA kernels never execute.
|
# where the CUDA kernels never execute.
|
||||||
proc = InputLogprobProcessor()
|
proc = InputLogprobProcessor(vocab_size=VOCAB)
|
||||||
batch = _build_batch([(4, 1), (3, 0)], torch.bfloat16)
|
batch = _build_batch([(4, 1), (3, 0)], torch.bfloat16)
|
||||||
got, _ = _run(proc, batch, True, None)
|
got, _ = _run(proc, batch, True, None)
|
||||||
self.assertEqual(got.token_logprobs.dtype, torch.float32)
|
self.assertEqual(got.token_logprobs.dtype, torch.float32)
|
||||||
@@ -323,7 +323,7 @@ class TestFastInputLogprobs(CustomTestCase):
|
|||||||
# (the CPU sweeps only cover the torch fallbacks), including the
|
# (the CPU sweeps only cover the torch fallbacks), including the
|
||||||
# k > FUSED_TOPK_MAX_K fallback.
|
# k > FUSED_TOPK_MAX_K fallback.
|
||||||
torch.manual_seed(0)
|
torch.manual_seed(0)
|
||||||
proc = InputLogprobProcessor()
|
proc = InputLogprobProcessor(vocab_size=64)
|
||||||
for k_override in (None, 20):
|
for k_override in (None, 20):
|
||||||
# k=20 exceeds FUSED_TOPK_MAX_K, exercising the torch fallback;
|
# k=20 exceeds FUSED_TOPK_MAX_K, exercising the torch fallback;
|
||||||
# it needs a vocab that can supply 20 entries.
|
# it needs a vocab that can supply 20 entries.
|
||||||
|
|||||||
+7
@@ -812,6 +812,10 @@ class TestStartupWeightLoadSchedulerRouting(CustomTestCase):
|
|||||||
stream=stream_context, Stream=lambda priority: schedule_stream
|
stream=stream_context, Stream=lambda priority: schedule_stream
|
||||||
),
|
),
|
||||||
),
|
),
|
||||||
|
patch(
|
||||||
|
"sglang.srt.managers.scheduler.prewarm_graph_pool_borrow",
|
||||||
|
side_effect=lambda: trace.append("borrow_prewarm"),
|
||||||
|
),
|
||||||
self.assertRaisesRegex(RuntimeError, "stop after startup"),
|
self.assertRaisesRegex(RuntimeError, "stop after startup"),
|
||||||
):
|
):
|
||||||
scheduler.init_model_worker()
|
scheduler.init_model_worker()
|
||||||
@@ -830,6 +834,7 @@ class TestStartupWeightLoadSchedulerRouting(CustomTestCase):
|
|||||||
"attention",
|
"attention",
|
||||||
"capture",
|
"capture",
|
||||||
"stream_enter",
|
"stream_enter",
|
||||||
|
"borrow_prewarm",
|
||||||
"prewarm",
|
"prewarm",
|
||||||
"stream_exit",
|
"stream_exit",
|
||||||
"resize",
|
"resize",
|
||||||
@@ -857,6 +862,7 @@ class TestStartupWeightLoadSchedulerRouting(CustomTestCase):
|
|||||||
"attention",
|
"attention",
|
||||||
"capture",
|
"capture",
|
||||||
"stream_enter",
|
"stream_enter",
|
||||||
|
"borrow_prewarm",
|
||||||
"prewarm",
|
"prewarm",
|
||||||
"stream_exit",
|
"stream_exit",
|
||||||
"resize",
|
"resize",
|
||||||
@@ -872,6 +878,7 @@ class TestStartupWeightLoadSchedulerRouting(CustomTestCase):
|
|||||||
"attention",
|
"attention",
|
||||||
"capture",
|
"capture",
|
||||||
"stream_enter",
|
"stream_enter",
|
||||||
|
"borrow_prewarm",
|
||||||
"draft_prewarm",
|
"draft_prewarm",
|
||||||
"stream_exit",
|
"stream_exit",
|
||||||
"resize",
|
"resize",
|
||||||
|
|||||||
@@ -251,11 +251,16 @@ class TestGraphPoolBorrow(CustomTestCase):
|
|||||||
torch.cuda.synchronize()
|
torch.cuda.synchronize()
|
||||||
|
|
||||||
device_id = torch.cuda.current_device()
|
device_id = torch.cuda.current_device()
|
||||||
reserved_before = torch.cuda.memory_reserved(device_id)
|
borrow_stream = torch.cuda.Stream()
|
||||||
with (
|
with (
|
||||||
envs.SGLANG_ENABLE_GRAPH_POOL_BORROW.override(True),
|
envs.SGLANG_ENABLE_GRAPH_POOL_BORROW.override(True),
|
||||||
patch.object(pool, "get_global_graph_memory_pool", return_value=handle),
|
patch.object(pool, "get_global_graph_memory_pool", return_value=handle),
|
||||||
|
torch.cuda.stream(borrow_stream),
|
||||||
):
|
):
|
||||||
|
lhs = torch.ones((32, 16), dtype=torch.bfloat16, device="cuda")
|
||||||
|
copied_product = torch.empty((32, 32), dtype=torch.float32, device="cuda")
|
||||||
|
pool.prewarm_graph_pool_borrow()
|
||||||
|
reserved_before = torch.cuda.memory_reserved(device_id)
|
||||||
runs = pool.find_free_graph_pool_runs(handle)
|
runs = pool.find_free_graph_pool_runs(handle)
|
||||||
self.assertGreaterEqual(len(runs), 2)
|
self.assertGreaterEqual(len(runs), 2)
|
||||||
largest_run_bytes = runs[0][1]
|
largest_run_bytes = runs[0][1]
|
||||||
@@ -301,6 +306,13 @@ class TestGraphPoolBorrow(CustomTestCase):
|
|||||||
self.assertEqual(reused.data_ptr(), recycled_address)
|
self.assertEqual(reused.data_ptr(), recycled_address)
|
||||||
del reused
|
del reused
|
||||||
|
|
||||||
|
# The first GEMM on this stream must not cache its workspace
|
||||||
|
# in borrowed storage, which replay would overwrite.
|
||||||
|
product = torch.mm(lhs, lhs.T, out_dtype=torch.float32)
|
||||||
|
self.assertTrue(on_a_run(product))
|
||||||
|
copied_product.copy_(product)
|
||||||
|
del product
|
||||||
|
|
||||||
# Captures retire the persistent borrow pool. Its storage aliases
|
# Captures retire the persistent borrow pool. Its storage aliases
|
||||||
# existing graph-pool runs, so the reserved footprint is unchanged.
|
# existing graph-pool runs, so the reserved footprint is unchanged.
|
||||||
pool._teardown_borrow_pool()
|
pool._teardown_borrow_pool()
|
||||||
@@ -315,6 +327,9 @@ class TestGraphPoolBorrow(CustomTestCase):
|
|||||||
graph.replay()
|
graph.replay()
|
||||||
torch.cuda.synchronize()
|
torch.cuda.synchronize()
|
||||||
self.assertTrue(torch.equal(y, torch.ones_like(y)))
|
self.assertTrue(torch.equal(y, torch.ones_like(y)))
|
||||||
|
self.assertTrue(
|
||||||
|
torch.equal(copied_product, torch.full_like(copied_product, 16))
|
||||||
|
)
|
||||||
|
|
||||||
self.assertEqual(torch.cuda.memory_reserved(device_id), reserved_before)
|
self.assertEqual(torch.cuda.memory_reserved(device_id), reserved_before)
|
||||||
del graph, y
|
del graph, y
|
||||||
|
|||||||
Reference in New Issue
Block a user