Support encoder_decoder on cpu_graph_runner (#10950)

This commit is contained in:
Cao E
2026-06-08 10:30:11 +08:00
committed by GitHub
parent 303757ccd8
commit 2d1856bf45
2 changed files with 183 additions and 59 deletions
@@ -29,6 +29,7 @@ import tqdm
from sglang.srt.distributed import get_tensor_model_parallel_rank
from sglang.srt.distributed.parallel_state import GroupCoordinator
from sglang.srt.layers.logits_processor import LogitsProcessorOutput
from sglang.srt.model_executor.cuda_graph_runner import model_capture_mode
from sglang.srt.model_executor.forward_batch_info import (
CaptureHiddenMode,
ForwardBatch,
@@ -38,6 +39,7 @@ from sglang.srt.model_executor.forward_batch_info import (
)
from sglang.srt.model_executor.forward_context import ForwardContext, forward_context
from sglang.srt.utils import (
empty_context,
log_info_on_rank0,
require_attn_tp_gather,
require_gathered_buffer,
@@ -48,6 +50,34 @@ from sglang.srt.utils.patch_torch import monkey_patch_torch_compile
logger = logging.getLogger(__name__)
# ---------------------------------------------------------------------------
# skip_cross_attention capture-mode helpers (CPU graph only)
# ---------------------------------------------------------------------------
# When CPUGraphRunner captures two graphs per batch size (one with cross-
# attention, one without), it uses this context variable so that
# encoder-decoder models (e.g. mllama) receive a compile-time-constant value
# for skip_cross_attention instead of a data-dependent branch to avoid recompiles.
_capture_skip_cross_attention: Optional[bool] = None
def get_capture_skip_cross_attention() -> Optional[bool]:
"""Return the active skip_cross_attention override, or None if not set."""
return _capture_skip_cross_attention
@contextmanager
def capture_with_skip_cross_attention(skip: bool):
"""Pin skip_cross_attention to *skip* for the duration of the context."""
global _capture_skip_cross_attention
previous = _capture_skip_cross_attention
_capture_skip_cross_attention = skip
try:
yield
finally:
_capture_skip_cross_attention = previous
if TYPE_CHECKING:
from sglang.srt.model_executor.model_runner import ModelRunner
@@ -484,7 +514,10 @@ class CPUGraphRunner:
# Parse args
self.model_runner = model_runner
self.device = model_runner.device
# bs -> compiled fn (text-only / skip_cross_attention=True)
self.graphs = {}
# bs -> compiled fn (cross-attention / skip_cross_attention=False, enc-dec only)
self.graphs_cross = {}
self.output_buffers = {}
self.enable_torch_compile = model_runner.server_args.enable_torch_compile
self.disable_padding = model_runner.server_args.disable_cuda_graph_padding
@@ -530,17 +563,17 @@ class CPUGraphRunner:
assert (
model_runner.spec_algorithm.is_none()
), "CPUGraphRunner does not support speculative inference yet."
# TODO add compile support for encoder-decoder models
assert (
not self.is_encoder_decoder
), "CPUGraphRunner does not support encoder-decoder models yet."
assert self.dp_size == 1, "CPUGraphRunner does not support DP yet."
assert self.pp_size == 1, "CPUGraphRunner does not support PP yet."
# Batch sizes to capture
self.capture_bs = get_batch_sizes_to_capture(model_runner)
log_info_on_rank0(logger, f"Capture cpu graph bs {self.capture_bs}")
# bs -> ForwardBatch (text-only / skip_cross_attention=True)
self.captured_forward_batches = {}
# bs -> ForwardBatch (cross-attention / skip=False, enc-dec only)
self.captured_forward_batches_cross = {}
# Attention backend
self.max_bs = max(self.capture_bs)
self.max_num_token = self.max_bs * self.num_tokens_per_bs
@@ -548,6 +581,7 @@ class CPUGraphRunner:
self.max_bs, self.max_num_token
)
self.encoder_len_fill_value = 0
self.seq_len_fill_value = (
self.model_runner.attn_backend.get_cpu_graph_seq_len_fill_value()
)
@@ -575,15 +609,41 @@ class CPUGraphRunner:
dtype=torch.bool,
device=self.device,
)
if self.is_encoder_decoder:
self.encoder_lens = torch.full(
(self.max_bs,), self.encoder_len_fill_value, dtype=torch.int64
)
else:
self.encoder_lens = None
# Capture
try:
self.capture()
# use model_capture_mode for encoder-decoder models to
# set skip_cross_attention to avoid
# "Graph Break Reason: Data-dependent branching" caused by
# skip_cross_attention = forward_batch.encoder_lens.max() == 0
capture_context = (
model_capture_mode if self.is_encoder_decoder else empty_context
)
with capture_context():
self.capture()
except RuntimeError as e:
raise Exception(
f"Capture CPU graph failed: {e}\n{CPU_GRAPH_CAPTURE_FAILED_MSG}"
)
def _get_skip_cross_attention(self, forward_batch: ForwardBatch) -> bool:
"""Return True when cross-attention layers should be skipped.
Non-encoder-decoder models have no cross-attention at all, so they
always use self.graphs (the skip=True / text-only graph dict).
For encoder-decoder models, skip when no request in the batch has
encoder output (i.e. no images).
"""
if not self.is_encoder_decoder:
return True
return bool(forward_batch.encoder_lens.max() == 0)
def can_run(self, forward_batch: ForwardBatch):
is_bs_supported = (
forward_batch.batch_size in self.graphs
@@ -626,12 +686,18 @@ class CPUGraphRunner:
num_tokens=bs * self.num_tokens_per_bs,
tp_group=self.model_runner.tp_group,
) as forward:
(
graph,
output_buffers,
) = self.capture_one_batch_size(bs, forward)
graph, output_buffers = self.capture_one_batch_size(
bs, forward, skip_cross_attention=True
)
self.graphs[bs] = graph
self.output_buffers[bs] = output_buffers
if self.is_encoder_decoder:
# Capture a second graph with cross-attention enabled
# (used when the batch contains images).
graph_cross, _ = self.capture_one_batch_size(
bs, forward, skip_cross_attention=False
)
self.graphs_cross[bs] = graph_cross
# Re-init states for qwen3-next as
# torch.compile may change the states
@@ -656,7 +722,9 @@ class CPUGraphRunner:
for v in vars(mamba_cache).values():
_zero_nested(v)
def capture_one_batch_size(self, bs: int, forward: Callable):
def capture_one_batch_size(
self, bs: int, forward: Callable, skip_cross_attention: bool = False
):
num_tokens = bs * self.num_tokens_per_bs
# Graph inputs
@@ -667,6 +735,10 @@ class CPUGraphRunner:
positions = self.positions[:num_tokens]
mrope_positions = self.mrope_positions[:, :num_tokens]
self.num_token_non_padded[...] = num_tokens
if self.is_encoder_decoder:
encoder_lens = self.encoder_lens[:bs]
else:
encoder_lens = None
spec_info = self.get_spec_info(num_tokens)
if self.capture_hidden_mode != CaptureHiddenMode.FULL:
@@ -682,6 +754,8 @@ class CPUGraphRunner:
seq_lens=seq_lens,
out_cache_loc=out_cache_loc,
seq_lens_sum=seq_lens.sum().item(),
encoder_lens=encoder_lens,
encoder_lens_cpu=encoder_lens,
return_logprob=False,
positions=positions,
mrope_positions=mrope_positions,
@@ -691,46 +765,58 @@ class CPUGraphRunner:
num_token_non_padded=self.num_token_non_padded,
global_forward_mode=self.capture_forward_mode,
)
with forward_context(
ForwardContext(attn_backend=self.model_runner.attn_backend)
):
self.model_runner.attn_backend.init_forward_metadata_capture_cpu_graph(
bs,
num_tokens,
req_pool_indices,
seq_lens,
None,
forward_batch.forward_mode,
forward_batch.spec_info,
)
with torch.no_grad():
self.model_runner.tp_group.barrier()
self.model_runner.model.forward(
forward_batch.input_ids,
forward_batch.positions,
forward_batch,
# Wrap all forward calls with capture_with_skip_cross_attention so that
# mllama (and any other encoder-decoder model) sees the correct compile-
# time constant for skip_cross_attention during tracing.
skip_ctx = (
capture_with_skip_cross_attention(skip_cross_attention)
if self.is_encoder_decoder
else empty_context()
)
with skip_ctx:
with forward_context(
ForwardContext(attn_backend=self.model_runner.attn_backend)
):
self.model_runner.attn_backend.init_forward_metadata_capture_cpu_graph(
bs,
num_tokens,
req_pool_indices,
seq_lens,
None,
forward_batch.forward_mode,
forward_batch.spec_info,
)
# Run and capture
def run_once():
# Clean intermediate result cache for DP attention
forward_batch.dp_local_start_pos = forward_batch.dp_local_num_tokens = (
None
)
logits_output_or_pp_proxy_tensors = forward(
forward_batch.input_ids,
forward_batch.positions,
forward_batch,
)
return logits_output_or_pp_proxy_tensors
with torch.no_grad():
for _ in range(2):
with torch.no_grad():
self.model_runner.tp_group.barrier()
out = run_once()
# Save the captured forward_batch
self.captured_forward_batches[bs] = forward_batch
return forward, out
self.model_runner.model.forward(
forward_batch.input_ids,
forward_batch.positions,
forward_batch,
)
# Run and capture
def run_once():
# Clean intermediate result cache for DP attention
forward_batch.dp_local_start_pos = (
forward_batch.dp_local_num_tokens
) = None
logits_output_or_pp_proxy_tensors = forward(
forward_batch.input_ids,
forward_batch.positions,
forward_batch,
)
return logits_output_or_pp_proxy_tensors
with torch.no_grad():
for _ in range(2):
self.model_runner.tp_group.barrier()
out = run_once()
# Save the captured forward_batch in the appropriate dict
if skip_cross_attention:
self.captured_forward_batches[bs] = forward_batch
else:
self.captured_forward_batches_cross[bs] = forward_batch
return forward, out
def recapture_if_needed(self, forward_batch: ForwardBatch):
@@ -766,11 +852,25 @@ class CPUGraphRunner:
def prepare_replay(
self,
forward_batch: ForwardBatch,
skip: bool = False,
):
self.recapture_if_needed(forward_batch)
graphs = self.graphs_cross if not skip else self.graphs
cfbs = (
self.captured_forward_batches_cross
if not skip
else self.captured_forward_batches
)
raw_bs = forward_batch.batch_size
if raw_bs in self.graphs:
if raw_bs in graphs:
# Keep encoder_out_cache_loc consistent with the captured graph (None).
if self.is_encoder_decoder:
# encoder_out_cache_loc is never accessed during decode (k/v are
# None so the KV-write path is skipped in the kernel). Use None
# consistently at both capture time and runtime.
forward_batch.encoder_out_cache_loc = None
self.model_runner.attn_backend.init_forward_metadata(forward_batch)
return forward_batch
@@ -782,7 +882,7 @@ class CPUGraphRunner:
self.raw_num_token = raw_num_token
self.bs = bs
captured_forward_batch = self.captured_forward_batches[bs]
captured_forward_batch = cfbs[bs]
assert captured_forward_batch is not None
captured_forward_batch.seq_lens.fill_(self.seq_len_fill_value)
captured_forward_batch.out_cache_loc.zero_()
@@ -805,6 +905,7 @@ class CPUGraphRunner:
captured_forward_batch.encoder_lens[:raw_bs].copy_(
forward_batch.encoder_lens
)
captured_forward_batch.encoder_out_cache_loc = None
if enable_num_token_non_padded():
captured_forward_batch.num_token_non_padded.copy_(
forward_batch.num_token_non_padded
@@ -822,13 +923,27 @@ class CPUGraphRunner:
pp_proxy_tensors is None
), "PPProxyTensors is not supported in CPUGraphRunner yet."
prepared_forward_batch = self.prepare_replay(forward_batch)
output = self.graphs[prepared_forward_batch.batch_size](
prepared_forward_batch.input_ids,
prepared_forward_batch.positions,
prepared_forward_batch,
replay_context = (
model_capture_mode if self.is_encoder_decoder else empty_context
)
if forward_batch.batch_size in self.graphs:
# Determine which compiled graph to use and pin skip_cross_attention so
# that any torch.compile re-tracing sees the same compile-time constant.
skip = self._get_skip_cross_attention(forward_batch)
graphs = self.graphs_cross if not skip else self.graphs
skip_ctx = (
capture_with_skip_cross_attention(skip)
if self.is_encoder_decoder
else empty_context()
)
with replay_context():
with skip_ctx:
prepared_forward_batch = self.prepare_replay(forward_batch, skip=skip)
output = graphs[prepared_forward_batch.batch_size](
prepared_forward_batch.input_ids,
prepared_forward_batch.positions,
prepared_forward_batch,
)
if forward_batch.batch_size in graphs:
return output
assert isinstance(output, LogitsProcessorOutput)
+12 -3
View File
@@ -968,9 +968,18 @@ class MllamaForConditionalGeneration(nn.Module):
cross_attention_states = None
if get_is_capture_mode():
# NOTE: when doing cuda graph capture, we do not want to skip cross attention
# Make is a constant value to avoid cuda graph capture issue
skip_cross_attention = False
# NOTE: during graph capture/replay skip_cross_attention must be a
# compile-time constant to avoid graph breaks from data-dependent
# branching. CPUGraphRunner captures two graphs per batch size (one
# with skip=True, one with skip=False) and sets the override via
# capture_with_skip_cross_attention(); fall back to False for CUDA
# graph capture which only captures the no-skip variant.
from sglang.srt.model_executor.cpu_graph_runner import (
get_capture_skip_cross_attention,
)
_override = get_capture_skip_cross_attention()
skip_cross_attention = _override if _override is not None else False
else:
# NOTE: we do not need image_inputs when prefill
assert len(forward_batch.encoder_lens) == len(forward_batch.seq_lens)