diff --git a/python/sglang/multimodal_gen/runtime/breakable_cuda_graph/runner.py b/python/sglang/multimodal_gen/runtime/breakable_cuda_graph/runner.py index f65ca2f7f..c335015ad 100644 --- a/python/sglang/multimodal_gen/runtime/breakable_cuda_graph/runner.py +++ b/python/sglang/multimodal_gen/runtime/breakable_cuda_graph/runner.py @@ -344,14 +344,65 @@ class BaseBreakableCudaGraphRunner: "[Diffusion BCG] differing fields (serving vs captured): %s", diffs[:8], ) - logger.warning( - "[Diffusion BCG] hint: graphs replay only for the exact shapes " - "captured at warmup. A ``hidden_states`` difference above means " - "the request resolution was never captured (the auto-derived " - "warmup resolution is the model default, which can differ from " - "the resolutions you actually serve) -- declare every served " - "resolution explicitly, e.g. --warmup-resolutions 1024x1024." - ) + temporal_miss = self._has_temporal_shape_miss(key) + if temporal_miss: + logger.warning( + "[Diffusion BCG] hint: the ``hidden_states`` temporal (frame) " + "dimension differs between serving and the captured graph. " + "Breakable CUDA graphs replay only the exact frame count " + "captured at warmup; the warmup frame count defaults to the " + "model sampling default, which can differ from the frames you " + "actually serve -- declare it explicitly, e.g. " + "--warmup-num-frames 17 (and --warmup-resolutions WxH)." + ) + else: + logger.warning( + "[Diffusion BCG] hint: graphs replay only for the exact shapes " + "captured at warmup. A ``hidden_states`` difference above means " + "the request resolution was never captured (the auto-derived " + "warmup resolution is the model default, which can differ from " + "the resolutions you actually serve) -- declare every served " + "resolution explicitly, e.g. --warmup-resolutions 1024x1024." + ) + + def _has_temporal_shape_miss(self, key: tuple) -> bool: + """True when a captured ``hidden_states`` differs only in its temporal dim. + + Video latents are shaped ``[B, C, F, H, W]``; a mismatch in ``F`` (not + ``H``/``W``) points to a frame-count gap rather than a resolution gap. + """ + + def _hidden_shape(k: tuple) -> tuple | None: + for name, leaf in k: + if name != "hidden_states": + continue + # tensor leaf: ("tensor", shape, dtype) + if ( + isinstance(leaf, tuple) + and len(leaf) == 3 + and leaf[0] == "tensor" + and isinstance(leaf[1], tuple) + and len(leaf[1]) == 5 + ): + return leaf[1] + return None + + serving = _hidden_shape(key) + if serving is None: + return False + for captured_key in self.entries: + cap = _hidden_shape(captured_key) + if cap is None: + continue + same_spatial = ( + serving[0] == cap[0] + and serving[1] == cap[1] + and serving[3] == cap[3] + and serving[4] == cap[4] + ) + if same_spatial and serving[2] != cap[2]: + return True + return False def replay(self, entry: _CaptureEntry, kwargs: dict[str, Any]) -> Any: live_leaves = _flatten_kwargs(kwargs) diff --git a/python/sglang/multimodal_gen/runtime/server_args/server_args.py b/python/sglang/multimodal_gen/runtime/server_args/server_args.py index b5fe29fe3..04d3e897c 100644 --- a/python/sglang/multimodal_gen/runtime/server_args/server_args.py +++ b/python/sglang/multimodal_gen/runtime/server_args/server_args.py @@ -680,6 +680,15 @@ class ServerArgs(DisaggServerArgsMixin): "model default warmup resolution. Requests at other " "resolutions run eager." ) + if self._is_video_gen_task() and self.warmup_num_frames is None: + default_frames = self._bcg_default_warmup_num_frames() + logger.info( + "[Diffusion BCG] --warmup-num-frames unset; capturing the " + "model default warmup frame count (%s). Requests with a " + "different frame count run eager. Pass --warmup-num-frames N " + "matching your served frame count.", + default_frames, + ) if self.bcg_text_buckets is not None and not any( int(b) > 0 for b in self.bcg_text_buckets ): @@ -687,6 +696,29 @@ class ServerArgs(DisaggServerArgsMixin): "--bcg-text-buckets must contain at least one positive integer." ) + def _is_video_gen_task(self) -> bool: + pipeline_config = getattr(self, "pipeline_config", None) + task_type = getattr(pipeline_config, "task_type", None) + is_video_gen = getattr(task_type, "is_video_gen", None) + return bool(is_video_gen()) if callable(is_video_gen) else False + + def _bcg_default_warmup_num_frames(self): + """Best-effort preview of the warmup frame count BCG will capture.""" + try: + from sglang.multimodal_gen.runtime.warmup_request_builder import ( + _resolve_warmup_num_frames, + get_model_sampling_defaults, + ) + + sampling_defaults = get_model_sampling_defaults(self) + return _resolve_warmup_num_frames( + self, + sampling_defaults, + server_based_warmup=True, + ) + except Exception: # pragma: no cover - defensive + return None + def _adjust_breakable_cuda_graph_support(self): if not self.enable_breakable_cuda_graph: return