diff --git a/src/diffusers/models/controlnets/controlnet_z_image.py b/src/diffusers/models/controlnets/controlnet_z_image.py index a4800b255ef0..38b610d50d97 100644 --- a/src/diffusers/models/controlnets/controlnet_z_image.py +++ b/src/diffusers/models/controlnets/controlnet_z_image.py @@ -29,6 +29,7 @@ from ..attention_dispatch import dispatch_attention_fn from ..controlnets.controlnet import zero_module from ..modeling_utils import ModelMixin +from ..transformers.transformer_z_image import COMPLEX_LESS_ROPE_BACKENDS, apply_rotary_emb_zimage ADALN_EMBED_DIM = 256 @@ -113,16 +114,9 @@ def __call__( key = attn.norm_k(key) # Apply RoPE - def apply_rotary_emb(x_in: torch.Tensor, freqs_cis: torch.Tensor) -> torch.Tensor: - with torch.amp.autocast("cuda", enabled=False): - x = torch.view_as_complex(x_in.float().reshape(*x_in.shape[:-1], -1, 2)) - freqs_cis = freqs_cis.unsqueeze(2) - x_out = torch.view_as_real(x * freqs_cis).flatten(3) - return x_out.type_as(x_in) # todo - if freqs_cis is not None: - query = apply_rotary_emb(query, freqs_cis) - key = apply_rotary_emb(key, freqs_cis) + query = apply_rotary_emb_zimage(query, freqs_cis) + key = apply_rotary_emb_zimage(key, freqs_cis) # Cast to correct dtype dtype = query.dtype @@ -317,20 +311,25 @@ def __call__(self, ids: torch.Tensor): assert ids.ndim == 2 assert ids.shape[-1] == len(self.axes_dims) device = ids.device + # On complex-less backends the complex64 tables cannot be moved to or indexed on device: keep them on + # CPU, index there, and cross the device boundary as reals via torch.view_as_real. + complex_less = device.type in COMPLEX_LESS_ROPE_BACKENDS if self.freqs_cis is None: self.freqs_cis = self.precompute_freqs_cis(self.axes_dims, self.axes_lens, theta=self.theta) - self.freqs_cis = [freqs_cis.to(device) for freqs_cis in self.freqs_cis] - else: - # Ensure freqs_cis are on the same device as ids - if self.freqs_cis[0].device != device: - self.freqs_cis = [freqs_cis.to(device) for freqs_cis in self.freqs_cis] + # Keep the tables on the indexing device: CPU on complex-less backends, ids' device otherwise + table_device = torch.device("cpu") if complex_less else device + if self.freqs_cis[0].device != table_device: + self.freqs_cis = [freqs_cis.to(table_device) for freqs_cis in self.freqs_cis] + index = ids.cpu() if complex_less else ids result = [] for i in range(len(self.axes_dims)): - index = ids[:, i] - result.append(self.freqs_cis[i][index]) - return torch.cat(result, dim=-1) + result.append(self.freqs_cis[i][index[:, i]]) + freqs = torch.cat(result, dim=-1) + if complex_less: + return torch.view_as_real(freqs).to(device) + return freqs @maybe_allow_in_graph diff --git a/src/diffusers/models/transformers/transformer_qwenimage.py b/src/diffusers/models/transformers/transformer_qwenimage.py index 5a242c8bd5c0..0bc672cd56c8 100644 --- a/src/diffusers/models/transformers/transformer_qwenimage.py +++ b/src/diffusers/models/transformers/transformer_qwenimage.py @@ -166,13 +166,18 @@ def apply_rotary_emb_qwen_neuron(x: torch.Tensor, freqs: torch.Tensor) -> torch. return (x.float() * cos + x_rotated.float() * sin).to(x.dtype) -# RoPE application is backend-dependent: the default path multiplies by a complex exponential, which Neuron cannot -# represent. Callers select by `device.type` and fall back to the default for any backend not listed here. +# RoPE application is backend-dependent: the default path multiplies by a complex exponential, which backends +# without complex64 kernels (Neuron, Ascend NPU) cannot represent. Callers select by `device.type` and fall back +# to the default for any backend not listed here. ROPE_PER_DEVICE = { "cuda": functools.partial(apply_rotary_emb_qwen, use_real=False), "neuron": apply_rotary_emb_qwen_neuron, + "npu": apply_rotary_emb_qwen_neuron, } +# Backends whose RoPE frequency tables are sent as rotation angles instead of complex exponentials. +COMPLEX_LESS_ROPE_BACKENDS = ("neuron", "npu") + def compute_text_seq_len_from_mask( encoder_hidden_states: torch.Tensor, encoder_hidden_states_mask: torch.Tensor | None @@ -268,8 +273,8 @@ def rope_params(self, index, dim, theta=10000): @lru_cache_unless_export(maxsize=None) def _get_device_freqs(self, device: torch.device) -> tuple[torch.Tensor, torch.Tensor]: """Return pos_freqs and neg_freqs on the given device.""" - if device is not None and device.type == "neuron": - # Neuron has no complex dtype, so send the rotation angles instead and let + if device is not None and device.type in COMPLEX_LESS_ROPE_BACKENDS: + # Neuron and Ascend NPU have no complex64 kernels, so send the rotation angles instead and let # `apply_rotary_emb_qwen_neuron` take cos/sin on device. `torch.angle` runs on CPU while the freqs are # still complex; wrapping into (-pi, pi] is harmless because only cos/sin of the angle are used. return torch.angle(self.pos_freqs).to(device), torch.angle(self.neg_freqs).to(device) @@ -397,8 +402,8 @@ def rope_params(self, index, dim, theta=10000): @lru_cache_unless_export(maxsize=None) def _get_device_freqs(self, device: torch.device) -> tuple[torch.Tensor, torch.Tensor]: """Return pos_freqs and neg_freqs on the given device.""" - if device is not None and device.type == "neuron": - # Neuron has no complex dtype, so send the rotation angles instead and let + if device is not None and device.type in COMPLEX_LESS_ROPE_BACKENDS: + # Neuron and Ascend NPU have no complex64 kernels, so send the rotation angles instead and let # `apply_rotary_emb_qwen_neuron` take cos/sin on device. `torch.angle` runs on CPU while the freqs are # still complex; wrapping into (-pi, pi] is harmless because only cos/sin of the angle are used. return torch.angle(self.pos_freqs).to(device), torch.angle(self.neg_freqs).to(device) diff --git a/src/diffusers/models/transformers/transformer_qwenimage21.py b/src/diffusers/models/transformers/transformer_qwenimage21.py index b6eabfd584bb..9088d198a561 100644 --- a/src/diffusers/models/transformers/transformer_qwenimage21.py +++ b/src/diffusers/models/transformers/transformer_qwenimage21.py @@ -133,6 +133,26 @@ def apply_rotary_emb_qwen( return x_out.type_as(x) +# Backends whose kernel libraries lack complex64 support (Ascend NPU, AWS Neuron, ...); see +# https://github.com/huggingface/diffusers/issues/12668. On these backends the RoPE frequency table stays on CPU, +# the per-request selection runs there, and frequencies cross the device boundary as reals (torch.view_as_real). +COMPLEX_LESS_ROPE_BACKENDS = ("neuron", "npu") + + +def apply_rotary_emb_qwen_real(x: torch.Tensor, freqs_cis: torch.Tensor) -> torch.Tensor: + """Real-valued variant of `apply_rotary_emb_qwen(..., use_real=False)` for backends without complex64 kernels. + + `freqs_cis` carries `torch.view_as_real` pairs of the complex exponentials, shape `[S, D//2, 2]`. The complex + multiply expands to the same four real multiply-adds, and the operands are bit-identical to the complex path's + (`torch.polar(1, theta)` stores exactly `(cos theta, sin theta)`), so the result matches the complex path to + within float32 multiply-add rounding — both paths compute in float32 via `x.float()`. + """ + x_real, x_imag = x.float().reshape(*x.shape[:-1], -1, 2).unbind(-1) # [B, S, H, D//2] + cos, sin = freqs_cis.unsqueeze(1).unbind(-1) # [S, 1, D//2] + out = torch.stack([x_real * cos - x_imag * sin, x_real * sin + x_imag * cos], dim=-1).flatten(3) + return out.type_as(x) + + class QwenImage21TemporalTimesteps(nn.Module): r"""Sinusoidal timestep embedding. `cos` occupies the first half of the channels and `sin` the second.""" @@ -345,8 +365,14 @@ def _qwenimage21_prepare_qkv( key = attn.norm_k(key).to(value.dtype) if rotary_emb is not None: - query = apply_rotary_emb_qwen(query, rotary_emb, use_real=False) - key = apply_rotary_emb_qwen(key, rotary_emb, use_real=False) + # Complex-capable backends carry complex64 exponentials; COMPLEX_LESS_ROPE_BACKENDS carry + # view_as_real pairs (see QwenImage21Rope.forward). The dtype picks the path. + if rotary_emb.is_complex(): + query = apply_rotary_emb_qwen(query, rotary_emb, use_real=False) + key = apply_rotary_emb_qwen(key, rotary_emb, use_real=False) + else: + query = apply_rotary_emb_qwen_real(query, rotary_emb) + key = apply_rotary_emb_qwen_real(key, rotary_emb) if layer_cache is not None: if kv_cache_mode == "extract" and cache_write_slice is not None: @@ -677,7 +703,13 @@ def rope_params(self, index: torch.Tensor, dim: int, theta: int = 10000) -> torc def forward( self, img_shapes: list[tuple[int, int, int]], image_pad_mask: torch.Tensor, device: torch.device ) -> torch.Tensor: - self.freqs = [freq.to(device) for freq in self.freqs] + # On complex-less backends the complex64 table cannot be moved to or indexed on device: keep it on + # CPU, do the selection there, and cross the device boundary as reals via torch.view_as_real. + complex_less = device.type in COMPLEX_LESS_ROPE_BACKENDS + index_device = torch.device("cpu") if complex_less else device + self.freqs = [freq.to(index_device) for freq in self.freqs] + if complex_less: + image_pad_mask = image_pad_mask.cpu() frame_index, height_index, width_index = [], [], [] image_height_index, image_width_index = [], [] @@ -701,13 +733,18 @@ def forward( if cursor < total_len: frame_index.extend(range(position, position + total_len - cursor)) - frame_index = torch.tensor(frame_index, dtype=torch.long, device=device) + frame_index = torch.tensor(frame_index, dtype=torch.long, device=index_device) height_index = frame_index.clone() width_index = frame_index.clone() - height_index[image_pad_mask] = torch.tensor(image_height_index, dtype=torch.long, device=device) - width_index[image_pad_mask] = torch.tensor(image_width_index, dtype=torch.long, device=device) + height_index[image_pad_mask] = torch.tensor(image_height_index, dtype=torch.long, device=index_device) + width_index[image_pad_mask] = torch.tensor(image_width_index, dtype=torch.long, device=index_device) - return torch.cat([self.freqs[0][frame_index], self.freqs[1][height_index], self.freqs[2][width_index]], dim=-1) + freqs = torch.cat( + [self.freqs[0][frame_index], self.freqs[1][height_index], self.freqs[2][width_index]], dim=-1 + ) + if complex_less: + return torch.view_as_real(freqs).to(device) + return freqs class QwenImage21Transformer2DModel( diff --git a/src/diffusers/models/transformers/transformer_z_image.py b/src/diffusers/models/transformers/transformer_z_image.py index 4cea745e5ed5..f8c9030526d3 100644 --- a/src/diffusers/models/transformers/transformer_z_image.py +++ b/src/diffusers/models/transformers/transformer_z_image.py @@ -72,6 +72,31 @@ def forward(self, t): return t_emb +# Backends whose kernel libraries lack complex64 support (Ascend NPU, AWS Neuron, ...); see +# https://github.com/huggingface/diffusers/issues/12668. On these backends RopeEmbedder keeps its frequency +# tables on CPU and returns view_as_real pairs instead of complex64. +COMPLEX_LESS_ROPE_BACKENDS = ("neuron", "npu") + + +def apply_rotary_emb_zimage(x_in: torch.Tensor, freqs_cis: torch.Tensor) -> torch.Tensor: + """Apply rotary embeddings to [B, S, H, D] with freqs broadcast over heads. + + `freqs_cis` is complex64 `[B, S, D//2]` on complex-capable backends and `torch.view_as_real` pairs + `[B, S, D//2, 2]` on complex-less ones (see RopeEmbedder.__call__); the dtype picks the path. The two paths + expand to the same real multiply-adds, so they match to within float32 rounding (both compute in float32). + """ + with torch.amp.autocast("cuda", enabled=False): + if freqs_cis.is_complex(): + x = torch.view_as_complex(x_in.float().reshape(*x_in.shape[:-1], -1, 2)) + freqs_cis = freqs_cis.unsqueeze(2) + x_out = torch.view_as_real(x * freqs_cis).flatten(3) + return x_out.type_as(x_in) # todo + x_real, x_imag = x_in.float().reshape(*x_in.shape[:-1], -1, 2).unbind(-1) + cos, sin = freqs_cis.unsqueeze(2).unbind(-1) # [B, S, 1, D//2] + out = torch.stack([x_real * cos - x_imag * sin, x_real * sin + x_imag * cos], dim=-1).flatten(3) + return out.type_as(x_in) + + class ZSingleStreamAttnProcessor: """ Processor for Z-Image single stream attention that adapts the existing Attention class to match the behavior of the @@ -110,16 +135,9 @@ def __call__( key = attn.norm_k(key) # Apply RoPE - def apply_rotary_emb(x_in: torch.Tensor, freqs_cis: torch.Tensor) -> torch.Tensor: - with torch.amp.autocast("cuda", enabled=False): - x = torch.view_as_complex(x_in.float().reshape(*x_in.shape[:-1], -1, 2)) - freqs_cis = freqs_cis.unsqueeze(2) - x_out = torch.view_as_real(x * freqs_cis).flatten(3) - return x_out.type_as(x_in) # todo - if freqs_cis is not None: - query = apply_rotary_emb(query, freqs_cis) - key = apply_rotary_emb(key, freqs_cis) + query = apply_rotary_emb_zimage(query, freqs_cis) + key = apply_rotary_emb_zimage(key, freqs_cis) # Cast to correct dtype dtype = query.dtype @@ -340,20 +358,25 @@ def __call__(self, ids: torch.Tensor): assert ids.ndim == 2 assert ids.shape[-1] == len(self.axes_dims) device = ids.device + # On complex-less backends the complex64 tables cannot be moved to or indexed on device: keep them on + # CPU, index there, and cross the device boundary as reals via torch.view_as_real. + complex_less = device.type in COMPLEX_LESS_ROPE_BACKENDS if self.freqs_cis is None: self.freqs_cis = self.precompute_freqs_cis(self.axes_dims, self.axes_lens, theta=self.theta) - self.freqs_cis = [freqs_cis.to(device) for freqs_cis in self.freqs_cis] - else: - # Ensure freqs_cis are on the same device as ids - if self.freqs_cis[0].device != device: - self.freqs_cis = [freqs_cis.to(device) for freqs_cis in self.freqs_cis] + # Keep the tables on the indexing device: CPU on complex-less backends, ids' device otherwise + table_device = torch.device("cpu") if complex_less else device + if self.freqs_cis[0].device != table_device: + self.freqs_cis = [freqs_cis.to(table_device) for freqs_cis in self.freqs_cis] + index = ids.cpu() if complex_less else ids result = [] for i in range(len(self.axes_dims)): - index = ids[:, i] - result.append(self.freqs_cis[i][index]) - return torch.cat(result, dim=-1) + result.append(self.freqs_cis[i][index[:, i]]) + freqs = torch.cat(result, dim=-1) + if complex_less: + return torch.view_as_real(freqs).to(device) + return freqs class ZImageTransformer2DModel(ModelMixin, ConfigMixin, PeftAdapterMixin, FromOriginalModelMixin): diff --git a/tests/models/transformers/test_models_transformer_qwenimage.py b/tests/models/transformers/test_models_transformer_qwenimage.py index 7a03a8fe2353..78e5a07b4811 100644 --- a/tests/models/transformers/test_models_transformer_qwenimage.py +++ b/tests/models/transformers/test_models_transformer_qwenimage.py @@ -506,3 +506,23 @@ class TestQwenImageTransformerTorchAo(QwenImageTransformerQuantTesterConfig, Tor @property def torch_dtype(self): return torch.bfloat16 + + +class TestQwenImageRopeComplexLess: + def test_npu_registered_and_neuron_path_matches_complex(self): + from diffusers.models.transformers.transformer_qwenimage import ( + ROPE_PER_DEVICE, + apply_rotary_emb_qwen, + apply_rotary_emb_qwen_neuron, + ) + + assert ROPE_PER_DEVICE.get("npu") is apply_rotary_emb_qwen_neuron + + torch.manual_seed(0) + x = torch.randn(2, 65, 8, 128) + angles = (torch.rand(65, 64) * 2 - 1) * torch.pi + freqs = torch.polar(torch.ones_like(angles), angles).to(torch.complex64) + ref = apply_rotary_emb_qwen(x, freqs, use_real=False) + cand = apply_rotary_emb_qwen_neuron(x, torch.angle(freqs)) + # The atan2 round-trip in torch.angle adds up to ~1e-7 on top of multiply-add rounding. + torch.testing.assert_close(ref, cand, rtol=1e-5, atol=1e-6) diff --git a/tests/models/transformers/test_models_transformer_qwenimage21.py b/tests/models/transformers/test_models_transformer_qwenimage21.py index f6d0a0004d0c..415885a37512 100644 --- a/tests/models/transformers/test_models_transformer_qwenimage21.py +++ b/tests/models/transformers/test_models_transformer_qwenimage21.py @@ -380,3 +380,34 @@ def pretrained_model_kwargs(self): @property def torch_dtype(self): return torch.bfloat16 + + +class TestQwenImage21RopeComplexLess: + """The real-valued RoPE path for COMPLEX_LESS_ROPE_BACKENDS matches the complex reference.""" + + def _inputs(self): + img_shapes = [(1, 32, 32), (1, 16, 48)] + text_len = 77 + mask = torch.zeros(text_len + sum(h * w for _, h, w in img_shapes), dtype=torch.bool) + mask[text_len:] = True + return img_shapes, mask + + def test_real_fallback_matches_complex(self, monkeypatch): + from diffusers.models.transformers import transformer_qwenimage21 as t21 + + img_shapes, mask = self._inputs() + ref = t21.QwenImage21Rope(theta=10000, axes_dim=[16, 56, 56])(img_shapes, mask, torch.device("cpu")) + assert ref.is_complex() + + monkeypatch.setattr(t21, "COMPLEX_LESS_ROPE_BACKENDS", ("cpu",)) + cand = t21.QwenImage21Rope(theta=10000, axes_dim=[16, 56, 56])(img_shapes, mask, torch.device("cpu")) + assert cand.dtype == torch.float32 and cand.shape == ref.shape + (2,) + + # The frequency selection crosses the boundary bit-exact. + torch.testing.assert_close(torch.view_as_real(ref), cand, rtol=0, atol=0) + + # The rotation matches to within float32 multiply-add rounding. + x = torch.randn(1, mask.shape[0], 8, 128, generator=torch.Generator().manual_seed(0)) + ref_out = t21.apply_rotary_emb_qwen(x, ref, use_real=False) + cand_out = t21.apply_rotary_emb_qwen_real(x, cand) + torch.testing.assert_close(ref_out, cand_out, rtol=1e-5, atol=1e-6) diff --git a/tests/models/transformers/test_models_transformer_z_image.py b/tests/models/transformers/test_models_transformer_z_image.py index 35bafa5702ae..5e4068afcfac 100644 --- a/tests/models/transformers/test_models_transformer_z_image.py +++ b/tests/models/transformers/test_models_transformer_z_image.py @@ -369,3 +369,34 @@ def pretrained_model_name_or_path(self): @property def pretrained_model_kwargs(self): return {"subfolder": "transformer"} + + +class TestZImageRopeComplexLess: + """The real-valued RoPE path for COMPLEX_LESS_ROPE_BACKENDS matches the complex reference.""" + + def test_real_fallback_matches_complex(self, monkeypatch): + from diffusers.models.transformers import transformer_z_image as tzi + + ids = torch.stack( + [ + torch.randint(0, 64, (37,), generator=torch.Generator().manual_seed(0)), + torch.randint(0, 128, (37,), generator=torch.Generator().manual_seed(1)), + torch.randint(0, 128, (37,), generator=torch.Generator().manual_seed(2)), + ], + dim=-1, + ) + ref = tzi.RopeEmbedder(theta=256.0, axes_dims=(16, 56, 56), axes_lens=(64, 128, 128))(ids) + assert ref.is_complex() + + monkeypatch.setattr(tzi, "COMPLEX_LESS_ROPE_BACKENDS", ("cpu",)) + cand = tzi.RopeEmbedder(theta=256.0, axes_dims=(16, 56, 56), axes_lens=(64, 128, 128))(ids) + assert cand.dtype == torch.float32 and cand.shape == ref.shape + (2,) + + # The frequency selection crosses the boundary bit-exact. + torch.testing.assert_close(torch.view_as_real(ref), cand, rtol=0, atol=0) + + # Same rotation to within float32 rounding, in the processor's [B, S, H, D] shape. + x = torch.randn(2, 37, 6, 128, generator=torch.Generator().manual_seed(3)) + ref_out = tzi.apply_rotary_emb_zimage(x, ref.unsqueeze(0)) + cand_out = tzi.apply_rotary_emb_zimage(x, cand.unsqueeze(0)) + torch.testing.assert_close(ref_out, cand_out, rtol=1e-5, atol=1e-6)