diff --git a/src/diffusers/models/attention_dispatch.py b/src/diffusers/models/attention_dispatch.py index f45c9e0d1107..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, diff --git a/src/diffusers/models/transformers/transformer_qwenimage21.py b/src/diffusers/models/transformers/transformer_qwenimage21.py index b6eabfd584bb..c18c2e757c8c 100644 --- a/src/diffusers/models/transformers/transformer_qwenimage21.py +++ b/src/diffusers/models/transformers/transformer_qwenimage21.py @@ -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 @@ -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) @@ -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: @@ -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): @@ -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] @@ -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) @@ -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) @@ -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__( diff --git a/tests/models/test_attention_dispatch.py b/tests/models/test_attention_dispatch.py index 2143707aac40..b01bf29b2caf 100644 --- a/tests/models/test_attention_dispatch.py +++ b/tests/models/test_attention_dispatch.py @@ -25,6 +25,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 +from diffusers.models.transformers.transformer_flux2 import Flux2KVLayerCache +from diffusers.models.transformers.transformer_qwenimage21 import QwenImage21KVLayerCache from ..testing_utils import ( is_attention, @@ -42,6 +44,135 @@ 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_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)) + 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, 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=ulysses_anything) + cp.setup(rank, world_size, torch.device("cpu"), mesh) + parallel = ParallelConfig(context_parallel_config=cp) + torch.manual_seed(4) + 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() + + valid = torch.ones(2, 8, dtype=torch.bool) + valid[-1, 1] = False + mask = torch.ones(8, 8, dtype=torch.bool).tril() + 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) + 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, + ) + 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() + + +@is_context_parallel +@pytest.mark.skipif(not dist.is_gloo_available(), reason="Gloo is required") +class TestContextParallelKVCache: + @pytest.mark.parametrize("world_size", [2, 4]) + @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(), ulysses_anything), + 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 f6d0a0004d0c..57697a59a722 100644 --- a/tests/models/transformers/test_models_transformer_qwenimage21.py +++ b/tests/models/transformers/test_models_transformer_qwenimage21.py @@ -15,22 +15,36 @@ 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 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, SingleFileTesterMixin, TaylorSeerCacheTesterMixin, TrainingTesterMixin, ) +from ..testing_utils.parallelism import DEVICE_CONFIG, _find_free_port enable_full_determinism() @@ -298,6 +312,120 @@ 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=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): + 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) + + @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 + + class TestQwenImage21TransformerTaylorSeerCache(QwenImage21TransformerTesterConfig, TaylorSeerCacheTesterMixin): """TaylorSeerCache tests for QwenImage 2.1 Transformer."""