[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:
co-authored by
BBuf
Claude Opus 4.8
BBuf
parent
c9303a08da
commit
33c3dfd7e0
@@ -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)
|
||||||
+102
@@ -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.
|
||||||
|
|||||||
+14
@@ -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,
|
||||||
|
|||||||
+18
-3
@@ -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
|
||||||
)
|
)
|
||||||
|
|||||||
+13
@@ -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",
|
||||||
|
]
|
||||||
|
|||||||
+16
-350
@@ -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
|
|
||||||
|
|||||||
+12
-43
@@ -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"
|
|
||||||
)
|
|
||||||
|
|||||||
+7
-32
@@ -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:]
|
|
||||||
|
|||||||
Reference in New Issue
Block a user