Skip to content
87 changes: 87 additions & 0 deletions src/diffusers/models/attention_dispatch.py
Original file line number Diff line number Diff line change
Expand Up @@ -410,6 +410,10 @@ def dispatch_attention_fn(
*,
backend: AttentionBackendName | None = None,
parallel_config: "ParallelConfig" | None = None,
kv_cache: Any | None = None,
kv_cache_mode: str | None = None,
cache_write_slice: slice | None = None,
block_causal_segments: list[tuple[int, int, bool]] | None = None,
) -> torch.Tensor:
attention_kwargs = attention_kwargs or {}

Expand All @@ -421,6 +425,89 @@ def dispatch_attention_fn(
backend_name = AttentionBackendName(backend)
backend_fn = _AttentionBackendRegistry._backends.get(backend_name)

cp_config = None if parallel_config is None else parallel_config.context_parallel_config
if cp_config is not None and (kv_cache is not None or block_causal_segments is not None):
if cp_config.ring_degree > 1:
raise NotImplementedError("KV caching and block-causal attention currently support Ulysses only.")
if enable_gqa:
raise ValueError("GQA is not yet supported for context-parallel attention.")

def forward_op(_ctx, query, key, value, *args, **kwargs):
return dispatch_attention_fn(
query,
key,
value,
attn_mask=attn_mask,
dropout_p=dropout_p,
is_causal=is_causal,
scale=scale,
enable_gqa=enable_gqa,
attention_kwargs=attention_kwargs,
backend=backend_name,
kv_cache=kv_cache,
kv_cache_mode=kv_cache_mode,
cache_write_slice=cache_write_slice,
block_causal_segments=block_causal_segments,
)

if cp_config.ulysses_anything:
return TemplatedUlyssesAnythingAttention.apply(
query,
key,
value,
None,
dropout_p,
is_causal,
scale,
enable_gqa,
False,
forward_op,
None,
parallel_config,
)
group = cp_config._ulysses_mesh.get_group()
query, key, value = (SeqAllToAllDim.apply(group, tensor, 2, 1) for tensor in (query, key, value))
hidden_states = forward_op(None, query, key, value)
return SeqAllToAllDim.apply(group, hidden_states, 1, 2)

if kv_cache is not None:
if kv_cache_mode == "extract" and cache_write_slice is not None:
# `clone()`, not `contiguous()`: at batch size 1 the prefix slice already counts as contiguous
# (size-1 dims are ignored), so `contiguous()` returns the same view and the cache would pin the
# whole prefill K/V for every step of the denoising loop.
kv_cache.store(key[:, cache_write_slice].clone(), value[:, cache_write_slice].clone())
elif kv_cache_mode == "cached":
cached_key, cached_value = kv_cache.get()
key = torch.cat([cached_key, key], dim=1)
value = torch.cat([cached_value, value], dim=1)

if block_causal_segments is not None:
prefix_len = block_causal_segments[-1][1] if block_causal_segments else 0
outputs = []
for start, end, is_segment_causal in [*block_causal_segments, (prefix_len, query.shape[1], False)]:
key_end = end if end <= prefix_len else key.shape[1]
segment_mask = None if attn_mask is None else attn_mask[..., :key_end]
if is_segment_causal:
positions = torch.arange(start, end, device=query.device)
causal_mask = positions[:, None] >= torch.arange(key_end, device=query.device)[None, :]
segment_mask = causal_mask if segment_mask is None else segment_mask & causal_mask
outputs.append(
dispatch_attention_fn(
query[:, start:end],
key[:, :key_end],
value[:, :key_end],
attn_mask=segment_mask,
dropout_p=dropout_p,
is_causal=is_causal,
scale=scale,
enable_gqa=enable_gqa,
attention_kwargs=attention_kwargs,
backend=backend_name,
parallel_config=parallel_config,
)
)
return torch.cat(outputs, dim=1)

kwargs = {
"query": query,
"key": key,
Expand Down
122 changes: 42 additions & 80 deletions src/diffusers/models/transformers/transformer_qwenimage21.py
Original file line number Diff line number Diff line change
Expand Up @@ -24,6 +24,7 @@
from ...utils import logging
from ...utils.peft_utils import apply_lora_scale
from ...utils.torch_utils import maybe_allow_in_graph
from .._modeling_parallel import ContextParallelInput, ContextParallelOutput
from ..attention import AttentionMixin, AttentionModuleMixin
from ..attention_dispatch import dispatch_attention_fn
from ..cache_utils import CacheMixin
Expand Down Expand Up @@ -328,11 +329,7 @@ def _qwenimage21_prepare_qkv(
attn: "QwenImage21Attention",
hidden_states: torch.Tensor,
rotary_emb: torch.Tensor | None,
layer_cache: QwenImage21KVLayerCache | None,
kv_cache_mode: str | None,
cache_write_slice: slice | None,
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, int]:
"""Shared QKV projection, norm, RoPE and KV-cache bookkeeping for both processors."""
query = attn.to_q(hidden_states)
key = attn.to_k(hidden_states)
value = attn.to_v(hidden_states)
Expand All @@ -348,22 +345,7 @@ def _qwenimage21_prepare_qkv(
query = apply_rotary_emb_qwen(query, rotary_emb, use_real=False)
key = apply_rotary_emb_qwen(key, rotary_emb, use_real=False)

if layer_cache is not None:
if kv_cache_mode == "extract" and cache_write_slice is not None:
# `clone()`, not `contiguous()`: at batch size 1 the prefix slice already counts as contiguous
# (size-1 dims are ignored), so `contiguous()` returns the same view and the cache would pin the
# whole prefill K/V for every step of the denoising loop.
layer_cache.store(
key[:, cache_write_slice].clone(),
value[:, cache_write_slice].clone(),
)
elif kv_cache_mode == "cached":
cached_k, cached_v = layer_cache.get()
key = torch.cat([cached_k, key], dim=1)
value = torch.cat([cached_v, value], dim=1)

seq_len_q = query.shape[1]
return query, key, value, seq_len_q
return query, key, value, query.shape[1]


class QwenImage21FlexAttnProcessor:
Expand Down Expand Up @@ -401,9 +383,12 @@ def __call__(
segments: list[tuple[int, int, bool]] | None = None,
key_valid: torch.Tensor | None = None,
) -> torch.Tensor:
query, key, value, seq_len_q = _qwenimage21_prepare_qkv(
attn, hidden_states, rotary_emb, layer_cache, kv_cache_mode, cache_write_slice
)
if self._parallel_config is not None:
raise NotImplementedError(
"Context parallelism is not implemented for QwenImage21FlexAttnProcessor. "
"Use QwenImage21AttnProcessor instead."
)
query, key, value, seq_len_q = _qwenimage21_prepare_qkv(attn, hidden_states, rotary_emb)

seq_len_kv = key.shape[1]
if isinstance(attention_mask, BlockMask):
Expand Down Expand Up @@ -440,6 +425,9 @@ def __call__(
dropout_p=0.0,
backend="flex",
parallel_config=self._parallel_config,
kv_cache=layer_cache,
kv_cache_mode=kv_cache_mode,
cache_write_slice=cache_write_slice,
)
else:
# decode: full attention over [cached prefix, target]
Expand All @@ -451,6 +439,9 @@ def __call__(
dropout_p=0.0,
backend=self._attention_backend,
parallel_config=self._parallel_config,
kv_cache=layer_cache,
kv_cache_mode=kv_cache_mode,
cache_write_slice=cache_write_slice,
)
hidden_states = hidden_states[:, :seq_len_q]
hidden_states = hidden_states.flatten(2, 3).type_as(query)
Expand Down Expand Up @@ -483,64 +474,22 @@ def __call__(
segments: list[tuple[int, int, bool]] | None = None,
key_valid: torch.Tensor | None = None,
) -> torch.Tensor:
query, key, value, seq_len_q = _qwenimage21_prepare_qkv(
attn, hidden_states, rotary_emb, layer_cache, kv_cache_mode, cache_write_slice
query, key, value, seq_len_q = _qwenimage21_prepare_qkv(attn, hidden_states, rotary_emb)
if segments is not None:
attention_mask = None if key_valid is None else key_valid[:, None, None, :]
hidden_states = dispatch_attention_fn(
query,
key,
value,
attn_mask=attention_mask,
dropout_p=0.0,
backend=self._attention_backend if segments is None else None,
parallel_config=self._parallel_config,
kv_cache=layer_cache,
kv_cache_mode=kv_cache_mode,
cache_write_slice=cache_write_slice,
block_causal_segments=segments,
)

if segments is None:
# decode: full attention over [cached prefix, target]
hidden_states = dispatch_attention_fn(
query,
key,
value,
attn_mask=attention_mask,
dropout_p=0.0,
backend=self._attention_backend,
parallel_config=self._parallel_config,
)
else:
# prefill: every segment attends to the keys `[0, end)` (everything before it plus its own block); text
# segments additionally get a causal triangle over their own keys; padded text keys are dropped.
# `attention_mask` may hold the flex `BlockMask` of the same structure, which is not used here.
prefix_len = segments[-1][1] if segments else 0
outputs = []
for start, end, is_text in segments:
seg_mask = None
if is_text:
seg_len = end - start
seg_mask = torch.cat(
[
torch.ones(seg_len, start, dtype=torch.bool, device=query.device),
torch.tril(torch.ones(seg_len, seg_len, dtype=torch.bool, device=query.device)),
],
dim=1,
)[None, None]
if key_valid is not None:
seg_key_valid = key_valid[:, None, None, :end]
seg_mask = seg_key_valid if seg_mask is None else (seg_mask & seg_key_valid)
outputs.append(
dispatch_attention_fn(
query[:, start:end],
key[:, :end],
value[:, :end],
attn_mask=seg_mask,
dropout_p=0.0,
backend=None,
parallel_config=self._parallel_config,
)
)
outputs.append(
dispatch_attention_fn(
query[:, prefix_len:],
key,
value,
attn_mask=None if key_valid is None else key_valid[:, None, None, :],
dropout_p=0.0,
backend=None,
parallel_config=self._parallel_config,
)
)
hidden_states = torch.cat(outputs, dim=1)
hidden_states = hidden_states[:, :seq_len_q]
hidden_states = hidden_states.flatten(2, 3).type_as(query)

Expand Down Expand Up @@ -758,6 +707,19 @@ class QwenImage21Transformer2DModel(
_skip_layerwise_casting_patterns = ["pos_embed", "norm"]
_repeated_blocks = ["QwenImage21TransformerBlock"]
_skip_keys = ["kv_cache"]
_cp_plan = {
"transformer_blocks.0": {
"hidden_states": ContextParallelInput(split_dim=1, expected_dims=3, split_output=False),
},
"transformer_blocks.*": {
"rotary_emb": ContextParallelInput(split_dim=0, expected_dims=2, split_output=False),
"target_token_mask": ContextParallelInput(split_dim=0, expected_dims=1, split_output=False),
},
"norm_out": {
"target_token_mask": ContextParallelInput(split_dim=0, expected_dims=1, split_output=False),
},
"proj_out": ContextParallelOutput(gather_dim=1, expected_dims=3),
}

@register_to_config
def __init__(
Expand Down
Loading
Loading