[PP] Support prefill CUDA graph proxy tensors (#36248)
Co-authored-by: github-actions[bot] <41898282+github-actions[bot]@users.noreply.github.com>
This commit is contained in:
co-authored by
github-actions[bot]
parent
5d92e60783
commit
b77cac06a9
@@ -115,6 +115,24 @@ def apply_cuda_graph_compatibility(server_args: Any):
|
|||||||
if (Phase.PREFILL, "backend") in server_args._cuda_graph_config_locked:
|
if (Phase.PREFILL, "backend") in server_args._cuda_graph_config_locked:
|
||||||
return
|
return
|
||||||
|
|
||||||
|
# PP prefill graph replay is opt-in. It is most useful for small
|
||||||
|
# aggregate forwards, while enabling it implicitly would also capture
|
||||||
|
# large buckets that can be slower than eager. An explicit backend
|
||||||
|
# selection bypasses this default policy.
|
||||||
|
if cfg.pp_size > 1 and cfg.cuda_graph_config.prefill.backend == Backend.BREAKABLE:
|
||||||
|
logger.info(
|
||||||
|
"Disabling breakable prefill CUDA graph by default for pipeline "
|
||||||
|
"parallelism. Set --cuda-graph-backend-prefill=breakable to opt in."
|
||||||
|
)
|
||||||
|
declare_resolution(
|
||||||
|
server_args,
|
||||||
|
"_apply_cuda_graph_compatibility",
|
||||||
|
cuda_graph_config=with_phase(
|
||||||
|
cfg.cuda_graph_config, Phase.PREFILL, backend=Backend.DISABLED
|
||||||
|
),
|
||||||
|
)
|
||||||
|
return
|
||||||
|
|
||||||
# Breakable is the CUDA default but not multimodal-compatible;
|
# Breakable is the CUDA default but not multimodal-compatible;
|
||||||
# piecewise-allowlisted archs run their validated decoder prefill
|
# piecewise-allowlisted archs run their validated decoder prefill
|
||||||
# there instead. Archs also on the breakable allowlist keep it --
|
# there instead. Archs also on the breakable allowlist keep it --
|
||||||
|
|||||||
@@ -16,10 +16,12 @@ from sglang.srt.arg_groups.overrides import (
|
|||||||
use_mla_backend,
|
use_mla_backend,
|
||||||
)
|
)
|
||||||
from sglang.srt.environ import envs
|
from sglang.srt.environ import envs
|
||||||
from sglang.srt.model_executor.cuda_graph_config import Backend
|
from sglang.srt.model_executor.cuda_graph_config import Backend, Phase
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
_DEFAULT_PP_PREFILL_CUDA_GRAPH_MAX_TOKENS = 8192
|
||||||
|
|
||||||
|
|
||||||
def handle_gpu_memory_settings(server_args: Any, gpu_mem):
|
def handle_gpu_memory_settings(server_args: Any, gpu_mem):
|
||||||
"""
|
"""
|
||||||
@@ -189,6 +191,18 @@ def handle_gpu_memory_settings(server_args: Any, gpu_mem):
|
|||||||
else:
|
else:
|
||||||
prefill_cuda_graph_config.max_bs = 2048
|
prefill_cuda_graph_config.max_bs = 2048
|
||||||
|
|
||||||
|
# For opt-in PP breakable graphs, capture small aggregate-token
|
||||||
|
# buckets by default and leave larger forwards on the eager path.
|
||||||
|
# Explicit max_bs or bs settings retain their existing semantics.
|
||||||
|
if (
|
||||||
|
cfg.pp_size > 1
|
||||||
|
and prefill_cuda_graph_config.backend == Backend.BREAKABLE
|
||||||
|
and (Phase.PREFILL, "bs") not in server_args._cuda_graph_config_locked
|
||||||
|
and prefill_cuda_graph_config.max_bs
|
||||||
|
> _DEFAULT_PP_PREFILL_CUDA_GRAPH_MAX_TOKENS
|
||||||
|
):
|
||||||
|
prefill_cuda_graph_config.max_bs = _DEFAULT_PP_PREFILL_CUDA_GRAPH_MAX_TOKENS
|
||||||
|
|
||||||
# If max_total_tokens is set, cap prefill max_bs to not exceed max_total_tokens.
|
# If max_total_tokens is set, cap prefill max_bs to not exceed max_total_tokens.
|
||||||
if cfg.max_total_tokens is not None:
|
if cfg.max_total_tokens is not None:
|
||||||
prefill_cuda_graph_config.max_bs = min(
|
prefill_cuda_graph_config.max_bs = min(
|
||||||
|
|||||||
@@ -53,6 +53,7 @@ def supports_prefill_cp_bcg(server_args: ServerArgs) -> bool:
|
|||||||
prefill_attention_backend, _ = attention_backends_of(resolved_view(server_args))
|
prefill_attention_backend, _ = attention_backends_of(resolved_view(server_args))
|
||||||
return (
|
return (
|
||||||
cfg.enable_prefill_cp
|
cfg.enable_prefill_cp
|
||||||
|
and cfg.pp_size == 1
|
||||||
and resolved.attn_cp_size == cfg.tp_size
|
and resolved.attn_cp_size == cfg.tp_size
|
||||||
and cfg.cp_strategy == "zigzag"
|
and cfg.cp_strategy == "zigzag"
|
||||||
and prefill_attention_backend == "trtllm_mha"
|
and prefill_attention_backend == "trtllm_mha"
|
||||||
|
|||||||
@@ -953,7 +953,10 @@ def build_prefill_registry(
|
|||||||
def _pp_source(key):
|
def _pp_source(key):
|
||||||
def _fn(_fb, ctx):
|
def _fn(_fb, ctx):
|
||||||
ppx = ctx.pp_proxy_tensors
|
ppx = ctx.pp_proxy_tensors
|
||||||
return None if ppx is None else ppx.tensors[key]
|
# Proxy contracts vary by model. The capture buffers are a
|
||||||
|
# stable-address superset; only copy fields present in the
|
||||||
|
# live proxy for this model.
|
||||||
|
return None if ppx is None else ppx.tensors.get(key)
|
||||||
|
|
||||||
return _fn
|
return _fn
|
||||||
|
|
||||||
|
|||||||
@@ -45,6 +45,7 @@ from sglang.srt.runtime_context import (
|
|||||||
get_disagg,
|
get_disagg,
|
||||||
get_exec,
|
get_exec,
|
||||||
get_flags,
|
get_flags,
|
||||||
|
get_parallel,
|
||||||
get_schedule,
|
get_schedule,
|
||||||
get_spec,
|
get_spec,
|
||||||
)
|
)
|
||||||
@@ -72,6 +73,8 @@ def _align_pipeline_layers(layers: list, layer_model) -> list:
|
|||||||
f"invalid pipeline layer range [{start_layer}, {end_layer}) for "
|
f"invalid pipeline layer range [{start_layer}, {end_layer}) for "
|
||||||
f"{len(layer_model.layers)} layers"
|
f"{len(layer_model.layers)} layers"
|
||||||
)
|
)
|
||||||
|
if len(layers) == len(layer_model.layers):
|
||||||
|
return layers
|
||||||
assert (
|
assert (
|
||||||
len(layers) <= end_layer - start_layer
|
len(layers) <= end_layer - start_layer
|
||||||
), f"found {len(layers)} layers in PP range [{start_layer}, {end_layer})"
|
), f"found {len(layers)} layers in PP range [{start_layer}, {end_layer})"
|
||||||
@@ -365,6 +368,17 @@ def capture_prefill_graph(
|
|||||||
|
|
||||||
prefill_config = get_exec().graph.cuda_graph_config.prefill
|
prefill_config = get_exec().graph.cuda_graph_config.prefill
|
||||||
prefill_backend = prefill_config.backend
|
prefill_backend = prefill_config.backend
|
||||||
|
parallel = get_parallel()
|
||||||
|
if (
|
||||||
|
prefill_backend == Backend.BREAKABLE
|
||||||
|
and parallel.enable_prefill_cp
|
||||||
|
and parallel.pp_size > 1
|
||||||
|
):
|
||||||
|
logger.warning(
|
||||||
|
"Disable prefill CUDA graph because pipeline parallelism combined "
|
||||||
|
"with prefill context parallelism is not validated."
|
||||||
|
)
|
||||||
|
return result(eager_runner)
|
||||||
context_length = model_runner.model_config.context_len
|
context_length = model_runner.model_config.context_len
|
||||||
if prefill_backend == Backend.FULL:
|
if prefill_backend == Backend.FULL:
|
||||||
max_capture_requests = prefill_config.full_prefill_max_req
|
max_capture_requests = prefill_config.full_prefill_max_req
|
||||||
|
|||||||
@@ -62,12 +62,11 @@ def compute_attention_and_moe_layers(layer_model: Any) -> AttentionAndMoeLayers:
|
|||||||
# Mamba layer with split op support - store the layer itself
|
# Mamba layer with split op support - store the layer itself
|
||||||
attn_layer = layer
|
attn_layer = layer
|
||||||
|
|
||||||
if attn_layer is not None:
|
# Keep these lists aligned with global layer ids. Pipeline-parallel
|
||||||
|
# models retain placeholders outside the local stage, while real
|
||||||
|
# attention modules use their global layer_id during graph replay.
|
||||||
attention_layers.append(attn_layer)
|
attention_layers.append(attn_layer)
|
||||||
mha_companion_layers.append(mha_companion_layer)
|
mha_companion_layers.append(mha_companion_layer)
|
||||||
elif hasattr(layer, "mixer"):
|
|
||||||
attention_layers.append(None)
|
|
||||||
mha_companion_layers.append(None)
|
|
||||||
|
|
||||||
moe_block = None
|
moe_block = None
|
||||||
moe_fusion = None
|
moe_fusion = None
|
||||||
|
|||||||
@@ -185,6 +185,23 @@ def _resolve_transformer_layer_model(model: torch.nn.Module) -> torch.nn.Module:
|
|||||||
return layer_model
|
return layer_model
|
||||||
|
|
||||||
|
|
||||||
|
def _build_layer_model_forward_kwargs(
|
||||||
|
layer_model: torch.nn.Module,
|
||||||
|
forward_batch: ForwardBatch,
|
||||||
|
pp_proxy_tensors: Optional[PPProxyTensors],
|
||||||
|
) -> Dict[str, Any]:
|
||||||
|
"""Bind optional transformer inputs by name across model signatures."""
|
||||||
|
parameters = inspect.signature(layer_model.forward).parameters
|
||||||
|
kwargs = {}
|
||||||
|
for embeds_name in ("input_embeds", "inputs_embeds"):
|
||||||
|
if embeds_name in parameters:
|
||||||
|
kwargs[embeds_name] = forward_batch.input_embeds
|
||||||
|
break
|
||||||
|
if pp_proxy_tensors is not None and "pp_proxy_tensors" in parameters:
|
||||||
|
kwargs["pp_proxy_tensors"] = pp_proxy_tensors
|
||||||
|
return kwargs
|
||||||
|
|
||||||
|
|
||||||
def _slice_output_rows(output: Any, num_tokens: int) -> Any:
|
def _slice_output_rows(output: Any, num_tokens: int) -> Any:
|
||||||
"""Slice every tensor leaf in a transformer-body output by token rows.
|
"""Slice every tensor leaf in a transformer-body output by token rows.
|
||||||
|
|
||||||
@@ -333,6 +350,7 @@ class PrefillCudaGraphRunner(BaseCudaGraphRunner):
|
|||||||
dtype=self.model_runner.dtype,
|
dtype=self.model_runner.dtype,
|
||||||
enable_mamba_track=self.mamba_track_enabled,
|
enable_mamba_track=self.mamba_track_enabled,
|
||||||
pp_size=self.pp_size,
|
pp_size=self.pp_size,
|
||||||
|
is_first_pp_rank=self.model_runner.pp_group.is_first_rank,
|
||||||
hc_hidden_size=getattr(
|
hc_hidden_size=getattr(
|
||||||
self.model_runner.model_config, "hc_hidden_size", None
|
self.model_runner.model_config, "hc_hidden_size", None
|
||||||
),
|
),
|
||||||
@@ -653,12 +671,12 @@ class PrefillCudaGraphRunner(BaseCudaGraphRunner):
|
|||||||
|
|
||||||
return forward_batch.positions
|
return forward_batch.positions
|
||||||
|
|
||||||
def _static_pp_proxy_tensors(self, num_tokens: int) -> Optional[PPProxyTensors]:
|
def _capture_pp_proxy_tensors(self, num_tokens: int) -> Optional[PPProxyTensors]:
|
||||||
buffers = self.buffers.pp_proxy_tensors
|
buffers = self.buffers.pp_proxy_tensors
|
||||||
if buffers is None:
|
if buffers is None or self.model_runner.pp_group.is_first_rank:
|
||||||
return None
|
return None
|
||||||
return PPProxyTensors(
|
return PPProxyTensors(
|
||||||
{key: value[:num_tokens] for key, value in buffers.items()}
|
{name: buffer[:num_tokens] for name, buffer in buffers.items()}
|
||||||
)
|
)
|
||||||
|
|
||||||
@contextmanager
|
@contextmanager
|
||||||
@@ -717,20 +735,28 @@ class PrefillCudaGraphRunner(BaseCudaGraphRunner):
|
|||||||
set_is_extend_in_batch(False)
|
set_is_extend_in_batch(False)
|
||||||
|
|
||||||
with self._prefill_forward_context(forward_batch):
|
with self._prefill_forward_context(forward_batch):
|
||||||
pp_kwargs = self.model_runner._pp_kwargs(
|
pp_proxy_tensors = self._capture_pp_proxy_tensors(num_tokens)
|
||||||
self._static_pp_proxy_tensors(num_tokens)
|
|
||||||
)
|
|
||||||
if self._uses_eager_prefill_tail():
|
if self._uses_eager_prefill_tail():
|
||||||
# BCG / Full: capture the transformer body only.
|
# BCG / Full: capture the transformer body only.
|
||||||
positions = self._get_layer_model_positions(forward_batch)
|
positions = self._get_layer_model_positions(forward_batch)
|
||||||
|
input_ids = forward_batch.input_ids
|
||||||
|
kwargs = _build_layer_model_forward_kwargs(
|
||||||
|
self.layer_model, forward_batch, pp_proxy_tensors
|
||||||
|
)
|
||||||
|
if pp_proxy_tensors is not None:
|
||||||
|
input_ids = None
|
||||||
|
for embeds_name in ("input_embeds", "inputs_embeds"):
|
||||||
|
if embeds_name in kwargs:
|
||||||
|
kwargs[embeds_name] = None
|
||||||
|
break
|
||||||
return self.layer_model.forward(
|
return self.layer_model.forward(
|
||||||
forward_batch.input_ids,
|
input_ids,
|
||||||
positions,
|
positions,
|
||||||
forward_batch,
|
forward_batch,
|
||||||
forward_batch.input_embeds,
|
**kwargs,
|
||||||
**pp_kwargs,
|
|
||||||
)
|
)
|
||||||
# tc_piecewise: compile/capture the outer model.forward path.
|
# tc_piecewise: compile/capture the outer model.forward path.
|
||||||
|
pp_kwargs = self.model_runner._pp_kwargs(pp_proxy_tensors)
|
||||||
return self.model_runner.model.forward(
|
return self.model_runner.model.forward(
|
||||||
forward_batch.input_ids,
|
forward_batch.input_ids,
|
||||||
forward_batch.positions,
|
forward_batch.positions,
|
||||||
@@ -1765,6 +1791,10 @@ class PrefillCudaGraphRunner(BaseCudaGraphRunner):
|
|||||||
# uses real request metadata instead of padded slots. BCG has no
|
# uses real request metadata instead of padded slots. BCG has no
|
||||||
# request-slot padding, so static_forward_batch is already the serving batch.
|
# request-slot padding, so static_forward_batch is already the serving batch.
|
||||||
tail_batch = forward_batch if full_path else static_forward_batch
|
tail_batch = forward_batch if full_path else static_forward_batch
|
||||||
|
if not full_path:
|
||||||
|
# MTP consumes the target model's live multimodal embeddings in its
|
||||||
|
# eager wrapper before the captured transformer body is replayed.
|
||||||
|
tail_batch.mm_input_embeds = forward_batch.mm_input_embeds
|
||||||
try:
|
try:
|
||||||
with self._prefill_forward_context(
|
with self._prefill_forward_context(
|
||||||
static_forward_batch,
|
static_forward_batch,
|
||||||
@@ -1836,7 +1866,7 @@ class PrefillCudaGraphRunner(BaseCudaGraphRunner):
|
|||||||
if isinstance(output, EmbeddingPoolerOutput):
|
if isinstance(output, EmbeddingPoolerOutput):
|
||||||
return output
|
return output
|
||||||
assert isinstance(output, PPProxyTensors)
|
assert isinstance(output, PPProxyTensors)
|
||||||
return output[: self.raw_num_tokens]
|
return _slice_output_rows(output, self.raw_num_tokens)
|
||||||
|
|
||||||
def _validate_capture_hidden_mode(self, forward_batch: ForwardBatch) -> None:
|
def _validate_capture_hidden_mode(self, forward_batch: ForwardBatch) -> None:
|
||||||
if self.capture_hidden_mode < forward_batch.capture_hidden_mode:
|
if self.capture_hidden_mode < forward_batch.capture_hidden_mode:
|
||||||
|
|||||||
@@ -60,6 +60,38 @@ def _grouped_foreach_copy_(dsts: List[torch.Tensor], srcs: List[torch.Tensor]) -
|
|||||||
foreach_copy(group_dsts, group_srcs)
|
foreach_copy(group_dsts, group_srcs)
|
||||||
|
|
||||||
|
|
||||||
|
def _allocate_pp_proxy_tensors(
|
||||||
|
*,
|
||||||
|
max_num_tokens: int,
|
||||||
|
max_hidden_tokens: int,
|
||||||
|
hidden_size: int,
|
||||||
|
dtype: torch.dtype,
|
||||||
|
hc_hidden_size: Optional[int] = None,
|
||||||
|
pp_proxy_topk_size: Optional[int] = None,
|
||||||
|
pp_proxy_residual_num_blocks: Optional[int] = None,
|
||||||
|
) -> Dict[str, torch.Tensor]:
|
||||||
|
"""Allocate the stable buffers consumed by an incoming PP proxy."""
|
||||||
|
is_mhc = hc_hidden_size is not None
|
||||||
|
pp_hidden_size = hc_hidden_size if is_mhc else hidden_size
|
||||||
|
pp_proxy_tensors = {
|
||||||
|
"hidden_states": torch.zeros((max_hidden_tokens, pp_hidden_size), dtype=dtype),
|
||||||
|
}
|
||||||
|
if not is_mhc:
|
||||||
|
# Only Kimi K3 supplies num_blocks: its PP bank is token-major
|
||||||
|
# [T, blocks, H]. Other models use the phase-specific hidden-token bound.
|
||||||
|
residual_shape = (
|
||||||
|
(max_num_tokens, pp_proxy_residual_num_blocks, hidden_size)
|
||||||
|
if pp_proxy_residual_num_blocks is not None
|
||||||
|
else (max_hidden_tokens, hidden_size)
|
||||||
|
)
|
||||||
|
pp_proxy_tensors["residual"] = torch.zeros(residual_shape, dtype=dtype)
|
||||||
|
if pp_proxy_topk_size is not None:
|
||||||
|
pp_proxy_tensors["topk_indices"] = torch.zeros(
|
||||||
|
(max_num_tokens, pp_proxy_topk_size), dtype=torch.int32
|
||||||
|
)
|
||||||
|
return pp_proxy_tensors
|
||||||
|
|
||||||
|
|
||||||
@dataclass
|
@dataclass
|
||||||
class DecodeInputBuffers(ForwardInputBuffers):
|
class DecodeInputBuffers(ForwardInputBuffers):
|
||||||
input_ids: torch.Tensor
|
input_ids: torch.Tensor
|
||||||
@@ -129,29 +161,19 @@ class DecodeInputBuffers(ForwardInputBuffers):
|
|||||||
torch.zeros((max_bs,), dtype=torch.bool) if enable_mamba_track else None
|
torch.zeros((max_bs,), dtype=torch.bool) if enable_mamba_track else None
|
||||||
)
|
)
|
||||||
|
|
||||||
if pp_size > 1:
|
pp_proxy_tensors = (
|
||||||
is_mhc = hc_hidden_size is not None
|
_allocate_pp_proxy_tensors(
|
||||||
hs = hc_hidden_size if is_mhc else hidden_size
|
max_num_tokens=max_num_token,
|
||||||
pp_proxy_tensors = {
|
max_hidden_tokens=max_num_token,
|
||||||
"hidden_states": torch.zeros((max_num_token, hs), dtype=dtype),
|
hidden_size=hidden_size,
|
||||||
}
|
dtype=dtype,
|
||||||
if not is_mhc:
|
hc_hidden_size=hc_hidden_size,
|
||||||
# Only Kimi K3 supplies num_blocks: its PP bank is token-major
|
pp_proxy_topk_size=pp_proxy_topk_size,
|
||||||
# [T, blocks, H]. Other models keep the legacy [max_bs, H].
|
pp_proxy_residual_num_blocks=pp_proxy_residual_num_blocks,
|
||||||
residual_shape = (
|
|
||||||
(max_num_token, pp_proxy_residual_num_blocks, hidden_size)
|
|
||||||
if pp_proxy_residual_num_blocks is not None
|
|
||||||
else (max_num_token, hidden_size)
|
|
||||||
)
|
)
|
||||||
pp_proxy_tensors["residual"] = torch.zeros(
|
if pp_size > 1
|
||||||
residual_shape, dtype=dtype
|
else None
|
||||||
)
|
)
|
||||||
if pp_proxy_topk_size is not None:
|
|
||||||
pp_proxy_tensors["topk_indices"] = torch.zeros(
|
|
||||||
(max_num_token, pp_proxy_topk_size), dtype=torch.int32
|
|
||||||
)
|
|
||||||
else:
|
|
||||||
pp_proxy_tensors = None
|
|
||||||
|
|
||||||
if is_encoder_decoder:
|
if is_encoder_decoder:
|
||||||
encoder_lens = torch.full(
|
encoder_lens = torch.full(
|
||||||
@@ -357,6 +379,7 @@ class PrefillInputBuffers(ForwardInputBuffers):
|
|||||||
dtype: torch.dtype,
|
dtype: torch.dtype,
|
||||||
enable_mamba_track: bool,
|
enable_mamba_track: bool,
|
||||||
pp_size: int = 1,
|
pp_size: int = 1,
|
||||||
|
is_first_pp_rank: bool = False,
|
||||||
hc_hidden_size: Optional[int] = None,
|
hc_hidden_size: Optional[int] = None,
|
||||||
pp_proxy_topk_size: Optional[int] = None,
|
pp_proxy_topk_size: Optional[int] = None,
|
||||||
pp_proxy_residual_num_blocks: Optional[int] = None,
|
pp_proxy_residual_num_blocks: Optional[int] = None,
|
||||||
@@ -387,29 +410,19 @@ class PrefillInputBuffers(ForwardInputBuffers):
|
|||||||
input_embeds = None
|
input_embeds = None
|
||||||
mrope_positions = None
|
mrope_positions = None
|
||||||
|
|
||||||
if pp_size > 1:
|
pp_proxy_tensors = (
|
||||||
is_mhc = hc_hidden_size is not None
|
_allocate_pp_proxy_tensors(
|
||||||
pp_hidden_size = hc_hidden_size if is_mhc else hidden_size
|
max_num_tokens=max_num_tokens,
|
||||||
pp_proxy_tensors = {
|
max_hidden_tokens=max_num_tokens,
|
||||||
"hidden_states": torch.zeros(
|
hidden_size=hidden_size,
|
||||||
(max_num_tokens, pp_hidden_size), dtype=dtype
|
dtype=dtype,
|
||||||
|
hc_hidden_size=hc_hidden_size,
|
||||||
|
pp_proxy_topk_size=pp_proxy_topk_size,
|
||||||
|
pp_proxy_residual_num_blocks=pp_proxy_residual_num_blocks,
|
||||||
)
|
)
|
||||||
}
|
if pp_size > 1 and not is_first_pp_rank
|
||||||
if not is_mhc:
|
else None
|
||||||
residual_shape = (
|
|
||||||
(max_num_tokens, pp_proxy_residual_num_blocks, hidden_size)
|
|
||||||
if pp_proxy_residual_num_blocks is not None
|
|
||||||
else (max_num_tokens, hidden_size)
|
|
||||||
)
|
)
|
||||||
pp_proxy_tensors["residual"] = torch.zeros(
|
|
||||||
residual_shape, dtype=dtype
|
|
||||||
)
|
|
||||||
if pp_proxy_topk_size is not None:
|
|
||||||
pp_proxy_tensors["topk_indices"] = torch.zeros(
|
|
||||||
(max_num_tokens, pp_proxy_topk_size), dtype=torch.int32
|
|
||||||
)
|
|
||||||
else:
|
|
||||||
pp_proxy_tensors = None
|
|
||||||
|
|
||||||
return cls(
|
return cls(
|
||||||
input_ids=input_ids,
|
input_ids=input_ids,
|
||||||
|
|||||||
@@ -31,6 +31,20 @@ class TestComputeAttentionAndMoeLayers(unittest.TestCase):
|
|||||||
self.assertEqual(mha_companion_layers, [attn_mha])
|
self.assertEqual(mha_companion_layers, [attn_mha])
|
||||||
self.assertNotIn("_pcg_mha_companion", vars(attn_mqa))
|
self.assertNotIn("_pcg_mha_companion", vars(attn_mqa))
|
||||||
|
|
||||||
|
def test_pipeline_placeholders_preserve_global_layer_ids(self):
|
||||||
|
local_attention = SimpleNamespace()
|
||||||
|
layer_model = SimpleNamespace(
|
||||||
|
layers=[SimpleNamespace(), SimpleNamespace()]
|
||||||
|
+ [SimpleNamespace(self_attn=SimpleNamespace(attn=local_attention))]
|
||||||
|
)
|
||||||
|
|
||||||
|
attention_layers, _, _, _, mha_companion_layers = (
|
||||||
|
compute_attention_and_moe_layers(layer_model)
|
||||||
|
)
|
||||||
|
|
||||||
|
self.assertEqual(attention_layers, [None, None, local_attention])
|
||||||
|
self.assertEqual(mha_companion_layers, [None, None, None])
|
||||||
|
|
||||||
|
|
||||||
if __name__ == "__main__":
|
if __name__ == "__main__":
|
||||||
unittest.main()
|
unittest.main()
|
||||||
|
|||||||
@@ -100,6 +100,8 @@ class TestPrefillCudaGraphRunnerChunkedPrefix(CustomTestCase):
|
|||||||
# out of the bags.
|
# out of the bags.
|
||||||
override = get_context().override_server_args(
|
override = get_context().override_server_args(
|
||||||
enable_lora=False,
|
enable_lora=False,
|
||||||
|
enable_prefill_cp=False,
|
||||||
|
pp_size=1,
|
||||||
cuda_graph_config=SimpleNamespace(
|
cuda_graph_config=SimpleNamespace(
|
||||||
prefill=SimpleNamespace(bs=[1], backend=Backend.BREAKABLE)
|
prefill=SimpleNamespace(bs=[1], backend=Backend.BREAKABLE)
|
||||||
),
|
),
|
||||||
|
|||||||
@@ -1,13 +1,19 @@
|
|||||||
"""Unit tests for prefill CUDA graph wrapper helpers."""
|
"""Unit tests for prefill CUDA graph wrapper helpers."""
|
||||||
|
|
||||||
import unittest
|
import unittest
|
||||||
|
from contextlib import nullcontext
|
||||||
|
from functools import partial
|
||||||
from types import SimpleNamespace
|
from types import SimpleNamespace
|
||||||
|
|
||||||
import torch
|
import torch
|
||||||
|
|
||||||
|
from sglang.srt.model_executor.cuda_graph_buffer_registry import (
|
||||||
|
build_prefill_registry,
|
||||||
|
)
|
||||||
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.prefill_cuda_graph_runner import (
|
from sglang.srt.model_executor.runner.prefill_cuda_graph_runner import (
|
||||||
PrefillCudaGraphRunner,
|
PrefillCudaGraphRunner,
|
||||||
|
_build_layer_model_forward_kwargs,
|
||||||
_resolve_transformer_layer_model,
|
_resolve_transformer_layer_model,
|
||||||
)
|
)
|
||||||
from sglang.srt.model_executor.runner_utils.buffers import PrefillInputBuffers
|
from sglang.srt.model_executor.runner_utils.buffers import PrefillInputBuffers
|
||||||
@@ -26,7 +32,95 @@ class _LayerModel:
|
|||||||
return input_embeds
|
return input_embeds
|
||||||
|
|
||||||
|
|
||||||
|
def _make_pp_buffers_and_registry():
|
||||||
|
base = torch.zeros(3, dtype=torch.int64)
|
||||||
|
buffers = SimpleNamespace(
|
||||||
|
**{name: base.clone() for name in ("input_ids", "positions", "out_cache_loc")},
|
||||||
|
pp_proxy_tensors={
|
||||||
|
key: torch.zeros((3, 2)) for key in ("hidden_states", "residual")
|
||||||
|
},
|
||||||
|
)
|
||||||
|
registry = build_prefill_registry(
|
||||||
|
device=base.device,
|
||||||
|
max_bs=1,
|
||||||
|
max_num_token=len(base),
|
||||||
|
cache_loc_dtype=torch.int64,
|
||||||
|
share_pool=False,
|
||||||
|
source=buffers,
|
||||||
|
)
|
||||||
|
return buffers, registry
|
||||||
|
|
||||||
|
|
||||||
class TestPrefillCudaGraphRunnerHelpers(CustomTestCase):
|
class TestPrefillCudaGraphRunnerHelpers(CustomTestCase):
|
||||||
|
def test_pp_proxy_stable_buffers_accept_full_and_hidden_only_contracts(self):
|
||||||
|
buffers, registry = _make_pp_buffers_and_registry()
|
||||||
|
full_proxy = PPProxyTensors(
|
||||||
|
{
|
||||||
|
"hidden_states": torch.full((3, 2), 2.0),
|
||||||
|
"residual": torch.full((3, 2), 3.0),
|
||||||
|
}
|
||||||
|
)
|
||||||
|
values = torch.arange(3)
|
||||||
|
fill = partial(
|
||||||
|
registry.fill_from,
|
||||||
|
SimpleNamespace(input_ids=values, positions=values, out_cache_loc=values),
|
||||||
|
raw_bs=1,
|
||||||
|
padded_bs=1,
|
||||||
|
raw_num_tokens=3,
|
||||||
|
padded_num_tokens=3,
|
||||||
|
)
|
||||||
|
fill(pp_proxy_tensors=full_proxy)
|
||||||
|
|
||||||
|
runner = PrefillCudaGraphRunner.__new__(PrefillCudaGraphRunner)
|
||||||
|
runner.buffers = buffers
|
||||||
|
runner.model_runner = SimpleNamespace(
|
||||||
|
pp_group=SimpleNamespace(is_first_rank=False)
|
||||||
|
)
|
||||||
|
capture_proxy = runner._capture_pp_proxy_tensors(3)
|
||||||
|
torch.testing.assert_close(capture_proxy.tensors, full_proxy.tensors)
|
||||||
|
self.assertEqual(
|
||||||
|
capture_proxy["hidden_states"].data_ptr(),
|
||||||
|
buffers.pp_proxy_tensors["hidden_states"].data_ptr(),
|
||||||
|
)
|
||||||
|
|
||||||
|
hidden_only_proxy = PPProxyTensors({"hidden_states": torch.full((3, 2), 4.0)})
|
||||||
|
fill(pp_proxy_tensors=hidden_only_proxy)
|
||||||
|
torch.testing.assert_close(
|
||||||
|
buffers.pp_proxy_tensors["hidden_states"][:3],
|
||||||
|
hidden_only_proxy["hidden_states"],
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_layer_model_kwargs_bind_optional_inputs_by_signature(self):
|
||||||
|
def proxy_before_embeds(a, b, c, pp_proxy_tensors=None, inputs_embeds=None):
|
||||||
|
pass
|
||||||
|
|
||||||
|
cases = (
|
||||||
|
(_LayerModel(), {"input_embeds": "embeds"}),
|
||||||
|
(
|
||||||
|
SimpleNamespace(forward=proxy_before_embeds),
|
||||||
|
{"inputs_embeds": "embeds", "pp_proxy_tensors": "proxy"},
|
||||||
|
),
|
||||||
|
)
|
||||||
|
forward_batch = SimpleNamespace(input_embeds="embeds")
|
||||||
|
for layer_model, expected in cases:
|
||||||
|
with self.subTest(signature=layer_model.forward.__name__):
|
||||||
|
kwargs = _build_layer_model_forward_kwargs(
|
||||||
|
layer_model, forward_batch, "proxy"
|
||||||
|
)
|
||||||
|
self.assertEqual(kwargs, expected)
|
||||||
|
layer_model.forward(None, None, forward_batch, **kwargs)
|
||||||
|
|
||||||
|
def test_finalize_pp_proxy_trims_padded_token_rows(self):
|
||||||
|
runner = PrefillCudaGraphRunner.__new__(PrefillCudaGraphRunner)
|
||||||
|
runner.raw_num_tokens = 3
|
||||||
|
output = PPProxyTensors({"hidden_states": torch.arange(10).reshape(5, 2)})
|
||||||
|
trimmed = runner._finalize_execute_output(output)
|
||||||
|
self.assertIsInstance(trimmed, PPProxyTensors)
|
||||||
|
self.assertEqual(tuple(trimmed["hidden_states"].shape), (3, 2))
|
||||||
|
torch.testing.assert_close(
|
||||||
|
trimmed["hidden_states"][-1], output["hidden_states"][2]
|
||||||
|
)
|
||||||
|
|
||||||
def test_resolve_layer_model_from_language_model_wrapper(self):
|
def test_resolve_layer_model_from_language_model_wrapper(self):
|
||||||
layer_model = _LayerModel()
|
layer_model = _LayerModel()
|
||||||
model = SimpleNamespace(language_model=SimpleNamespace(model=layer_model))
|
model = SimpleNamespace(language_model=SimpleNamespace(model=layer_model))
|
||||||
@@ -74,6 +168,7 @@ class TestPrefillCudaGraphRunnerHelpers(CustomTestCase):
|
|||||||
dtype=torch.bfloat16,
|
dtype=torch.bfloat16,
|
||||||
enable_mamba_track=False,
|
enable_mamba_track=False,
|
||||||
pp_size=2,
|
pp_size=2,
|
||||||
|
is_first_pp_rank=False,
|
||||||
pp_proxy_residual_num_blocks=3,
|
pp_proxy_residual_num_blocks=3,
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -93,6 +188,38 @@ class TestPrefillCudaGraphRunnerHelpers(CustomTestCase):
|
|||||||
finalized = runner._finalize_execute_output(output)
|
finalized = runner._finalize_execute_output(output)
|
||||||
self.assertEqual(finalized["hidden_states"].shape, (3, 8))
|
self.assertEqual(finalized["hidden_states"].shape, (3, 8))
|
||||||
|
|
||||||
|
def test_bcg_eager_tail_uses_live_multimodal_embeddings(self):
|
||||||
|
live_embeds = object()
|
||||||
|
live_batch = SimpleNamespace(mm_input_embeds=live_embeds)
|
||||||
|
static_batch = SimpleNamespace(
|
||||||
|
input_ids=None,
|
||||||
|
positions=None,
|
||||||
|
mm_input_embeds=None,
|
||||||
|
)
|
||||||
|
|
||||||
|
runner = PrefillCudaGraphRunner.__new__(PrefillCudaGraphRunner)
|
||||||
|
runner._is_full_backend = False
|
||||||
|
runner._input_embeds_arg_idx = None
|
||||||
|
runner.buffer_registry = SimpleNamespace(has_slot=lambda _name: False)
|
||||||
|
runner.backend = SimpleNamespace(replay=lambda *_args, **_kwargs: None)
|
||||||
|
runner.layer_model = SimpleNamespace(forward=lambda *_args, **_kwargs: None)
|
||||||
|
runner.model_runner = SimpleNamespace(
|
||||||
|
model=SimpleNamespace(
|
||||||
|
forward=lambda _ids, _positions, batch, **_kwargs: batch.mm_input_embeds
|
||||||
|
)
|
||||||
|
)
|
||||||
|
runner._prefill_forward_context = lambda *_args, **_kwargs: nullcontext()
|
||||||
|
|
||||||
|
output = runner._execute_body_capture(
|
||||||
|
live_batch,
|
||||||
|
static_batch,
|
||||||
|
static_num_tokens=1,
|
||||||
|
raw_num_tokens=1,
|
||||||
|
shape_key=object(),
|
||||||
|
)
|
||||||
|
|
||||||
|
self.assertIs(output, live_embeds)
|
||||||
|
|
||||||
|
|
||||||
if __name__ == "__main__":
|
if __name__ == "__main__":
|
||||||
unittest.main()
|
unittest.main()
|
||||||
|
|||||||
@@ -15,6 +15,7 @@ from sglang.srt.arg_groups.attention_hook import (
|
|||||||
handle_deterministic_inference,
|
handle_deterministic_inference,
|
||||||
)
|
)
|
||||||
from sglang.srt.arg_groups.cuda_graph_hook import (
|
from sglang.srt.arg_groups.cuda_graph_hook import (
|
||||||
|
apply_cuda_graph_compatibility,
|
||||||
disable_tc_piecewise_cudagraph_if_incompatible,
|
disable_tc_piecewise_cudagraph_if_incompatible,
|
||||||
handle_cuda_graph_config,
|
handle_cuda_graph_config,
|
||||||
)
|
)
|
||||||
@@ -31,6 +32,7 @@ from sglang.srt.arg_groups.kv_cache_hook import (
|
|||||||
validate_prefill_only_disable_kv_cache_args,
|
validate_prefill_only_disable_kv_cache_args,
|
||||||
)
|
)
|
||||||
from sglang.srt.arg_groups.mamba_hook import handle_mamba_backend
|
from sglang.srt.arg_groups.mamba_hook import handle_mamba_backend
|
||||||
|
from sglang.srt.arg_groups.memory_hook import handle_gpu_memory_settings
|
||||||
from sglang.srt.arg_groups.model_path_hook import handle_load_format
|
from sglang.srt.arg_groups.model_path_hook import handle_load_format
|
||||||
from sglang.srt.arg_groups.moe_hook import (
|
from sglang.srt.arg_groups.moe_hook import (
|
||||||
handle_a2a_moe,
|
handle_a2a_moe,
|
||||||
@@ -1852,6 +1854,58 @@ class TestCudaGraphConfigDataclassAccess(CustomTestCase):
|
|||||||
self.assertEqual(config.compiler, "eager")
|
self.assertEqual(config.compiler, "eager")
|
||||||
|
|
||||||
|
|
||||||
|
class TestPipelineParallelPrefillCudaGraphPolicy(CustomTestCase):
|
||||||
|
def test_pp_prefill_graph_is_opt_in(self):
|
||||||
|
cases = (
|
||||||
|
(set(), Backend.DISABLED),
|
||||||
|
({(Phase.PREFILL, "backend")}, Backend.BREAKABLE),
|
||||||
|
)
|
||||||
|
for locked, expected in cases:
|
||||||
|
with self.subTest(locked=locked):
|
||||||
|
args = ServerArgs(
|
||||||
|
model_path="dummy",
|
||||||
|
pp_size=4,
|
||||||
|
cuda_graph_config=CudaGraphConfig(
|
||||||
|
prefill=PhaseConfig(backend=Backend.BREAKABLE)
|
||||||
|
),
|
||||||
|
)
|
||||||
|
args._cuda_graph_config_locked = locked
|
||||||
|
apply_cuda_graph_compatibility(args)
|
||||||
|
self.assertEqual(
|
||||||
|
resolution_result(args, "cuda_graph_config").prefill.backend,
|
||||||
|
expected,
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_pp_prefill_capture_limit_policy(self):
|
||||||
|
cases = (
|
||||||
|
(4096, None, 4096),
|
||||||
|
(32768, None, 8192),
|
||||||
|
(32768, 16384, 16384),
|
||||||
|
)
|
||||||
|
for chunked_prefill_size, max_bs, expected in cases:
|
||||||
|
with self.subTest(chunked_prefill_size=chunked_prefill_size, max_bs=max_bs):
|
||||||
|
args = ServerArgs(
|
||||||
|
model_path="dummy",
|
||||||
|
pp_size=4,
|
||||||
|
chunked_prefill_size=chunked_prefill_size,
|
||||||
|
mem_fraction_static=0.8,
|
||||||
|
cuda_graph_config=CudaGraphConfig(
|
||||||
|
decode=PhaseConfig(backend=Backend.DISABLED, max_bs=1, bs=[1]),
|
||||||
|
prefill=PhaseConfig(backend=Backend.BREAKABLE, max_bs=max_bs),
|
||||||
|
),
|
||||||
|
)
|
||||||
|
args._cuda_graph_config_locked = {(Phase.PREFILL, "backend")} | (
|
||||||
|
{(Phase.PREFILL, "max_bs")} if max_bs is not None else set()
|
||||||
|
)
|
||||||
|
with patch(
|
||||||
|
"sglang.srt.arg_groups.memory_hook.use_mla_backend",
|
||||||
|
return_value=False,
|
||||||
|
):
|
||||||
|
handle_gpu_memory_settings(args, gpu_mem=None)
|
||||||
|
prefill = resolution_result(args, "cuda_graph_config").prefill
|
||||||
|
self.assertEqual((prefill.max_bs, prefill.bs[-1]), (expected, expected))
|
||||||
|
|
||||||
|
|
||||||
class TestCudaGraphDisaggregationRoles(CustomTestCase):
|
class TestCudaGraphDisaggregationRoles(CustomTestCase):
|
||||||
def _handled_args(self, **overrides):
|
def _handled_args(self, **overrides):
|
||||||
args = ServerArgs(model_path="dummy", **overrides)
|
args = ServerArgs(model_path="dummy", **overrides)
|
||||||
|
|||||||
Reference in New Issue
Block a user