From 3d1ae640419f62ceb5ae99a63625859375b97035 Mon Sep 17 00:00:00 2001 From: Edson Date: Thu, 1 Oct 2026 01:42:15 -0400 Subject: [PATCH] fix(z-image): handle missing unified mask in regional guidance forward diffusers 0.40 returns attn_mask=None from _build_unified_sequence when all batch items have the same length, which is always the case for the single-item call in z_image_denoise. The regional forward called unified_mask.bool() and crashed on every Z-Image regional guidance generation. Treat a None mask as all tokens valid. Closes #9612 --- .../z_image/z_image_transformer_patch.py | 9 ++- .../z_image/test_z_image_transformer_patch.py | 75 +++++++++++++++++++ 2 files changed, 83 insertions(+), 1 deletion(-) create mode 100644 tests/backend/z_image/test_z_image_transformer_patch.py diff --git a/invokeai/backend/z_image/z_image_transformer_patch.py b/invokeai/backend/z_image/z_image_transformer_patch.py index a6707fb7f50..49d7a894ba7 100644 --- a/invokeai/backend/z_image/z_image_transformer_patch.py +++ b/invokeai/backend/z_image/z_image_transformer_patch.py @@ -155,9 +155,16 @@ def regional_forward( # token is valid, -inf where it is padding) so non-matching items behave normally. neg_inf = torch.finfo(unified.dtype).min zero = torch.zeros((), dtype=unified.dtype, device=device) + # diffusers returns ``None`` when every item has the same length (always the case for + # bsz == 1), meaning no token is padding. + valid_mask = ( + unified_mask.bool() + if unified_mask is not None + else torch.ones((bsz, unified_seqlen), dtype=torch.bool, device=device) + ) float_mask = ( torch.where( - unified_mask.bool().unsqueeze(1).unsqueeze(1), # (bsz, 1, 1, S) + valid_mask.unsqueeze(1).unsqueeze(1), # (bsz, 1, 1, S) zero, torch.full((), neg_inf, dtype=unified.dtype, device=device), ) diff --git a/tests/backend/z_image/test_z_image_transformer_patch.py b/tests/backend/z_image/test_z_image_transformer_patch.py new file mode 100644 index 00000000000..ab089a0c69c --- /dev/null +++ b/tests/backend/z_image/test_z_image_transformer_patch.py @@ -0,0 +1,75 @@ +import diffusers +import torch + +from invokeai.backend.z_image.z_image_transformer_patch import patch_transformer_for_regional_prompting + + +def _tiny_z_image_transformer() -> diffusers.ZImageTransformer2DModel: + torch.manual_seed(0) + return diffusers.ZImageTransformer2DModel( + all_patch_size=(1,), + all_f_patch_size=(1,), + in_channels=4, + dim=8, + n_layers=2, + n_refiner_layers=1, + n_heads=1, + n_kv_heads=1, + cap_feat_dim=8, + axes_dims=[2, 2, 4], + axes_lens=[64, 64, 64], + ).eval() + + +def _run( + model: diffusers.ZImageTransformer2DModel, + x: list[torch.Tensor], + cap_feats: torch.Tensor, + regional_mask: torch.Tensor, + img_len: int, +) -> tuple[torch.Tensor, torch.Tensor]: + t = torch.tensor([0.5]) + with torch.no_grad(): + baseline = model(x, t, cap_feats=[cap_feats], patch_size=1)[0][0] + with patch_transformer_for_regional_prompting( + model, regional_mask, img_len, positive_cap_feats=cap_feats + ) as patched: + out = patched(x, t, cap_feats=[cap_feats], patch_size=1)[0][0] + return baseline, out + + +def test_regional_forward_single_item_batch() -> None: + """Regression test for #9612. + + z_image_denoise always calls the transformer with a single item. With diffusers 0.40 the unified attention mask + is ``None`` when all items have the same length, which must not crash the regional forward. + """ + model = _tiny_z_image_transformer() + x = [torch.randn(4, 1, 4, 4)] # 16 image tokens + cap_feats = torch.randn(5, 8) + img_len, txt_len = 16, 5 + regional_mask = torch.ones(img_len + txt_len, img_len + txt_len, dtype=torch.bool) + regional_mask[:8, img_len:] = False # top half of the image does not attend to the prompt + + baseline, out = _run(model, x, cap_feats, regional_mask, img_len) + + assert out.shape == baseline.shape + assert torch.isfinite(out).all() + + +def test_regional_forward_single_item_batch_matches_unpatched_without_padding() -> None: + """With no SEQ_MULTI_OF padding, an all-True regional mask must reproduce the unpatched forward, and a + restrictive one must still change the output.""" + model = _tiny_z_image_transformer() + x = [torch.randn(4, 1, 8, 8)] # 64 image tokens, no padding + cap_feats = torch.randn(32, 8) # 32 text tokens, no padding + img_len, txt_len = 64, 32 + seq_len = img_len + txt_len + + baseline, out = _run(model, x, cap_feats, torch.ones(seq_len, seq_len, dtype=torch.bool), img_len) + torch.testing.assert_close(out, baseline) + + restrictive = torch.ones(seq_len, seq_len, dtype=torch.bool) + restrictive[: img_len // 2, img_len:] = False + _, out_restricted = _run(model, x, cap_feats, restrictive, img_len) + assert not torch.allclose(out_restricted, baseline)