[diffusion] Enable breakable CUDA graph (BCG) for diffusion DiTs (#27436)

Co-authored-by: BBuf <xiaoyu.zhang@radixark.net>
Co-authored-by: Claude Opus 4.8 <noreply@anthropic.com>
Co-authored-by: BBuf <bbuf@sglang.local>
This commit is contained in:
Xiaoyu Zhang
2026-07-08 14:45:48 +08:00
committed by GitHub
co-authored by BBuf Claude Opus 4.8 BBuf
parent c9303a08da
commit 33c3dfd7e0
31 changed files with 2952 additions and 503 deletions
@@ -58,26 +58,30 @@ class GlmImagePipelineConfig(SpatialImagePipelineConfig):
return cos, sin return cos, sin
def prepare_pos_cond_kwargs(self, batch, device, rotary_emb, dtype): def prepare_pos_cond_kwargs(self, batch, device, rotary_emb, dtype):
return { kwargs = {
"prior_token_id": batch.prior_token_id, "prior_token_id": batch.prior_token_id,
"prior_token_drop": batch.prior_token_drop_cond, "prior_token_drop": batch.prior_token_drop_cond,
"crop_coords": batch.crop_coords, "crop_coords": batch.crop_coords,
"target_size": batch.target_size, "target_size": batch.target_size,
"kv_caches": batch.kv_caches,
"kv_caches_mode": "read",
"freqs_cis": self.get_freqs_cis(batch, device, rotary_emb, dtype), "freqs_cis": self.get_freqs_cis(batch, device, rotary_emb, dtype),
} }
if getattr(batch, "prior_token_image_ids", None) is not None:
kwargs["kv_caches"] = batch.kv_caches
kwargs["kv_caches_mode"] = "read"
return kwargs
def prepare_neg_cond_kwargs(self, batch, device, rotary_emb, dtype): def prepare_neg_cond_kwargs(self, batch, device, rotary_emb, dtype):
return { kwargs = {
"prior_token_id": batch.prior_token_id, "prior_token_id": batch.prior_token_id,
"prior_token_drop": batch.prior_token_drop_uncond, "prior_token_drop": batch.prior_token_drop_uncond,
"crop_coords": batch.crop_coords, "crop_coords": batch.crop_coords,
"target_size": batch.target_size, "target_size": batch.target_size,
"kv_caches": batch.kv_caches,
"kv_caches_mode": "skip",
"freqs_cis": self.get_freqs_cis(batch, device, rotary_emb, dtype), "freqs_cis": self.get_freqs_cis(batch, device, rotary_emb, dtype),
} }
if getattr(batch, "prior_token_image_ids", None) is not None:
kwargs["kv_caches"] = batch.kv_caches
kwargs["kv_caches_mode"] = "skip"
return kwargs
def get_decode_scale_and_shift(self, device, dtype, vae): def get_decode_scale_and_shift(self, device, dtype, vae):
latents_mean = ( latents_mean = (
@@ -0,0 +1 @@
"""Diffusion breakable CUDA graph runtime helpers."""
@@ -0,0 +1 @@
"""Model-specific prompt padders for diffusion breakable CUDA graph."""
@@ -0,0 +1,131 @@
# Copyright 2023-2026 SGLang Team
# Licensed under the Apache License, Version 2.0
# ==============================================================================
"""Ideogram-4 breakable CUDA graph (BCG) prompt padding."""
from __future__ import annotations
from typing import Any
import torch
from sglang.multimodal_gen.runtime.breakable_cuda_graph import (
prompt_padding as bcg_utils,
)
from sglang.multimodal_gen.runtime.layers.attention import DynamicVarlenMaskMeta
_SEQUENCE_PADDING_INDICATOR = -1
_OUTPUT_IMAGE_INDICATOR = 2
_LLM_TOKEN_INDICATOR = 3
_DYNAMIC_MASK_META_ATTR = "_sglang_bcg_ideogram_attn_mask_meta"
def is_ideogram_transformer(current_model: Any, call_kwargs: dict) -> bool:
return (
bcg_utils.transformer_class_name_matches(current_model, "ideogram")
and "llm_features" in call_kwargs
and "x" in call_kwargs
and "indicator" in call_kwargs
and "position_ids" in call_kwargs
)
def _unwrap_model(current_model: Any) -> Any:
for attr in ("module", "_orig_mod"):
wrapped = getattr(current_model, attr, None)
if wrapped is not None:
current_model = wrapped
return current_model
def _dynamic_mask_meta(current_model: Any) -> DynamicVarlenMaskMeta:
model = _unwrap_model(current_model)
meta = getattr(model, _DYNAMIC_MASK_META_ATTR, None)
if not isinstance(meta, DynamicVarlenMaskMeta):
meta = DynamicVarlenMaskMeta()
setattr(model, _DYNAMIC_MASK_META_ATTR, meta)
return meta
def _first_indicator(call_kwargs: dict) -> torch.Tensor | None:
indicator = bcg_utils.first_tensor(call_kwargs.get("indicator"))
if not torch.is_tensor(indicator) or indicator.dim() < 2:
return None
return indicator
def _text_and_image_lengths(indicator: torch.Tensor) -> tuple[int, int] | None:
row = indicator[0]
if not torch.any(row == _LLM_TOKEN_INDICATOR):
return None
image_positions = (row == _OUTPUT_IMAGE_INDICATOR).nonzero(as_tuple=False)
if image_positions.numel() == 0:
return None
text_seq = int(image_positions[0].item())
if text_seq <= 0:
return None
image_seq = int(row.numel()) - text_seq
if image_seq <= 0:
return None
return text_seq, image_seq
def _pad_total_dim(obj: Any, *, source: int, target: int, value: float = 0) -> Any:
return bcg_utils.pad_nested_dim(
obj, dim=1, source=source, target=target, value=value
)
def pad_ideogram_prompt_kwargs(
call_kwargs: dict, current_model: Any, buckets: tuple[int, ...]
) -> dict:
indicator = _first_indicator(call_kwargs)
if indicator is None:
return call_kwargs
lengths = _text_and_image_lengths(indicator)
if lengths is None:
return call_kwargs
text_seq, image_seq = lengths
bucket = bcg_utils.select_text_bucket(text_seq, buckets)
if bucket is None:
return call_kwargs
source_total = text_seq + image_seq
target_total = bucket + image_seq
out = dict(call_kwargs)
if source_total < target_total:
for key in ("llm_features", "x"):
if key in out and out[key] is not None:
out[key] = _pad_total_dim(
out[key], source=source_total, target=target_total
)
if out.get("position_ids") is not None:
out["position_ids"] = _pad_total_dim(
out["position_ids"], source=source_total, target=target_total
)
if out.get("segment_ids") is not None:
out["segment_ids"] = _pad_total_dim(
out["segment_ids"],
source=source_total,
target=target_total,
value=_SEQUENCE_PADDING_INDICATOR,
)
if out.get("indicator") is not None:
out["indicator"] = _pad_total_dim(
out["indicator"], source=source_total, target=target_total
)
if out.get("attn_mask") is not None:
out["attn_mask"] = _pad_total_dim(
out["attn_mask"], source=source_total, target=target_total
)
if out.get("attn_mask") is not None:
out["attn_mask_meta"] = _dynamic_mask_meta(current_model)
return out
bcg_utils.register_prompt_padder(is_ideogram_transformer, pad_ideogram_prompt_kwargs)
@@ -0,0 +1,102 @@
# Copyright 2023-2026 SGLang Team
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
# ==============================================================================
"""Qwen-Image breakable CUDA graph (BCG) prompt padding.
Qwen-Image / Qwen-Image-Edit carry text length on dim 1 of
``encoder_hidden_states`` and a separate ``freqs_cis`` text-rope cache plus
``txt_seq_lens``; they may not pass an explicit prompt mask, so this padder
synthesizes one. Registered with the base denoising stage's padder registry.
"""
from __future__ import annotations
from typing import Any
import torch
from sglang.multimodal_gen.runtime.breakable_cuda_graph import (
prompt_padding as bcg_utils,
)
def is_qwen_transformer(current_model: Any, call_kwargs: dict) -> bool:
return (
bcg_utils.transformer_class_name_matches(current_model, "qwen")
and "txt_seq_lens" in call_kwargs
and "freqs_cis" in call_kwargs
)
def pad_qwen_prompt_kwargs(
call_kwargs: dict, current_model: Any, buckets: tuple[int, ...]
) -> dict:
ehs = call_kwargs.get("encoder_hidden_states")
ehs_tensor = bcg_utils.first_tensor(ehs)
if not torch.is_tensor(ehs_tensor) or ehs_tensor.dim() < 2:
return call_kwargs
seq = ehs_tensor.shape[1]
bucket = bcg_utils.select_text_bucket(seq, buckets)
if bucket is None:
return call_kwargs
out = dict(call_kwargs)
if seq < bucket:
out["encoder_hidden_states"] = bcg_utils.pad_nested_dim(
ehs, dim=1, source=seq, target=bucket
)
if (
"encoder_hidden_states_2" in out
and out["encoder_hidden_states_2"] is not None
):
out["encoder_hidden_states_2"] = bcg_utils.pad_nested_dim(
out["encoder_hidden_states_2"], dim=1, source=seq, target=bucket
)
mask = out.get("encoder_hidden_states_mask")
if mask is None:
mask = torch.ones(
ehs_tensor.shape[:2],
device=ehs_tensor.device,
dtype=torch.bool,
)
if mask is not None:
out["encoder_hidden_states_mask"] = bcg_utils.pad_nested_dim(
mask, dim=1, source=seq, target=bucket
)
if "encoder_attention_mask" in out and out["encoder_attention_mask"] is not None:
out["encoder_attention_mask"] = bcg_utils.pad_nested_dim(
out["encoder_attention_mask"], dim=1, source=seq, target=bucket
)
freqs_cis = out.get("freqs_cis")
if isinstance(freqs_cis, tuple) and len(freqs_cis) == 2:
img_cache, txt_cache = freqs_cis
txt_cache = bcg_utils.pad_nested_dim(
txt_cache, dim=0, source=seq, target=bucket
)
out["freqs_cis"] = (img_cache, txt_cache)
elif isinstance(freqs_cis, list) and len(freqs_cis) == 2:
img_cache, txt_cache = freqs_cis
txt_cache = bcg_utils.pad_nested_dim(
txt_cache, dim=0, source=seq, target=bucket
)
out["freqs_cis"] = [img_cache, txt_cache]
out["txt_seq_lens"] = bcg_utils.bucket_txt_seq_lens(out.get("txt_seq_lens"), bucket)
return out
bcg_utils.register_prompt_padder(is_qwen_transformer, pad_qwen_prompt_kwargs)
@@ -0,0 +1,169 @@
# Copyright 2023-2026 SGLang Team
# Licensed under the Apache License, Version 2.0
# ==============================================================================
"""Z-Image breakable CUDA graph (BCG) prompt padding."""
from __future__ import annotations
from typing import Any
import torch
from sglang.multimodal_gen.runtime.breakable_cuda_graph import (
prompt_padding as bcg_utils,
)
def is_zimage_transformer(current_model: Any, call_kwargs: dict) -> bool:
return (
bcg_utils.transformer_class_name_matches(current_model, "zimage")
and "encoder_hidden_states" in call_kwargs
and "freqs_cis" in call_kwargs
)
def _first_caption_tensor(encoder_hidden_states: Any) -> torch.Tensor | None:
tensor = bcg_utils.first_tensor(encoder_hidden_states)
if not torch.is_tensor(tensor):
return None
if tensor.dim() == 2:
return tensor
if tensor.dim() == 3:
return tensor[0]
return None
def _caption_seq_len(tensor: torch.Tensor) -> int:
if tensor.dim() == 2:
return int(tensor.shape[0])
if tensor.dim() == 3:
return int(tensor.shape[1])
raise ValueError("Z-Image caption tensor must have rank 2 or 3")
def _pad_caption(obj: Any, *, target: int) -> Any:
if torch.is_tensor(obj):
if obj.dim() == 2:
return bcg_utils.pad_tensor_dim(obj, 0, target)
if obj.dim() == 3:
return bcg_utils.pad_tensor_dim(obj, 1, target)
return obj
if isinstance(obj, list):
return [_pad_caption(item, target=target) for item in obj]
if isinstance(obj, tuple):
return tuple(_pad_caption(item, target=target) for item in obj)
return obj
def _unwrap_model(current_model: Any) -> Any:
for attr in ("module", "_orig_mod"):
wrapped = getattr(current_model, attr, None)
if wrapped is not None:
current_model = wrapped
return current_model
def _build_caption_freqs(current_model: Any, *, target: int, device: torch.device):
rotary_emb = getattr(_unwrap_model(current_model), "rotary_emb", None)
if rotary_emb is None:
return None
axes = [
torch.arange(1, target + 1, dtype=torch.int32, device=device),
torch.zeros(target, dtype=torch.int32, device=device),
torch.zeros(target, dtype=torch.int32, device=device),
]
cap_pos_ids = torch.stack(axes, dim=-1)
return rotary_emb(cap_pos_ids)
def _pad_caption_freqs(freqs_cis: Any, current_model: Any, *, target: int) -> Any:
if not isinstance(freqs_cis, (tuple, list)) or len(freqs_cis) != 2:
return freqs_cis
cap_cache, image_cache = freqs_cis
cap_tensor = bcg_utils.first_tensor(cap_cache)
if torch.is_tensor(cap_tensor) and cap_tensor.dim() >= 1:
cap_freqs = _build_caption_freqs(
current_model, target=target, device=cap_tensor.device
)
if cap_freqs is not None:
cap_cache = cap_freqs
if isinstance(freqs_cis, tuple):
return (cap_cache, image_cache)
return [cap_cache, image_cache]
def _caption_mask(
call_kwargs: dict, *, caption: torch.Tensor, seq: int, bucket: int
) -> torch.Tensor:
mask = bcg_utils.first_tensor(call_kwargs.get("encoder_hidden_states_mask"))
if not torch.is_tensor(mask):
mask = bcg_utils.first_tensor(call_kwargs.get("encoder_attention_mask"))
if torch.is_tensor(mask):
if mask.dim() == 1:
mask = mask[:seq].unsqueeze(0)
elif mask.dim() >= 2:
mask = mask[:, :seq]
mask = mask.to(device=caption.device, dtype=torch.bool)
else:
batch = int(caption.shape[0]) if caption.dim() == 3 else 1
mask = torch.ones((batch, seq), device=caption.device, dtype=torch.bool)
return bcg_utils.pad_tensor_dim(mask, 1, bucket)
def pad_zimage_prompt_kwargs(
call_kwargs: dict, current_model: Any, buckets: tuple[int, ...]
) -> dict:
caption = _first_caption_tensor(call_kwargs.get("encoder_hidden_states"))
if caption is None:
return call_kwargs
seq = _caption_seq_len(caption)
cap_freq = None
freqs_cis = call_kwargs.get("freqs_cis")
if isinstance(freqs_cis, (tuple, list)) and len(freqs_cis) == 2:
cap_freq = bcg_utils.first_tensor(freqs_cis[0])
cap_freq_len = int(cap_freq.shape[0]) if torch.is_tensor(cap_freq) else seq
bucket = bcg_utils.select_text_bucket(max(seq, cap_freq_len), buckets)
if bucket is None:
return call_kwargs
out = {
key: value
for key, value in call_kwargs.items()
if key
in {
"hidden_states",
"timestep",
"guidance",
"encoder_hidden_states",
"encoder_attention_mask",
"encoder_hidden_states_mask",
"freqs_cis",
"image_seq_len_target",
"patch_size",
"f_patch_size",
}
}
if seq < bucket:
out["encoder_hidden_states"] = _pad_caption(
out["encoder_hidden_states"], target=bucket
)
caption_mask = _caption_mask(call_kwargs, caption=caption, seq=seq, bucket=bucket)
out["encoder_hidden_states_mask"] = caption_mask
out["caption_valid_lens"] = caption_mask.sum(dim=1).to(dtype=torch.long)
out["_use_caption_valid_mask"] = True
if out.get("encoder_attention_mask") is not None:
out["encoder_attention_mask"] = out["encoder_hidden_states_mask"]
out["freqs_cis"] = _pad_caption_freqs(
out.get("freqs_cis"), current_model, target=bucket
)
return out
bcg_utils.register_prompt_padder(is_zimage_transformer, pad_zimage_prompt_kwargs)
@@ -0,0 +1,306 @@
# Copyright 2023-2026 SGLang Team
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
# ==============================================================================
"""Utilities for breakable CUDA graph (BCG) prompt padding.
These helpers bucket prompt-conditioning inputs by sequence length so diffusion
DiT forward calls with different prompt lengths can reuse captured CUDA graphs.
Model-specific padders can register custom handling under
``breakable_cuda_graph.model_padders``.
"""
from __future__ import annotations
import logging
from typing import Any, Callable
import torch
logger = logging.getLogger(__name__)
# Prompt-conditioning kwarg keys, grouped by which dim carries the text length.
PROMPT_MASK_KEYS = (
"encoder_attention_mask",
"encoder_hidden_states_mask",
"attention_mask",
"text_mask",
"prompt_attention_mask",
"negative_attention_mask",
"prompt_embeds_mask",
"negative_prompt_embeds_mask",
)
TEXT_DIM1_KEYS = (
"encoder_hidden_states",
"encoder_hidden_states_2",
"encoder_attention_mask",
"encoder_hidden_states_mask",
"attention_mask",
"text_mask",
"text_ids",
"text_pos_ids",
"txt_ids",
"prompt_embeds",
"negative_prompt_embeds",
"prompt_attention_mask",
"negative_attention_mask",
"prompt_embeds_mask",
"negative_prompt_embeds_mask",
"audio_encoder_hidden_states",
"audio_encoder_attention_mask",
)
TEXT_DIM0_KEYS = (
"txt_freqs_cis",
"text_freqs_cis",
)
TEXT_SEQ_LEN_KEYS = (
"txt_seq_lens",
"text_seq_lens",
)
def first_tensor(obj: Any) -> torch.Tensor | None:
"""First tensor leaf found by depth-first traversal (dicts in sorted-key
order), or ``None``."""
if torch.is_tensor(obj):
return obj
if isinstance(obj, (list, tuple)):
for item in obj:
tensor = first_tensor(item)
if tensor is not None:
return tensor
if isinstance(obj, dict):
for key in sorted(obj):
tensor = first_tensor(obj[key])
if tensor is not None:
return tensor
return None
def select_text_bucket(seq: int, buckets: tuple[int, ...]) -> int | None:
"""Smallest bucket that fits ``seq``; ``None`` (and a warning) when ``seq``
exceeds the largest bucket so the caller runs that length eagerly."""
for bucket in buckets:
if seq <= bucket:
return bucket
logger.warning(
"[Diffusion BCG] text length %d exceeds max bucket %d; not padding "
"(this length captures its own graph). Raise --bcg-text-buckets.",
seq,
buckets[-1],
)
return None
def pad_tensor_dim(tensor: Any, dim: int, target: int, value: float = 0) -> Any:
if not torch.is_tensor(tensor) or tensor.dim() <= dim:
return tensor
seq = tensor.shape[dim]
if seq >= target:
return tensor
pad = [0, 0] * tensor.dim()
pad_index = 2 * (tensor.dim() - dim - 1) + 1
pad[pad_index] = target - seq
return torch.nn.functional.pad(tensor, tuple(pad), value=value)
def pad_nested_dim(
obj: Any,
*,
dim: int,
source: int,
target: int,
value: float = 0,
) -> Any:
if torch.is_tensor(obj):
if obj.dim() > dim and obj.shape[dim] == source:
return pad_tensor_dim(obj, dim, target, value)
return obj
if isinstance(obj, list):
return [
pad_nested_dim(item, dim=dim, source=source, target=target, value=value)
for item in obj
]
if isinstance(obj, tuple):
return tuple(
pad_nested_dim(item, dim=dim, source=source, target=target, value=value)
for item in obj
)
return obj
def bucket_txt_seq_lens(txt_seq_lens: Any, bucket: int) -> Any:
if txt_seq_lens is None:
return txt_seq_lens
if torch.is_tensor(txt_seq_lens):
return torch.full_like(txt_seq_lens, bucket)
if isinstance(txt_seq_lens, list):
return [bucket_txt_seq_lens(seq_len, bucket) for seq_len in txt_seq_lens]
if isinstance(txt_seq_lens, tuple):
return tuple(bucket_txt_seq_lens(seq_len, bucket) for seq_len in txt_seq_lens)
if isinstance(txt_seq_lens, int):
return bucket
return txt_seq_lens
def prompt_seq_and_dim(call_kwargs: dict) -> tuple[int, int] | None:
"""Return ``(text_seq_len, seq_dim)`` inferred from the prompt embeddings or
a prompt mask, or ``None`` when no text conditioning is present."""
ehs_tensor = first_tensor(call_kwargs.get("encoder_hidden_states"))
if torch.is_tensor(ehs_tensor) and ehs_tensor.dim() >= 2:
if ehs_tensor.dim() == 2:
return int(ehs_tensor.shape[0]), 0
return int(ehs_tensor.shape[1]), 1
for key in PROMPT_MASK_KEYS:
tensor = first_tensor(call_kwargs.get(key))
if torch.is_tensor(tensor) and tensor.dim() >= 2:
if tensor.shape[0] == 1:
return int(tensor.shape[1]), 1
return int(tensor.shape[0]), 0
return None
def pad_nested_text_dim(
obj: Any,
*,
source: int,
target: int,
preferred_dim: int,
) -> Any:
if torch.is_tensor(obj):
if obj.dim() > preferred_dim and obj.shape[preferred_dim] == source:
return pad_tensor_dim(obj, preferred_dim, target)
for dim in (1, 0):
if dim != preferred_dim and obj.dim() > dim and obj.shape[dim] == source:
return pad_tensor_dim(obj, dim, target)
return obj
if isinstance(obj, list):
return [
pad_nested_text_dim(
item, source=source, target=target, preferred_dim=preferred_dim
)
for item in obj
]
if isinstance(obj, tuple):
return tuple(
pad_nested_text_dim(
item, source=source, target=target, preferred_dim=preferred_dim
)
for item in obj
)
if isinstance(obj, dict):
return {
key: pad_nested_text_dim(
value, source=source, target=target, preferred_dim=preferred_dim
)
for key, value in obj.items()
}
return obj
def bucket_text_seq_lens(obj: Any, *, target: int) -> Any:
if isinstance(obj, int) and not isinstance(obj, bool):
return target
if isinstance(obj, list):
return [bucket_text_seq_lens(item, target=target) for item in obj]
if isinstance(obj, tuple):
return tuple(bucket_text_seq_lens(item, target=target) for item in obj)
return obj
def pad_masked_prompt_kwargs(call_kwargs: dict, buckets: tuple[int, ...]) -> dict:
"""Generic, model-agnostic prompt padding for models that pass a prompt
attention mask alongside their text embeddings."""
seq_and_dim = prompt_seq_and_dim(call_kwargs)
if seq_and_dim is None:
return call_kwargs
seq, seq_dim = seq_and_dim
has_mask = any(
first_tensor(call_kwargs.get(key)) is not None for key in PROMPT_MASK_KEYS
)
if not has_mask:
return call_kwargs
bucket = select_text_bucket(seq, buckets)
if bucket is None or seq == bucket:
return call_kwargs
out = dict(call_kwargs)
for key in TEXT_DIM1_KEYS:
if key in out and out[key] is not None:
out[key] = pad_nested_text_dim(
out[key], source=seq, target=bucket, preferred_dim=seq_dim
)
for key in TEXT_DIM0_KEYS:
if key in out and out[key] is not None:
out[key] = pad_nested_dim(out[key], dim=0, source=seq, target=bucket)
for key in TEXT_SEQ_LEN_KEYS:
if key in out and out[key] is not None:
out[key] = bucket_text_seq_lens(out[key], target=bucket)
return out
def transformer_class_name_matches(current_model: Any, needle: str) -> bool:
"""True when ``current_model`` (or its ``module`` / ``_orig_mod`` wrapper)
is a transformer whose qualified class name contains ``needle``."""
candidates = [current_model]
for attr in ("module", "_orig_mod"):
wrapped = getattr(current_model, attr, None)
if wrapped is not None:
candidates.append(wrapped)
for candidate in candidates:
cls = type(candidate)
name = f"{cls.__module__}.{cls.__qualname__}".lower()
if needle in name:
return True
return False
# --- Model-specific prompt-padder registry ------------------------------- #
# Each model that needs custom prompt padding registers a (predicate, padder)
# pair from its own module in ``model_specific_stages`` so the base denoising
# stage stays model-agnostic. ``padder(call_kwargs, current_model, buckets)``
# returns the padded kwargs.
PromptPadder = Callable[[dict, Any, tuple], dict]
_PROMPT_PADDERS: list[tuple[Callable[[Any, dict], bool], PromptPadder]] = []
def register_prompt_padder(
predicate: Callable[[Any, dict], bool], padder: PromptPadder
) -> None:
_PROMPT_PADDERS.append((predicate, padder))
def select_prompt_padder(current_model: Any, call_kwargs: dict) -> PromptPadder | None:
"""Return the registered model-specific padder for ``current_model``, or
``None`` to fall back to :func:`pad_masked_prompt_kwargs`."""
_ensure_model_padders_registered()
for predicate, padder in _PROMPT_PADDERS:
if predicate(current_model, call_kwargs):
return padder
return None
_model_padders_registered = False
def _ensure_model_padders_registered() -> None:
"""Import the model-specific padder modules once so they register."""
global _model_padders_registered
if _model_padders_registered:
return
_model_padders_registered = True
from sglang.multimodal_gen.runtime.breakable_cuda_graph.model_padders import ( # noqa: F401
ideogram,
qwen_image,
zimage,
)
@@ -0,0 +1,468 @@
# Copyright 2023-2026 SGLang Team
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
# ==============================================================================
"""Breakable CUDA graph (BCG) runner for diffusion DiT transformers.
A runner wraps a callable ``nn.Module`` and turns it into an *eager runner* that
transparently proxies every attribute to the wrapped module and, when called,
replays a previously captured graph for the input signature — or runs the
module eagerly when no graph was captured for that signature. Capture is an
explicit, idempotent ``capture()`` call (driven at warmup) so that serving never
triggers a fresh capture.
This file is intentionally local to ``multimodal_gen``: diffusion reuses the
low-level SRT BCG primitives, but the capture/replay runner owns diffusion DiT
signature handling, static tensor buffers, prompt-bucket warmup, and fallback
behavior.
"""
from __future__ import annotations
import logging
import os
from dataclasses import dataclass
from typing import Any
import torch
import torch.nn as nn
from sglang.srt.model_executor.runner_backend_utils.breakable_cuda_graph.breakable_cuda_graph import (
BreakableCUDAGraph,
BreakableCUDAGraphCapture,
)
from sglang.srt.model_executor.runner_backend_utils.breakable_cuda_graph.context import (
enable_breakable_cuda_graph,
)
# Log under the multimodal_gen namespace so the diffusion server's logging
# config surfaces the "[Diffusion BCG] captured ..." lines.
logger = logging.getLogger(__name__)
def _env_int(name: str, default: int) -> int:
raw = os.environ.get(name)
if raw is None:
return default
try:
return int(raw)
except ValueError:
logger.warning("[BCG] ignoring invalid integer %s=%r", name, raw)
return default
def _env_float(name: str, default: float) -> float:
raw = os.environ.get(name)
if raw is None:
return default
try:
return float(raw)
except ValueError:
logger.warning("[BCG] ignoring invalid float %s=%r", name, raw)
return default
def _map_tensors(obj, fn):
"""Rebuild ``obj`` applying ``fn`` to every tensor leaf, recursing into
list/tuple/dict containers; everything else passes through unchanged."""
if torch.is_tensor(obj):
return fn(obj)
if isinstance(obj, tuple):
return tuple(_map_tensors(o, fn) for o in obj)
if isinstance(obj, list):
return [_map_tensors(o, fn) for o in obj]
if isinstance(obj, dict):
return {k: _map_tensors(v, fn) for k, v in obj.items()}
return obj
def _flatten_tensors(obj, out: list):
"""Depth-first collect every tensor leaf into ``out`` (deterministic order:
dicts traversed in sorted-key order to match across calls)."""
if torch.is_tensor(obj):
out.append(obj)
elif isinstance(obj, (list, tuple)):
for o in obj:
_flatten_tensors(o, out)
elif isinstance(obj, dict):
for k in sorted(obj):
_flatten_tensors(obj[k], out)
def _flatten_kwargs(kwargs: dict[str, Any]) -> list[torch.Tensor]:
out: list[torch.Tensor] = []
for name in sorted(kwargs):
_flatten_tensors(kwargs[name], out)
return out
def _signature_leaf(obj: Any) -> Any:
if torch.is_tensor(obj):
return ("tensor", tuple(obj.shape), str(obj.dtype))
if isinstance(obj, tuple):
return ("tuple", tuple(_signature_leaf(o) for o in obj))
if isinstance(obj, list):
return ("list", tuple(_signature_leaf(o) for o in obj))
if isinstance(obj, dict):
return (
"dict",
tuple((k, _signature_leaf(obj[k])) for k in sorted(obj)),
)
if obj is None or isinstance(obj, (bool, int, float, str)):
return ("const", obj)
return ("object", type(obj).__module__, type(obj).__qualname__, id(obj))
def _signature_kwargs(kwargs: dict[str, Any]) -> tuple:
return tuple((name, _signature_leaf(kwargs[name])) for name in sorted(kwargs))
def _signature_summary_leaf(sig: Any, *, depth: int = 0) -> Any:
if not isinstance(sig, tuple) or not sig:
return sig
tag = sig[0]
if tag == "tensor":
return sig
if tag == "const":
value = sig[1]
if isinstance(value, str) and len(value) > 64:
value = value[:61] + "..."
return (tag, value)
if tag == "object":
return sig[:3]
if depth >= 2:
return (tag, "...")
if tag in ("tuple", "list"):
items = sig[1]
preview = tuple(
_signature_summary_leaf(item, depth=depth + 1) for item in items[:4]
)
if len(items) > 4:
preview += (("...", len(items) - 4),)
return (tag, len(items), preview)
if tag == "dict":
items = sig[1]
preview = tuple(
(key, _signature_summary_leaf(value, depth=depth + 1))
for key, value in items[:4]
)
if len(items) > 4:
preview += (("...", len(items) - 4),)
return (tag, len(items), preview)
return sig
def _signature_summary(key: tuple) -> tuple:
return tuple((name, _signature_summary_leaf(value)) for name, value in key[:16]) + (
(("...", len(key) - 16),) if len(key) > 16 else ()
)
def _clone_output(out: Any) -> Any:
if torch.is_tensor(out):
return out.clone()
if isinstance(out, tuple):
return tuple(_clone_output(o) for o in out)
if isinstance(out, list):
return [_clone_output(o) for o in out]
return out
@dataclass
class _CaptureEntry:
graph: BreakableCUDAGraph
# full captured kwargs with persistent static buffers at every tensor leaf
static_kwargs: dict[str, Any]
# the same static buffers, flattened in _flatten_kwargs order (replay copies
# live tensors into these positionally)
static_leaves: list[torch.Tensor]
output: Any
num_segments: int
class _CaptureRejected(RuntimeError):
pass
class BaseBreakableCudaGraphRunner:
"""Eager runner around ``transformer`` with an explicit capture/replay API.
The capture/replay contract:
* :meth:`capture` captures a BCG graph for the given input signature, once
(idempotent). It is intended to be driven at warmup so that every
signature served later is already captured.
* :meth:`replay` copies live inputs into the captured static buffers and
replays the graph, returning a clone of the captured output.
* :meth:`__call__` is the *eager runner*: it replays when a graph exists for
the signature and otherwise runs ``transformer`` eagerly. It never
captures, so serving never pays a capture cost.
Any attribute not defined on the runner is proxied to ``transformer`` so the
runner can stand in for the wrapped module ("other functions directly
pass").
"""
def __init__(
self,
transformer: nn.Module,
device: torch.device,
pool=None,
) -> None:
self.transformer = transformer
self.device = device
self.device_module = torch.get_device_module(device)
# One shared mempool across all captured graphs/segments so per-block
# intermediates can be reclaimed and weak-ref'd safely.
self._pool = (
pool if pool is not None else self.device_module.graph_pool_handle()
)
self._capture_stream = self.device_module.Stream(device=device)
self.entries: dict[tuple, _CaptureEntry] = {}
# Signatures we have given up capturing (capture raised); run eager.
self._blocked: set[tuple] = set()
self._disabled_reason: str | None = None
self.max_entries = max(0, _env_int("SGLANG_DIFFUSION_BCG_MAX_ENTRIES", 32))
self.max_segments = max(0, _env_int("SGLANG_DIFFUSION_BCG_MAX_SEGMENTS", 128))
def __getattr__(self, name: str) -> Any:
# Only reached for attributes the runner itself does not define; proxy
# them to the wrapped transformer so callers can treat the runner as a
# transparent stand-in. Use __dict__ to avoid recursing through
# __getattr__ before ``transformer`` is assigned in __init__.
try:
transformer = self.__dict__["transformer"]
except KeyError as e: # pragma: no cover - during/ before __init__
raise AttributeError(name) from e
return getattr(transformer, name)
# ------------------------------------------------------------------ #
# Public capture / replay API
# ------------------------------------------------------------------ #
@torch.no_grad()
def capture(self, **kwargs) -> bool:
"""Capture a graph for ``kwargs``'s signature if not already captured.
Idempotent: returns ``True`` when a graph is available for the
signature afterwards (already captured or newly captured), ``False``
when capture is disabled/blocked or failed (the caller then runs eager).
"""
if self._disabled_reason is not None:
return False
key = self._signature(kwargs)
if key in self._blocked:
return False
if key in self.entries:
return True
try:
entry = self._capture(kwargs, key)
except Exception as e: # noqa: BLE001 — never break generation on capture
logger.warning(
"[Diffusion BCG] capture failed for signature %s (%s); "
"this signature will run eager.",
_signature_summary(key),
e,
)
self._blocked.add(key)
return False
self.entries[key] = entry
self._evict_entries_if_needed()
return True
def _should_capture_on_call(self, key: tuple) -> bool:
"""Whether ``__call__`` may lazily capture an unseen signature.
Base runners only ever capture through the explicit :meth:`capture`
API, so this returns ``False``: serving never records a fresh graph.
Subclasses gate lazy capture on a warmup window (see the diffusion
runner) so warmup can capture by simply driving the forward as usual.
"""
return False
@torch.no_grad()
def __call__(self, **kwargs) -> Any:
"""Eager runner: replay a captured graph, else run ``transformer``.
While serving this never captures, so no new graph is recorded once
warmup is over. During the warmup window subclasses opt into lazy
capture via :meth:`_should_capture_on_call`.
"""
if self._disabled_reason is not None:
return self.transformer(**kwargs)
key = self._signature(kwargs)
entry = self.entries.get(key)
if entry is None:
if not self._should_capture_on_call(key):
return self.transformer(**kwargs)
if not self.capture(**kwargs):
return self.transformer(**kwargs)
entry = self.entries[key]
return self.replay(entry, kwargs)
def replay(self, entry: _CaptureEntry, kwargs: dict[str, Any]) -> Any:
live_leaves = _flatten_kwargs(kwargs)
if len(live_leaves) != len(entry.static_leaves):
# Structure changed under a matching shape key — should not happen;
# fall back to eager rather than copy mismatched buffers.
return self.transformer(**kwargs)
for buf, live in zip(entry.static_leaves, live_leaves):
buf.copy_(live, non_blocking=True)
entry.graph.replay()
# Clone so the caller can hold the result across the next replay / the
# other CFG branch (which shares this static output buffer when shapes
# match). The clone is one cheap DtoD copy relative to the full DiT.
return _clone_output(entry.output)
# ------------------------------------------------------------------ #
# Internals
# ------------------------------------------------------------------ #
def _signature(self, kwargs: dict[str, Any]) -> tuple:
"""Capture key for tensor leaves and non-tensor control values.
Tensor leaves are keyed by shape+dtype so their values can change per
replay. Non-tensor leaves are baked into the captured Python control
flow, so simple constants must be part of the key as well. Mutable
objects are keyed by identity to avoid replaying a graph whose eager
break points still reference a previous request's state object.
"""
return _signature_kwargs(kwargs)
def _empty_cache(self) -> None:
empty_cache = getattr(self.device_module, "empty_cache", None)
if callable(empty_cache):
empty_cache()
@staticmethod
def _drop_entry(entry: _CaptureEntry) -> None:
entry.graph._break_fns.clear()
entry.graph._segments.clear()
entry.static_kwargs.clear()
entry.static_leaves.clear()
entry.output = None
def reset(self, *, disabled_reason: str | None = None) -> None:
for entry in self.entries.values():
self._drop_entry(entry)
self.entries.clear()
self._blocked.clear()
self._pool = None
self._empty_cache()
if disabled_reason is not None:
self._disabled_reason = disabled_reason
def _capture_limit_reason(self, entry: _CaptureEntry) -> str | None:
if self.max_segments and entry.num_segments > self.max_segments:
return (
f"captured {entry.num_segments} segments, above "
f"SGLANG_DIFFUSION_BCG_MAX_SEGMENTS={self.max_segments}"
)
return None
def _evict_entries_if_needed(self) -> None:
if not self.max_entries:
return
while len(self.entries) > self.max_entries:
evicted_key = next(iter(self.entries))
entry = self.entries.pop(evicted_key)
self._drop_entry(entry)
logger.info(
"[Diffusion BCG] evicted oldest capture for signature %s "
"(SGLANG_DIFFUSION_BCG_MAX_ENTRIES=%d)",
_signature_summary(evicted_key),
self.max_entries,
)
self._empty_cache()
def _capture(self, kwargs: dict[str, Any], key: tuple) -> _CaptureEntry:
if self._pool is None:
self._pool = self.device_module.graph_pool_handle()
# Persistent static buffers at every tensor leaf; bake non-tensors.
def _to_static(t: torch.Tensor) -> torch.Tensor:
# Static buffers live on the capture device. A CPU input (e.g. a
# scalar timestep/sigma or an index tensor built on the host)
# would otherwise force a CPU->CUDA copy inside the captured
# region, which is illegal; place its buffer on the device so the
# only host->device copy happens here, before capture, and replay
# is device-to-device.
if t.device.type == "cpu":
buf = torch.empty(t.shape, dtype=t.dtype, device=self.device)
else:
buf = torch.empty_like(t)
buf.copy_(t)
return buf
static_kwargs = {
name: _map_tensors(v, _to_static) for name, v in kwargs.items()
}
static_leaves = _flatten_kwargs(static_kwargs)
# Warm up on the capture stream so cuBLAS/cuDNN/Triton workspaces and
# any lazy JIT are materialized before capture (mirrors the LLM runner
# and torch.cuda.make_graphed_callables).
self.device_module.synchronize()
with self.device_module.stream(self._capture_stream):
for _ in range(2):
self.transformer(**static_kwargs)
self._capture_stream.synchronize()
self.device_module.synchronize()
graph = BreakableCUDAGraph()
with enable_breakable_cuda_graph():
with BreakableCUDAGraphCapture(
cuda_graph=graph, pool=self._pool, stream=self._capture_stream
):
output = self.transformer(**static_kwargs)
self.device_module.synchronize()
logger.info(
"[Diffusion BCG] captured %d segment(s), %d tensor input(s) for "
"signature %s",
len(graph._segments),
len(static_leaves),
_signature_summary(key),
)
entry = _CaptureEntry(
graph=graph,
static_kwargs=static_kwargs,
static_leaves=static_leaves,
output=output,
num_segments=len(graph._segments),
)
limit_reason = self._capture_limit_reason(entry)
if limit_reason is not None:
self._drop_entry(entry)
self.reset(disabled_reason=limit_reason)
raise _CaptureRejected(
f"{limit_reason}; disabling this BCG runner and using eager"
)
return entry
class DiffusionBreakableCudaGraphRunner(BaseBreakableCudaGraphRunner):
"""Capture/replay a diffusion DiT ``transformer`` with BCG.
Unknown attributes proxy to the wrapped transformer, so the runner can
stand in for the module while only intercepting ``forward`` calls.
"""
def _should_capture_on_call(self, key) -> bool:
try:
from sglang.multimodal_gen.runtime.managers.forward_context import (
get_forward_context,
)
forward_batch = get_forward_context().forward_batch
except Exception:
return False
return bool(getattr(forward_batch, "is_warmup", False))
@@ -8,6 +8,7 @@ from sglang.multimodal_gen.runtime.layers.attention.backends.attention_backend i
AttentionMetadataBuilder, AttentionMetadataBuilder,
) )
from sglang.multimodal_gen.runtime.layers.attention.layer import ( from sglang.multimodal_gen.runtime.layers.attention.layer import (
DynamicVarlenMaskMeta,
LocalAttention, LocalAttention,
UlyssesAttention, UlyssesAttention,
UlyssesAttention_VSA, UlyssesAttention_VSA,
@@ -22,6 +23,7 @@ from sglang.multimodal_gen.runtime.layers.attention.turbo_layer import MinimalA2
__all__ = [ __all__ = [
"USPAttention", "USPAttention",
"LocalAttention", "LocalAttention",
"DynamicVarlenMaskMeta",
"UlyssesAttention", "UlyssesAttention",
"UlyssesAttention_VSA", "UlyssesAttention_VSA",
"MinimalA2AAttnOp", "MinimalA2AAttnOp",
@@ -1,6 +1,7 @@
# Copied and adapted from: https://github.com/hao-ai-lab/FastVideo # Copied and adapted from: https://github.com/hao-ai-lab/FastVideo
# SPDX-License-Identifier: Apache-2.0 # SPDX-License-Identifier: Apache-2.0
import functools
import os import os
from collections.abc import Sequence from collections.abc import Sequence
from contextlib import nullcontext from contextlib import nullcontext
@@ -52,6 +53,11 @@ from sglang.multimodal_gen.runtime.managers.forward_context import (
) )
from sglang.multimodal_gen.runtime.platforms import AttentionBackendEnum from sglang.multimodal_gen.runtime.platforms import AttentionBackendEnum
from sglang.multimodal_gen.utils import get_compute_dtype from sglang.multimodal_gen.utils import get_compute_dtype
from sglang.srt.breakable_cuda_graph import (
eager_on_graph,
get_current_replay_token,
is_in_breakable_cuda_graph,
)
_PYTORCH_DEFAULT_CUDA_SDP_BACKENDS = [ _PYTORCH_DEFAULT_CUDA_SDP_BACKENDS = [
SDPBackend.CUDNN_ATTENTION, SDPBackend.CUDNN_ATTENTION,
@@ -171,6 +177,38 @@ def build_varlen_mask_meta_from_ranges(
} }
class DynamicVarlenMaskMeta:
"""Replay-local builder for varlen attention metadata.
BCG attention break points capture Python kwargs once. Passing a plain
``attn_mask_meta`` dict would replay stale cu_seqlens/indices when the same
graph bucket is reused for a different prompt length. This helper keeps only
replay-local metadata and rebuilds it from the current ``attn_mask`` tensor
on the first attention block of each graph replay.
"""
def __init__(self) -> None:
self._cache_key = None
self._meta = None
def resolve(self, attn_mask: torch.Tensor | None) -> dict | None:
if attn_mask is None:
self._cache_key = None
self._meta = None
return None
replay_token = get_current_replay_token()
if replay_token is None:
cache_key = ("capture", id(attn_mask), tuple(attn_mask.shape))
else:
cache_key = ("replay", replay_token, tuple(attn_mask.shape))
if cache_key != self._cache_key:
self._meta = build_varlen_mask_meta(attn_mask)
self._cache_key = cache_key
return self._meta
class UlyssesAttention(nn.Module): class UlyssesAttention(nn.Module):
"""Ulysses-style SequenceParallelism attention layer.""" """Ulysses-style SequenceParallelism attention layer."""
@@ -612,6 +650,9 @@ class USPAttention(nn.Module):
effective_skip_sp = ( effective_skip_sp = (
self.skip_sequence_parallel or skip_sequence_parallel_override self.skip_sequence_parallel or skip_sequence_parallel_override
) )
if isinstance(attn_mask_meta, DynamicVarlenMaskMeta):
attn_mask_meta = attn_mask_meta.resolve(attn_mask)
# Tail-pad meta alone (sp_shard.tail_attn_meta; mask derivable from the # Tail-pad meta alone (sp_shard.tail_attn_meta; mask derivable from the
# pad span) also opts into the masked SP branch. gap_* = legacy alias. # pad span) also opts into the masked SP branch. gap_* = legacy alias.
meta_pad_start = meta_pad_end = None meta_pad_start = meta_pad_end = None
@@ -1134,3 +1175,38 @@ class USPAttention(nn.Module):
) )
out_rep, out_shard = out[:, :num_rep], out[:, num_rep:] out_rep, out_shard = out[:, :num_rep], out[:, num_rep:]
return torch.cat([out_shard, out_rep], dim=1) return torch.cat([out_shard, out_rep], dim=1)
def _make_breakable_attention_forward(forward_method):
"""Wrap a DiT attention module's ``forward`` so it becomes a breakable
CUDA graph (BCG) break point.
During BCG capture the whole attention forward runs eagerly between
captured graph segments -- the sequence-parallel all-to-all collectives,
varlen packing, and dynamic/sparse attention kernels that live here
cannot (or should not) be captured into a static CUDA graph. When BCG is
disabled this is a transparent pass-through to the original method.
"""
bcg_forward = eager_on_graph(True)(forward_method)
@functools.wraps(forward_method)
def forward(self, *args, **kwargs):
if is_in_breakable_cuda_graph():
return bcg_forward(self, *args, **kwargs)
return forward_method(self, *args, **kwargs)
return forward
# Install the break points on every DiT attention entry point. All diffusion
# models route attention through one of these modules (e.g. FLUX -> USPAttention),
# so wrapping here gives universal, model-agnostic BCG break points without
# touching individual model files.
for _attn_cls in (
UlyssesAttention,
UlyssesAttention_VSA,
LocalAttention,
USPAttention,
):
_attn_cls.forward = _make_breakable_attention_forward(_attn_cls.forward)
del _attn_cls
@@ -908,7 +908,7 @@ class GlmImageTransformer2DModel(CachableDiT, LayerwiseOffloadableModuleMixin):
batch_size, num_channels, height, width = hidden_states.shape batch_size, num_channels, height, width = hidden_states.shape
timestep -= 1.0 timestep = timestep - 1.0
if isinstance(encoder_hidden_states, list): if isinstance(encoder_hidden_states, list):
encoder_hidden_states = encoder_hidden_states[0] encoder_hidden_states = encoder_hidden_states[0]
@@ -925,7 +925,7 @@ class GlmImageTransformer2DModel(CachableDiT, LayerwiseOffloadableModuleMixin):
hidden_states = self.image_projector(hidden_states) hidden_states = self.image_projector(hidden_states)
encoder_hidden_states = self.glyph_projector(encoder_hidden_states) encoder_hidden_states = self.glyph_projector(encoder_hidden_states)
prior_embedding = self.prior_token_embedding(prior_token_id) prior_embedding = self.prior_token_embedding(prior_token_id)
prior_embedding[prior_token_drop] *= 0.0 prior_embedding = prior_embedding.masked_fill(prior_token_drop.unsqueeze(-1), 0)
prior_hidden_states = self.prior_projector(prior_embedding) prior_hidden_states = self.prior_projector(prior_embedding)
# SP: when latents are H-sharded, hidden_states has fewer patches than prior_hidden_states. # SP: when latents are H-sharded, hidden_states has fewer patches than prior_hidden_states.
# Shard prior_hidden_states along seq dim to match (prior is row-major, same as latent patches). # Shard prior_hidden_states along seq dim to match (prior is row-major, same as latent patches).
@@ -31,6 +31,7 @@ from sglang.multimodal_gen.runtime.distributed.sp_shard_utils import (
tail_attn_meta, tail_attn_meta,
) )
from sglang.multimodal_gen.runtime.layers.attention import ( from sglang.multimodal_gen.runtime.layers.attention import (
DynamicVarlenMaskMeta,
USPAttention, USPAttention,
build_varlen_mask_meta, build_varlen_mask_meta,
) )
@@ -67,6 +68,9 @@ from sglang.multimodal_gen.runtime.managers.memory_managers.layerwise_offload im
from sglang.multimodal_gen.runtime.models.dits.base import CachableDiT from sglang.multimodal_gen.runtime.models.dits.base import CachableDiT
from sglang.multimodal_gen.runtime.platforms import AttentionBackendEnum from sglang.multimodal_gen.runtime.platforms import AttentionBackendEnum
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
from sglang.srt.model_executor.runner_backend_utils.breakable_cuda_graph import (
is_in_breakable_cuda_graph,
)
logger = init_logger(__name__) # pylint: disable=invalid-name logger = init_logger(__name__) # pylint: disable=invalid-name
@@ -1002,6 +1006,38 @@ class QwenImageTransformerBlock(nn.Module):
self.img_mlp = NunchakuFeedForward(self.img_mlp, **nunchaku_kwargs) self.img_mlp = NunchakuFeedForward(self.img_mlp, **nunchaku_kwargs)
self.txt_mlp = NunchakuFeedForward(self.txt_mlp, **nunchaku_kwargs) self.txt_mlp = NunchakuFeedForward(self.txt_mlp, **nunchaku_kwargs)
def _norm_scale_shift(
self,
norm_module: LayerNormScaleShift,
x: torch.Tensor,
shift: torch.Tensor,
scale: torch.Tensor,
) -> torch.Tensor:
return norm_module(x=x, shift=shift, scale=scale)
def _scale_residual_norm_scale_shift(
self,
norm_module: ScaleResidualLayerNormScaleShift,
*,
residual: torch.Tensor,
x: torch.Tensor,
gate: torch.Tensor | int,
shift: torch.Tensor,
scale: torch.Tensor,
) -> tuple[torch.Tensor, torch.Tensor]:
return norm_module(
residual=residual,
x=x,
gate=gate,
shift=shift,
scale=scale,
)
def _mul_add(
self, a: torch.Tensor, b: torch.Tensor, c: torch.Tensor, k: int = 0
) -> torch.Tensor:
return self.fuse_mul_add(a, b, c, k)
def _modulate( def _modulate(
self, self,
x: torch.Tensor, x: torch.Tensor,
@@ -1010,6 +1046,7 @@ class QwenImageTransformerBlock(nn.Module):
index: Optional[torch.Tensor] = None, index: Optional[torch.Tensor] = None,
gate_x: Optional[torch.Tensor] = None, gate_x: Optional[torch.Tensor] = None,
residual_x: Optional[torch.Tensor] = None, residual_x: Optional[torch.Tensor] = None,
use_bcg_helpers: bool = False,
) -> Union[ ) -> Union[
Tuple[torch.Tensor, torch.Tensor], Tuple[torch.Tensor, torch.Tensor],
Tuple[torch.Tensor, torch.Tensor, torch.Tensor], Tuple[torch.Tensor, torch.Tensor, torch.Tensor],
@@ -1071,6 +1108,16 @@ class QwenImageTransformerBlock(nn.Module):
scale_result = scale.unsqueeze(1) scale_result = scale.unsqueeze(1)
gate_result = gate.unsqueeze(1) gate_result = gate.unsqueeze(1)
if is_scale_residual: if is_scale_residual:
if use_bcg_helpers:
modulated, residual_out = self._scale_residual_norm_scale_shift(
norm_module,
residual=residual_x,
x=x,
gate=gate_x,
shift=shift_result,
scale=scale_result,
)
else:
modulated, residual_out = norm_module( modulated, residual_out = norm_module(
residual=residual_x, residual=residual_x,
x=x, x=x,
@@ -1079,6 +1126,11 @@ class QwenImageTransformerBlock(nn.Module):
scale=scale_result, scale=scale_result,
) )
return modulated, residual_out, gate_result return modulated, residual_out, gate_result
else:
if use_bcg_helpers:
modulated = self._norm_scale_shift(
norm_module, x=x, shift=shift_result, scale=scale_result
)
else: else:
modulated = norm_module(x=x, shift=shift_result, scale=scale_result) modulated = norm_module(x=x, shift=shift_result, scale=scale_result)
return modulated, gate_result return modulated, gate_result
@@ -1119,13 +1171,26 @@ class QwenImageTransformerBlock(nn.Module):
# Split modulation parameters for norm1 and norm2 # Split modulation parameters for norm1 and norm2
img_mod1, img_mod2 = img_mod_params.chunk(2, dim=-1) # Each [B, 3*dim] img_mod1, img_mod2 = img_mod_params.chunk(2, dim=-1) # Each [B, 3*dim]
txt_mod1, txt_mod2 = txt_mod_params.chunk(2, dim=-1) # Each [B, 3*dim] txt_mod1, txt_mod2 = txt_mod_params.chunk(2, dim=-1) # Each [B, 3*dim]
use_bcg_helpers = is_in_breakable_cuda_graph()
# Process image stream - norm1 + modulation # Process image stream - norm1 + modulation
img_modulated, img_gate1 = self._modulate( img_modulated, img_gate1 = self._modulate(
hidden_states, img_mod1, self.img_norm1, modulate_index hidden_states,
img_mod1,
self.img_norm1,
modulate_index,
use_bcg_helpers=use_bcg_helpers,
) )
# Process text stream - norm1 + modulation # Process text stream - norm1 + modulation
txt_shift1, txt_scale1, txt_gate1_raw = txt_mod1.chunk(3, dim=-1) txt_shift1, txt_scale1, txt_gate1_raw = txt_mod1.chunk(3, dim=-1)
if use_bcg_helpers:
txt_modulated = self._norm_scale_shift(
self.txt_norm1,
encoder_hidden_states,
shift=txt_shift1,
scale=txt_scale1,
)
else:
txt_modulated = self.txt_norm1( txt_modulated = self.txt_norm1(
encoder_hidden_states, shift=txt_shift1, scale=txt_scale1 encoder_hidden_states, shift=txt_shift1, scale=txt_scale1
) )
@@ -1158,15 +1223,32 @@ class QwenImageTransformerBlock(nn.Module):
modulate_index, modulate_index,
gate_x=img_gate1, gate_x=img_gate1,
residual_x=hidden_states, residual_x=hidden_states,
use_bcg_helpers=use_bcg_helpers,
) )
img_mlp_output = self.img_mlp(img_modulated2) img_mlp_output = self.img_mlp(img_modulated2)
if img_mlp_output.dim() == 2: if img_mlp_output.dim() == 2:
img_mlp_output = img_mlp_output.unsqueeze(0) img_mlp_output = img_mlp_output.unsqueeze(0)
if use_bcg_helpers:
hidden_states = self._mul_add(img_mlp_output, img_gate2, hidden_states)
else:
hidden_states = self.fuse_mul_add(img_mlp_output, img_gate2, hidden_states) hidden_states = self.fuse_mul_add(img_mlp_output, img_gate2, hidden_states)
# Process text stream - norm2 + MLP # Process text stream - norm2 + MLP
txt_shift2, txt_scale2, txt_gate2_raw = txt_mod2.chunk(3, dim=-1) txt_shift2, txt_scale2, txt_gate2_raw = txt_mod2.chunk(3, dim=-1)
if use_bcg_helpers:
(
txt_modulated2,
encoder_hidden_states,
) = self._scale_residual_norm_scale_shift(
self.txt_norm2,
residual=encoder_hidden_states,
x=txt_attn_output,
gate=txt_gate1,
shift=txt_shift2,
scale=txt_scale2,
)
else:
txt_modulated2, encoder_hidden_states = self.txt_norm2( txt_modulated2, encoder_hidden_states = self.txt_norm2(
residual=encoder_hidden_states, residual=encoder_hidden_states,
x=txt_attn_output, x=txt_attn_output,
@@ -1179,6 +1261,11 @@ class QwenImageTransformerBlock(nn.Module):
if txt_mlp_output.dim() == 2: if txt_mlp_output.dim() == 2:
txt_mlp_output = txt_mlp_output.unsqueeze(0) txt_mlp_output = txt_mlp_output.unsqueeze(0)
if use_bcg_helpers:
encoder_hidden_states = self._mul_add(
txt_mlp_output, txt_gate2, encoder_hidden_states
)
else:
encoder_hidden_states = self.fuse_mul_add( encoder_hidden_states = self.fuse_mul_add(
txt_mlp_output, txt_gate2, encoder_hidden_states txt_mlp_output, txt_gate2, encoder_hidden_states
) )
@@ -1430,8 +1517,15 @@ class QwenImageTransformer2DModel(CachableDiT, LayerwiseOffloadableModuleMixin):
) )
joint_mask = torch.cat([encoder_hidden_states_mask, image_mask], dim=1) joint_mask = torch.cat([encoder_hidden_states_mask, image_mask], dim=1)
block_attention_kwargs["attn_mask"] = joint_mask block_attention_kwargs["attn_mask"] = joint_mask
# Precompute varlen metadata once per request so every block reuses if is_in_breakable_cuda_graph():
# the same cu_seqlens / indices instead of rebuilding. # Qwen/FireRed BCG buckets text inputs so different prompt
# lengths can share a graph. Attention break kwargs are captured
# once, so build varlen metadata replay-locally from the current
# static mask instead of closing over stale cu_seqlens/indices.
block_attention_kwargs["attn_mask_meta"] = DynamicVarlenMaskMeta()
else:
# Precompute varlen metadata once per request so every block
# reuses the same cu_seqlens / indices instead of rebuilding.
block_attention_kwargs["attn_mask_meta"] = build_varlen_mask_meta( block_attention_kwargs["attn_mask_meta"] = build_varlen_mask_meta(
joint_mask joint_mask
) )
@@ -949,6 +949,7 @@ class ZImageTransformer2DModel(CachableDiT, LayerwiseOffloadableModuleMixin):
f_patch_size: int, f_patch_size: int,
image_seq_len_target: int | None = None, image_seq_len_target: int | None = None,
caption_valid_lens: torch.Tensor | None = None, caption_valid_lens: torch.Tensor | None = None,
caption_valid_mask: torch.Tensor | None = None,
): ):
"""Patchify images and pad image/caption tokens to batch targets. """Patchify images and pad image/caption tokens to batch targets.
@@ -963,6 +964,10 @@ class ZImageTransformer2DModel(CachableDiT, LayerwiseOffloadableModuleMixin):
) )
if not all_image: if not all_image:
raise ValueError("Z-Image batch must contain at least one image latent") raise ValueError("Z-Image batch must contain at least one image latent")
if caption_valid_mask is not None and caption_valid_mask.shape[0] != len(
all_cap_feats
):
raise ValueError("caption_valid_mask must have one row per Z-Image caption")
pH = pW = patch_size pH = pW = patch_size
pF = f_patch_size pF = f_patch_size
@@ -971,6 +976,7 @@ class ZImageTransformer2DModel(CachableDiT, LayerwiseOffloadableModuleMixin):
all_cap_feats_out = [] all_cap_feats_out = []
all_image_valid_lens = [] all_image_valid_lens = []
all_cap_valid_lens = [] all_cap_valid_lens = []
all_cap_valid_masks = []
all_image_attn_lens = [] all_image_attn_lens = []
all_cap_attn_lens = [] all_cap_attn_lens = []
image_records = [] image_records = []
@@ -994,6 +1000,21 @@ class ZImageTransformer2DModel(CachableDiT, LayerwiseOffloadableModuleMixin):
dim=0, dim=0,
) )
all_cap_feats_out.append(cap_padded_feat) all_cap_feats_out.append(cap_padded_feat)
if caption_valid_mask is not None:
mask_row = caption_valid_mask[idx].to(
device=cap_feat.device, dtype=torch.bool
)
if mask_row.dim() != 1:
mask_row = mask_row.reshape(-1)
if mask_row.shape[0] > cap_seq_len_target:
mask_row = mask_row[:cap_seq_len_target]
elif mask_row.shape[0] < cap_seq_len_target:
mask_row = torch.nn.functional.pad(
mask_row,
(0, cap_seq_len_target - mask_row.shape[0]),
value=0,
)
all_cap_valid_masks.append(mask_row)
if caption_valid_lens is None: if caption_valid_lens is None:
all_cap_valid_lens.append(cap_ori_len) all_cap_valid_lens.append(cap_ori_len)
else: else:
@@ -1045,6 +1066,11 @@ class ZImageTransformer2DModel(CachableDiT, LayerwiseOffloadableModuleMixin):
cap_valid_lens_out, cap_valid_lens_out,
all_image_attn_lens, all_image_attn_lens,
all_cap_attn_lens, all_cap_attn_lens,
(
torch.stack(all_cap_valid_masks, dim=0)
if caption_valid_mask is not None
else None
),
) )
def _build_single_sample_freqs_cis( def _build_single_sample_freqs_cis(
@@ -1339,6 +1365,45 @@ class ZImageTransformer2DModel(CachableDiT, LayerwiseOffloadableModuleMixin):
return cap_feats return cap_feats
return cap_feats return cap_feats
@staticmethod
def _caption_valid_mask_from_mask(
mask, *, batch_size: int, max_seq_len: int
) -> torch.Tensor | None:
if mask is None:
return None
if isinstance(mask, (list, tuple)):
if not mask:
return None
if len(mask) == 1:
return ZImageTransformer2DModel._caption_valid_mask_from_mask(
mask[0], batch_size=batch_size, max_seq_len=max_seq_len
)
rows = []
for item in mask:
item_mask = ZImageTransformer2DModel._caption_valid_mask_from_mask(
item, batch_size=1, max_seq_len=max_seq_len
)
if item_mask is None:
return None
rows.append(item_mask[0])
return torch.stack(rows, dim=0) if len(rows) == batch_size else None
if not torch.is_tensor(mask):
return None
mask = mask.to(dtype=torch.bool)
if mask.ndim == 1:
if batch_size != 1:
return None
mask = mask[:max_seq_len].unsqueeze(0)
elif mask.ndim == 2 and mask.shape[0] == batch_size:
mask = mask[:, :max_seq_len]
elif mask.ndim == 2 and batch_size == 1 and mask.shape[0] == 1:
mask = mask[:, :max_seq_len]
else:
return None
return mask
@staticmethod @staticmethod
def _replace_padding_with_token( def _replace_padding_with_token(
tensor: torch.Tensor, tensor: torch.Tensor,
@@ -1346,20 +1411,43 @@ class ZImageTransformer2DModel(CachableDiT, LayerwiseOffloadableModuleMixin):
pad_token: torch.Tensor, pad_token: torch.Tensor,
) -> torch.Tensor: ) -> torch.Tensor:
"""Replace padded token rows after each valid sequence length.""" """Replace padded token rows after each valid sequence length."""
if not ZImageTransformer2DModel._has_padding(valid_lens, tensor.shape[1]): if not torch.is_tensor(valid_lens) and all(
int(length) >= tensor.shape[1] for length in valid_lens
):
return tensor return tensor
positions = torch.arange(tensor.shape[1], device=tensor.device).unsqueeze(0) positions = torch.arange(tensor.shape[1], device=tensor.device).unsqueeze(0)
if torch.is_tensor(valid_lens): if torch.is_tensor(valid_lens):
lengths = valid_lens.to(device=tensor.device, dtype=torch.long) lengths = valid_lens.to(device=tensor.device, dtype=torch.long)
else: else:
lengths = torch.tensor(valid_lens, device=tensor.device) lengths = torch.tensor(valid_lens, device=tensor.device)
if lengths.ndim == 0:
lengths = lengths.reshape(1)
lengths = lengths.unsqueeze(1) lengths = lengths.unsqueeze(1)
pad_mask = positions >= lengths pad_mask = positions >= lengths
tensor = tensor.clone() tensor = tensor.clone()
tensor[pad_mask] = pad_token.to(device=tensor.device, dtype=tensor.dtype) tensor[pad_mask] = pad_token.to(device=tensor.device, dtype=tensor.dtype)
return tensor return tensor
@staticmethod
def _replace_padding_with_token_mask(
tensor: torch.Tensor,
valid_mask: torch.Tensor,
pad_token: torch.Tensor,
) -> torch.Tensor:
"""Replace padded token rows using a fixed-shape tensor mask."""
seq_len = tensor.shape[1]
valid_mask = valid_mask.to(device=tensor.device, dtype=torch.bool)
if valid_mask.shape[1] > seq_len:
valid_mask = valid_mask[:, :seq_len]
elif valid_mask.shape[1] < seq_len:
valid_mask = torch.nn.functional.pad(
valid_mask,
(0, seq_len - valid_mask.shape[1]),
value=0,
)
pad_value = pad_token.to(device=tensor.device, dtype=tensor.dtype)
return torch.where(valid_mask.unsqueeze(-1), tensor, pad_value.view(1, 1, -1))
def forward( def forward(
self, self,
hidden_states: List[torch.Tensor], hidden_states: List[torch.Tensor],
@@ -1370,6 +1458,7 @@ class ZImageTransformer2DModel(CachableDiT, LayerwiseOffloadableModuleMixin):
f_patch_size=1, f_patch_size=1,
freqs_cis=None, freqs_cis=None,
image_seq_len_target: int | None = None, image_seq_len_target: int | None = None,
encoder_hidden_states_mask=None,
caption_valid_lens: torch.Tensor | None = None, caption_valid_lens: torch.Tensor | None = None,
**kwargs, **kwargs,
): ):
@@ -1380,7 +1469,13 @@ class ZImageTransformer2DModel(CachableDiT, LayerwiseOffloadableModuleMixin):
cap_feats = self._as_caption_list(encoder_hidden_states) cap_feats = self._as_caption_list(encoder_hidden_states)
input_images = x input_images = x
input_cap_feats = cap_feats input_cap_feats = cap_feats
caption_valid_mask = None
if kwargs.pop("_use_caption_valid_mask", False):
caption_valid_mask = self._caption_valid_mask_from_mask(
encoder_hidden_states_mask,
batch_size=len(cap_feats),
max_seq_len=max(cap_feat.shape[0] for cap_feat in cap_feats),
)
timestep = 1000.0 - timestep timestep = 1000.0 - timestep
t = timestep t = timestep
t = self.t_embedder(t) t = self.t_embedder(t)
@@ -1393,6 +1488,7 @@ class ZImageTransformer2DModel(CachableDiT, LayerwiseOffloadableModuleMixin):
cap_valid_lens, cap_valid_lens,
x_attn_lens, x_attn_lens,
cap_attn_lens, cap_attn_lens,
cap_valid_mask,
) = self.patchify_and_embed( ) = self.patchify_and_embed(
x, x,
cap_feats, cap_feats,
@@ -1400,6 +1496,7 @@ class ZImageTransformer2DModel(CachableDiT, LayerwiseOffloadableModuleMixin):
f_patch_size, f_patch_size,
image_seq_len_target=image_seq_len_target, image_seq_len_target=image_seq_len_target,
caption_valid_lens=caption_valid_lens, caption_valid_lens=caption_valid_lens,
caption_valid_mask=caption_valid_mask,
) )
x, _ = self.all_x_embedder[f"{patch_size}-{f_patch_size}"](x) x, _ = self.all_x_embedder[f"{patch_size}-{f_patch_size}"](x)
@@ -1435,6 +1532,11 @@ class ZImageTransformer2DModel(CachableDiT, LayerwiseOffloadableModuleMixin):
) )
cap_feats, _ = self.cap_embedder(cap_feats) cap_feats, _ = self.cap_embedder(cap_feats)
if cap_valid_mask is not None:
cap_feats = self._replace_padding_with_token_mask(
cap_feats, cap_valid_mask, self.cap_pad_token
)
else:
cap_feats = self._replace_padding_with_token( cap_feats = self._replace_padding_with_token(
cap_feats, cap_valid_lens, self.cap_pad_token cap_feats, cap_valid_lens, self.cap_pad_token
) )
@@ -27,6 +27,9 @@ from sglang.multimodal_gen.configs.pipeline_configs.flux import (
FluxPipelineConfig, FluxPipelineConfig,
) )
from sglang.multimodal_gen.configs.pipeline_configs.zimage import ZImagePipelineConfig from sglang.multimodal_gen.configs.pipeline_configs.zimage import ZImagePipelineConfig
from sglang.multimodal_gen.runtime.breakable_cuda_graph import (
prompt_padding as bcg_utils,
)
from sglang.multimodal_gen.runtime.cache.cache_dit_integration import ( from sglang.multimodal_gen.runtime.cache.cache_dit_integration import (
CacheDitConfig, CacheDitConfig,
enable_cache_on_dual_transformer, enable_cache_on_dual_transformer,
@@ -214,6 +217,8 @@ class DenoisingStage(PipelineStage, RolloutDenoisingMixin):
self._cache_dit_enabled = False self._cache_dit_enabled = False
self._cached_num_steps = None self._cached_num_steps = None
self._torch_compile_registry = CompiledModuleRegistry() self._torch_compile_registry = CompiledModuleRegistry()
# Breakable CUDA graph runners, one per transformer module (lazy).
self._bcg_runners: dict[int, Any] = {}
hidden_size = self.server_args.pipeline_config.dit_config.hidden_size hidden_size = self.server_args.pipeline_config.dit_config.hidden_size
num_attention_heads = ( num_attention_heads = (
@@ -370,9 +375,13 @@ class DenoisingStage(PipelineStage, RolloutDenoisingMixin):
Compile a module with torch.compile, and enable inductor overlap tweak if available. Compile a module with torch.compile, and enable inductor overlap tweak if available.
No-op if torch compile is disabled or the object is not a nn.Module. No-op if torch compile is disabled or the object is not a nn.Module.
""" """
if not self.server_args.enable_torch_compile or not isinstance( if self.server_args.enable_breakable_cuda_graph:
module, nn.Module # BCG captures the eager kernel stream itself; compiling first
): # would capture inductor's own cudagraph trees / guards.
return
if not getattr(
self.server_args, "enable_torch_compile", False
) or not isinstance(module, nn.Module):
return return
if envs.SGLANG_CACHE_DIT_ENABLED and not self._cache_dit_enabled: if envs.SGLANG_CACHE_DIT_ENABLED and not self._cache_dit_enabled:
logger.debug("Deferring torch.compile until cache-dit is enabled") logger.debug("Deferring torch.compile until cache-dit is enabled")
@@ -564,6 +573,10 @@ class DenoisingStage(PipelineStage, RolloutDenoisingMixin):
transformers with (potentially) different configurations. transformers with (potentially) different configurations.
""" """
if self.server_args.enable_breakable_cuda_graph:
# Cache-DiT wraps transformer.forward with step-skipping control
# flow that must not be baked into a captured CUDA graph.
return
# NOTE: When a new request arrives, we need to refresh the cache-dit context. # NOTE: When a new request arrives, we need to refresh the cache-dit context.
if self._cache_dit_enabled: if self._cache_dit_enabled:
primary_num_steps, secondary_num_steps = self._cache_dit_step_counts( primary_num_steps, secondary_num_steps = self._cache_dit_step_counts(
@@ -596,7 +609,9 @@ class DenoisingStage(PipelineStage, RolloutDenoisingMixin):
# warmup to mount cache-dit before Dynamo traces the transformer. # warmup to mount cache-dit before Dynamo traces the transformer.
if not envs.SGLANG_CACHE_DIT_ENABLED: if not envs.SGLANG_CACHE_DIT_ENABLED:
return return
if batch.is_warmup and not self.server_args.enable_torch_compile: if batch.is_warmup and not getattr(
self.server_args, "enable_torch_compile", False
):
return return
primary_num_steps, secondary_num_steps = self._cache_dit_step_counts( primary_num_steps, secondary_num_steps = self._cache_dit_step_counts(
@@ -876,8 +891,13 @@ class DenoisingStage(PipelineStage, RolloutDenoisingMixin):
) )
else: else:
reserved_frames_mask_sp, z_sp = ( reserved_frames_mask_sp, z_sp = (
reserved_frames_masks[0] if reserved_frames_masks is not None else None (
), z reserved_frames_masks[0]
if reserved_frames_masks is not None
else None
),
z,
)
guidance = self.get_or_build_guidance( guidance = self.get_or_build_guidance(
# TODO: replace with raw_latent_shape? # TODO: replace with raw_latent_shape?
@@ -900,7 +920,11 @@ class DenoisingStage(PipelineStage, RolloutDenoisingMixin):
{ {
"encoder_hidden_states_2": batch.clip_embedding_pos, "encoder_hidden_states_2": batch.clip_embedding_pos,
"encoder_attention_mask": batch.prompt_attention_mask, "encoder_attention_mask": batch.prompt_attention_mask,
"encoder_hidden_states_mask": batch.prompt_attention_mask, "encoder_hidden_states_mask": (
batch.prompt_embeds_mask
if batch.prompt_embeds_mask is not None
else batch.prompt_attention_mask
),
} }
| server_args.pipeline_config.prepare_pos_cond_kwargs( | server_args.pipeline_config.prepare_pos_cond_kwargs(
batch, batch,
@@ -921,7 +945,11 @@ class DenoisingStage(PipelineStage, RolloutDenoisingMixin):
{ {
"encoder_hidden_states_2": batch.clip_embedding_neg, "encoder_hidden_states_2": batch.clip_embedding_neg,
"encoder_attention_mask": batch.negative_attention_mask, "encoder_attention_mask": batch.negative_attention_mask,
"encoder_hidden_states_mask": batch.negative_attention_mask, "encoder_hidden_states_mask": (
batch.negative_prompt_embeds_mask
if batch.negative_prompt_embeds_mask is not None
else batch.negative_attention_mask
),
} }
| server_args.pipeline_config.prepare_neg_cond_kwargs( | server_args.pipeline_config.prepare_neg_cond_kwargs(
batch, batch,
@@ -1989,14 +2017,109 @@ class DenoisingStage(PipelineStage, RolloutDenoisingMixin):
getattr(current_model, "forward", current_model), getattr(current_model, "forward", current_model),
{"guidance": guidance}, {"guidance": guidance},
) )
model_output = current_model( call_kwargs = dict(
hidden_states=latent_model_input, hidden_states=latent_model_input,
timestep=timestep, timestep=timestep,
**guidance_kwargs, **guidance_kwargs,
**kwargs, **kwargs,
) )
runner = self._maybe_get_bcg_runner(current_model)
if runner is not None:
model_output = self._bcg_run(runner, call_kwargs, current_model)
else:
model_output = current_model(**call_kwargs)
return _ensure_tensor_model_output(model_output) return _ensure_tensor_model_output(model_output)
@staticmethod
def _bcg_is_warmup() -> bool:
"""True when the current forward is a warmup request."""
from sglang.multimodal_gen.runtime.managers.forward_context import (
get_forward_context,
)
try:
forward_batch = get_forward_context().forward_batch
except Exception:
return False
return bool(getattr(forward_batch, "is_warmup", False))
def _bcg_run(self, runner, call_kwargs: dict, current_model):
"""Run the DiT through the BCG runner.
During warmup we proactively capture one graph per text bucket (in
addition to the request's own bucket) so that serving never records a
fresh graph for a different prompt length — every bucket is already
captured. Serving just replays (or runs eager for an uncaptured
signature, never capturing).
"""
if self._bcg_is_warmup():
for bucket in self._bcg_text_buckets():
runner.capture(
**self._bcg_pad_prompt_kwargs(
call_kwargs, current_model=current_model, force_bucket=bucket
)
)
return runner(
**self._bcg_pad_prompt_kwargs(call_kwargs, current_model=current_model)
)
@staticmethod
def _bcg_text_buckets() -> tuple[int, ...]:
"""Prompt sequence-length buckets, from --bcg-text-buckets."""
from sglang.multimodal_gen.runtime.server_args import (
DEFAULT_BCG_TEXT_BUCKETS,
get_global_server_args,
)
try:
resolver = get_global_server_args().resolved_bcg_text_buckets
return resolver()
except Exception:
return DEFAULT_BCG_TEXT_BUCKETS
def _bcg_pad_prompt_kwargs(
self, call_kwargs: dict, current_model=None, force_bucket: int | None = None
):
"""Bucket prompt-conditioning inputs so BCG signatures ignore prompt length.
Generic padding lives in ``breakable_cuda_graph.prompt_padding``;
model-specific padders register from ``breakable_cuda_graph.model_padders``.
``force_bucket`` pads to exactly that bucket (used by warmup to capture
every bucket); a prompt already longer than ``force_bucket`` is left
unchanged, exactly as the normal bucket selection would do.
"""
buckets = (
(force_bucket,) if force_bucket is not None else self._bcg_text_buckets()
)
padder = bcg_utils.select_prompt_padder(current_model, call_kwargs)
if padder is not None:
return padder(call_kwargs, current_model, buckets)
return bcg_utils.pad_masked_prompt_kwargs(call_kwargs, buckets)
def _maybe_get_bcg_runner(self, current_model):
"""Return (lazily creating) the breakable CUDA graph runner for
``current_model``, or ``None`` if BCG is disabled / inapplicable.
"""
if not self.server_args.enable_breakable_cuda_graph:
return None
if not isinstance(current_model, nn.Module):
return None
key = id(current_model)
runner = self._bcg_runners.get(key)
if runner is None:
from sglang.multimodal_gen.runtime.breakable_cuda_graph.runner import (
DiffusionBreakableCudaGraphRunner,
)
# DenoisingStage can switch between transformer and transformer_2;
# each module owns separate graph state and static input buffers.
runner = DiffusionBreakableCudaGraphRunner(
current_model, get_local_torch_device()
)
self._bcg_runners[key] = runner
return runner
def prepare_sta_param(self, batch: Req, server_args: ServerArgs): def prepare_sta_param(self, batch: Req, server_args: ServerArgs):
""" """
Prepare Sliding Tile Attention (STA) parameters and settings. Prepare Sliding Tile Attention (STA) parameters and settings.
@@ -308,6 +308,20 @@ class GlmImageAR(PipelineStage):
width = width or ar_condition_images[0].width width = width or ar_condition_images[0].width
time_start = time.time() time_start = time.time()
seed = getattr(batch, "seed", None)
if seed is None:
prior_token_id, prior_token_image_ids = self.generate_prior_tokens(
prompt=prompt,
image=ar_condition_images,
height=height,
width=width,
)
else:
rng_devices = []
if device.type == "cuda":
rng_devices.append(torch.cuda.current_device())
with torch.random.fork_rng(devices=rng_devices, enabled=True):
torch.manual_seed(int(seed))
prior_token_id, prior_token_image_ids = self.generate_prior_tokens( prior_token_id, prior_token_image_ids = self.generate_prior_tokens(
prompt=prompt, prompt=prompt,
image=ar_condition_images, image=ar_condition_images,
@@ -314,6 +314,14 @@ class Ideogram4DenoisingStage(DenoisingStage):
return return
super()._manage_dit_use_site(current_model, current_phase, batch) super()._manage_dit_use_site(current_model, current_phase, batch)
def _run_ideogram_transformer(
self, current_model: torch.nn.Module, call_kwargs: dict
) -> torch.Tensor:
runner = self._maybe_get_bcg_runner(current_model)
if runner is not None:
return self._bcg_run(runner, call_kwargs, current_model)
return current_model(**call_kwargs)
def _preprocess_sp_latents(self, batch: Req, server_args: ServerArgs): def _preprocess_sp_latents(self, batch: Req, server_args: ServerArgs):
batch.did_sp_shard_latents = False batch.did_sp_shard_latents = False
@@ -417,6 +425,7 @@ class Ideogram4DenoisingStage(DenoisingStage):
z = ctx.latents.to(dtype=torch.float32) z = ctx.latents.to(dtype=torch.float32)
llm_features = batch.prompt_embeds[0] llm_features = batch.prompt_embeds[0]
max_text_tokens = data["max_text_tokens"] max_text_tokens = data["max_text_tokens"]
num_image_tokens = data["num_image_tokens"]
schedule_values = ctx.extra["ideogram4_schedule_values"] schedule_values = ctx.extra["ideogram4_schedule_values"]
schedule_deltas = ctx.extra["ideogram4_schedule_deltas"] schedule_deltas = ctx.extra["ideogram4_schedule_deltas"]
guidance_schedule = ctx.extra["ideogram4_guidance_schedule"] guidance_schedule = ctx.extra["ideogram4_guidance_schedule"]
@@ -433,7 +442,9 @@ class Ideogram4DenoisingStage(DenoisingStage):
attn_metadata=step.attn_metadata, attn_metadata=step.attn_metadata,
forward_batch=batch, forward_batch=batch,
): ):
pos_out = step.current_model( pos_out = self._run_ideogram_transformer(
step.current_model,
dict(
llm_features=llm_features, llm_features=llm_features,
x=pos_z, x=pos_z,
t=t, t=t,
@@ -442,8 +453,9 @@ class Ideogram4DenoisingStage(DenoisingStage):
indicator=data["indicator"], indicator=data["indicator"],
attn_mask=ctx.extra["ideogram4_attn_mask"], attn_mask=ctx.extra["ideogram4_attn_mask"],
attn_mask_meta=ctx.extra["ideogram4_attn_mask_meta"], attn_mask_meta=ctx.extra["ideogram4_attn_mask_meta"],
),
) )
pos_v = pos_out[:, max_text_tokens:] pos_v = pos_out[:, max_text_tokens : max_text_tokens + num_image_tokens]
self._manage_unconditional_transformer_use_site(batch) self._manage_unconditional_transformer_use_site(batch)
with set_forward_context( with set_forward_context(
@@ -451,7 +463,9 @@ class Ideogram4DenoisingStage(DenoisingStage):
attn_metadata=step.attn_metadata, attn_metadata=step.attn_metadata,
forward_batch=batch, forward_batch=batch,
): ):
neg_v = self.unconditional_transformer( neg_v = self._run_ideogram_transformer(
self.unconditional_transformer,
dict(
llm_features=ctx.extra["ideogram4_neg_llm_features"], llm_features=ctx.extra["ideogram4_neg_llm_features"],
x=z, x=z,
t=t, t=t,
@@ -460,6 +474,7 @@ class Ideogram4DenoisingStage(DenoisingStage):
indicator=ctx.extra["ideogram4_neg_indicator"], indicator=ctx.extra["ideogram4_neg_indicator"],
attn_mask=ctx.extra["ideogram4_neg_attn_mask"], attn_mask=ctx.extra["ideogram4_neg_attn_mask"],
attn_mask_meta=ctx.extra["ideogram4_neg_attn_mask_meta"], attn_mask_meta=ctx.extra["ideogram4_neg_attn_mask_meta"],
),
) )
with maybe_nvtx_range("scheduler_step", use_nvtx): with maybe_nvtx_range("scheduler_step", use_nvtx):
@@ -121,6 +121,55 @@ class Backend(str, Enum):
WARMUP_MODES = ("off", "request", "server") WARMUP_MODES = ("off", "request", "server")
# Default prompt sequence-length buckets for breakable CUDA graph (BCG) padding.
# Prompt-conditioning is padded up to the smallest bucket that fits so prompts
# of different lengths share one captured graph.
DEFAULT_BCG_TEXT_BUCKETS = (64, 128, 256, 512, 1024)
BREAKABLE_CUDA_GRAPH_SUPPORTED_MODEL_IDS = frozenset(
{
"comfy-org/ideogram-4",
"glm-image",
"ideogram-4",
"ideogram-4-fp8",
"ideogram-4-nf4",
"ideogram-ai/ideogram-4-fp8",
"ideogram-ai/ideogram-4-nf4",
"qwen/qwen-image",
"qwen/qwen-image-2512",
"qwen-image",
"qwen-image-2512",
"tongyi-mai/z-image",
"tongyi-mai/z-image-turbo",
"zai-org/glm-image",
"z-image",
"z-image-turbo",
}
)
BREAKABLE_CUDA_GRAPH_SUPPORTED_PIPELINE_CONFIGS = frozenset(
{
"GlmImagePipelineConfig",
"Ideogram4PipelineConfig",
"QwenImagePipelineConfig",
"ZImagePipelineConfig",
}
)
def _normalized_bcg_model_refs(model_ref: str | None) -> set[str]:
if not model_ref:
return set()
normalized = str(model_ref).strip().rstrip("/").lower()
refs = {normalized, os.path.basename(normalized)}
if "models--" in normalized:
hf_cache_name = normalized.split("models--", 1)[1].split("/", 1)[0]
refs.add(hf_cache_name.replace("--", "/"))
return refs
@dataclasses.dataclass @dataclasses.dataclass
class ServerArgs(DisaggServerArgsMixin): class ServerArgs(DisaggServerArgsMixin):
@@ -230,6 +279,22 @@ class ServerArgs(DisaggServerArgsMixin):
# Compilation # Compilation
enable_torch_compile: bool = False enable_torch_compile: bool = False
# Breakable CUDA graph (BCG): capture the DiT forward as CUDA-graph
# segments split at attention modules (SP all-to-all / dynamic attention
# stay eager). Mutually exclusive with --enable-torch-compile and
# Cache-DiT; BCG takes priority when more than one is requested.
#
# BCG graphs are resolution-specific, so --warmup-resolutions is required
# when BCG is enabled: every requested resolution is captured at warmup so
# serving never triggers a fresh capture.
enable_breakable_cuda_graph: bool = False
# Text/prompt sequence-length padding budget for BCG. Prompt-conditioning
# inputs are padded up to the smallest bucket that fits, so prompts of
# different lengths reuse one captured graph. Warmup captures one graph per
# bucket; a prompt longer than the largest bucket falls back to eager.
# ``None`` resolves to DEFAULT_BCG_TEXT_BUCKETS.
bcg_text_buckets: list[int] = None
# NVTX profiling # NVTX profiling
enable_layerwise_nvtx_marker: bool = False enable_layerwise_nvtx_marker: bool = False
@@ -382,6 +447,7 @@ class ServerArgs(DisaggServerArgsMixin):
auto_tuner.maybe_replace_cpu_offloaded_components_with_layerwise() auto_tuner.maybe_replace_cpu_offloaded_components_with_layerwise()
self._adjust_path() self._adjust_path()
self._adjust_quant_config() self._adjust_quant_config()
self._adjust_breakable_cuda_graph_support()
self._adjust_warmup() self._adjust_warmup()
self._adjust_network_ports() self._adjust_network_ports()
# adjust parallelism before attention backend # adjust parallelism before attention backend
@@ -416,6 +482,65 @@ class ServerArgs(DisaggServerArgsMixin):
self._validate_parallelism() self._validate_parallelism()
self._validate_cfg_parallel() self._validate_cfg_parallel()
self._validate_batching() self._validate_batching()
self._validate_breakable_cuda_graph()
def resolved_bcg_text_buckets(self) -> tuple[int, ...]:
"""Sorted, de-duplicated, positive BCG text buckets.
Falls back to :data:`DEFAULT_BCG_TEXT_BUCKETS` when ``--bcg-text-buckets``
is unset, so both prompt padding and warmup capture share one source of
truth instead of the legacy ``SGLANG_BCG_TEXT_BUCKETS`` env var.
"""
raw = self.bcg_text_buckets
if not raw:
return DEFAULT_BCG_TEXT_BUCKETS
buckets = sorted({int(b) for b in raw if int(b) > 0})
return tuple(buckets) or DEFAULT_BCG_TEXT_BUCKETS
def _validate_breakable_cuda_graph(self):
if not self.enable_breakable_cuda_graph:
return
# BCG graphs are captured per resolution and only replay for that exact
# latent shape, so the user must declare the resolutions up front. We
# capture every one of them at warmup; serving then never re-captures.
if not self.warmup_resolutions:
raise ValueError(
"--enable-breakable-cuda-graph requires --warmup-resolutions: "
"diffusion CUDA graphs only replay for a fixed resolution, so "
"every served resolution must be declared and captured at "
"warmup, e.g. --warmup-resolutions 1024x1024 1328x1328."
)
if self.bcg_text_buckets is not None and not any(
int(b) > 0 for b in self.bcg_text_buckets
):
raise ValueError(
"--bcg-text-buckets must contain at least one positive integer."
)
def _adjust_breakable_cuda_graph_support(self):
if not self.enable_breakable_cuda_graph:
return
pipeline_config = getattr(self, "pipeline_config", None)
pipeline_config_name = type(pipeline_config).__name__
if (
pipeline_config_name in BREAKABLE_CUDA_GRAPH_SUPPORTED_PIPELINE_CONFIGS
and self._is_breakable_cuda_graph_supported_model()
):
return
logger.warning(
"[Diffusion BCG] disabled for %s: only Ideogram-4, Qwen/Qwen-Image, "
"Qwen/Qwen-Image-2512, Tongyi-MAI/Z-Image/Z-Image-Turbo, "
"and zai-org/GLM-Image are currently supported.",
pipeline_config_name,
)
self.enable_breakable_cuda_graph = False
def _is_breakable_cuda_graph_supported_model(self) -> bool:
refs = _normalized_bcg_model_refs(self.model_id)
refs.update(_normalized_bcg_model_refs(self.model_path))
return bool(refs & BREAKABLE_CUDA_GRAPH_SUPPORTED_MODEL_IDS)
def _adjust_save_paths(self): def _adjust_save_paths(self):
"""Normalize empty-string save paths to None (disabled).""" """Normalize empty-string save paths to None (disabled)."""
@@ -771,6 +896,14 @@ class ServerArgs(DisaggServerArgsMixin):
"to disable this behavior." "to disable this behavior."
) )
# BCG captures every graph during a synthetic warmup forward at startup
# so that serving never records a fresh graph. That requires
# server-based warmup (a real warmup request issued at startup), not
# request-based warmup which runs no forward until the first request.
if self.enable_breakable_cuda_graph and self.disagg_role == RoleType.MONOLITHIC:
self.warmup = True
self.server_warmup = True
if self.disagg_role != RoleType.MONOLITHIC: if self.disagg_role != RoleType.MONOLITHIC:
self.server_warmup = False self.server_warmup = False
@@ -1343,6 +1476,28 @@ class ServerArgs(DisaggServerArgsMixin):
default=ServerArgs.offload_during_compile, default=ServerArgs.offload_during_compile,
help="Offload components during the torch.compile warmup (the DiT layerwise) so max-autotune fits on tighter-memory GPUs, then restore the configured residency for serving. Skipped when the DiT is already layerwise-offloaded, or under cache-dit / FSDP.", help="Offload components during the torch.compile warmup (the DiT layerwise) so max-autotune fits on tighter-memory GPUs, then restore the configured residency for serving. Skipped when the DiT is already layerwise-offloaded, or under cache-dit / FSDP.",
) )
parser.add_argument(
"--enable-breakable-cuda-graph",
action=StoreBoolean,
default=ServerArgs.enable_breakable_cuda_graph,
help="Capture the DiT forward as breakable CUDA graph segments "
"(split at attention; SP all-to-all / dynamic attention stay "
"eager) to cut per-kernel launch overhead. Mutually exclusive "
"with --enable-torch-compile and Cache-DiT (BCG takes priority). "
"Requires --warmup-resolutions; all of them are captured at warmup.",
)
parser.add_argument(
"--bcg-text-buckets",
type=int,
nargs="+",
default=ServerArgs.bcg_text_buckets,
help="Prompt sequence-length padding budget for breakable CUDA "
"graph. Prompt-conditioning is padded up to the smallest bucket "
"that fits so different prompt lengths reuse one captured graph; "
"warmup captures one graph per bucket. Defaults to "
f"{' '.join(map(str, DEFAULT_BCG_TEXT_BUCKETS))}. "
"Replaces the legacy SGLANG_BCG_TEXT_BUCKETS env var.",
)
parser.add_argument( parser.add_argument(
"--enable-layerwise-nvtx-marker", "--enable-layerwise-nvtx-marker",
@@ -259,6 +259,17 @@ def _resolve_warmup_steps(
server_based_warmup: bool, server_based_warmup: bool,
) -> int: ) -> int:
warmup_steps = server_args.warmup_steps warmup_steps = server_args.warmup_steps
default_steps = sampling_defaults.num_inference_steps
# Breakable CUDA graph captures one graph per step-branch at warmup so that
# serving never records a fresh graph. Run the model's full recommended
# steps (uncapped) so every step-branch signature is captured up front.
if (
getattr(server_args, "enable_breakable_cuda_graph", False) is True
and default_steps
):
return max(int(default_steps), warmup_steps)
if not server_based_warmup: if not server_based_warmup:
return warmup_steps return warmup_steps
@@ -288,6 +299,8 @@ def should_include_warmup_image(
return False return False
if task_type.requires_image_input(): if task_type.requires_image_input():
return True return True
if type(server_args.pipeline_config).__name__ == "GlmImagePipelineConfig":
return False
if server_based_warmup: if server_based_warmup:
return task_type in (ModelTaskType.TI2I, ModelTaskType.TI2V) return task_type in (ModelTaskType.TI2I, ModelTaskType.TI2V)
return True return True
@@ -38,6 +38,7 @@ def _make_unit_server_args():
comfyui_mode=False, comfyui_mode=False,
disable_autocast=False, disable_autocast=False,
enable_cfg_parallel=False, enable_cfg_parallel=False,
enable_breakable_cuda_graph=False,
enable_layerwise_nvtx_marker=False, enable_layerwise_nvtx_marker=False,
enable_torch_compile=False, enable_torch_compile=False,
model_loaded={}, model_loaded={},
@@ -0,0 +1,391 @@
import unittest
from types import SimpleNamespace
from unittest.mock import patch
import torch
from sglang.multimodal_gen.runtime.breakable_cuda_graph.runner import (
DiffusionBreakableCudaGraphRunner,
_CaptureEntry,
_signature_kwargs,
)
from sglang.multimodal_gen.runtime.layers.attention import (
DynamicVarlenMaskMeta,
build_varlen_mask_meta,
)
from sglang.multimodal_gen.runtime.pipelines_core.stages.denoising import (
DenoisingStage,
)
from sglang.multimodal_gen.runtime.server_args import (
BREAKABLE_CUDA_GRAPH_SUPPORTED_MODEL_IDS,
BREAKABLE_CUDA_GRAPH_SUPPORTED_PIPELINE_CONFIGS,
)
class QwenImageTransformer2DModel(torch.nn.Module):
pass
class OtherTransformer2DModel(torch.nn.Module):
pass
class Ideogram4Transformer2DModel(torch.nn.Module):
pass
class ZImageTransformer2DModel(torch.nn.Module):
def rotary_emb(self, pos_ids):
return torch.zeros(pos_ids.shape[0], 8, device=pos_ids.device)
class TestDiffusionBCGPadding(unittest.TestCase):
def setUp(self):
self.stage = DenoisingStage.__new__(DenoisingStage)
self.qwen_model = QwenImageTransformer2DModel()
self.ideogram_model = Ideogram4Transformer2DModel()
self.zimage_model = ZImageTransformer2DModel()
self.other_model = OtherTransformer2DModel()
def _patch_buckets(self, *buckets: int):
resolved = tuple(sorted({b for b in buckets if b > 0}))
return patch.object(
DenoisingStage,
"_bcg_text_buckets",
staticmethod(lambda: resolved),
)
def _qwen_kwargs(self, seq_len: int, *, fill: float = 1.0):
return {
"hidden_states": torch.zeros(1, 4096, 64),
"timestep": torch.zeros(1),
"encoder_hidden_states": [
torch.full((1, seq_len, 3584), fill, dtype=torch.float32)
],
"encoder_hidden_states_mask": None,
"txt_seq_lens": [seq_len],
"freqs_cis": (
torch.zeros(4096, 128, dtype=torch.float32),
torch.ones(seq_len, 128, dtype=torch.float32),
),
"img_shapes": [[(1, 64, 64)]],
}
def test_qwen_prompt_lengths_share_bucket_signature(self):
with self._patch_buckets(256, 512, 1024):
short = self.stage._bcg_pad_prompt_kwargs(
self._qwen_kwargs(19), current_model=self.qwen_model
)
longer = self.stage._bcg_pad_prompt_kwargs(
self._qwen_kwargs(47), current_model=self.qwen_model
)
self.assertEqual(short["encoder_hidden_states"][0].shape, (1, 256, 3584))
self.assertEqual(longer["encoder_hidden_states"][0].shape, (1, 256, 3584))
self.assertEqual(short["encoder_hidden_states_mask"].shape, (1, 256))
self.assertTrue(short["encoder_hidden_states_mask"][0, :19].all())
self.assertFalse(short["encoder_hidden_states_mask"][0, 19:].any())
self.assertTrue(longer["encoder_hidden_states_mask"][0, :47].all())
self.assertFalse(longer["encoder_hidden_states_mask"][0, 47:].any())
self.assertEqual(short["freqs_cis"][1].shape, (256, 128))
self.assertEqual(short["txt_seq_lens"], [256])
self.assertEqual(longer["txt_seq_lens"], [256])
self.assertEqual(_signature_kwargs(short), _signature_kwargs(longer))
def test_qwen_prompt_content_changes_do_not_change_signature(self):
with self._patch_buckets(256, 512, 1024):
first = self.stage._bcg_pad_prompt_kwargs(
self._qwen_kwargs(47, fill=1.0), current_model=self.qwen_model
)
second = self.stage._bcg_pad_prompt_kwargs(
self._qwen_kwargs(47, fill=2.0), current_model=self.qwen_model
)
self.assertFalse(
torch.equal(
first["encoder_hidden_states"][0],
second["encoder_hidden_states"][0],
)
)
self.assertEqual(_signature_kwargs(first), _signature_kwargs(second))
def test_qwen_default_bucket_preserves_mask(self):
def kwargs(valid_len: int):
mask = torch.zeros(1, 64, dtype=torch.bool)
mask[:, :valid_len] = True
out = self._qwen_kwargs(64)
out["encoder_hidden_states_mask"] = mask
out["txt_seq_lens"] = [valid_len]
return out
first = self.stage._bcg_pad_prompt_kwargs(
kwargs(19), current_model=self.qwen_model
)
second = self.stage._bcg_pad_prompt_kwargs(
kwargs(47), current_model=self.qwen_model
)
self.assertEqual(first["encoder_hidden_states"][0].shape[1], 64)
self.assertEqual(first["txt_seq_lens"], [64])
self.assertEqual(second["txt_seq_lens"], [64])
self.assertTrue(first["encoder_hidden_states_mask"][0, :19].all())
self.assertFalse(first["encoder_hidden_states_mask"][0, 19:].any())
self.assertTrue(second["encoder_hidden_states_mask"][0, :47].all())
self.assertFalse(second["encoder_hidden_states_mask"][0, 47:].any())
self.assertEqual(_signature_kwargs(first), _signature_kwargs(second))
def test_non_qwen_kwargs_do_not_take_qwen_padding_path(self):
kwargs = self._qwen_kwargs(47)
with self._patch_buckets(256, 512, 1024):
out = self.stage._bcg_pad_prompt_kwargs(
kwargs, current_model=self.other_model
)
self.assertIs(out, kwargs)
self.assertIsNone(out["encoder_hidden_states_mask"])
self.assertEqual(out["encoder_hidden_states"][0].shape[1], 47)
self.assertEqual(out["txt_seq_lens"], [47])
def _zimage_kwargs(self, seq_len: int, *, fill: float = 1.0):
return {
"hidden_states": [torch.zeros(16, 1, 4, 4)],
"timestep": torch.zeros(1),
"guidance": torch.zeros(1),
"encoder_hidden_states": [
torch.full((seq_len, 16), fill, dtype=torch.float32)
],
"encoder_hidden_states_mask": torch.ones(1, seq_len, dtype=torch.bool),
"freqs_cis": (
torch.zeros(seq_len, 8, dtype=torch.float32),
torch.zeros(256, 8, dtype=torch.float32),
),
"image_seq_len_target": 256,
}
def test_zimage_prompt_lengths_share_bucket_signature(self):
with self._patch_buckets(64, 128):
short = self.stage._bcg_pad_prompt_kwargs(
self._zimage_kwargs(19), current_model=self.zimage_model
)
longer = self.stage._bcg_pad_prompt_kwargs(
self._zimage_kwargs(47), current_model=self.zimage_model
)
self.assertEqual(short["encoder_hidden_states"][0].shape, (64, 16))
self.assertEqual(longer["encoder_hidden_states"][0].shape, (64, 16))
self.assertEqual(short["encoder_hidden_states_mask"].shape, (1, 64))
self.assertEqual(short["caption_valid_lens"].shape, (1,))
self.assertEqual(short["caption_valid_lens"].item(), 19)
self.assertEqual(longer["caption_valid_lens"].item(), 47)
self.assertTrue(short["_use_caption_valid_mask"])
self.assertTrue(longer["_use_caption_valid_mask"])
self.assertFalse(short["encoder_hidden_states_mask"][0, 19:].any())
self.assertFalse(longer["encoder_hidden_states_mask"][0, 47:].any())
self.assertEqual(short["freqs_cis"][0].shape, (64, 8))
self.assertEqual(_signature_kwargs(short), _signature_kwargs(longer))
def _ideogram_kwargs(self, text_seq: int, *, image_seq: int = 4):
total_seq = text_seq + image_seq
indicator = torch.zeros(1, total_seq, dtype=torch.long)
if text_seq:
indicator[:, :text_seq] = 3
indicator[:, text_seq:] = 2
segment_ids = torch.ones(1, total_seq, dtype=torch.long)
if text_seq:
segment_ids[:, :text_seq] = 1
return {
"llm_features": torch.ones(1, total_seq, 8),
"x": torch.zeros(1, total_seq, 16),
"t": torch.zeros(1),
"position_ids": torch.zeros(1, total_seq, 3, dtype=torch.long),
"segment_ids": segment_ids,
"indicator": indicator,
"attn_mask": segment_ids > 0,
"attn_mask_meta": build_varlen_mask_meta(segment_ids > 0),
}
def test_ideogram_prompt_lengths_share_bucket_signature(self):
with self._patch_buckets(64, 128):
short = self.stage._bcg_pad_prompt_kwargs(
self._ideogram_kwargs(19), current_model=self.ideogram_model
)
longer = self.stage._bcg_pad_prompt_kwargs(
self._ideogram_kwargs(47), current_model=self.ideogram_model
)
self.assertEqual(short["llm_features"].shape, (1, 68, 8))
self.assertEqual(longer["llm_features"].shape, (1, 68, 8))
self.assertEqual(short["x"].shape, (1, 68, 16))
self.assertEqual(short["position_ids"].shape, (1, 68, 3))
self.assertEqual(short["segment_ids"][0, 23:].tolist(), [-1] * 45)
self.assertFalse(short["attn_mask"][0, 23:].any())
self.assertIsInstance(short["attn_mask_meta"], DynamicVarlenMaskMeta)
self.assertIs(short["attn_mask_meta"], longer["attn_mask_meta"])
self.assertEqual(_signature_kwargs(short), _signature_kwargs(longer))
def test_ideogram_image_only_kwargs_are_not_prompt_padded(self):
kwargs = self._ideogram_kwargs(0)
with self._patch_buckets(64, 128):
out = self.stage._bcg_pad_prompt_kwargs(
kwargs, current_model=self.ideogram_model
)
self.assertIs(out, kwargs)
self.assertEqual(out["x"].shape, (1, 4, 16))
self.assertIsInstance(out["attn_mask_meta"], dict)
def test_ideogram_is_registered_as_bcg_supported(self):
self.assertIn(
"ideogram-ai/ideogram-4-fp8",
BREAKABLE_CUDA_GRAPH_SUPPORTED_MODEL_IDS,
)
self.assertIn(
"comfy-org/ideogram-4",
BREAKABLE_CUDA_GRAPH_SUPPORTED_MODEL_IDS,
)
self.assertIn(
"Ideogram4PipelineConfig",
BREAKABLE_CUDA_GRAPH_SUPPORTED_PIPELINE_CONFIGS,
)
def test_dynamic_varlen_mask_meta_rebuilds_once_per_replay_token(self):
builder = DynamicVarlenMaskMeta()
mask = torch.tensor([[True, True, False, False]])
calls = []
def fake_build(current_mask):
calls.append(current_mask.clone())
return {"valid": int(current_mask.sum().item())}
with (
patch(
"sglang.multimodal_gen.runtime.layers.attention.layer."
"build_varlen_mask_meta",
side_effect=fake_build,
),
patch(
"sglang.multimodal_gen.runtime.layers.attention.layer."
"get_current_replay_token",
side_effect=[1, 1, 2],
),
):
first = builder.resolve(mask)
mask[0, 2] = True
second = builder.resolve(mask)
third = builder.resolve(mask)
self.assertEqual(first, {"valid": 2})
self.assertIs(second, first)
self.assertEqual(third, {"valid": 3})
self.assertEqual(len(calls), 2)
def test_disabled_bcg_flag_skips_runner(self):
self.stage.server_args = SimpleNamespace(
enable_breakable_cuda_graph=False,
enable_torch_compile=False,
)
self.stage._bcg_runners = {}
self.stage._cache_dit_enabled = False
self.assertIsNone(self.stage._maybe_get_bcg_runner(self.qwen_model))
self.stage._maybe_torch_compile(self.qwen_model)
self.stage._maybe_enable_cache_dit(1, SimpleNamespace(is_warmup=True))
self.assertEqual(self.stage._bcg_runners, {})
def test_bcg_runner_cache_is_per_model_module(self):
self.stage.server_args = SimpleNamespace(enable_breakable_cuda_graph=True)
self.stage._bcg_runners = {}
def fake_runner(model, device):
return SimpleNamespace(model=model, device=device)
with (
patch(
"sglang.multimodal_gen.runtime.breakable_cuda_graph.runner."
"DiffusionBreakableCudaGraphRunner",
side_effect=fake_runner,
),
patch(
"sglang.multimodal_gen.runtime.pipelines_core.stages.denoising."
"get_local_torch_device",
return_value=torch.device("cpu"),
),
):
first = self.stage._maybe_get_bcg_runner(self.qwen_model)
second = self.stage._maybe_get_bcg_runner(self.other_model)
first_again = self.stage._maybe_get_bcg_runner(self.qwen_model)
self.assertIs(first_again, first)
self.assertIsNot(first, second)
self.assertIs(first.model, self.qwen_model)
self.assertIs(second.model, self.other_model)
self.assertEqual(len(self.stage._bcg_runners), 2)
def test_bcg_runner_rejects_too_many_segments(self):
runner = object.__new__(DiffusionBreakableCudaGraphRunner)
runner.max_segments = 2
entry = _CaptureEntry(
graph=SimpleNamespace(_break_fns=[], _segments=[object()] * 3),
static_kwargs={},
static_leaves=[],
output=None,
num_segments=3,
)
self.assertIn("captured 3 segments", runner._capture_limit_reason(entry))
def test_bcg_runner_lazy_capture_only_during_warmup(self):
runner = object.__new__(DiffusionBreakableCudaGraphRunner)
with patch(
"sglang.multimodal_gen.runtime.managers.forward_context.get_forward_context",
return_value=SimpleNamespace(forward_batch=SimpleNamespace(is_warmup=True)),
):
self.assertTrue(runner._should_capture_on_call(("sig",)))
with patch(
"sglang.multimodal_gen.runtime.managers.forward_context.get_forward_context",
return_value=SimpleNamespace(
forward_batch=SimpleNamespace(is_warmup=False)
),
):
self.assertFalse(runner._should_capture_on_call(("sig",)))
def test_bcg_runner_reset_drops_entries_and_marks_disabled(self):
runner = object.__new__(DiffusionBreakableCudaGraphRunner)
runner.device_module = SimpleNamespace(empty_cache=lambda: None)
entry = _CaptureEntry(
graph=SimpleNamespace(_break_fns=[lambda: None], _segments=[object()]),
static_kwargs={"x": torch.zeros(1)},
static_leaves=[torch.zeros(1)],
output=torch.zeros(1),
num_segments=1,
)
runner.entries = {("sig",): entry}
runner._blocked = {("sig",)}
runner.reset(disabled_reason="too much memory")
self.assertEqual(runner.entries, {})
self.assertEqual(runner._blocked, set())
self.assertEqual(entry.graph._break_fns, [])
self.assertEqual(entry.graph._segments, [])
self.assertIsNone(entry.output)
self.assertEqual(runner._disabled_reason, "too much memory")
def test_bcg_runner_allows_unlimited_segments(self):
runner = object.__new__(DiffusionBreakableCudaGraphRunner)
runner.max_segments = 0
entry = _CaptureEntry(
graph=SimpleNamespace(_break_fns=[], _segments=[object()]),
static_kwargs={},
static_leaves=[],
output=None,
num_segments=1,
)
self.assertIsNone(runner._capture_limit_reason(entry))
if __name__ == "__main__":
unittest.main()
@@ -70,6 +70,7 @@ class _GlobalStageArgsMixin:
server_args = SimpleNamespace( server_args = SimpleNamespace(
comfyui_mode=False, comfyui_mode=False,
enable_torch_compile=False, enable_torch_compile=False,
enable_breakable_cuda_graph=False,
enable_cfg_parallel=False, enable_cfg_parallel=False,
attention_backend=None, attention_backend=None,
**kwargs, **kwargs,
@@ -63,7 +63,11 @@ from sglang.multimodal_gen.runtime.pipelines.ideogram import (
_resolve_ideogram4_unconditional_transformer_weights_path, _resolve_ideogram4_unconditional_transformer_weights_path,
) )
from sglang.multimodal_gen.runtime.pipelines_core.schedule_batch import Req from sglang.multimodal_gen.runtime.pipelines_core.schedule_batch import Req
from sglang.multimodal_gen.runtime.pipelines_core.stages.denoising import DenoisingStage from sglang.multimodal_gen.runtime.pipelines_core.stages.denoising import (
DenoisingContext,
DenoisingStage,
DenoisingStepState,
)
from sglang.multimodal_gen.runtime.pipelines_core.stages.model_specific_stages.ideogram import ( from sglang.multimodal_gen.runtime.pipelines_core.stages.model_specific_stages.ideogram import (
IMAGE_POSITION_OFFSET, IMAGE_POSITION_OFFSET,
LLM_TOKEN_INDICATOR, LLM_TOKEN_INDICATOR,
@@ -153,6 +157,7 @@ def _fake_server_args(cfg=None):
pipeline_config=cfg or Ideogram4PipelineConfig(), pipeline_config=cfg or Ideogram4PipelineConfig(),
comfyui_mode=False, comfyui_mode=False,
enable_torch_compile=False, enable_torch_compile=False,
enable_breakable_cuda_graph=False,
attention_backend="torch_sdpa", attention_backend="torch_sdpa",
enable_layerwise_nvtx_marker=False, enable_layerwise_nvtx_marker=False,
model_loaded={"transformer": True}, model_loaded={"transformer": True},
@@ -1117,6 +1122,127 @@ class TestIdeogram4(unittest.TestCase):
self.assertEqual(tuple(decoded.output.shape), (1, 3, 2, 2)) self.assertEqual(tuple(decoded.output.shape), (1, 3, 2, 2))
def test_ideogram_bcg_padded_positive_output_is_cropped(self):
import sglang.multimodal_gen.runtime.server_args as server_args_module
cfg = Ideogram4PipelineConfig()
args = _fake_server_args(cfg)
device = get_local_torch_device()
prev_args = server_args_module._global_server_args
try:
set_global_server_args(args)
transformer = FakeIdeogramTransformer()
unconditional_transformer = FakeIdeogramTransformer()
stage = Ideogram4DenoisingStage(
transformer=transformer,
unconditional_transformer=unconditional_transformer,
pipeline=_fake_ideogram_pipeline(
transformer, unconditional_transformer
),
)
batch = Req(
sampling_params=Ideogram4SamplingParams(
prompt="11 12",
height=256,
width=512,
preset="V4_TURBO_12",
suppress_logs=True,
)
)
batch.prompt_embeds = [torch.zeros(1, 3, 8, device=device)]
batch.extra["ideogram4"] = {
"max_text_tokens": 1,
"num_image_tokens": 2,
"position_ids": torch.zeros(1, 3, 3, dtype=torch.long, device=device),
"segment_ids": torch.ones(1, 3, dtype=torch.long, device=device),
"indicator": torch.tensor(
[
[
LLM_TOKEN_INDICATOR,
OUTPUT_IMAGE_INDICATOR,
OUTPUT_IMAGE_INDICATOR,
]
],
dtype=torch.long,
device=device,
),
}
ctx = DenoisingContext(
scheduler=None,
extra_step_kwargs={},
target_dtype=torch.float32,
autocast_enabled=False,
timesteps=torch.tensor([0], device=device),
num_inference_steps=1,
num_warmup_steps=0,
image_kwargs={},
pos_cond_kwargs={},
neg_cond_kwargs={},
latents=torch.zeros(1, 2, 128, device=device),
boundary_timestep=None,
z=None,
reserved_frames_mask=None,
seq_len=None,
guidance=torch.ones(1, device=device),
is_warmup=False,
extra={
"ideogram4_schedule_values": torch.tensor(
[1.0, 0.0], device=device
),
"ideogram4_schedule_deltas": torch.tensor([1.0], device=device),
"ideogram4_guidance_schedule": torch.tensor([1.0], device=device),
"ideogram4_text_z_padding": torch.zeros(1, 1, 128, device=device),
"ideogram4_attn_mask": torch.ones(
1, 3, dtype=torch.bool, device=device
),
"ideogram4_attn_mask_meta": None,
"ideogram4_neg_position_ids": torch.zeros(
1, 2, 3, dtype=torch.long, device=device
),
"ideogram4_neg_segment_ids": torch.ones(
1, 2, dtype=torch.long, device=device
),
"ideogram4_neg_indicator": torch.full(
(1, 2),
OUTPUT_IMAGE_INDICATOR,
dtype=torch.long,
device=device,
),
"ideogram4_neg_attn_mask": torch.ones(
1, 2, dtype=torch.bool, device=device
),
"ideogram4_neg_attn_mask_meta": None,
"ideogram4_neg_llm_features": torch.zeros(1, 2, 8, device=device),
},
)
step = DenoisingStepState(
step_index=0,
t_host=torch.tensor(0),
t_device=torch.tensor(0, device=device),
t_int=0,
current_model=transformer,
current_guidance_scale=None,
attn_metadata=None,
)
def fake_run(current_model, call_kwargs):
if current_model is transformer:
out = torch.zeros(1, 5, 128, device=device)
out[:, 1:3] = 4.0
out[:, 3:] = 99.0
return out
return torch.ones(1, 2, 128, device=device)
with patch.object(stage, "_run_ideogram_transformer", side_effect=fake_run):
stage._run_denoising_step(ctx, step, batch, args)
finally:
set_global_server_args(prev_args)
self.assertEqual(tuple(ctx.latents.shape), (1, 2, 128))
self.assertTrue(
torch.allclose(ctx.latents, torch.full((1, 2, 128), 4.0, device=device))
)
def test_text_input_builder_matches_official_layout(self): def test_text_input_builder_matches_official_layout(self):
prev_args = None prev_args = None
import sglang.multimodal_gen.runtime.server_args as server_args_module import sglang.multimodal_gen.runtime.server_args as server_args_module
@@ -0,0 +1,44 @@
# Copyright 2023-2026 SGLang Team
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
# ==============================================================================
"""Model-agnostic breakable CUDA graph (BCG) primitives.
Shared by the LLM runtime (``sglang.srt.model_executor``) and the diffusion
runtime (``sglang.multimodal_gen``). Capture a forward region as a sequence of
``torch.cuda.CUDAGraph`` segments separated by eager break points inserted via
:func:`eager_on_graph`-decorated callables.
"""
from sglang.srt.breakable_cuda_graph.breakable_cuda_graph import (
BreakableCUDAGraph,
BreakableCUDAGraphCapture,
break_graph,
eager_on_graph,
get_current_replay_token,
)
from sglang.srt.breakable_cuda_graph.context import (
BCG_FAILURE_HINT,
enable_breakable_cuda_graph,
is_in_breakable_cuda_graph,
)
__all__ = [
"BreakableCUDAGraph",
"BreakableCUDAGraphCapture",
"break_graph",
"eager_on_graph",
"get_current_replay_token",
"BCG_FAILURE_HINT",
"enable_breakable_cuda_graph",
"is_in_breakable_cuda_graph",
]
@@ -0,0 +1,389 @@
# Copyright 2023-2026 SGLang Team
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
# ==============================================================================
"""Breakable CUDA Graph: capture a region as a sequence of
``torch.cuda.CUDAGraph`` segments separated by eager break points.
Each segment is a real ``torch.cuda.CUDAGraph``. Its destructor calls
``releasePool`` on the shared mempool, so the pool's ``use_count`` tracks how
many segments are alive; the pool stays pinned as long as any segment graph
is alive. This lets ``weak_ref_tensor`` views of intermediate pool-allocated
tensors remain valid across replays — we don't need Python-managed bridge
buffers to keep break-point tensors at stable addresses.
This module is model-agnostic. The LLM runtime (``sglang.srt``) breaks at
radix-attention / mamba; the diffusion runtime (``sglang.multimodal_gen``)
breaks at the DiT attention modules, where sequence-parallel all-to-all and
dynamic/varlen/sparse attention kernels must run eagerly between captured
segments. Break-point callables may return a single tensor, a tuple/list of
tensors, or an object/dict of tensors — see :func:`_copy_output`.
"""
import itertools
import logging
import threading
from contextvars import ContextVar
from typing import Any, Callable
import torch
try:
from cuda.bindings import runtime as rt
except ImportError:
rt = None
from sglang.srt.breakable_cuda_graph.cuda_utils import checkCudaErrors
from sglang.srt.utils import is_hip
logger = logging.getLogger(__name__)
__all__ = [
"eager_on_graph",
"BreakableCUDAGraph",
"BreakableCUDAGraphCapture",
"break_graph",
"get_current_replay_token",
]
def _check_cuda_bindings():
if rt is None:
raise ImportError(
"Breakable CUDA graph requires the 'cuda-python' package. "
"Install it with: pip install cuda-python"
)
# Active BreakableCUDAGraphCapture context for the currently-capturing thread.
# eager_on_graph's wrapper uses this to split the current torch.cuda.CUDAGraph
# at break points.
_current_capture_var: ContextVar["BreakableCUDAGraphCapture | None"] = ContextVar(
"current_capture", default=None
)
_current_stream_var: ContextVar[torch.cuda.Stream | None] = ContextVar(
"current_stream", default=None
)
_current_replay_token_var: ContextVar[int | None] = ContextVar(
"current_replay_token", default=None
)
_forked_streams_var: ContextVar[set[torch.cuda.Stream] | None] = ContextVar(
"forked_streams", default=None
)
_replay_token_counter = itertools.count(1)
def get_current_stream(device: torch.device | None = None) -> torch.cuda.Stream:
stream = _current_stream_var.get()
if stream is None:
return torch.cuda.current_stream(device)
return stream
def get_current_replay_token() -> int | None:
"""Return a unique token for the current BCG replay, or ``None``.
Eager break-point code can use this to cache metadata within a single replay
while still rebuilding it for the next replay when static buffers change.
This was added for diffusion model adaptation, where Qwen Image rebuilds
replay-local varlen attention metadata from the current prompt mask.
"""
return _current_replay_token_var.get()
def _capture_status(stream_ptr: int) -> "rt.cudaStreamCaptureStatus":
_check_cuda_bindings()
status, *_ = checkCudaErrors(rt.cudaStreamGetCaptureInfo(stream_ptr))
return status
def _is_stream_capturing(stream: torch.cuda.Stream) -> bool:
# On ROCm/HIP, cuda-python is unavailable, so use the portable torch API
# (which maps to the HIP runtime). On NVIDIA, keep querying the CUDA runtime
# directly via cuda-python: torch.cuda.is_current_stream_capturing() has
# proven unreliable there, so we preserve the original behavior.
if is_hip():
with torch.cuda.stream(stream):
return torch.cuda.is_current_stream_capturing()
return (
_capture_status(stream.cuda_stream)
== rt.cudaStreamCaptureStatus.cudaStreamCaptureStatusActive
)
# Hook torch.cuda.Stream.wait_stream to track side-stream forks/joins that happen
# during breakable capture. We need this because capture_end() on a torch
# CUDAGraph fails if there are still side streams participating in the capture
# — so before ending each segment we auto-join any forked-but-not-rejoined streams.
_original_wait_stream: Callable | None = None
_hook_lock = threading.Lock()
_hook_refcount = 0
def _hooked_wait_stream(self: torch.cuda.Stream, other: torch.cuda.Stream):
assert _original_wait_stream is not None
forked = _forked_streams_var.get()
if forked is None:
_original_wait_stream(self, other)
return
capturing = _current_stream_var.get()
if capturing is None:
_original_wait_stream(self, other)
return
cap_ptr = capturing.cuda_stream
is_self_cap = self is capturing or self.cuda_stream == cap_ptr
is_other_cap = other is capturing or other.cuda_stream == cap_ptr
if is_self_cap and not is_other_cap:
if not _is_stream_capturing(other):
return
_original_wait_stream(self, other)
forked.discard(other)
elif is_other_cap and not is_self_cap:
_original_wait_stream(self, other)
forked.add(self)
else:
_original_wait_stream(self, other)
def _install_wait_stream_hook():
global _original_wait_stream, _hook_refcount
with _hook_lock:
if _hook_refcount == 0:
_original_wait_stream = torch.cuda.Stream.wait_stream
torch.cuda.Stream.wait_stream = _hooked_wait_stream # type: ignore[assignment]
_hook_refcount += 1
def _uninstall_wait_stream_hook():
global _original_wait_stream, _hook_refcount
with _hook_lock:
_hook_refcount -= 1
if _hook_refcount == 0:
assert _original_wait_stream is not None, "wait_stream hook not installed"
torch.cuda.Stream.wait_stream = _original_wait_stream # type: ignore[assignment]
_original_wait_stream = None
def _weak_ref_if_tensor(x):
"""Return a weak-ref tensor view (shared storage, no refcount) for tensors;
recurse into tuples/lists; pass-through for everything else. Weak-ref'ing
captured args/outputs lets the shared mempool reclaim per-layer
intermediates between segments — storage stays alive for each segment
CUDAGraph's lifetime via its pool use_count.
``weak_ref_tensors`` is imported lazily: the module hard-raises on
non-CUDA/NPU platforms, and we only reach this code during an active
BCG capture (which can't happen on CPU-only runners anyway)."""
if torch.is_tensor(x):
from sglang.srt.compilation.weak_ref_tensor import weak_ref_tensors
return weak_ref_tensors(x)
if isinstance(x, tuple):
return tuple(_weak_ref_if_tensor(e) for e in x)
if isinstance(x, list):
return [_weak_ref_if_tensor(e) for e in x]
return x
def _copy_output(dst: Any, src: Any) -> Any:
"""Copy src output into dst in-place where possible.
Handles plain tensors, tuples/lists of tensors, dataclass/object with
tensor attributes, and dicts of tensors. Returns dst if in-place copy
succeeded, otherwise returns src.
The in-place copy is what keeps a break point's output at a stable address
across replays: ``dst`` is the weak-ref'd capture-time output (pinned by the
segment mempool), and the downstream captured segment reads from that
address, so each replay must write fresh data back into ``dst`` rather than
return a freshly-allocated tensor.
"""
if torch.is_tensor(dst) and torch.is_tensor(src):
dst.copy_(src)
return dst
if (
isinstance(dst, (tuple, list))
and isinstance(src, (tuple, list))
and len(dst) == len(src)
):
copied = [_copy_output(d, s) for d, s in zip(dst, src)]
return tuple(copied) if isinstance(dst, tuple) else copied
if hasattr(dst, "__dict__") and hasattr(src, "__dict__"):
for key, src_val in src.__dict__.items():
dst_val = getattr(dst, key, None)
if torch.is_tensor(dst_val) and torch.is_tensor(src_val):
dst_val.copy_(src_val)
else:
setattr(dst, key, src_val)
return dst
if isinstance(dst, dict) and isinstance(src, dict):
for key, src_val in src.items():
dst_val = dst.get(key)
if torch.is_tensor(dst_val) and torch.is_tensor(src_val):
dst_val.copy_(src_val)
else:
dst[key] = src_val
return dst
return src
def eager_on_graph(enable: bool):
def decorator(inner: Callable):
if not enable:
return inner
def wrapper(*args, **kwargs):
capture = _current_capture_var.get()
if capture is None:
return inner(*args, **kwargs)
logger.debug("Break graph due to function: %s", inner.__name__)
# End the segment that captured up to this break point.
capture._end_current_segment()
# Run the eager function once so it allocates its outputs and
# writes real data into them.
output = inner(*args, **kwargs)
# Weak-ref the closure state. Storage lives with the segment
# CUDAGraphs' mempool pin; Python refs don't need to prevent
# pool reuse across layers.
captured_inner = inner
captured_args = tuple(_weak_ref_if_tensor(a) for a in args)
captured_kwargs = {k: _weak_ref_if_tensor(v) for k, v in kwargs.items()}
captured_output = _weak_ref_if_tensor(output)
def replay_fn():
new_out = captured_inner(*captured_args, **captured_kwargs)
return _copy_output(captured_output, new_out)
capture.cuda_graph._break_fns.append(replay_fn)
# Start a fresh CUDAGraph segment for the remainder of the forward.
capture._begin_new_segment()
return output
return wrapper
return decorator
class BreakableCUDAGraph:
"""Container holding one ``torch.cuda.CUDAGraph`` per segment plus an
eager break function between consecutive segments."""
def __init__(self) -> None:
self._segments: list[torch.cuda.CUDAGraph] = []
self._break_fns: list[Callable[[], Any]] = []
def replay(self) -> None:
stream = torch.cuda.current_stream()
stream_token = _current_stream_var.set(stream)
replay_token = _current_replay_token_var.set(next(_replay_token_counter))
try:
for i, seg in enumerate(self._segments):
seg.replay()
if i < len(self._break_fns):
self._break_fns[i]()
finally:
_current_replay_token_var.reset(replay_token)
_current_stream_var.reset(stream_token)
class BreakableCUDAGraphCapture:
"""Context manager that captures the enclosed code as one or more
``torch.cuda.CUDAGraph`` segments separated by eager break points.
Each segment shares the supplied ``pool`` (``MempoolId_t`` tuple) so
pool-allocated intermediates can be reused across segments. While any
segment is alive, its ``beginAllocateToPool`` call keeps the mempool's
``use_count`` > 0, which makes ``weak_ref_tensor`` of segment-allocated
tensors safe across subsequent replays.
"""
def __init__(
self,
cuda_graph: BreakableCUDAGraph,
pool=None,
stream: torch.cuda.Stream | None = None,
capture_error_mode: str = "global",
):
assert isinstance(
cuda_graph, BreakableCUDAGraph
), "cuda_graph must be a BreakableCUDAGraph"
self.cuda_graph = cuda_graph
self._pool = pool if pool is not None else (0, 0)
self._stream = stream
self._capture_error_mode = capture_error_mode
self._stream_ctx = None
self._capture_token = None
self._stream_token = None
self._forked_token = None
def __enter__(self):
_install_wait_stream_hook()
if self._stream is not None:
self._stream_ctx = torch.cuda.stream(self._stream)
self._stream_ctx.__enter__()
self._capture_token = _current_capture_var.set(self)
self._stream_token = _current_stream_var.set(
self._stream or torch.cuda.current_stream()
)
self._forked_token = _forked_streams_var.set(set())
self._begin_new_segment()
return self
def __exit__(self, *args: object):
try:
self._end_current_segment()
finally:
_forked_streams_var.reset(self._forked_token)
_current_stream_var.reset(self._stream_token)
_current_capture_var.reset(self._capture_token)
if self._stream_ctx is not None:
self._stream_ctx.__exit__(*args)
self._stream_ctx = None
_uninstall_wait_stream_hook()
return False
def _begin_new_segment(self) -> None:
graph = torch.cuda.CUDAGraph()
graph.capture_begin(
pool=self._pool, capture_error_mode=self._capture_error_mode
)
self.cuda_graph._segments.append(graph)
def _end_current_segment(self) -> None:
# Auto-join any side streams forked during this segment but not joined.
main_stream = get_current_stream()
forked = _forked_streams_var.get()
if forked:
assert _original_wait_stream is not None
for side in list(forked):
if _is_stream_capturing(side):
_original_wait_stream(main_stream, side)
forked.clear()
self.cuda_graph._segments[-1].capture_end()
@eager_on_graph(True)
def break_graph() -> None:
"""Insert a graph break. The @eager_on_graph decorator does the actual
segment split; this function body intentionally does nothing."""
pass
@@ -0,0 +1,50 @@
# Copyright 2023-2026 SGLang Team
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
# ==============================================================================
"""Runtime state for the breakable CUDA graph (BCG) runner.
Kept intentionally separate from ``compilation/piecewise_context_manager.py``:
BCG no longer inherits from the torch.compile-based PCG path, so its
capture/replay lifecycle is managed on its own.
This module is model-agnostic: it is shared by the LLM runtime
(``sglang.srt``) and the diffusion runtime (``sglang.multimodal_gen``).
"""
from __future__ import annotations
from contextlib import contextmanager
_in_breakable_cuda_graph = False
BCG_FAILURE_HINT = (
"1. change to tc_piecewise by --cuda-graph-backend-prefill=tc_piecewise\n"
"2. disable the prefill CUDA graph by --cuda-graph-backend-prefill=disabled\n"
"3. if it is an OOM problem, set --mem-fraction-static to a smaller value "
"(e.g., 0.8 or 0.7) or set --cuda-graph-max-bs-prefill to a smaller value "
"(e.g., 2048)\n"
)
def is_in_breakable_cuda_graph() -> bool:
return _in_breakable_cuda_graph
@contextmanager
def enable_breakable_cuda_graph():
global _in_breakable_cuda_graph
_in_breakable_cuda_graph = True
try:
yield
finally:
_in_breakable_cuda_graph = False
@@ -0,0 +1,48 @@
# Copyright 2023-2026 SGLang Team
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
# ==============================================================================
"""CUDA runtime binding utilities."""
try:
from cuda.bindings import runtime as rt
except ImportError:
rt = None
def _cudaGetErrorString(error):
if rt is None:
return "<cuda.bindings not available>"
err, msg = rt.cudaGetErrorString(error)
if err != rt.cudaError_t.cudaSuccess:
return "<unknown>"
if isinstance(msg, bytes):
return msg.decode("utf-8", "replace")
return str(msg)
def checkCudaErrors(result):
if rt is None:
raise RuntimeError(
"cuda.bindings is not available. "
"Install it with: pip install cuda-python"
)
if result[0] != rt.cudaError_t.cudaSuccess:
raise RuntimeError(
f"CUDA error {int(result[0])}({_cudaGetErrorString(result[0])})"
)
if len(result) == 1:
return None
elif len(result) == 2:
return result[1]
else:
return result[1:]
@@ -117,7 +117,7 @@ class BreakableCudaGraphBackend(DedupedCudaGraphMixin, BaseCudaGraphBackend):
if post_warmup_hook is not None: if post_warmup_hook is not None:
post_warmup_hook() post_warmup_hook()
graph = BreakableCUDAGraph(self.deduped_cuda_graph) graph = BreakableCUDAGraph()
captured_fn = ( captured_fn = (
eager_on_graph(True)(forward_fn) if self._debug_eager else forward_fn eager_on_graph(True)(forward_fn) if self._debug_eager else forward_fn
) )
@@ -14,8 +14,21 @@ from sglang.srt.model_executor.runner_backend_utils.breakable_cuda_graph.breakab
BreakableCUDAGraphCapture, BreakableCUDAGraphCapture,
break_graph, break_graph,
eager_on_graph, eager_on_graph,
get_current_replay_token,
) )
from sglang.srt.model_executor.runner_backend_utils.breakable_cuda_graph.context import ( # noqa: F401 from sglang.srt.model_executor.runner_backend_utils.breakable_cuda_graph.context import ( # noqa: F401
BCG_FAILURE_HINT,
enable_breakable_cuda_graph, enable_breakable_cuda_graph,
is_in_breakable_cuda_graph, is_in_breakable_cuda_graph,
) )
__all__ = [
"BreakableCUDAGraph",
"BreakableCUDAGraphCapture",
"break_graph",
"eager_on_graph",
"get_current_replay_token",
"BCG_FAILURE_HINT",
"enable_breakable_cuda_graph",
"is_in_breakable_cuda_graph",
]
@@ -11,364 +11,30 @@
# See the License for the specific language governing permissions and # See the License for the specific language governing permissions and
# limitations under the License. # limitations under the License.
# ============================================================================== # ==============================================================================
"""Breakable CUDA Graph: capture a region as a sequence of """Backward-compatible re-export shim.
torch.cuda.CUDAGraph segments separated by eager break points.
Each segment is a real torch.cuda.CUDAGraph. Its destructor calls The breakable CUDA graph primitives moved to the model-agnostic package
releasePool on the shared mempool, so the pool's use_count tracks how :mod:`sglang.srt.breakable_cuda_graph` so the diffusion runtime
many segments are alive; the pool stays pinned as long as any segment graph (``sglang.multimodal_gen``) can share them with the LLM runtime. This module
is alive. This lets weak_ref_tensor views of intermediate pool-allocated preserves the historical import path.
tensors remain valid across replays — we don't need Python-managed bridge
buffers to keep break-point tensors at stable addresses.
""" """
import logging from sglang.srt.breakable_cuda_graph.breakable_cuda_graph import ( # noqa: F401
import threading BreakableCUDAGraph,
from contextvars import ContextVar BreakableCUDAGraphCapture,
from typing import Any, Callable _copy_output,
break_graph,
import torch eager_on_graph,
get_current_replay_token,
try: get_current_stream,
from cuda.bindings import runtime as rt
except ImportError:
rt = None
from sglang.srt.model_executor.runner_backend_utils.breakable_cuda_graph.cuda_utils import (
checkCudaErrors,
) )
from sglang.srt.utils import is_hip
logger = logging.getLogger(__name__)
__all__ = [ __all__ = [
"eager_on_graph", "eager_on_graph",
"BreakableCUDAGraph", "BreakableCUDAGraph",
"BreakableCUDAGraphCapture", "BreakableCUDAGraphCapture",
"_copy_output",
"break_graph", "break_graph",
"get_current_stream",
"get_current_replay_token",
] ]
def _check_cuda_bindings():
if rt is None:
raise ImportError(
"Breakable CUDA graph on NVIDIA requires the 'cuda-python' package. "
"Install it with: pip install cuda-python"
)
# Active BreakableCUDAGraphCapture context for the currently-capturing thread.
# eager_on_graph's wrapper uses this to split the current torch.cuda.CUDAGraph
# at break points.
_current_capture_var: ContextVar["BreakableCUDAGraphCapture | None"] = ContextVar(
"current_capture", default=None
)
_current_stream_var: ContextVar[torch.cuda.Stream | None] = ContextVar(
"current_stream", default=None
)
_forked_streams_var: ContextVar[set[torch.cuda.Stream] | None] = ContextVar(
"forked_streams", default=None
)
def get_current_stream(device: torch.device | None = None) -> torch.cuda.Stream:
stream = _current_stream_var.get()
if stream is None:
return torch.cuda.current_stream(device)
return stream
def _capture_status(stream_ptr: int) -> "rt.cudaStreamCaptureStatus":
_check_cuda_bindings()
status, *_ = checkCudaErrors(rt.cudaStreamGetCaptureInfo(stream_ptr))
return status
def _is_stream_capturing(stream: torch.cuda.Stream) -> bool:
# On ROCm/HIP, cuda-python is unavailable, so use the portable torch API
# (which maps to the HIP runtime). On NVIDIA, keep querying the CUDA runtime
# directly via cuda-python: torch.cuda.is_current_stream_capturing() has
# proven unreliable there, so we preserve the original behavior.
if is_hip():
with torch.cuda.stream(stream):
return torch.cuda.is_current_stream_capturing()
return (
_capture_status(stream.cuda_stream)
== rt.cudaStreamCaptureStatus.cudaStreamCaptureStatusActive
)
# Hook torch.cuda.Stream.wait_stream to track side-stream forks/joins that happen
# during breakable capture. We need this because capture_end() on a torch
# CUDAGraph fails if there are still side streams participating in the capture
# — so before ending each segment we auto-join any forked-but-not-rejoined streams.
_original_wait_stream: Callable | None = None
_hook_lock = threading.Lock()
_hook_refcount = 0
def _hooked_wait_stream(self: torch.cuda.Stream, other: torch.cuda.Stream):
assert _original_wait_stream is not None
forked = _forked_streams_var.get()
if forked is None:
_original_wait_stream(self, other)
return
capturing = _current_stream_var.get()
if capturing is None:
_original_wait_stream(self, other)
return
cap_ptr = capturing.cuda_stream
is_self_cap = self is capturing or self.cuda_stream == cap_ptr
is_other_cap = other is capturing or other.cuda_stream == cap_ptr
if is_self_cap and not is_other_cap:
if not _is_stream_capturing(other):
return
_original_wait_stream(self, other)
forked.discard(other)
elif is_other_cap and not is_self_cap:
_original_wait_stream(self, other)
forked.add(self)
else:
_original_wait_stream(self, other)
def _install_wait_stream_hook():
global _original_wait_stream, _hook_refcount
with _hook_lock:
if _hook_refcount == 0:
_original_wait_stream = torch.cuda.Stream.wait_stream
torch.cuda.Stream.wait_stream = _hooked_wait_stream # type: ignore[assignment]
_hook_refcount += 1
def _uninstall_wait_stream_hook():
global _original_wait_stream, _hook_refcount
with _hook_lock:
_hook_refcount -= 1
if _hook_refcount == 0:
assert _original_wait_stream is not None, "wait_stream hook not installed"
torch.cuda.Stream.wait_stream = _original_wait_stream # type: ignore[assignment]
_original_wait_stream = None
def _weak_ref_if_tensor(x):
"""Return a weak-ref tensor view (shared storage, no refcount) for tensors;
pass-through for non-tensors. Weak-ref'ing captured args lets the shared
mempool reclaim per-layer intermediates between segments — storage stays
alive for each segment CUDAGraph's lifetime via its pool use_count.
weak_ref_tensors is imported lazily because it hard-raises on
platforms without a CUDA/HIP/NPU backend; we only reach this code during
an active Breakable capture, which runs only on those backends."""
if torch.is_tensor(x):
from sglang.srt.compilation.weak_ref_tensor import weak_ref_tensors
return weak_ref_tensors(x)
return x
def _copy_output(dst: Any, src: Any) -> Any:
"""Copy src output into dst in-place where possible.
Handles plain tensors, dataclass/object with tensor attributes,
and dicts of tensors. Returns dst if in-place copy succeeded,
otherwise returns src.
"""
if torch.is_tensor(dst) and torch.is_tensor(src):
dst.copy_(src)
return dst
if hasattr(dst, "__dict__") and hasattr(src, "__dict__"):
for key, src_val in src.__dict__.items():
dst_val = getattr(dst, key, None)
if torch.is_tensor(dst_val) and torch.is_tensor(src_val):
dst_val.copy_(src_val)
else:
setattr(dst, key, src_val)
return dst
if isinstance(dst, dict) and isinstance(src, dict):
for key, src_val in src.items():
dst_val = dst.get(key)
if torch.is_tensor(dst_val) and torch.is_tensor(src_val):
dst_val.copy_(src_val)
else:
dst[key] = src_val
return dst
return src
def eager_on_graph(enable: bool):
def decorator(inner: Callable):
if not enable:
return inner
def wrapper(*args, **kwargs):
capture = _current_capture_var.get()
if capture is None:
return inner(*args, **kwargs)
logger.debug("Break graph due to function: %s", inner.__name__)
# End the segment that captured up to this break point.
capture._end_current_segment()
# Run the eager function once so it allocates its outputs and
# writes real data into them.
output = inner(*args, **kwargs)
# Weak-ref the closure state. Storage lives with the segment
# CUDAGraphs' mempool pin; Python refs don't need to prevent
# pool reuse across layers.
captured_inner = inner
captured_args = tuple(_weak_ref_if_tensor(a) for a in args)
captured_kwargs = {k: _weak_ref_if_tensor(v) for k, v in kwargs.items()}
captured_output = _weak_ref_if_tensor(output)
def replay_fn():
new_out = captured_inner(*captured_args, **captured_kwargs)
return _copy_output(captured_output, new_out)
capture.cuda_graph._break_fns.append(replay_fn)
# Start a fresh CUDAGraph segment for the remainder of the forward.
capture._begin_new_segment()
return output
return wrapper
return decorator
class BreakableCUDAGraph:
"""Container holding one torch.cuda.CUDAGraph per segment plus an
eager break function between consecutive segments."""
def __init__(self, deduped_cuda_graph=None) -> None:
self._segments: list[Any] = []
self._break_fns: list[Callable[[], Any]] = []
self._deduped_cuda_graph = deduped_cuda_graph
def replay(self) -> None:
stream = torch.cuda.current_stream()
token = _current_stream_var.set(stream)
try:
for i, seg in enumerate(self._segments):
seg.replay()
if i < len(self._break_fns):
self._break_fns[i]()
finally:
_current_stream_var.reset(token)
def _append_segment(
self, graph: torch.cuda.CUDAGraph, needs_instantiate: bool
) -> None:
if self._deduped_cuda_graph is not None:
self._segments.append(self._deduped_cuda_graph.register(graph))
return
if needs_instantiate:
graph.instantiate()
self._segments.append(graph)
class BreakableCUDAGraphCapture:
"""Context manager that captures the enclosed code as one or more
torch.cuda.CUDAGraph segments separated by eager break points.
Each segment shares the supplied pool (MempoolId_t tuple) so
pool-allocated intermediates can be reused across segments. While any
segment is alive, its beginAllocateToPool call keeps the mempool's
use_count > 0, which makes weak_ref_tensor of segment-allocated
tensors safe across subsequent replays.
"""
def __init__(
self,
cuda_graph: BreakableCUDAGraph,
pool=None,
stream: torch.cuda.Stream | None = None,
capture_error_mode: str = "global",
):
assert isinstance(
cuda_graph, BreakableCUDAGraph
), "cuda_graph must be a BreakableCUDAGraph"
self.cuda_graph = cuda_graph
self._pool = pool if pool is not None else (0, 0)
self._stream = stream
self._capture_error_mode = capture_error_mode
self._stream_ctx = None
self._capture_token = None
self._stream_token = None
self._forked_token = None
self._current_graph: torch.cuda.CUDAGraph | None = None
self._current_graph_needs_instantiate = False
def __enter__(self):
_install_wait_stream_hook()
if self._stream is not None:
self._stream_ctx = torch.cuda.stream(self._stream)
self._stream_ctx.__enter__()
self._capture_token = _current_capture_var.set(self)
self._stream_token = _current_stream_var.set(
self._stream or torch.cuda.current_stream()
)
self._forked_token = _forked_streams_var.set(set())
self._begin_new_segment()
return self
def __exit__(self, *args: object):
try:
self._end_current_segment()
finally:
_forked_streams_var.reset(self._forked_token)
_current_stream_var.reset(self._stream_token)
_current_capture_var.reset(self._capture_token)
if self._stream_ctx is not None:
self._stream_ctx.__exit__(*args)
self._stream_ctx = None
_uninstall_wait_stream_hook()
return False
def _begin_new_segment(self) -> None:
# keep_graph retains the raw graph for dedup; skip it on the plain path.
if self.cuda_graph._deduped_cuda_graph is not None:
try:
graph = torch.cuda.CUDAGraph(keep_graph=True)
self._current_graph_needs_instantiate = True
except TypeError:
graph = torch.cuda.CUDAGraph()
self._current_graph_needs_instantiate = False
else:
graph = torch.cuda.CUDAGraph()
self._current_graph_needs_instantiate = False
graph.capture_begin(
pool=self._pool, capture_error_mode=self._capture_error_mode
)
self._current_graph = graph
def _end_current_segment(self) -> None:
# Auto-join any side streams forked during this segment but not joined.
main_stream = get_current_stream()
forked = _forked_streams_var.get()
if forked:
assert _original_wait_stream is not None
for side in list(forked):
if _is_stream_capturing(side):
_original_wait_stream(main_stream, side)
forked.clear()
graph = self._current_graph
assert graph is not None
graph.capture_end()
self.cuda_graph._append_segment(graph, self._current_graph_needs_instantiate)
self._current_graph = None
self._current_graph_needs_instantiate = False
@eager_on_graph(True)
def break_graph() -> None:
"""Insert a graph break. The @eager_on_graph decorator does the actual
segment split; this function body intentionally does nothing."""
pass
@@ -11,50 +11,19 @@
# See the License for the specific language governing permissions and # See the License for the specific language governing permissions and
# limitations under the License. # limitations under the License.
# ============================================================================== # ==============================================================================
"""Runtime state for the breakable CUDA graph runner.""" """Backward-compatible re-export shim for the moved BCG context helpers.
from __future__ import annotations See :mod:`sglang.srt.breakable_cuda_graph.context`.
"""
import logging from sglang.srt.breakable_cuda_graph.context import ( # noqa: F401
from contextlib import contextmanager BCG_FAILURE_HINT,
enable_breakable_cuda_graph,
from sglang.srt.model_executor.cuda_graph_config import Backend is_in_breakable_cuda_graph,
from sglang.srt.model_executor.runner_backend_utils import (
PREFILL_CUDA_GRAPH_CAPTURE_FAILED_MSG,
) )
logger = logging.getLogger(__name__) __all__ = [
"BCG_FAILURE_HINT",
_in_breakable_cuda_graph = False "enable_breakable_cuda_graph",
"is_in_breakable_cuda_graph",
]
def is_in_breakable_cuda_graph() -> bool:
return _in_breakable_cuda_graph
@contextmanager
def enable_breakable_cuda_graph():
"""Mark the enclosed scope as inside a BCG capture/replay. Any exception
raised inside is logged with the BCG-specific failure hint, then re-raised
for the caller to handle."""
global _in_breakable_cuda_graph
_in_breakable_cuda_graph = True
try:
yield
except Exception as exc:
msg = PREFILL_CUDA_GRAPH_CAPTURE_FAILED_MSG.format(
backend=Backend.BREAKABLE, suggestions=BCG_FAILURE_HINT
)
logger.error(f"{type(exc).__name__}: {exc}\n{msg}")
raise
finally:
_in_breakable_cuda_graph = False
BCG_FAILURE_HINT = (
"1. change to tc_piecewise by --cuda-graph-backend-prefill=tc_piecewise\n"
"2. disable the prefill CUDA graph by --cuda-graph-backend-prefill=disabled\n"
"3. if it is an OOM problem, set --mem-fraction-static to a smaller value "
"(e.g., 0.8 or 0.7) or set --cuda-graph-max-bs-prefill to a smaller value "
"(e.g., 2048)\n"
)
@@ -11,38 +11,13 @@
# See the License for the specific language governing permissions and # See the License for the specific language governing permissions and
# limitations under the License. # limitations under the License.
# ============================================================================== # ==============================================================================
"""CUDA runtime binding utilities.""" """Backward-compatible re-export shim for the moved CUDA runtime utilities.
try: See :mod:`sglang.srt.breakable_cuda_graph.cuda_utils`.
from cuda.bindings import runtime as rt """
except ImportError:
rt = None
from sglang.srt.breakable_cuda_graph.cuda_utils import ( # noqa: F401
checkCudaErrors,
)
def _cudaGetErrorString(error): __all__ = ["checkCudaErrors"]
if rt is None:
return "<cuda.bindings not available>"
err, msg = rt.cudaGetErrorString(error)
if err != rt.cudaError_t.cudaSuccess:
return "<unknown>"
if isinstance(msg, bytes):
return msg.decode("utf-8", "replace")
return str(msg)
def checkCudaErrors(result):
if rt is None:
raise RuntimeError(
"cuda.bindings is not available. "
"Install it with: pip install cuda-python"
)
if result[0] != rt.cudaError_t.cudaSuccess:
raise RuntimeError(
f"CUDA error {int(result[0])}({_cudaGetErrorString(result[0])})"
)
if len(result) == 1:
return None
elif len(result) == 2:
return result[1]
else:
return result[1:]