Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
1 change: 1 addition & 0 deletions invokeai/app/invocations/z_image_denoise.py
Original file line number Diff line number Diff line change
Expand Up @@ -558,6 +558,7 @@ def _run_diffusion(self, context: InvocationContext) -> torch.Tensor:
transformer=transformer,
regional_attn_mask=regional_extension.regional_attn_mask,
img_seq_len=img_seq_len,
positive_cap_feats=pos_prompt_embeds,
)
)

Expand Down
296 changes: 151 additions & 145 deletions invokeai/backend/z_image/z_image_transformer_patch.py
Original file line number Diff line number Diff line change
Expand Up @@ -4,25 +4,38 @@
from typing import Callable, List, Optional, Tuple

import torch
from torch.nn.utils.rnn import pad_sequence


def create_regional_forward(
original_forward: Callable,
regional_attn_mask: torch.Tensor,
img_seq_len: int,
positive_cap_feats: torch.Tensor,
) -> Callable:
"""Create a modified forward function that uses a regional attention mask.

The regional attention mask replaces the internally computed padding mask,
allowing for regional prompting where different image regions attend to
different text prompts.
The regional attention mask replaces the internally computed padding mask on the
main transformer layers (alternating with the plain padding mask), allowing for
regional prompting where different image regions attend to different text prompts.

This delegates to the model's own helper methods (``patchify_and_embed``,
``_prepare_sequence``, ``_build_unified_sequence``) so it stays in sync with the
upstream diffusers ``ZImageTransformer2DModel.forward`` implementation. Only the
main-layer attention mask is overridden.

Args:
original_forward: The original forward method of ZImageTransformer2DModel.
regional_attn_mask: Attention mask of shape (seq_len, seq_len) where
seq_len = img_seq_len + txt_seq_len.
img_seq_len: Number of image tokens in the sequence.
original_forward: The original forward method of ZImageTransformer2DModel
(kept for signature compatibility; not used directly).
regional_attn_mask: Boolean attention mask of shape (seq_len, seq_len) where
seq_len = img_seq_len + txt_seq_len, ordered [img, txt].
img_seq_len: Number of (unpadded) image tokens in the sequence.
positive_cap_feats: The exact caption-embedding tensor the regional mask was
built for (the conditioned/positive pass). The regional mask is applied only
to forward calls whose ``cap_feats`` is this same object; the negative/CFG
pass supplies a different tensor and is left to run with the plain padding
mask. Identity is used instead of a token-length heuristic so the positive
and negative passes can never be confused even when their padded lengths
coincide.

Returns:
A modified forward function with regional attention support.
Expand All @@ -38,160 +51,146 @@ def regional_forward(
) -> Tuple[List[torch.Tensor], dict]:
"""Modified forward with regional attention mask injection.

This is based on the original ZImageTransformer2DModel.forward but
replaces the padding-based attention mask with a regional attention mask.
Mirrors the basic (non-omni) path of ZImageTransformer2DModel.forward but
injects a regional attention mask into the main transformer layers.
"""
assert patch_size in self.all_patch_size
assert f_patch_size in self.all_f_patch_size

bsz = len(x)
device = x[0].device
t_scaled = t * self.t_scale
t_emb = self.t_embedder(t_scaled)

SEQ_MULTI_OF = 32 # From diffusers transformer_z_image.py
# Identify which caption inputs belong to the conditioned (positive) pass the regional
# mask was built for. Capture this before patchify_and_embed reassigns ``cap_feats``.
# The negative/CFG pass supplies a different tensor, so object identity distinguishes the
# passes regardless of token length (avoids the positive mask leaking into the uncond
# prediction when prompt lengths happen to pad to the same multiple).
is_positive_pass = [ci is positive_cap_feats for ci in cap_feats]

# Single adaLN embedding for all tokens (basic mode).
adaln_input = self.t_embedder(t * self.t_scale).type_as(x[0])

# Patchify and embed (reusing the original method)
# Patchify & embed (basic mode: single image per batch item).
(
x,
cap_feats,
x_size,
x_pos_ids,
cap_pos_ids,
x_inner_pad_mask,
cap_inner_pad_mask,
x_pad_mask,
cap_pad_mask,
) = self.patchify_and_embed(x, cap_feats, patch_size, f_patch_size)

# x embed & refine
x_item_seqlens = [len(_) for _ in x]
assert all(_ % SEQ_MULTI_OF == 0 for _ in x_item_seqlens)
x_max_item_seqlen = max(x_item_seqlens)

x_cat = torch.cat(x, dim=0)
x_cat = self.all_x_embedder[f"{patch_size}-{f_patch_size}"](x_cat)

adaln_input = t_emb.type_as(x_cat)
x_cat[torch.cat(x_inner_pad_mask)] = self.x_pad_token
x_list = list(x_cat.split(x_item_seqlens, dim=0))
x_freqs_cis = list(self.rope_embedder(torch.cat(x_pos_ids, dim=0)).split(x_item_seqlens, dim=0))

x_padded = pad_sequence(x_list, batch_first=True, padding_value=0.0)
x_freqs_cis_padded = pad_sequence(x_freqs_cis, batch_first=True, padding_value=0.0)
x_attn_mask = torch.zeros((bsz, x_max_item_seqlen), dtype=torch.bool, device=device)
for i, seq_len in enumerate(x_item_seqlens):
x_attn_mask[i, :seq_len] = 1

# Process through noise_refiner
if torch.is_grad_enabled() and self.gradient_checkpointing:
for layer in self.noise_refiner:
x_padded = self._gradient_checkpointing_func(
layer, x_padded, x_attn_mask, x_freqs_cis_padded, adaln_input
)
else:
for layer in self.noise_refiner:
x_padded = layer(x_padded, x_attn_mask, x_freqs_cis_padded, adaln_input)

# cap embed & refine
cap_item_seqlens = [len(_) for _ in cap_feats]
assert all(_ % SEQ_MULTI_OF == 0 for _ in cap_item_seqlens)
cap_max_item_seqlen = max(cap_item_seqlens)

cap_cat = torch.cat(cap_feats, dim=0)
cap_cat = self.cap_embedder(cap_cat)
cap_cat[torch.cat(cap_inner_pad_mask)] = self.cap_pad_token
cap_list = list(cap_cat.split(cap_item_seqlens, dim=0))
cap_freqs_cis = list(self.rope_embedder(torch.cat(cap_pos_ids, dim=0)).split(cap_item_seqlens, dim=0))

cap_padded = pad_sequence(cap_list, batch_first=True, padding_value=0.0)
cap_freqs_cis_padded = pad_sequence(cap_freqs_cis, batch_first=True, padding_value=0.0)
cap_attn_mask = torch.zeros((bsz, cap_max_item_seqlen), dtype=torch.bool, device=device)
for i, seq_len in enumerate(cap_item_seqlens):
cap_attn_mask[i, :seq_len] = 1

# Process through context_refiner
if torch.is_grad_enabled() and self.gradient_checkpointing:
for layer in self.context_refiner:
cap_padded = self._gradient_checkpointing_func(layer, cap_padded, cap_attn_mask, cap_freqs_cis_padded)
else:
for layer in self.context_refiner:
cap_padded = layer(cap_padded, cap_attn_mask, cap_freqs_cis_padded)

# Unified sequence: [img_tokens, txt_tokens]
unified = []
unified_freqs_cis = []
for i in range(bsz):
x_len = x_item_seqlens[i]
cap_len = cap_item_seqlens[i]
unified.append(torch.cat([x_padded[i][:x_len], cap_padded[i][:cap_len]]))
unified_freqs_cis.append(torch.cat([x_freqs_cis_padded[i][:x_len], cap_freqs_cis_padded[i][:cap_len]]))

unified_item_seqlens = [a + b for a, b in zip(cap_item_seqlens, x_item_seqlens, strict=False)]
assert unified_item_seqlens == [len(_) for _ in unified]
unified_max_item_seqlen = max(unified_item_seqlens)

unified_padded = pad_sequence(unified, batch_first=True, padding_value=0.0)
unified_freqs_cis_padded = pad_sequence(unified_freqs_cis, batch_first=True, padding_value=0.0)
# X embed & refine.
x_seqlens = [len(xi) for xi in x]
x = self.all_x_embedder[f"{patch_size}-{f_patch_size}"](torch.cat(x, dim=0))
x, x_freqs, x_mask, _, _ = self._prepare_sequence(
list(x.split(x_seqlens, dim=0)), x_pos_ids, x_pad_mask, self.x_pad_token, None, device
)
for layer in self.noise_refiner:
x = layer(x, x_mask, x_freqs, adaln_input, None, None, None)

# Cap embed & refine.
cap_seqlens = [len(ci) for ci in cap_feats]
cap_feats = self.cap_embedder(torch.cat(cap_feats, dim=0))
cap_feats, cap_freqs, cap_mask, _, _ = self._prepare_sequence(
list(cap_feats.split(cap_seqlens, dim=0)), cap_pos_ids, cap_pad_mask, self.cap_pad_token, None, device
)
for layer in self.context_refiner:
cap_feats = layer(cap_feats, cap_mask, cap_freqs)

# Unified sequence: basic mode order [x, cap].
unified, unified_freqs, unified_mask, _ = self._build_unified_sequence(
x,
x_freqs,
x_seqlens,
None,
cap_feats,
cap_freqs,
cap_seqlens,
None,
None,
None,
None,
None,
False, # omni_mode
device,
)

bsz = unified.shape[0]
unified_seqlen = unified.shape[1]

# --- REGIONAL ATTENTION MASK INJECTION ---
# Instead of using the padding mask, we use the regional attention mask
# The regional mask is (seq_len, seq_len), we need to expand it to (batch, seq_len, seq_len)
# and then add the batch dimension for broadcasting: (batch, 1, seq_len, seq_len)

# Expand regional mask to match the actual sequence length (may include padding)
if regional_attn_mask.shape[0] != unified_max_item_seqlen:
# Pad the regional mask to match unified sequence length
padded_regional_mask = torch.zeros(
(unified_max_item_seqlen, unified_max_item_seqlen),
dtype=regional_attn_mask.dtype,
device=device,
# The regional mask is (S, S) with S = img_seq_len + txt_seq_len, ordered [img, txt],
# using the *unpadded* image and text token counts. In the unified sequence, however,
# both the image block and the caption block are individually padded to a multiple of
# SEQ_MULTI_OF, so the real layout per item is:
# [ img_real | img_pad | txt_real | txt_pad ]
# We therefore scatter the four regional sub-blocks (img-img, img-txt, txt-img, txt-txt)
# into their padding-aware positions instead of assuming a contiguous top-left block.
#
# The patched forward also runs for the negative/CFG pass (a different prompt). The
# regional mask was built for the positive prompt only, so we apply it only to the
# conditioned items and fall back to the plain padding mask otherwise.
regional = regional_attn_mask.to(device=device, dtype=torch.bool)
txt_seq_len = regional.shape[0] - img_seq_len

# Decide per item whether the regional mask applies, using only cheap scalar checks, so
# that on passes that never match (e.g. every negative/CFG pass) we avoid materializing
# the (bsz, 1, S, S) float mask at all.
applied_regional = [
is_positive_pass[i]
and txt_seq_len > 0
and img_seq_len <= x_seqlens[i]
and x_seqlens[i] + cap_seqlens[i] <= unified_seqlen
for i in range(bsz)
]

# Main transformer layers: alternate regional mask (even) with plain padding mask (odd).
# If no item matched the positive pass, skip regional injection entirely.
use_regional = any(applied_regional)

float_mask = None
if use_regional:
# Build a per-item additive float mask. Start from the plain padding mask (0 where a
# 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)
float_mask = (
torch.where(
unified_mask.bool().unsqueeze(1).unsqueeze(1), # (bsz, 1, 1, S)
zero,
torch.full((), neg_inf, dtype=unified.dtype, device=device),
)
.expand(bsz, 1, unified_seqlen, unified_seqlen)
.clone()
)
mask_size = min(regional_attn_mask.shape[0], unified_max_item_seqlen)
padded_regional_mask[:mask_size, :mask_size] = regional_attn_mask[:mask_size, :mask_size]
else:
padded_regional_mask = regional_attn_mask.to(device)

# Convert boolean mask to additive float mask for attention
# True (attend) -> 0.0, False (block) -> -inf
# This is required because the attention backend expects additive masks for 4D inputs
# Use bfloat16 to match the transformer's query dtype
float_mask = torch.zeros_like(padded_regional_mask, dtype=torch.bfloat16)
float_mask[~padded_regional_mask] = float("-inf")

# Expand to (batch, 1, seq_len, seq_len) for attention
unified_attn_mask = float_mask.unsqueeze(0).unsqueeze(0).expand(bsz, 1, -1, -1)

# Process through main layers with regional attention mask
if torch.is_grad_enabled() and self.gradient_checkpointing:
for layer_idx, layer in enumerate(self.layers):
# Alternate between regional mask and full attention
if layer_idx % 2 == 0:
unified_padded = self._gradient_checkpointing_func(
layer, unified_padded, unified_attn_mask, unified_freqs_cis_padded, adaln_input
)
else:
# Use padding mask only for odd layers (allows global coherence)
padding_mask = torch.zeros((bsz, unified_max_item_seqlen), dtype=torch.bool, device=device)
for i, seq_len in enumerate(unified_item_seqlens):
padding_mask[i, :seq_len] = 1
unified_padded = self._gradient_checkpointing_func(
layer, unified_padded, padding_mask, unified_freqs_cis_padded, adaln_input
)
else:
for layer_idx, layer in enumerate(self.layers):
# Alternate between regional mask and full attention
if layer_idx % 2 == 0:
unified_padded = layer(unified_padded, unified_attn_mask, unified_freqs_cis_padded, adaln_input)
else:
# Use padding mask only for odd layers (allows global coherence)
padding_mask = torch.zeros((bsz, unified_max_item_seqlen), dtype=torch.bool, device=device)
for i, seq_len in enumerate(unified_item_seqlens):
padding_mask[i, :seq_len] = 1
unified_padded = layer(unified_padded, padding_mask, unified_freqs_cis_padded, adaln_input)

# Final layer
unified_out = self.all_final_layer[f"{patch_size}-{f_patch_size}"](unified_padded, adaln_input)
unified_list = list(unified_out.unbind(dim=0))
x_out = self.unpatchify(unified_list, x_size, patch_size, f_patch_size)

for i in range(bsz):
if not applied_regional[i]:
continue
x_len = x_seqlens[i]

ii, it = slice(0, img_seq_len), slice(img_seq_len, img_seq_len + txt_seq_len)
ui = slice(0, img_seq_len) # real image positions in unified item
ut = slice(x_len, x_len + txt_seq_len) # real text positions in unified item

# Reset the masked region so only regional rules apply to real img/txt tokens;
# their rows start fully blocked and we open the allowed sub-blocks below.
float_mask[i, 0, ui, :] = neg_inf
float_mask[i, 0, ut, :] = neg_inf

float_mask[i, 0, ui, ui] = torch.where(regional[ii, ii], zero, neg_inf) # img -> img
float_mask[i, 0, ui, ut] = torch.where(regional[ii, it], zero, neg_inf) # img -> txt
float_mask[i, 0, ut, ui] = torch.where(regional[it, ii], zero, neg_inf) # txt -> img
float_mask[i, 0, ut, ut] = torch.where(regional[it, it], zero, neg_inf) # txt -> txt

for layer_idx, layer in enumerate(self.layers):
attn_mask = float_mask if (use_regional and layer_idx % 2 == 0) else unified_mask
unified = layer(unified, attn_mask, unified_freqs, adaln_input, None, None, None)

# Final layer + unpatchify.
unified = self.all_final_layer[f"{patch_size}-{f_patch_size}"](unified, c=adaln_input)
x_out = self.unpatchify(list(unified.unbind(dim=0)), x_size, patch_size, f_patch_size)

return x_out, {}

Expand All @@ -203,6 +202,7 @@ def patch_transformer_for_regional_prompting(
transformer,
regional_attn_mask: Optional[torch.Tensor],
img_seq_len: int,
positive_cap_feats: Optional[torch.Tensor] = None,
):
"""Context manager to temporarily patch the transformer for regional prompting.

Expand All @@ -211,6 +211,9 @@ def patch_transformer_for_regional_prompting(
regional_attn_mask: Regional attention mask of shape (seq_len, seq_len).
If None, the transformer is not patched.
img_seq_len: Number of image tokens.
positive_cap_feats: The caption-embedding tensor the regional mask was built for.
Required when ``regional_attn_mask`` is provided; the mask is applied only to
forward calls whose ``cap_feats`` is this exact object (the conditioned pass).

Yields:
The (possibly patched) transformer.
Expand All @@ -220,11 +223,14 @@ def patch_transformer_for_regional_prompting(
yield transformer
return

if positive_cap_feats is None:
raise ValueError("positive_cap_feats is required when regional_attn_mask is provided")

# Store original forward
original_forward = transformer.forward

# Create and bind the regional forward
regional_fwd = create_regional_forward(original_forward, regional_attn_mask, img_seq_len)
regional_fwd = create_regional_forward(original_forward, regional_attn_mask, img_seq_len, positive_cap_feats)
transformer.forward = lambda *args, **kwargs: regional_fwd(transformer, *args, **kwargs)

try:
Expand Down
Loading