diff --git a/python/sglang/multimodal_gen/runtime/breakable_cuda_graph/model_padders/zimage.py b/python/sglang/multimodal_gen/runtime/breakable_cuda_graph/model_padders/zimage.py index 088983e05..7517ee327 100644 --- a/python/sglang/multimodal_gen/runtime/breakable_cuda_graph/model_padders/zimage.py +++ b/python/sglang/multimodal_gen/runtime/breakable_cuda_graph/model_padders/zimage.py @@ -127,9 +127,17 @@ def pad_zimage_prompt_kwargs( 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 + # Z-Image attends its caption slots UNMASKED: the pipeline pads captions + # to the native length (a multiple of 32) with learned pad-token + # embeddings that act as attended registers, and the DiT derives the + # attention length from the full padded tensor (`lens == target` -> + # mask=None). Padding further to a text bucket therefore changes how many + # registers every token attends -- a materially different (not bit-exact) + # forward, cascading over the few-step distilled sampler. Capture at the + # native length instead: signatures stay bounded because the pipeline + # already quantizes caption lengths, and unseen lengths fall back to + # eager at serving time. + bucket = max(seq, cap_freq_len) out = { key: value 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 32fca00ee..f65ca2f7f 100644 --- a/python/sglang/multimodal_gen/runtime/breakable_cuda_graph/runner.py +++ b/python/sglang/multimodal_gen/runtime/breakable_cuda_graph/runner.py @@ -344,6 +344,14 @@ 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." + ) def replay(self, entry: _CaptureEntry, kwargs: dict[str, Any]) -> Any: live_leaves = _flatten_kwargs(kwargs) diff --git a/python/sglang/multimodal_gen/runtime/models/dits/zimage.py b/python/sglang/multimodal_gen/runtime/models/dits/zimage.py index 16ca5c720..ebb2d4be9 100644 --- a/python/sglang/multimodal_gen/runtime/models/dits/zimage.py +++ b/python/sglang/multimodal_gen/runtime/models/dits/zimage.py @@ -41,6 +41,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.platforms import current_platform 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, +) try: from nunchaku.models.attention import NunchakuFeedForward # type: ignore[import] @@ -1198,7 +1201,7 @@ class ZImageTransformer2DModel(CachableDiT, LayerwiseOffloadableModuleMixin): ) cached = getattr(self, "_cached_batched_freqs_cis", None) if cached is not None and cached[0] == cache_key: - return cached[1] + return self._pin_for_active_capture(cached[1]) freqs_cis = self._build_batched_freqs_cis( images, @@ -1209,7 +1212,32 @@ class ZImageTransformer2DModel(CachableDiT, LayerwiseOffloadableModuleMixin): cap_target_len=cap_target_len, ) self._cached_batched_freqs_cis = (cache_key, freqs_cis) - return freqs_cis + return self._pin_for_active_capture(freqs_cis) + + def _pin_for_active_capture(self, value): + """Keep cache values consumed under CUDA graph capture alive forever. + + The single-slot shape-keyed caches below hold tensors that are pure + functions of their cache key. Capturing a second signature (e.g. the + next BCG caption bucket, whose static input buffers change every + ``data_ptr()``-keyed entry) replaces the slot and frees the old + tensors -- but a previously captured graph baked their device + addresses, so replaying it dereferences freed memory (observed as an + illegal memory access or a hang at the first replayed segment). + Pinning every value a capture consumes keeps those addresses alive; + contents stay correct because a value never changes for its key. + Growth is bounded by O(cache sites x captured signatures) small + tensors, and nothing is pinned outside graph capture. + """ + if is_in_breakable_cuda_graph() or ( + _is_cuda and torch.cuda.is_current_stream_capturing() + ): + pinned = getattr(self, "_bcg_pinned_cache_values", None) + if pinned is None: + pinned = [] + self._bcg_pinned_cache_values = pinned + pinned.append(value) + return value def _get_rope_cache( self, @@ -1234,7 +1262,7 @@ class ZImageTransformer2DModel(CachableDiT, LayerwiseOffloadableModuleMixin): ) cached = getattr(self, cache_attr, None) if cached is not None and cached[0] == cache_key: - return cached[1] + return self._pin_for_active_capture(cached[1]) if cos.dim() == 3: batch_size, seq_len = cos.shape[:2] @@ -1262,7 +1290,7 @@ class ZImageTransformer2DModel(CachableDiT, LayerwiseOffloadableModuleMixin): rope_cache = (cos_sin_cache, positions) setattr(self, cache_attr, (cache_key, rope_cache)) - return rope_cache + return self._pin_for_active_capture(rope_cache) def _get_attn_mask_and_meta( self, cache_attr: str, lengths: list[int], target_len: int, device: torch.device @@ -1278,7 +1306,7 @@ class ZImageTransformer2DModel(CachableDiT, LayerwiseOffloadableModuleMixin): ) cached = getattr(self, cache_attr, None) if cached is not None and cached[0] == cache_key: - return cached[1] + return self._pin_for_active_capture(cached[1]) positions = torch.arange(target_len, device=device).unsqueeze(0) length_tensor = torch.as_tensor( @@ -1288,7 +1316,7 @@ class ZImageTransformer2DModel(CachableDiT, LayerwiseOffloadableModuleMixin): meta = build_varlen_mask_meta_from_lengths(length_key, target_len, device) result = (mask, meta) setattr(self, cache_attr, (cache_key, result)) - return result + return self._pin_for_active_capture(result) def _get_joint_attn_mask_and_meta( self, @@ -1314,7 +1342,7 @@ class ZImageTransformer2DModel(CachableDiT, LayerwiseOffloadableModuleMixin): ) cached = getattr(self, "_cached_joint_attn_mask_meta", None) if cached is not None and cached[0] == cache_key: - return cached[1] + return self._pin_for_active_capture(cached[1]) image_pos = torch.arange(image_target_len, device=device).unsqueeze(0) cap_pos = torch.arange(cap_target_len, device=device).unsqueeze(0) @@ -1341,7 +1369,7 @@ class ZImageTransformer2DModel(CachableDiT, LayerwiseOffloadableModuleMixin): ) result = (mask, meta) self._cached_joint_attn_mask_meta = (cache_key, result) - return result + return self._pin_for_active_capture(result) @staticmethod def _has_padding(valid_lens: list[int], target_len: int) -> bool: diff --git a/python/sglang/multimodal_gen/test/unit/test_diffusion_bcg_padding.py b/python/sglang/multimodal_gen/test/unit/test_diffusion_bcg_padding.py index c0086492b..7d774c209 100644 --- a/python/sglang/multimodal_gen/test/unit/test_diffusion_bcg_padding.py +++ b/python/sglang/multimodal_gen/test/unit/test_diffusion_bcg_padding.py @@ -175,27 +175,38 @@ class TestDiffusionBCGPadding(unittest.TestCase): "image_seq_len_target": 256, } - def test_zimage_prompt_lengths_share_bucket_signature(self): + def test_zimage_captures_at_native_caption_length(self): + """Z-Image attends its caption slots unmasked (learned pad tokens act + as attended registers), so padding to a shared text bucket changes how + many registers every token attends and drifts the output + (sgl-project/sglang#34183). The padder therefore captures at the + incoming native length: lengths are never extended, and distinct + native lengths intentionally do NOT share a graph signature (unseen + lengths fall back to eager at serving time).""" with self._patch_buckets(64, 128): short = self.stage._bcg_pad_prompt_kwargs( self._zimage_kwargs(19), current_model=self.zimage_model ) + short_again = 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,)) + # Captions keep their native length -- no bucket extension. + self.assertEqual(short["encoder_hidden_states"][0].shape, (19, 16)) + self.assertEqual(longer["encoder_hidden_states"][0].shape, (47, 16)) + self.assertEqual(short["encoder_hidden_states_mask"].shape, (1, 19)) 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)) + self.assertTrue(short["encoder_hidden_states_mask"].all()) + self.assertEqual(short["freqs_cis"][0].shape, (19, 8)) + # Same native length -> same signature (graph reuse works); different + # native lengths -> different signatures, by design. + self.assertEqual(_signature_kwargs(short), _signature_kwargs(short_again)) + self.assertNotEqual(_signature_kwargs(short), _signature_kwargs(longer)) def _minimax_h3_kwargs(self, text_seq: int): image_seq = 4