[diffusion] Z-Image single-GPU BCG: fix the replay crash and make output bit-exact vs eager (#34183) (#34210)
Co-authored-by: Claude Fable 5 <noreply@anthropic.com>
This commit is contained in:
co-authored by
Claude Fable 5
parent
fd3036523a
commit
f5f0c3ee7a
@@ -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
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user