[diffusion] Enable breakable CUDA graph for LTX-2 (H200 two-stage e2e 10.75 s -> 6.90 s, 1.56x) (#33885)
This commit is contained in:
@@ -237,7 +237,10 @@ class BaseBreakableCudaGraphRunner:
|
|||||||
self._blocked: set[tuple] = set()
|
self._blocked: set[tuple] = set()
|
||||||
self._disabled_reason: str | None = None
|
self._disabled_reason: str | None = None
|
||||||
self.max_entries = max(0, _env_int("SGLANG_DIFFUSION_BCG_MAX_ENTRIES", 32))
|
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))
|
# LTX-2 dual-tower blocks carry 6 attention break points each
|
||||||
|
# (video/audio self, video/audio prompt-cross, a2v, v2a), so 48 blocks
|
||||||
|
# capture ~289 segments; keep headroom above that.
|
||||||
|
self.max_segments = max(0, _env_int("SGLANG_DIFFUSION_BCG_MAX_SEGMENTS", 512))
|
||||||
|
|
||||||
def __getattr__(self, name: str) -> Any:
|
def __getattr__(self, name: str) -> Any:
|
||||||
# Only reached for attributes the runner itself does not define; proxy
|
# Only reached for attributes the runner itself does not define; proxy
|
||||||
@@ -307,12 +310,41 @@ class BaseBreakableCudaGraphRunner:
|
|||||||
entry = self.entries.get(key)
|
entry = self.entries.get(key)
|
||||||
if entry is None:
|
if entry is None:
|
||||||
if not self._should_capture_on_call(key):
|
if not self._should_capture_on_call(key):
|
||||||
|
self._log_signature_miss(key)
|
||||||
return self.transformer(**kwargs)
|
return self.transformer(**kwargs)
|
||||||
if not self.capture(**kwargs):
|
if not self.capture(**kwargs):
|
||||||
return self.transformer(**kwargs)
|
return self.transformer(**kwargs)
|
||||||
entry = self.entries[key]
|
entry = self.entries[key]
|
||||||
return self.replay(entry, kwargs)
|
return self.replay(entry, kwargs)
|
||||||
|
|
||||||
|
def _log_signature_miss(self, key: tuple) -> None:
|
||||||
|
"""One-shot diagnostic: serving signature missed every captured graph."""
|
||||||
|
if getattr(self, "_miss_logged", False) or not self.entries:
|
||||||
|
return
|
||||||
|
self._miss_logged = True
|
||||||
|
key_d = dict(key)
|
||||||
|
logger.warning(
|
||||||
|
"[Diffusion BCG] serving signature MISSED %d captured graph(s); "
|
||||||
|
"running eager.",
|
||||||
|
len(self.entries),
|
||||||
|
)
|
||||||
|
for captured_key in self.entries:
|
||||||
|
cap_d = dict(captured_key)
|
||||||
|
names = sorted(set(key_d) | set(cap_d))
|
||||||
|
diffs = [
|
||||||
|
(
|
||||||
|
n,
|
||||||
|
_signature_summary_leaf(key_d.get(n, "<absent>")),
|
||||||
|
_signature_summary_leaf(cap_d.get(n, "<absent>")),
|
||||||
|
)
|
||||||
|
for n in names
|
||||||
|
if key_d.get(n, "<absent>") != cap_d.get(n, "<absent>")
|
||||||
|
]
|
||||||
|
logger.warning(
|
||||||
|
"[Diffusion BCG] differing fields (serving vs captured): %s",
|
||||||
|
diffs[:8],
|
||||||
|
)
|
||||||
|
|
||||||
def replay(self, entry: _CaptureEntry, kwargs: dict[str, Any]) -> Any:
|
def replay(self, entry: _CaptureEntry, kwargs: dict[str, Any]) -> Any:
|
||||||
live_leaves = _flatten_kwargs(kwargs)
|
live_leaves = _flatten_kwargs(kwargs)
|
||||||
if len(live_leaves) != len(entry.static_leaves):
|
if len(live_leaves) != len(entry.static_leaves):
|
||||||
|
|||||||
+50
-2
@@ -1169,6 +1169,28 @@ class LTX2DenoisingStage(DenoisingStage):
|
|||||||
audio_latent_model_input,
|
audio_latent_model_input,
|
||||||
num_frames=audio_num_frames_latent,
|
num_frames=audio_num_frames_latent,
|
||||||
)
|
)
|
||||||
|
if server_args.enable_breakable_cuda_graph:
|
||||||
|
# The in-model RoPE coordinate construction builds host
|
||||||
|
# tensors (torch.tensor(list, device=cuda)), which is an
|
||||||
|
# unpinned H2D copy and therefore illegal inside CUDA graph
|
||||||
|
# capture. Build the coords outside the captured region with
|
||||||
|
# the exact same rope helpers (start_frame=0 == the sp<=1
|
||||||
|
# in-model path), so values are bit-identical.
|
||||||
|
if video_coords is None:
|
||||||
|
video_coords = step.current_model.rope.prepare_video_coords(
|
||||||
|
batch_size=int(latent_model_input.shape[0]),
|
||||||
|
num_frames=ctx.latent_num_frames_for_model,
|
||||||
|
height=ctx.latent_height,
|
||||||
|
width=ctx.latent_width,
|
||||||
|
device=latent_model_input.device,
|
||||||
|
fps=batch.fps,
|
||||||
|
)
|
||||||
|
if audio_coords is None:
|
||||||
|
audio_coords = step.current_model.audio_rope.prepare_audio_coords(
|
||||||
|
batch_size=int(audio_latent_model_input.shape[0]),
|
||||||
|
num_frames=audio_num_frames_latent,
|
||||||
|
device=audio_latent_model_input.device,
|
||||||
|
)
|
||||||
|
|
||||||
batch_size = int(latent_model_input.shape[0])
|
batch_size = int(latent_model_input.shape[0])
|
||||||
use_raw_sigma_timestep = ctx.use_ltx23_hq_timestep_semantics
|
use_raw_sigma_timestep = ctx.use_ltx23_hq_timestep_semantics
|
||||||
@@ -1469,6 +1491,28 @@ class LTX2DenoisingStage(DenoisingStage):
|
|||||||
):
|
):
|
||||||
yield
|
yield
|
||||||
|
|
||||||
|
def _ltx2_call_current_model(
|
||||||
|
self,
|
||||||
|
ctx: "LTX2DenoisingContext",
|
||||||
|
step: DenoisingStepState,
|
||||||
|
model_kwargs: dict,
|
||||||
|
):
|
||||||
|
"""Run the LTX-2 DiT forward, replaying a breakable CUDA graph when
|
||||||
|
one is captured for this input signature.
|
||||||
|
|
||||||
|
LTX-2 builds its model kwargs locally instead of going through the
|
||||||
|
generic ``predict_noise`` path, so BCG must be routed here. Capture is
|
||||||
|
driven explicitly from the warmup request (``ctx.is_warmup``); LTX-2
|
||||||
|
tokenizes prompts to a fixed max length, so every serving request
|
||||||
|
shares the warmup signatures and no text bucketing is needed.
|
||||||
|
"""
|
||||||
|
runner = self._maybe_get_bcg_runner(step.current_model)
|
||||||
|
if runner is None:
|
||||||
|
return step.current_model(**model_kwargs)
|
||||||
|
if ctx.is_warmup:
|
||||||
|
runner.capture(**model_kwargs)
|
||||||
|
return runner(**model_kwargs)
|
||||||
|
|
||||||
def _prepare_denoising_loop(
|
def _prepare_denoising_loop(
|
||||||
self,
|
self,
|
||||||
batch: Req,
|
batch: Req,
|
||||||
@@ -1752,7 +1796,9 @@ class LTX2DenoisingStage(DenoisingStage):
|
|||||||
)
|
)
|
||||||
|
|
||||||
with self._ltx2_model_forward_context(ctx, step):
|
with self._ltx2_model_forward_context(ctx, step):
|
||||||
model_video, model_audio = step.current_model(**model_kwargs)
|
model_video, model_audio = self._ltx2_call_current_model(
|
||||||
|
ctx, step, model_kwargs
|
||||||
|
)
|
||||||
|
|
||||||
model_video = model_video.float()
|
model_video = model_video.float()
|
||||||
model_audio = model_audio.float()
|
model_audio = model_audio.float()
|
||||||
@@ -1848,7 +1894,9 @@ class LTX2DenoisingStage(DenoisingStage):
|
|||||||
)
|
)
|
||||||
|
|
||||||
with self._ltx2_model_forward_context(ctx, step):
|
with self._ltx2_model_forward_context(ctx, step):
|
||||||
mid_v, mid_a = step.current_model(**model_kwargs_local)
|
mid_v, mid_a = self._ltx2_call_current_model(
|
||||||
|
ctx, step, model_kwargs_local
|
||||||
|
)
|
||||||
|
|
||||||
mid_v = mid_v.float()
|
mid_v = mid_v.float()
|
||||||
mid_a = mid_a.float()
|
mid_a = mid_a.float()
|
||||||
|
|||||||
@@ -143,6 +143,8 @@ BREAKABLE_CUDA_GRAPH_SUPPORTED_MODEL_IDS = frozenset(
|
|||||||
"ideogram-v4-instant",
|
"ideogram-v4-instant",
|
||||||
"ideogram-ai/ideogram-4-fp8",
|
"ideogram-ai/ideogram-4-fp8",
|
||||||
"ideogram-ai/ideogram-4-nf4",
|
"ideogram-ai/ideogram-4-nf4",
|
||||||
|
"lightricks/ltx-2",
|
||||||
|
"ltx-2",
|
||||||
"minimax-h3",
|
"minimax-h3",
|
||||||
"minimaxai/minimax-h3",
|
"minimaxai/minimax-h3",
|
||||||
"qwen/qwen-image",
|
"qwen/qwen-image",
|
||||||
@@ -161,6 +163,7 @@ BREAKABLE_CUDA_GRAPH_SUPPORTED_PIPELINE_CONFIGS = frozenset(
|
|||||||
{
|
{
|
||||||
"GlmImagePipelineConfig",
|
"GlmImagePipelineConfig",
|
||||||
"Ideogram4PipelineConfig",
|
"Ideogram4PipelineConfig",
|
||||||
|
"LTX2PipelineConfig",
|
||||||
"MiniMaxH3PipelineConfig",
|
"MiniMaxH3PipelineConfig",
|
||||||
"QwenImagePipelineConfig",
|
"QwenImagePipelineConfig",
|
||||||
"ZImagePipelineConfig",
|
"ZImagePipelineConfig",
|
||||||
@@ -556,7 +559,7 @@ class ServerArgs(DisaggServerArgsMixin):
|
|||||||
return
|
return
|
||||||
|
|
||||||
logger.warning(
|
logger.warning(
|
||||||
"[Diffusion BCG] disabled for %s: only Ideogram-4, MiniMax-H3, "
|
"[Diffusion BCG] disabled for %s: only Ideogram-4, Lightricks/LTX-2, MiniMax-H3, "
|
||||||
"Qwen/Qwen-Image, Qwen/Qwen-Image-2512, "
|
"Qwen/Qwen-Image, Qwen/Qwen-Image-2512, "
|
||||||
"Tongyi-MAI/Z-Image/Z-Image-Turbo, and zai-org/GLM-Image are "
|
"Tongyi-MAI/Z-Image/Z-Image-Turbo, and zai-org/GLM-Image are "
|
||||||
"currently supported.",
|
"currently supported.",
|
||||||
|
|||||||
@@ -240,6 +240,13 @@ def _resolve_warmup_num_frames(
|
|||||||
# use default num frames
|
# use default num frames
|
||||||
return num_frames
|
return num_frames
|
||||||
|
|
||||||
|
# Breakable CUDA graph replays only exact latent shapes: the warmup
|
||||||
|
# request must run the full serving frame count so its captured graphs
|
||||||
|
# match serving signatures (mirrors the uncapped-steps rule in
|
||||||
|
# _resolve_warmup_steps).
|
||||||
|
if getattr(server_args, "enable_breakable_cuda_graph", False) is True:
|
||||||
|
return num_frames
|
||||||
|
|
||||||
return min(num_frames, SERVER_WARMUP_MAX_VIDEO_FRAMES)
|
return min(num_frames, SERVER_WARMUP_MAX_VIDEO_FRAMES)
|
||||||
|
|
||||||
|
|
||||||
@@ -298,6 +305,14 @@ 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 getattr(server_args, "enable_breakable_cuda_graph", False) is True:
|
||||||
|
# BCG replays only exact input signatures. A synthetic warmup image
|
||||||
|
# flips optional-TI2V pipelines (e.g. LTX-2) into image-conditioned
|
||||||
|
# kwargs (denoise-mask -> per-token timestep) that pure T2V serving
|
||||||
|
# never produces, so every T2V request would miss the captured
|
||||||
|
# graphs and silently run eager. Capture the T2V signature instead;
|
||||||
|
# image-conditioned requests fall back to eager.
|
||||||
|
return False
|
||||||
if type(server_args.pipeline_config).__name__ == "GlmImagePipelineConfig":
|
if type(server_args.pipeline_config).__name__ == "GlmImagePipelineConfig":
|
||||||
return False
|
return False
|
||||||
if server_based_warmup:
|
if server_based_warmup:
|
||||||
|
|||||||
Reference in New Issue
Block a user