From 875e4876437845b013fe1b36058bde2058731451 Mon Sep 17 00:00:00 2001 From: Sayak Paul Date: Sat, 19 Sep 2026 04:00:30 +0000 Subject: [PATCH 1/3] feat: cp support in qwenimage 2.1 --- .../transformers/transformer_qwenimage21.py | 144 ++++++++++++++++-- .../test_models_transformer_qwenimage21.py | 48 +++++- 2 files changed, 182 insertions(+), 10 deletions(-) diff --git a/src/diffusers/models/transformers/transformer_qwenimage21.py b/src/diffusers/models/transformers/transformer_qwenimage21.py index b6eabfd584bb..f5887bb0c6b6 100644 --- a/src/diffusers/models/transformers/transformer_qwenimage21.py +++ b/src/diffusers/models/transformers/transformer_qwenimage21.py @@ -16,6 +16,7 @@ from typing import Any import torch +import torch.distributed as dist import torch.nn as nn import torch.nn.functional as F @@ -24,6 +25,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, gather_size_by_comm from ..attention import AttentionMixin, AttentionModuleMixin from ..attention_dispatch import dispatch_attention_fn from ..cache_utils import CacheMixin @@ -324,6 +326,37 @@ def _qwenimage21_prefix_segments(image_ids: torch.Tensor, prefix_len: int) -> li return segments +def _qwenimage21_dense_block_causal_mask( + segments: list[tuple[int, int, bool]], + seq_len: int, + key_valid: torch.Tensor | None, + batch_size: int, + device: torch.device, +) -> torch.Tensor: + attention_mask = torch.ones(seq_len, seq_len, dtype=torch.bool, device=device).tril() + for start, end, is_text in segments: + if not is_text: + attention_mask[start:end, start:end] = True + prefix_len = segments[-1][1] if segments else 0 + attention_mask[prefix_len:] = True + attention_mask = attention_mask.view(1, 1, seq_len, seq_len) + attention_mask = attention_mask.expand(batch_size, -1, -1, -1) + if key_valid is not None: + attention_mask = attention_mask & key_valid[:, None, None, :] + return attention_mask + + +def _qwenimage21_all_gather_sequence(tensor: torch.Tensor, group: dist.ProcessGroup) -> torch.Tensor: + local_sizes = gather_size_by_comm(tensor.shape[1], group) + max_local_size = max(local_sizes) + if tensor.shape[1] < max_local_size: + padding = tensor.new_zeros(tensor.shape[0], max_local_size - tensor.shape[1], *tensor.shape[2:]) + tensor = torch.cat([tensor, padding], dim=1) + gathered = [torch.empty_like(tensor) for _ in local_sizes] + dist.all_gather(gathered, tensor, group=group) + return torch.cat([value[:, :size] for value, size in zip(gathered, local_sizes)], dim=1) + + def _qwenimage21_prepare_qkv( attn: "QwenImage21Attention", hidden_states: torch.Tensor, @@ -331,6 +364,7 @@ def _qwenimage21_prepare_qkv( layer_cache: QwenImage21KVLayerCache | None, kv_cache_mode: str | None, cache_write_slice: slice | None, + parallel_config: Any | None = 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) @@ -350,13 +384,40 @@ def _qwenimage21_prepare_qkv( 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(), - ) + context_parallel_config = None if parallel_config is None else parallel_config.context_parallel_config + if context_parallel_config is 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(), + ) + else: + group = context_parallel_config._ulysses_mesh.get_group() + rank = dist.get_rank(group) + world_size = dist.get_world_size(group) + local_seq_lens = gather_size_by_comm(key.shape[1], group) + local_offset = sum(local_seq_lens[:rank]) + cache_start = 0 if cache_write_slice.start is None else cache_write_slice.start + cache_stop = sum(local_seq_lens) if cache_write_slice.stop is None else cache_write_slice.stop + local_cache_start = max(cache_start - local_offset, 0) + local_cache_stop = min(cache_stop - local_offset, key.shape[1]) + local_cache_start = min(local_cache_start, key.shape[1]) + local_cache_stop = max(local_cache_stop, local_cache_start) + + cached_key = _qwenimage21_all_gather_sequence(key[:, local_cache_start:local_cache_stop], group) + cached_value = _qwenimage21_all_gather_sequence(value[:, local_cache_start:local_cache_stop], group) + if not context_parallel_config.ulysses_anything and cached_key.shape[1] % world_size != 0: + raise ValueError( + "The cached prefix length must be divisible by the Ulysses degree. Enable " + "`ulysses_anything=True` to cache an uneven prefix." + ) + split_fn = torch.tensor_split if context_parallel_config.ulysses_anything else torch.chunk + layer_cache.store( + split_fn(cached_key, world_size, dim=1)[rank].clone(), + split_fn(cached_value, world_size, dim=1)[rank].clone(), + ) elif kv_cache_mode == "cached": cached_k, cached_v = layer_cache.get() key = torch.cat([cached_k, key], dim=1) @@ -401,8 +462,19 @@ def __call__( segments: list[tuple[int, int, bool]] | None = None, key_valid: torch.Tensor | None = None, ) -> torch.Tensor: + 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, layer_cache, kv_cache_mode, cache_write_slice + attn, + hidden_states, + rotary_emb, + layer_cache, + kv_cache_mode, + cache_write_slice, + self._parallel_config, ) seq_len_kv = key.shape[1] @@ -483,12 +555,34 @@ def __call__( segments: list[tuple[int, int, bool]] | None = None, key_valid: torch.Tensor | None = None, ) -> torch.Tensor: + context_parallel_config = ( + None if self._parallel_config is None else self._parallel_config.context_parallel_config + ) + if context_parallel_config is not None and context_parallel_config.ring_degree > 1: + raise NotImplementedError("QwenImage21AttnProcessor currently supports Ulysses context parallelism only.") + query, key, value, seq_len_q = _qwenimage21_prepare_qkv( - attn, hidden_states, rotary_emb, layer_cache, kv_cache_mode, cache_write_slice + attn, + hidden_states, + rotary_emb, + layer_cache, + kv_cache_mode, + cache_write_slice, + self._parallel_config, ) if segments is None: # decode: full attention over [cached prefix, target] + if context_parallel_config is not None and attention_mask is not None: + group = context_parallel_config._ulysses_mesh.get_group() + rank = dist.get_rank(group) + world_size = dist.get_world_size(group) + target_len = sum(gather_size_by_comm(seq_len_q, group)) + prefix_len = attention_mask.shape[-1] - target_len + split_fn = torch.tensor_split if context_parallel_config.ulysses_anything else torch.chunk + prefix_masks = split_fn(attention_mask[..., :prefix_len], world_size, dim=-1) + target_masks = split_fn(attention_mask[..., prefix_len:], world_size, dim=-1) + attention_mask = torch.cat([prefix_masks[rank], target_masks[rank]], dim=-1) hidden_states = dispatch_attention_fn( query, key, @@ -498,6 +592,25 @@ def __call__( backend=self._attention_backend, parallel_config=self._parallel_config, ) + elif context_parallel_config is not None: + group = context_parallel_config._ulysses_mesh.get_group() + global_seq_len = sum(gather_size_by_comm(seq_len_q, group)) + attention_mask = _qwenimage21_dense_block_causal_mask( + segments, + global_seq_len, + key_valid, + query.shape[0], + query.device, + ) + hidden_states = dispatch_attention_fn( + query, + key, + value, + attn_mask=attention_mask, + dropout_p=0.0, + backend=None, + 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. @@ -758,6 +871,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__( diff --git a/tests/models/transformers/test_models_transformer_qwenimage21.py b/tests/models/transformers/test_models_transformer_qwenimage21.py index cb23fc4ad1d2..b3de7b19bf31 100644 --- a/tests/models/transformers/test_models_transformer_qwenimage21.py +++ b/tests/models/transformers/test_models_transformer_qwenimage21.py @@ -21,10 +21,17 @@ from diffusers.models.transformers.transformer_qwenimage21 import build_qwenimage21_block_causal_mask from diffusers.utils.torch_utils import randn_tensor -from ...testing_utils import enable_full_determinism, torch_device +from ...testing_utils import ( + enable_full_determinism, + is_context_parallel, + require_torch_multi_accelerator, + torch_device, +) from ..testing_utils import ( AttentionTesterMixin, BaseModelTesterConfig, + ContextParallelAttentionBackendsTesterMixin, + ContextParallelTesterMixin, MemoryTesterMixin, ModelTesterMixin, TrainingTesterMixin, @@ -294,3 +301,42 @@ def test_gradient_checkpointing_is_applied(self): class TestQwenImage21TransformerAttention(QwenImage21TransformerTesterConfig, AttentionTesterMixin): pass + + +@is_context_parallel +@require_torch_multi_accelerator +class TestQwenImage21TransformerContextParallel(QwenImage21TransformerTesterConfig, ContextParallelTesterMixin): + @pytest.mark.parametrize("cp_type", ["ulysses_degree"], ids=["ulysses"]) + def test_context_parallel_inference(self, cp_type, batch_size: int = 1): + super().test_context_parallel_inference(cp_type, batch_size=batch_size) + + @pytest.mark.parametrize("cp_type", ["ulysses_degree"], ids=["ulysses"]) + def test_context_parallel_batch_inputs(self, cp_type): + super().test_context_parallel_inference(cp_type, batch_size=2) + + @pytest.mark.parametrize("cp_type", ["ulysses_degree"], ids=["ulysses"]) + def test_context_parallel_backward(self, cp_type, batch_size: int = 1): + super().test_context_parallel_backward(cp_type, batch_size=batch_size) + + @pytest.mark.parametrize("cp_type", ["ulysses_degree"], ids=["ulysses"]) + def test_context_parallel_backward_batch_inputs(self, cp_type): + super().test_context_parallel_backward(cp_type, batch_size=2) + + @pytest.mark.parametrize( + "cp_type,mesh_shape,mesh_dim_names", + [("ulysses_degree", (1, 2, 1), ("ring", "ulysses", "fsdp"))], + ids=["ulysses-3d-fsdp"], + ) + def test_context_parallel_custom_mesh(self, cp_type, mesh_shape, mesh_dim_names): + super().test_context_parallel_custom_mesh(cp_type, mesh_shape, mesh_dim_names) + + +class TestQwenImage21TransformerContextParallelAttnBackends( + QwenImage21TransformerTesterConfig, ContextParallelAttentionBackendsTesterMixin +): + unsupported_attn_backends = ["flash_hub", "flash_varlen_hub", "_flash_3_hub", "_flash_3_varlen_hub"] + + def get_dummy_inputs(self, batch_size: int = 1) -> dict[str, torch.Tensor]: + inputs = super().get_dummy_inputs(batch_size=batch_size) + inputs["encoder_hidden_states_mask"][:, 1] = 0 + return inputs From db37d305f0f8ff2f1280e28241fae1f5b469f513 Mon Sep 17 00:00:00 2001 From: Sayak Paul Date: Wed, 7 Oct 2026 10:43:38 +0000 Subject: [PATCH 2/3] first try at cp general. --- src/diffusers/models/attention_dispatch.py | 65 ++++++- src/diffusers/models/cache_utils.py | 76 ++++++++ .../transformers/transformer_qwenimage21.py | 177 +++--------------- tests/models/test_attention_dispatch.py | 88 ++++++++- .../test_models_transformer_qwenimage21.py | 84 ++++++++- 5 files changed, 332 insertions(+), 158 deletions(-) diff --git a/src/diffusers/models/attention_dispatch.py b/src/diffusers/models/attention_dispatch.py index f45c9e0d1107..d99ba523f5dd 100644 --- a/src/diffusers/models/attention_dispatch.py +++ b/src/diffusers/models/attention_dispatch.py @@ -449,6 +449,68 @@ def dispatch_attention_fn( return backend_fn(**kwargs) +def dispatch_block_causal_attention_fn( + query: torch.Tensor, + key: torch.Tensor, + value: torch.Tensor, + segments: list[tuple[int, int, bool]], + *, + key_valid: torch.Tensor | None = None, + parallel_config: ParallelConfig | None = None, +) -> torch.Tensor: + cp_config = None if parallel_config is None else parallel_config.context_parallel_config + prefix_len = segments[-1][1] if segments else 0 + if cp_config is not None: + if cp_config.ring_degree > 1: + raise NotImplementedError("Block-causal attention currently supports Ulysses context parallelism only.") + group = cp_config._ulysses_mesh.get_group() + seq_len = sum(gather_size_by_comm(query.shape[1], group)) + attention_mask = torch.ones(seq_len, seq_len, dtype=torch.bool, device=query.device).tril() + for start, end, is_causal in segments: + if not is_causal: + attention_mask[start:end, start:end] = True + attention_mask[prefix_len:] = True + attention_mask = attention_mask[None, None] + if key_valid is not None: + attention_mask = attention_mask & key_valid[:, None, None, :] + return dispatch_attention_fn(query, key, value, attn_mask=attention_mask, parallel_config=parallel_config) + + outputs = [] + for start, end, is_causal in segments: + seg_mask = None + if is_causal: + 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, + parallel_config=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, :], + parallel_config=parallel_config, + ) + ) + return torch.cat(outputs, dim=1) + + # ===== Checks ===== # A list of very simple functions to catch common errors quickly when debugging. @@ -2808,8 +2870,7 @@ def forward( attn_mask = F.pad(attn_mask, (0, max_local - attn_mask.shape[-1])) mask_list = [torch.empty_like(attn_mask) for _ in range(dist.get_world_size(group=group))] dist.all_gather(mask_list, attn_mask, group=group) - attn_mask = torch.cat(mask_list, dim=-1) - attn_mask = attn_mask[..., : sum(mask_local_sizes)] + attn_mask = torch.cat([mask[..., :size] for mask, size in zip(mask_list, mask_local_sizes)], dim=-1) out = forward_op( ctx, diff --git a/src/diffusers/models/cache_utils.py b/src/diffusers/models/cache_utils.py index 886ab6032bd4..5ef61b2b25a7 100644 --- a/src/diffusers/models/cache_utils.py +++ b/src/diffusers/models/cache_utils.py @@ -13,13 +13,89 @@ # limitations under the License. from contextlib import contextmanager +from typing import Any + +import torch +import torch.distributed as dist from ..utils.logging import get_logger +from ._modeling_parallel import ParallelConfig, gather_size_by_comm logger = get_logger(__name__) # pylint: disable=invalid-name +def apply_kv_cache( + key: torch.Tensor, + value: torch.Tensor, + layer_cache, + cache_mode: str | None, + cache_write_slice: slice | None, + *, + attention_mask: Any | None = None, + parallel_config: ParallelConfig | None = None, +) -> tuple[torch.Tensor, torch.Tensor, Any]: + if layer_cache is None: + return key, value, attention_mask + + cp_config = None if parallel_config is None else parallel_config.context_parallel_config + if cp_config is not None and cp_config.ring_degree > 1: + raise NotImplementedError("KV caching currently supports Ulysses context parallelism only.") + + if cache_mode == "extract" and cache_write_slice is not None: + if cp_config is 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()) + else: + group = cp_config._ulysses_mesh.get_group() + rank = dist.get_rank(group) + world_size = dist.get_world_size(group) + local_sizes = gather_size_by_comm(key.shape[1], group) + cache_start, cache_stop, cache_step = cache_write_slice.indices(sum(local_sizes)) + if cache_step != 1: + raise ValueError("Context-parallel KV caching requires a slice with step 1.") + cache_stop = max(cache_start, cache_stop) + cache_size = cache_stop - cache_start + if not cp_config.ulysses_anything and cache_size % world_size != 0: + raise ValueError( + "The cached sequence length must be divisible by the Ulysses degree. Enable " + "`ulysses_anything=True` to cache an uneven sequence." + ) + offsets = [sum(local_sizes[:index]) for index in range(world_size)] + starts = [min(max(cache_start - offset, 0), size) for offset, size in zip(offsets, local_sizes)] + stops = [min(max(cache_stop - offset, 0), size) for offset, size in zip(offsets, local_sizes)] + cache_sizes = [stop - start for start, stop in zip(starts, stops)] + max_size = max(cache_sizes) + cached = [] + for tensor in (key, value): + local = tensor[:, starts[rank] : stops[rank]].contiguous() + if local.shape[1] < max_size: + padding = tensor.new_zeros(tensor.shape[0], max_size - local.shape[1], *tensor.shape[2:]) + local = torch.cat([local, padding], dim=1) + gathered = [torch.empty_like(local) for _ in range(world_size)] + dist.all_gather(gathered, local, group=group) + full_cache = torch.cat([part[:, :size] for part, size in zip(gathered, cache_sizes)], dim=1) + cached.append(torch.tensor_split(full_cache, world_size, dim=1)[rank].clone()) + layer_cache.store(*cached) + elif cache_mode == "cached": + cached_key, cached_value = layer_cache.get() + if cp_config is not None and attention_mask is not None: + group = cp_config._ulysses_mesh.get_group() + rank = dist.get_rank(group) + world_size = dist.get_world_size(group) + target_len = sum(gather_size_by_comm(key.shape[1], group)) + prefix_len = attention_mask.shape[-1] - target_len + prefix_masks = torch.tensor_split(attention_mask[..., :prefix_len], world_size, dim=-1) + target_masks = torch.tensor_split(attention_mask[..., prefix_len:], world_size, dim=-1) + attention_mask = torch.cat([prefix_masks[rank], target_masks[rank]], dim=-1) + key = torch.cat([cached_key, key], dim=1) + value = torch.cat([cached_value, value], dim=1) + + return key, value, attention_mask + + class CacheMixin: r""" A class for enable/disabling caching techniques on diffusion models. diff --git a/src/diffusers/models/transformers/transformer_qwenimage21.py b/src/diffusers/models/transformers/transformer_qwenimage21.py index f5887bb0c6b6..7fefa2f3f599 100644 --- a/src/diffusers/models/transformers/transformer_qwenimage21.py +++ b/src/diffusers/models/transformers/transformer_qwenimage21.py @@ -16,7 +16,6 @@ from typing import Any import torch -import torch.distributed as dist import torch.nn as nn import torch.nn.functional as F @@ -25,10 +24,10 @@ 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, gather_size_by_comm +from .._modeling_parallel import ContextParallelInput, ContextParallelOutput from ..attention import AttentionMixin, AttentionModuleMixin -from ..attention_dispatch import dispatch_attention_fn -from ..cache_utils import CacheMixin +from ..attention_dispatch import dispatch_attention_fn, dispatch_block_causal_attention_fn +from ..cache_utils import CacheMixin, apply_kv_cache from ..embeddings import TimestepEmbedding from ..modeling_outputs import Transformer2DModelOutput from ..modeling_utils import ModelMixin @@ -326,37 +325,6 @@ def _qwenimage21_prefix_segments(image_ids: torch.Tensor, prefix_len: int) -> li return segments -def _qwenimage21_dense_block_causal_mask( - segments: list[tuple[int, int, bool]], - seq_len: int, - key_valid: torch.Tensor | None, - batch_size: int, - device: torch.device, -) -> torch.Tensor: - attention_mask = torch.ones(seq_len, seq_len, dtype=torch.bool, device=device).tril() - for start, end, is_text in segments: - if not is_text: - attention_mask[start:end, start:end] = True - prefix_len = segments[-1][1] if segments else 0 - attention_mask[prefix_len:] = True - attention_mask = attention_mask.view(1, 1, seq_len, seq_len) - attention_mask = attention_mask.expand(batch_size, -1, -1, -1) - if key_valid is not None: - attention_mask = attention_mask & key_valid[:, None, None, :] - return attention_mask - - -def _qwenimage21_all_gather_sequence(tensor: torch.Tensor, group: dist.ProcessGroup) -> torch.Tensor: - local_sizes = gather_size_by_comm(tensor.shape[1], group) - max_local_size = max(local_sizes) - if tensor.shape[1] < max_local_size: - padding = tensor.new_zeros(tensor.shape[0], max_local_size - tensor.shape[1], *tensor.shape[2:]) - tensor = torch.cat([tensor, padding], dim=1) - gathered = [torch.empty_like(tensor) for _ in local_sizes] - dist.all_gather(gathered, tensor, group=group) - return torch.cat([value[:, :size] for value, size in zip(gathered, local_sizes)], dim=1) - - def _qwenimage21_prepare_qkv( attn: "QwenImage21Attention", hidden_states: torch.Tensor, @@ -364,8 +332,9 @@ def _qwenimage21_prepare_qkv( layer_cache: QwenImage21KVLayerCache | None, kv_cache_mode: str | None, cache_write_slice: slice | None, + attention_mask: Any | None = None, parallel_config: Any | None = None, -) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, int]: +) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, int, Any]: """Shared QKV projection, norm, RoPE and KV-cache bookkeeping for both processors.""" query = attn.to_q(hidden_states) key = attn.to_k(hidden_states) @@ -382,49 +351,16 @@ 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: - context_parallel_config = None if parallel_config is None else parallel_config.context_parallel_config - if context_parallel_config is 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(), - ) - else: - group = context_parallel_config._ulysses_mesh.get_group() - rank = dist.get_rank(group) - world_size = dist.get_world_size(group) - local_seq_lens = gather_size_by_comm(key.shape[1], group) - local_offset = sum(local_seq_lens[:rank]) - cache_start = 0 if cache_write_slice.start is None else cache_write_slice.start - cache_stop = sum(local_seq_lens) if cache_write_slice.stop is None else cache_write_slice.stop - local_cache_start = max(cache_start - local_offset, 0) - local_cache_stop = min(cache_stop - local_offset, key.shape[1]) - local_cache_start = min(local_cache_start, key.shape[1]) - local_cache_stop = max(local_cache_stop, local_cache_start) - - cached_key = _qwenimage21_all_gather_sequence(key[:, local_cache_start:local_cache_stop], group) - cached_value = _qwenimage21_all_gather_sequence(value[:, local_cache_start:local_cache_stop], group) - if not context_parallel_config.ulysses_anything and cached_key.shape[1] % world_size != 0: - raise ValueError( - "The cached prefix length must be divisible by the Ulysses degree. Enable " - "`ulysses_anything=True` to cache an uneven prefix." - ) - split_fn = torch.tensor_split if context_parallel_config.ulysses_anything else torch.chunk - layer_cache.store( - split_fn(cached_key, world_size, dim=1)[rank].clone(), - split_fn(cached_value, world_size, dim=1)[rank].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 + key, value, attention_mask = apply_kv_cache( + key, + value, + layer_cache, + kv_cache_mode, + cache_write_slice, + attention_mask=attention_mask, + parallel_config=parallel_config, + ) + return query, key, value, query.shape[1], attention_mask class QwenImage21FlexAttnProcessor: @@ -467,13 +403,14 @@ def __call__( "Context parallelism is not implemented for QwenImage21FlexAttnProcessor. " "Use QwenImage21AttnProcessor instead." ) - query, key, value, seq_len_q = _qwenimage21_prepare_qkv( + query, key, value, seq_len_q, attention_mask = _qwenimage21_prepare_qkv( attn, hidden_states, rotary_emb, layer_cache, kv_cache_mode, cache_write_slice, + attention_mask, self._parallel_config, ) @@ -555,34 +492,19 @@ def __call__( segments: list[tuple[int, int, bool]] | None = None, key_valid: torch.Tensor | None = None, ) -> torch.Tensor: - context_parallel_config = ( - None if self._parallel_config is None else self._parallel_config.context_parallel_config - ) - if context_parallel_config is not None and context_parallel_config.ring_degree > 1: - raise NotImplementedError("QwenImage21AttnProcessor currently supports Ulysses context parallelism only.") - - query, key, value, seq_len_q = _qwenimage21_prepare_qkv( + query, key, value, seq_len_q, attention_mask = _qwenimage21_prepare_qkv( attn, hidden_states, rotary_emb, layer_cache, kv_cache_mode, cache_write_slice, + attention_mask, self._parallel_config, ) if segments is None: # decode: full attention over [cached prefix, target] - if context_parallel_config is not None and attention_mask is not None: - group = context_parallel_config._ulysses_mesh.get_group() - rank = dist.get_rank(group) - world_size = dist.get_world_size(group) - target_len = sum(gather_size_by_comm(seq_len_q, group)) - prefix_len = attention_mask.shape[-1] - target_len - split_fn = torch.tensor_split if context_parallel_config.ulysses_anything else torch.chunk - prefix_masks = split_fn(attention_mask[..., :prefix_len], world_size, dim=-1) - target_masks = split_fn(attention_mask[..., prefix_len:], world_size, dim=-1) - attention_mask = torch.cat([prefix_masks[rank], target_masks[rank]], dim=-1) hidden_states = dispatch_attention_fn( query, key, @@ -592,68 +514,15 @@ def __call__( backend=self._attention_backend, parallel_config=self._parallel_config, ) - elif context_parallel_config is not None: - group = context_parallel_config._ulysses_mesh.get_group() - global_seq_len = sum(gather_size_by_comm(seq_len_q, group)) - attention_mask = _qwenimage21_dense_block_causal_mask( - segments, - global_seq_len, - key_valid, - query.shape[0], - query.device, - ) - hidden_states = dispatch_attention_fn( + else: + hidden_states = dispatch_block_causal_attention_fn( query, key, value, - attn_mask=attention_mask, - dropout_p=0.0, - backend=None, + segments, + key_valid=key_valid, 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) diff --git a/tests/models/test_attention_dispatch.py b/tests/models/test_attention_dispatch.py index 2143707aac40..185d824c66b8 100644 --- a/tests/models/test_attention_dispatch.py +++ b/tests/models/test_attention_dispatch.py @@ -24,7 +24,9 @@ from diffusers.models._modeling_parallel import ContextParallelConfig, ParallelConfig from diffusers.models.attention_dispatch import attention_backend as attention_backend_ctx -from diffusers.models.attention_dispatch import dispatch_attention_fn +from diffusers.models.attention_dispatch import dispatch_attention_fn, dispatch_block_causal_attention_fn +from diffusers.models.cache_utils import apply_kv_cache +from diffusers.models.transformers.transformer_qwenimage21 import QwenImage21KVLayerCache from ..testing_utils import ( is_attention, @@ -42,6 +44,90 @@ GRAD_RTOL = 2e-2 +class TestBlockCausalAttention: + @pytest.mark.parametrize("batch_size", [1, 2]) + @pytest.mark.parametrize("pad_keys", [False, True]) + @pytest.mark.parametrize("has_prefix", [False, True]) + def test_matches_dense_attention(self, batch_size, pad_keys, has_prefix): + torch.manual_seed(0) + query, key, value = [torch.randn(batch_size, 12, 2, 16, requires_grad=True) for _ in range(3)] + segments = [(0, 3, True), (3, 5, False), (5, 7, False), (7, 8, True)] if has_prefix else [] + block_ids = torch.tensor([-1, -1, -1, 0, 0, 1, 1, -1, 2, 2, 2, 2] if has_prefix else [0] * 12) + positions = torch.arange(12) + same_image = (block_ids[:, None] == block_ids[None, :]) & (block_ids[:, None] >= 0) + mask = ((positions[:, None] >= positions[None, :]) | same_image)[None, None] + key_valid = None + if pad_keys: + key_valid = torch.ones(batch_size, 12, dtype=torch.bool) + key_valid[-1, 1] = False + mask = mask & key_valid[:, None, None, :] + expected = F.scaled_dot_product_attention( + query.transpose(1, 2), key.transpose(1, 2), value.transpose(1, 2), attn_mask=mask + ).transpose(1, 2) + output = dispatch_block_causal_attention_fn(query, key, value, segments, key_valid=key_valid) + torch.testing.assert_close(output, expected) + expected_grads = torch.autograd.grad(expected.sum(), (query, key, value)) + grads = torch.autograd.grad(output.sum(), (query, key, value)) + for grad, expected_grad in zip(grads, expected_grads): + torch.testing.assert_close(grad, expected_grad) + + +def _context_parallel_kv_cache_worker(rank, world_size, port): + dist.init_process_group("cpu:gloo", init_method=f"tcp://localhost:{port}", rank=rank, world_size=world_size) + try: + mesh = dist.device_mesh.init_device_mesh("cpu", (1, world_size), mesh_dim_names=("ring", "ulysses")) + cp = ContextParallelConfig(ulysses_degree=world_size, ulysses_anything=True) + cp.setup(rank, world_size, torch.device("cpu"), mesh) + parallel = ParallelConfig(context_parallel_config=cp) + torch.manual_seed(4) + q, k, v = [torch.randn(2, 8, 4, 16) for _ in range(3)] + + def shard(tensor): + return torch.tensor_split(tensor, world_size, dim=1)[rank].contiguous() + + cache = QwenImage21KVLayerCache() + apply_kv_cache(shard(k), shard(v), cache, "extract", slice(0, 3), parallel_config=parallel) + for cached, original in zip(cache.get(), (k, v)): + torch.testing.assert_close(cached, shard(original[:, :3])) + valid = torch.ones(2, 8, dtype=torch.bool) + valid[-1, 1] = False + mask = torch.ones(8, 8, dtype=torch.bool).tril() + mask[3:] = True + mask = mask[None, None] & valid[:, None, None] + ref = F.scaled_dot_product_attention( + q.transpose(1, 2), k.transpose(1, 2), v.transpose(1, 2), attn_mask=mask + ).transpose(1, 2) + out = dispatch_block_causal_attention_fn( + shard(q), shard(k), shard(v), [(0, 3, True)], key_valid=valid, parallel_config=parallel + ) + torch.testing.assert_close(out, shard(ref)) + qd, kd, vd = [torch.randn(2, 5, 4, 16) for _ in range(3)] + fullk, fullv = torch.cat([k[:, :3], kd], dim=1), torch.cat([v[:, :3], vd], dim=1) + ref = F.scaled_dot_product_attention( + qd.transpose(1, 2), fullk.transpose(1, 2), fullv.transpose(1, 2), attn_mask=valid[:, None, None] + ).transpose(1, 2) + ks, vs, mask = apply_kv_cache( + shard(kd), shard(vd), cache, "cached", None, attention_mask=valid[:, None, None], parallel_config=parallel + ) + out = dispatch_attention_fn(shard(qd), ks, vs, attn_mask=mask, parallel_config=parallel) + torch.testing.assert_close(out, shard(ref)) + finally: + dist.destroy_process_group() + + +@is_context_parallel +@pytest.mark.skipif(not dist.is_gloo_available(), reason="Gloo is required") +class TestContextParallelKVCache: + @pytest.mark.parametrize("world_size", [2, 4]) + def test_cached_attention_matches_full_attention(self, world_size): + mp.spawn( + _context_parallel_kv_cache_worker, + args=(world_size, _find_free_port()), + nprocs=world_size, + join=True, + ) + + def _attention_backward_parity_worker(rank, world_size, master_port, cp_dict, attention_backend, return_dict): """Op-level worker: check `dispatch_attention_fn` gradients against a single-process reference. diff --git a/tests/models/transformers/test_models_transformer_qwenimage21.py b/tests/models/transformers/test_models_transformer_qwenimage21.py index c1b62993689f..32adb3800726 100644 --- a/tests/models/transformers/test_models_transformer_qwenimage21.py +++ b/tests/models/transformers/test_models_transformer_qwenimage21.py @@ -15,10 +15,16 @@ import pytest import torch +import torch.distributed as dist +import torch.multiprocessing as mp from torch.nn.attention.flex_attention import create_mask from diffusers import QwenImage21Transformer2DModel -from diffusers.models.transformers.transformer_qwenimage21 import build_qwenimage21_block_causal_mask +from diffusers.models._modeling_parallel import ContextParallelConfig +from diffusers.models.transformers.transformer_qwenimage21 import ( + QwenImage21KVCache, + build_qwenimage21_block_causal_mask, +) from diffusers.utils.torch_utils import randn_tensor from ...testing_utils import ( @@ -38,6 +44,7 @@ TaylorSeerCacheTesterMixin, TrainingTesterMixin, ) +from ..testing_utils.parallelism import DEVICE_CONFIG, _find_free_port enable_full_determinism() @@ -305,9 +312,84 @@ class TestQwenImage21TransformerAttention(QwenImage21TransformerTesterConfig, At pass +def _qwenimage21_cached_context_parallel_worker(rank, world_size, port, ulysses_anything): + device_config = DEVICE_CONFIG[torch_device] + device_config["module"].set_device(rank) + device = torch.device(f"{torch_device}:{rank}") + dist.init_process_group( + device_config["backend"], init_method=f"tcp://localhost:{port}", rank=rank, world_size=world_size + ) + try: + config = QwenImage21TransformerTesterConfig() + cases = [(1, 4, 0, False), (2, 4, 2, True)] + if ulysses_anything: + cases.extend([(1, 1, 0, True), (2, 3, 2, True)]) + for batch_size, text_len, num_conditions, pad_prompt in cases: + torch.manual_seed(0) + init_dict = dict(config.get_init_dict(), num_attention_heads=4) + model = config.model_class(**init_dict).to(device).eval() + inputs = { + "hidden_states": torch.randn(batch_size, 4 * (num_conditions + 1), 4, device=device), + "encoder_hidden_states": torch.randn(batch_size, text_len + num_conditions, 8, device=device), + "encoder_hidden_states_mask": None, + "timestep": torch.full((batch_size,), 0.9, device=device), + "img_shapes": [[(1, 2, 2)] * (num_conditions + 1)] * batch_size, + "img_mask": torch.tensor( + [[False] * text_len + [True] * (num_conditions + 1)] * batch_size, device=device + ), + } + if pad_prompt: + inputs["encoder_hidden_states_mask"] = torch.ones( + batch_size, text_len + num_conditions, device=device, dtype=torch.long + ) + inputs["encoder_hidden_states_mask"][-1, text_len - 1] = 0 + steps = [inputs] + for timestep in (0.6, 0.3): + hidden_states = inputs["hidden_states"].clone() + hidden_states[:, -4:] = torch.randn_like(hidden_states[:, -4:]) + steps.append(dict(inputs, hidden_states=hidden_states, timestep=inputs["timestep"] * timestep)) + + reference_cache = QwenImage21KVCache(init_dict["num_layers"]) + with torch.no_grad(): + reference_prefill = model( + **inputs, kv_cache=reference_cache, kv_cache_mode="extract", return_dict=False + )[0] + references = [model(**step, return_dict=False)[0][:, -4:] for step in steps[1:]] + + model.enable_parallelism( + config=ContextParallelConfig(ulysses_degree=world_size, ulysses_anything=ulysses_anything) + ) + cache = QwenImage21KVCache(init_dict["num_layers"]) + with torch.no_grad(): + prefill = model(**inputs, kv_cache=cache, kv_cache_mode="extract", return_dict=False)[0] + torch.testing.assert_close(prefill, reference_prefill, atol=2e-5, rtol=2e-5) + for index in range(init_dict["num_layers"]): + for cached, full in zip(cache.get_layer(index).get(), reference_cache.get_layer(index).get()): + expected = torch.tensor_split(full, world_size, dim=1)[rank] + torch.testing.assert_close(cached, expected, atol=2e-5, rtol=2e-5) + assert cached.untyped_storage().nbytes() == cached.numel() * cached.element_size() + for step, reference in zip(steps[1:], references): + decoded = model(**step, kv_cache=cache, kv_cache_mode="cached", return_dict=False)[0] + torch.testing.assert_close(decoded, reference, atol=2e-5, rtol=2e-5) + finally: + dist.destroy_process_group() + + @is_context_parallel @require_torch_multi_accelerator class TestQwenImage21TransformerContextParallel(QwenImage21TransformerTesterConfig, ContextParallelTesterMixin): + @pytest.mark.parametrize("world_size", [2, 4]) + @pytest.mark.parametrize("ulysses_anything", [False, True]) + def test_context_parallel_kv_cache(self, world_size, ulysses_anything): + if DEVICE_CONFIG[torch_device]["module"].device_count() < world_size: + pytest.skip(f"Requires {world_size} devices") + mp.spawn( + _qwenimage21_cached_context_parallel_worker, + args=(world_size, _find_free_port(), ulysses_anything), + nprocs=world_size, + join=True, + ) + @pytest.mark.parametrize("cp_type", ["ulysses_degree"], ids=["ulysses"]) def test_context_parallel_inference(self, cp_type, batch_size: int = 1): super().test_context_parallel_inference(cp_type, batch_size=batch_size) From c6a9a3ce8aa334902607e464e2adbac3a24e96b8 Mon Sep 17 00:00:00 2001 From: Sayak Paul Date: Wed, 7 Oct 2026 12:14:10 +0000 Subject: [PATCH 3/3] generalize more. --- src/diffusers/models/attention_dispatch.py | 152 ++++++++++-------- src/diffusers/models/cache_utils.py | 76 --------- .../transformers/transformer_qwenimage21.py | 85 +++------- tests/models/test_attention_dispatch.py | 97 ++++++++--- .../test_models_transformer_qwenimage21.py | 2 +- 5 files changed, 187 insertions(+), 225 deletions(-) diff --git a/src/diffusers/models/attention_dispatch.py b/src/diffusers/models/attention_dispatch.py index d99ba523f5dd..025f5d496776 100644 --- a/src/diffusers/models/attention_dispatch.py +++ b/src/diffusers/models/attention_dispatch.py @@ -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 {} @@ -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, @@ -449,68 +536,6 @@ def dispatch_attention_fn( return backend_fn(**kwargs) -def dispatch_block_causal_attention_fn( - query: torch.Tensor, - key: torch.Tensor, - value: torch.Tensor, - segments: list[tuple[int, int, bool]], - *, - key_valid: torch.Tensor | None = None, - parallel_config: ParallelConfig | None = None, -) -> torch.Tensor: - cp_config = None if parallel_config is None else parallel_config.context_parallel_config - prefix_len = segments[-1][1] if segments else 0 - if cp_config is not None: - if cp_config.ring_degree > 1: - raise NotImplementedError("Block-causal attention currently supports Ulysses context parallelism only.") - group = cp_config._ulysses_mesh.get_group() - seq_len = sum(gather_size_by_comm(query.shape[1], group)) - attention_mask = torch.ones(seq_len, seq_len, dtype=torch.bool, device=query.device).tril() - for start, end, is_causal in segments: - if not is_causal: - attention_mask[start:end, start:end] = True - attention_mask[prefix_len:] = True - attention_mask = attention_mask[None, None] - if key_valid is not None: - attention_mask = attention_mask & key_valid[:, None, None, :] - return dispatch_attention_fn(query, key, value, attn_mask=attention_mask, parallel_config=parallel_config) - - outputs = [] - for start, end, is_causal in segments: - seg_mask = None - if is_causal: - 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, - parallel_config=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, :], - parallel_config=parallel_config, - ) - ) - return torch.cat(outputs, dim=1) - - # ===== Checks ===== # A list of very simple functions to catch common errors quickly when debugging. @@ -2870,7 +2895,8 @@ def forward( attn_mask = F.pad(attn_mask, (0, max_local - attn_mask.shape[-1])) mask_list = [torch.empty_like(attn_mask) for _ in range(dist.get_world_size(group=group))] dist.all_gather(mask_list, attn_mask, group=group) - attn_mask = torch.cat([mask[..., :size] for mask, size in zip(mask_list, mask_local_sizes)], dim=-1) + attn_mask = torch.cat(mask_list, dim=-1) + attn_mask = attn_mask[..., : sum(mask_local_sizes)] out = forward_op( ctx, diff --git a/src/diffusers/models/cache_utils.py b/src/diffusers/models/cache_utils.py index 5ef61b2b25a7..886ab6032bd4 100644 --- a/src/diffusers/models/cache_utils.py +++ b/src/diffusers/models/cache_utils.py @@ -13,89 +13,13 @@ # limitations under the License. from contextlib import contextmanager -from typing import Any - -import torch -import torch.distributed as dist from ..utils.logging import get_logger -from ._modeling_parallel import ParallelConfig, gather_size_by_comm logger = get_logger(__name__) # pylint: disable=invalid-name -def apply_kv_cache( - key: torch.Tensor, - value: torch.Tensor, - layer_cache, - cache_mode: str | None, - cache_write_slice: slice | None, - *, - attention_mask: Any | None = None, - parallel_config: ParallelConfig | None = None, -) -> tuple[torch.Tensor, torch.Tensor, Any]: - if layer_cache is None: - return key, value, attention_mask - - cp_config = None if parallel_config is None else parallel_config.context_parallel_config - if cp_config is not None and cp_config.ring_degree > 1: - raise NotImplementedError("KV caching currently supports Ulysses context parallelism only.") - - if cache_mode == "extract" and cache_write_slice is not None: - if cp_config is 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()) - else: - group = cp_config._ulysses_mesh.get_group() - rank = dist.get_rank(group) - world_size = dist.get_world_size(group) - local_sizes = gather_size_by_comm(key.shape[1], group) - cache_start, cache_stop, cache_step = cache_write_slice.indices(sum(local_sizes)) - if cache_step != 1: - raise ValueError("Context-parallel KV caching requires a slice with step 1.") - cache_stop = max(cache_start, cache_stop) - cache_size = cache_stop - cache_start - if not cp_config.ulysses_anything and cache_size % world_size != 0: - raise ValueError( - "The cached sequence length must be divisible by the Ulysses degree. Enable " - "`ulysses_anything=True` to cache an uneven sequence." - ) - offsets = [sum(local_sizes[:index]) for index in range(world_size)] - starts = [min(max(cache_start - offset, 0), size) for offset, size in zip(offsets, local_sizes)] - stops = [min(max(cache_stop - offset, 0), size) for offset, size in zip(offsets, local_sizes)] - cache_sizes = [stop - start for start, stop in zip(starts, stops)] - max_size = max(cache_sizes) - cached = [] - for tensor in (key, value): - local = tensor[:, starts[rank] : stops[rank]].contiguous() - if local.shape[1] < max_size: - padding = tensor.new_zeros(tensor.shape[0], max_size - local.shape[1], *tensor.shape[2:]) - local = torch.cat([local, padding], dim=1) - gathered = [torch.empty_like(local) for _ in range(world_size)] - dist.all_gather(gathered, local, group=group) - full_cache = torch.cat([part[:, :size] for part, size in zip(gathered, cache_sizes)], dim=1) - cached.append(torch.tensor_split(full_cache, world_size, dim=1)[rank].clone()) - layer_cache.store(*cached) - elif cache_mode == "cached": - cached_key, cached_value = layer_cache.get() - if cp_config is not None and attention_mask is not None: - group = cp_config._ulysses_mesh.get_group() - rank = dist.get_rank(group) - world_size = dist.get_world_size(group) - target_len = sum(gather_size_by_comm(key.shape[1], group)) - prefix_len = attention_mask.shape[-1] - target_len - prefix_masks = torch.tensor_split(attention_mask[..., :prefix_len], world_size, dim=-1) - target_masks = torch.tensor_split(attention_mask[..., prefix_len:], world_size, dim=-1) - attention_mask = torch.cat([prefix_masks[rank], target_masks[rank]], dim=-1) - key = torch.cat([cached_key, key], dim=1) - value = torch.cat([cached_value, value], dim=1) - - return key, value, attention_mask - - class CacheMixin: r""" A class for enable/disabling caching techniques on diffusion models. diff --git a/src/diffusers/models/transformers/transformer_qwenimage21.py b/src/diffusers/models/transformers/transformer_qwenimage21.py index 7fefa2f3f599..c18c2e757c8c 100644 --- a/src/diffusers/models/transformers/transformer_qwenimage21.py +++ b/src/diffusers/models/transformers/transformer_qwenimage21.py @@ -26,8 +26,8 @@ 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, dispatch_block_causal_attention_fn -from ..cache_utils import CacheMixin, apply_kv_cache +from ..attention_dispatch import dispatch_attention_fn +from ..cache_utils import CacheMixin from ..embeddings import TimestepEmbedding from ..modeling_outputs import Transformer2DModelOutput from ..modeling_utils import ModelMixin @@ -329,13 +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, - attention_mask: Any | None = None, - parallel_config: Any | None = None, -) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, int, Any]: - """Shared QKV projection, norm, RoPE and KV-cache bookkeeping for both processors.""" +) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, int]: query = attn.to_q(hidden_states) key = attn.to_k(hidden_states) value = attn.to_v(hidden_states) @@ -351,16 +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) - key, value, attention_mask = apply_kv_cache( - key, - value, - layer_cache, - kv_cache_mode, - cache_write_slice, - attention_mask=attention_mask, - parallel_config=parallel_config, - ) - return query, key, value, query.shape[1], attention_mask + return query, key, value, query.shape[1] class QwenImage21FlexAttnProcessor: @@ -403,16 +388,7 @@ def __call__( "Context parallelism is not implemented for QwenImage21FlexAttnProcessor. " "Use QwenImage21AttnProcessor instead." ) - query, key, value, seq_len_q, attention_mask = _qwenimage21_prepare_qkv( - attn, - hidden_states, - rotary_emb, - layer_cache, - kv_cache_mode, - cache_write_slice, - attention_mask, - self._parallel_config, - ) + 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): @@ -449,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] @@ -460,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) @@ -492,37 +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, attention_mask = _qwenimage21_prepare_qkv( - attn, - hidden_states, - rotary_emb, - layer_cache, - kv_cache_mode, - cache_write_slice, - attention_mask, - self._parallel_config, + 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: - hidden_states = dispatch_block_causal_attention_fn( - query, - key, - value, - segments, - key_valid=key_valid, - parallel_config=self._parallel_config, - ) hidden_states = hidden_states[:, :seq_len_q] hidden_states = hidden_states.flatten(2, 3).type_as(query) diff --git a/tests/models/test_attention_dispatch.py b/tests/models/test_attention_dispatch.py index 185d824c66b8..b01bf29b2caf 100644 --- a/tests/models/test_attention_dispatch.py +++ b/tests/models/test_attention_dispatch.py @@ -24,8 +24,8 @@ from diffusers.models._modeling_parallel import ContextParallelConfig, ParallelConfig from diffusers.models.attention_dispatch import attention_backend as attention_backend_ctx -from diffusers.models.attention_dispatch import dispatch_attention_fn, dispatch_block_causal_attention_fn -from diffusers.models.cache_utils import apply_kv_cache +from diffusers.models.attention_dispatch import dispatch_attention_fn +from diffusers.models.transformers.transformer_flux2 import Flux2KVLayerCache from diffusers.models.transformers.transformer_qwenimage21 import QwenImage21KVLayerCache from ..testing_utils import ( @@ -64,7 +64,13 @@ def test_matches_dense_attention(self, batch_size, pad_keys, has_prefix): expected = F.scaled_dot_product_attention( query.transpose(1, 2), key.transpose(1, 2), value.transpose(1, 2), attn_mask=mask ).transpose(1, 2) - output = dispatch_block_causal_attention_fn(query, key, value, segments, key_valid=key_valid) + output = dispatch_attention_fn( + query, + key, + value, + attn_mask=None if key_valid is None else key_valid[:, None, None, :], + block_causal_segments=segments, + ) torch.testing.assert_close(output, expected) expected_grads = torch.autograd.grad(expected.sum(), (query, key, value)) grads = torch.autograd.grad(output.sum(), (query, key, value)) @@ -72,45 +78,83 @@ def test_matches_dense_attention(self, batch_size, pad_keys, has_prefix): torch.testing.assert_close(grad, expected_grad) -def _context_parallel_kv_cache_worker(rank, world_size, port): +def _context_parallel_kv_cache_worker(rank, world_size, port, ulysses_anything): dist.init_process_group("cpu:gloo", init_method=f"tcp://localhost:{port}", rank=rank, world_size=world_size) try: mesh = dist.device_mesh.init_device_mesh("cpu", (1, world_size), mesh_dim_names=("ring", "ulysses")) - cp = ContextParallelConfig(ulysses_degree=world_size, ulysses_anything=True) + cp = ContextParallelConfig(ulysses_degree=world_size, ulysses_anything=ulysses_anything) cp.setup(rank, world_size, torch.device("cpu"), mesh) parallel = ParallelConfig(context_parallel_config=cp) torch.manual_seed(4) - q, k, v = [torch.randn(2, 8, 4, 16) for _ in range(3)] + prefix_len = 3 if ulysses_anything else 4 + heads = 7 if ulysses_anything else 4 + q, k, v = [torch.randn(2, 8, heads, 16) for _ in range(3)] def shard(tensor): return torch.tensor_split(tensor, world_size, dim=1)[rank].contiguous() - cache = QwenImage21KVLayerCache() - apply_kv_cache(shard(k), shard(v), cache, "extract", slice(0, 3), parallel_config=parallel) - for cached, original in zip(cache.get(), (k, v)): - torch.testing.assert_close(cached, shard(original[:, :3])) valid = torch.ones(2, 8, dtype=torch.bool) valid[-1, 1] = False mask = torch.ones(8, 8, dtype=torch.bool).tril() - mask[3:] = True + mask[prefix_len:] = True mask = mask[None, None] & valid[:, None, None] ref = F.scaled_dot_product_attention( q.transpose(1, 2), k.transpose(1, 2), v.transpose(1, 2), attn_mask=mask ).transpose(1, 2) - out = dispatch_block_causal_attention_fn( - shard(q), shard(k), shard(v), [(0, 3, True)], key_valid=valid, parallel_config=parallel - ) - torch.testing.assert_close(out, shard(ref)) - qd, kd, vd = [torch.randn(2, 5, 4, 16) for _ in range(3)] - fullk, fullv = torch.cat([k[:, :3], kd], dim=1), torch.cat([v[:, :3], vd], dim=1) - ref = F.scaled_dot_product_attention( - qd.transpose(1, 2), fullk.transpose(1, 2), fullv.transpose(1, 2), attn_mask=valid[:, None, None] - ).transpose(1, 2) - ks, vs, mask = apply_kv_cache( - shard(kd), shard(vd), cache, "cached", None, attention_mask=valid[:, None, None], parallel_config=parallel + uncached = dispatch_attention_fn(shard(q), shard(k), shard(v), attn_mask=mask, parallel_config=parallel) + torch.testing.assert_close(uncached, shard(ref)) + train_inputs = [shard(tensor).requires_grad_() for tensor in (q, k, v)] + output = dispatch_attention_fn( + *train_inputs, + attn_mask=valid[:, None, None], + block_causal_segments=[(0, prefix_len, True)], + parallel_config=parallel, ) - out = dispatch_attention_fn(shard(qd), ks, vs, attn_mask=mask, parallel_config=parallel) - torch.testing.assert_close(out, shard(ref)) + if ulysses_anything: + with pytest.raises(NotImplementedError, match="Backward pass for Ulysses Anything"): + output.sum().backward() + else: + output.sum().backward() + ref_inputs = [tensor.detach().requires_grad_() for tensor in (q, k, v)] + ref_output = F.scaled_dot_product_attention( + *(tensor.transpose(1, 2) for tensor in ref_inputs), attn_mask=mask + ) + ref_output.sum().backward() + for tensor, reference in zip(train_inputs, ref_inputs): + torch.testing.assert_close(tensor.grad, shard(reference.grad)) + for cache_class in (QwenImage21KVLayerCache, Flux2KVLayerCache): + cache = cache_class() + out = dispatch_attention_fn( + shard(q), + shard(k), + shard(v), + attn_mask=valid[:, None, None], + parallel_config=parallel, + kv_cache=cache, + kv_cache_mode="extract", + cache_write_slice=slice(0, prefix_len), + block_causal_segments=[(0, prefix_len, True)], + ) + torch.testing.assert_close(out, shard(ref)) + for cached, original in zip(cache.get(), (k, v)): + expected = torch.tensor_split(original[:, :prefix_len], world_size, dim=2)[rank] + torch.testing.assert_close(cached, expected) + assert cached.untyped_storage().nbytes() == cached.numel() * cached.element_size() + qd, kd, vd = [torch.randn(2, 8 - prefix_len, heads, 16) for _ in range(3)] + fullk, fullv = torch.cat([k[:, :prefix_len], kd], dim=1), torch.cat([v[:, :prefix_len], vd], dim=1) + decode_ref = F.scaled_dot_product_attention( + qd.transpose(1, 2), fullk.transpose(1, 2), fullv.transpose(1, 2), attn_mask=valid[:, None, None] + ).transpose(1, 2) + out = dispatch_attention_fn( + shard(qd), + shard(kd), + shard(vd), + attn_mask=valid[:, None, None], + parallel_config=parallel, + kv_cache=cache, + kv_cache_mode="cached", + ) + torch.testing.assert_close(out, shard(decode_ref)) finally: dist.destroy_process_group() @@ -119,10 +163,11 @@ def shard(tensor): @pytest.mark.skipif(not dist.is_gloo_available(), reason="Gloo is required") class TestContextParallelKVCache: @pytest.mark.parametrize("world_size", [2, 4]) - def test_cached_attention_matches_full_attention(self, world_size): + @pytest.mark.parametrize("ulysses_anything", [False, True]) + def test_cached_attention_matches_full_attention(self, world_size, ulysses_anything): mp.spawn( _context_parallel_kv_cache_worker, - args=(world_size, _find_free_port()), + args=(world_size, _find_free_port(), ulysses_anything), nprocs=world_size, join=True, ) diff --git a/tests/models/transformers/test_models_transformer_qwenimage21.py b/tests/models/transformers/test_models_transformer_qwenimage21.py index 32adb3800726..57697a59a722 100644 --- a/tests/models/transformers/test_models_transformer_qwenimage21.py +++ b/tests/models/transformers/test_models_transformer_qwenimage21.py @@ -365,7 +365,7 @@ def _qwenimage21_cached_context_parallel_worker(rank, world_size, port, ulysses_ torch.testing.assert_close(prefill, reference_prefill, atol=2e-5, rtol=2e-5) for index in range(init_dict["num_layers"]): for cached, full in zip(cache.get_layer(index).get(), reference_cache.get_layer(index).get()): - expected = torch.tensor_split(full, world_size, dim=1)[rank] + expected = torch.tensor_split(full, world_size, dim=2)[rank] torch.testing.assert_close(cached, expected, atol=2e-5, rtol=2e-5) assert cached.untyped_storage().nbytes() == cached.numel() * cached.element_size() for step, reference in zip(steps[1:], references):