Support encoder_decoder on cpu_graph_runner (#10950)
This commit is contained in:
@@ -29,6 +29,7 @@ import tqdm
|
|||||||
from sglang.srt.distributed import get_tensor_model_parallel_rank
|
from sglang.srt.distributed import get_tensor_model_parallel_rank
|
||||||
from sglang.srt.distributed.parallel_state import GroupCoordinator
|
from sglang.srt.distributed.parallel_state import GroupCoordinator
|
||||||
from sglang.srt.layers.logits_processor import LogitsProcessorOutput
|
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 (
|
from sglang.srt.model_executor.forward_batch_info import (
|
||||||
CaptureHiddenMode,
|
CaptureHiddenMode,
|
||||||
ForwardBatch,
|
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.model_executor.forward_context import ForwardContext, forward_context
|
||||||
from sglang.srt.utils import (
|
from sglang.srt.utils import (
|
||||||
|
empty_context,
|
||||||
log_info_on_rank0,
|
log_info_on_rank0,
|
||||||
require_attn_tp_gather,
|
require_attn_tp_gather,
|
||||||
require_gathered_buffer,
|
require_gathered_buffer,
|
||||||
@@ -48,6 +50,34 @@ from sglang.srt.utils.patch_torch import monkey_patch_torch_compile
|
|||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
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:
|
if TYPE_CHECKING:
|
||||||
from sglang.srt.model_executor.model_runner import ModelRunner
|
from sglang.srt.model_executor.model_runner import ModelRunner
|
||||||
|
|
||||||
@@ -484,7 +514,10 @@ class CPUGraphRunner:
|
|||||||
# Parse args
|
# Parse args
|
||||||
self.model_runner = model_runner
|
self.model_runner = model_runner
|
||||||
self.device = model_runner.device
|
self.device = model_runner.device
|
||||||
|
# bs -> compiled fn (text-only / skip_cross_attention=True)
|
||||||
self.graphs = {}
|
self.graphs = {}
|
||||||
|
# bs -> compiled fn (cross-attention / skip_cross_attention=False, enc-dec only)
|
||||||
|
self.graphs_cross = {}
|
||||||
self.output_buffers = {}
|
self.output_buffers = {}
|
||||||
self.enable_torch_compile = model_runner.server_args.enable_torch_compile
|
self.enable_torch_compile = model_runner.server_args.enable_torch_compile
|
||||||
self.disable_padding = model_runner.server_args.disable_cuda_graph_padding
|
self.disable_padding = model_runner.server_args.disable_cuda_graph_padding
|
||||||
@@ -530,17 +563,17 @@ class CPUGraphRunner:
|
|||||||
assert (
|
assert (
|
||||||
model_runner.spec_algorithm.is_none()
|
model_runner.spec_algorithm.is_none()
|
||||||
), "CPUGraphRunner does not support speculative inference yet."
|
), "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.dp_size == 1, "CPUGraphRunner does not support DP yet."
|
||||||
assert self.pp_size == 1, "CPUGraphRunner does not support PP yet."
|
assert self.pp_size == 1, "CPUGraphRunner does not support PP yet."
|
||||||
|
|
||||||
# Batch sizes to capture
|
# Batch sizes to capture
|
||||||
self.capture_bs = get_batch_sizes_to_capture(model_runner)
|
self.capture_bs = get_batch_sizes_to_capture(model_runner)
|
||||||
log_info_on_rank0(logger, f"Capture cpu graph bs {self.capture_bs}")
|
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 = {}
|
self.captured_forward_batches = {}
|
||||||
|
# bs -> ForwardBatch (cross-attention / skip=False, enc-dec only)
|
||||||
|
self.captured_forward_batches_cross = {}
|
||||||
# Attention backend
|
# Attention backend
|
||||||
self.max_bs = max(self.capture_bs)
|
self.max_bs = max(self.capture_bs)
|
||||||
self.max_num_token = self.max_bs * self.num_tokens_per_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.max_bs, self.max_num_token
|
||||||
)
|
)
|
||||||
|
|
||||||
|
self.encoder_len_fill_value = 0
|
||||||
self.seq_len_fill_value = (
|
self.seq_len_fill_value = (
|
||||||
self.model_runner.attn_backend.get_cpu_graph_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,
|
dtype=torch.bool,
|
||||||
device=self.device,
|
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
|
# Capture
|
||||||
try:
|
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:
|
except RuntimeError as e:
|
||||||
raise Exception(
|
raise Exception(
|
||||||
f"Capture CPU graph failed: {e}\n{CPU_GRAPH_CAPTURE_FAILED_MSG}"
|
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):
|
def can_run(self, forward_batch: ForwardBatch):
|
||||||
is_bs_supported = (
|
is_bs_supported = (
|
||||||
forward_batch.batch_size in self.graphs
|
forward_batch.batch_size in self.graphs
|
||||||
@@ -626,12 +686,18 @@ class CPUGraphRunner:
|
|||||||
num_tokens=bs * self.num_tokens_per_bs,
|
num_tokens=bs * self.num_tokens_per_bs,
|
||||||
tp_group=self.model_runner.tp_group,
|
tp_group=self.model_runner.tp_group,
|
||||||
) as forward:
|
) as forward:
|
||||||
(
|
graph, output_buffers = self.capture_one_batch_size(
|
||||||
graph,
|
bs, forward, skip_cross_attention=True
|
||||||
output_buffers,
|
)
|
||||||
) = self.capture_one_batch_size(bs, forward)
|
|
||||||
self.graphs[bs] = graph
|
self.graphs[bs] = graph
|
||||||
self.output_buffers[bs] = output_buffers
|
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
|
# Re-init states for qwen3-next as
|
||||||
# torch.compile may change the states
|
# torch.compile may change the states
|
||||||
@@ -656,7 +722,9 @@ class CPUGraphRunner:
|
|||||||
for v in vars(mamba_cache).values():
|
for v in vars(mamba_cache).values():
|
||||||
_zero_nested(v)
|
_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
|
num_tokens = bs * self.num_tokens_per_bs
|
||||||
|
|
||||||
# Graph inputs
|
# Graph inputs
|
||||||
@@ -667,6 +735,10 @@ class CPUGraphRunner:
|
|||||||
positions = self.positions[:num_tokens]
|
positions = self.positions[:num_tokens]
|
||||||
mrope_positions = self.mrope_positions[:, :num_tokens]
|
mrope_positions = self.mrope_positions[:, :num_tokens]
|
||||||
self.num_token_non_padded[...] = 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)
|
spec_info = self.get_spec_info(num_tokens)
|
||||||
if self.capture_hidden_mode != CaptureHiddenMode.FULL:
|
if self.capture_hidden_mode != CaptureHiddenMode.FULL:
|
||||||
@@ -682,6 +754,8 @@ class CPUGraphRunner:
|
|||||||
seq_lens=seq_lens,
|
seq_lens=seq_lens,
|
||||||
out_cache_loc=out_cache_loc,
|
out_cache_loc=out_cache_loc,
|
||||||
seq_lens_sum=seq_lens.sum().item(),
|
seq_lens_sum=seq_lens.sum().item(),
|
||||||
|
encoder_lens=encoder_lens,
|
||||||
|
encoder_lens_cpu=encoder_lens,
|
||||||
return_logprob=False,
|
return_logprob=False,
|
||||||
positions=positions,
|
positions=positions,
|
||||||
mrope_positions=mrope_positions,
|
mrope_positions=mrope_positions,
|
||||||
@@ -691,46 +765,58 @@ class CPUGraphRunner:
|
|||||||
num_token_non_padded=self.num_token_non_padded,
|
num_token_non_padded=self.num_token_non_padded,
|
||||||
global_forward_mode=self.capture_forward_mode,
|
global_forward_mode=self.capture_forward_mode,
|
||||||
)
|
)
|
||||||
with forward_context(
|
# Wrap all forward calls with capture_with_skip_cross_attention so that
|
||||||
ForwardContext(attn_backend=self.model_runner.attn_backend)
|
# mllama (and any other encoder-decoder model) sees the correct compile-
|
||||||
):
|
# time constant for skip_cross_attention during tracing.
|
||||||
self.model_runner.attn_backend.init_forward_metadata_capture_cpu_graph(
|
skip_ctx = (
|
||||||
bs,
|
capture_with_skip_cross_attention(skip_cross_attention)
|
||||||
num_tokens,
|
if self.is_encoder_decoder
|
||||||
req_pool_indices,
|
else empty_context()
|
||||||
seq_lens,
|
)
|
||||||
None,
|
with skip_ctx:
|
||||||
forward_batch.forward_mode,
|
with forward_context(
|
||||||
forward_batch.spec_info,
|
ForwardContext(attn_backend=self.model_runner.attn_backend)
|
||||||
)
|
):
|
||||||
with torch.no_grad():
|
self.model_runner.attn_backend.init_forward_metadata_capture_cpu_graph(
|
||||||
self.model_runner.tp_group.barrier()
|
bs,
|
||||||
self.model_runner.model.forward(
|
num_tokens,
|
||||||
forward_batch.input_ids,
|
req_pool_indices,
|
||||||
forward_batch.positions,
|
seq_lens,
|
||||||
forward_batch,
|
None,
|
||||||
|
forward_batch.forward_mode,
|
||||||
|
forward_batch.spec_info,
|
||||||
)
|
)
|
||||||
|
with torch.no_grad():
|
||||||
# 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()
|
self.model_runner.tp_group.barrier()
|
||||||
out = run_once()
|
self.model_runner.model.forward(
|
||||||
# Save the captured forward_batch
|
forward_batch.input_ids,
|
||||||
self.captured_forward_batches[bs] = forward_batch
|
forward_batch.positions,
|
||||||
return forward, out
|
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):
|
def recapture_if_needed(self, forward_batch: ForwardBatch):
|
||||||
|
|
||||||
@@ -766,11 +852,25 @@ class CPUGraphRunner:
|
|||||||
def prepare_replay(
|
def prepare_replay(
|
||||||
self,
|
self,
|
||||||
forward_batch: ForwardBatch,
|
forward_batch: ForwardBatch,
|
||||||
|
skip: bool = False,
|
||||||
):
|
):
|
||||||
self.recapture_if_needed(forward_batch)
|
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
|
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)
|
self.model_runner.attn_backend.init_forward_metadata(forward_batch)
|
||||||
return forward_batch
|
return forward_batch
|
||||||
|
|
||||||
@@ -782,7 +882,7 @@ class CPUGraphRunner:
|
|||||||
self.raw_num_token = raw_num_token
|
self.raw_num_token = raw_num_token
|
||||||
self.bs = bs
|
self.bs = bs
|
||||||
|
|
||||||
captured_forward_batch = self.captured_forward_batches[bs]
|
captured_forward_batch = cfbs[bs]
|
||||||
assert captured_forward_batch is not None
|
assert captured_forward_batch is not None
|
||||||
captured_forward_batch.seq_lens.fill_(self.seq_len_fill_value)
|
captured_forward_batch.seq_lens.fill_(self.seq_len_fill_value)
|
||||||
captured_forward_batch.out_cache_loc.zero_()
|
captured_forward_batch.out_cache_loc.zero_()
|
||||||
@@ -805,6 +905,7 @@ class CPUGraphRunner:
|
|||||||
captured_forward_batch.encoder_lens[:raw_bs].copy_(
|
captured_forward_batch.encoder_lens[:raw_bs].copy_(
|
||||||
forward_batch.encoder_lens
|
forward_batch.encoder_lens
|
||||||
)
|
)
|
||||||
|
captured_forward_batch.encoder_out_cache_loc = None
|
||||||
if enable_num_token_non_padded():
|
if enable_num_token_non_padded():
|
||||||
captured_forward_batch.num_token_non_padded.copy_(
|
captured_forward_batch.num_token_non_padded.copy_(
|
||||||
forward_batch.num_token_non_padded
|
forward_batch.num_token_non_padded
|
||||||
@@ -822,13 +923,27 @@ class CPUGraphRunner:
|
|||||||
pp_proxy_tensors is None
|
pp_proxy_tensors is None
|
||||||
), "PPProxyTensors is not supported in CPUGraphRunner yet."
|
), "PPProxyTensors is not supported in CPUGraphRunner yet."
|
||||||
|
|
||||||
prepared_forward_batch = self.prepare_replay(forward_batch)
|
replay_context = (
|
||||||
output = self.graphs[prepared_forward_batch.batch_size](
|
model_capture_mode if self.is_encoder_decoder else empty_context
|
||||||
prepared_forward_batch.input_ids,
|
|
||||||
prepared_forward_batch.positions,
|
|
||||||
prepared_forward_batch,
|
|
||||||
)
|
)
|
||||||
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
|
return output
|
||||||
|
|
||||||
assert isinstance(output, LogitsProcessorOutput)
|
assert isinstance(output, LogitsProcessorOutput)
|
||||||
|
|||||||
@@ -968,9 +968,18 @@ class MllamaForConditionalGeneration(nn.Module):
|
|||||||
cross_attention_states = None
|
cross_attention_states = None
|
||||||
|
|
||||||
if get_is_capture_mode():
|
if get_is_capture_mode():
|
||||||
# NOTE: when doing cuda graph capture, we do not want to skip cross attention
|
# NOTE: during graph capture/replay skip_cross_attention must be a
|
||||||
# Make is a constant value to avoid cuda graph capture issue
|
# compile-time constant to avoid graph breaks from data-dependent
|
||||||
skip_cross_attention = False
|
# 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:
|
else:
|
||||||
# NOTE: we do not need image_inputs when prefill
|
# NOTE: we do not need image_inputs when prefill
|
||||||
assert len(forward_batch.encoder_lens) == len(forward_batch.seq_lens)
|
assert len(forward_batch.encoder_lens) == len(forward_batch.seq_lens)
|
||||||
|
|||||||
Reference in New Issue
Block a user