From 2d1856bf45a613184f3916d0a8659c0083e286fb Mon Sep 17 00:00:00 2001 From: Cao E Date: Mon, 8 Jun 2026 10:30:11 +0800 Subject: [PATCH] Support encoder_decoder on cpu_graph_runner (#10950) --- .../srt/model_executor/cpu_graph_runner.py | 227 +++++++++++++----- python/sglang/srt/models/mllama.py | 15 +- 2 files changed, 183 insertions(+), 59 deletions(-) diff --git a/python/sglang/srt/model_executor/cpu_graph_runner.py b/python/sglang/srt/model_executor/cpu_graph_runner.py index 07028f04d..8c059f7f4 100644 --- a/python/sglang/srt/model_executor/cpu_graph_runner.py +++ b/python/sglang/srt/model_executor/cpu_graph_runner.py @@ -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) diff --git a/python/sglang/srt/models/mllama.py b/python/sglang/srt/models/mllama.py index ba50cf8eb..771ef54ac 100644 --- a/python/sglang/srt/models/mllama.py +++ b/python/sglang/srt/models/mllama.py @@ -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)