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

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
33 changes: 16 additions & 17 deletions src/diffusers/models/controlnets/controlnet_z_image.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -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
Expand Down
17 changes: 11 additions & 6 deletions src/diffusers/models/transformers/transformer_qwenimage.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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)
Expand Down Expand Up @@ -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)
Expand Down
51 changes: 44 additions & 7 deletions src/diffusers/models/transformers/transformer_qwenimage21.py
Original file line number Diff line number Diff line change
Expand Up @@ -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."""

Expand Down Expand Up @@ -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:
Expand Down Expand Up @@ -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 = [], []
Expand All @@ -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(
Expand Down
57 changes: 40 additions & 17 deletions src/diffusers/models/transformers/transformer_z_image.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -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):
Expand Down
20 changes: 20 additions & 0 deletions tests/models/transformers/test_models_transformer_qwenimage.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)
31 changes: 31 additions & 0 deletions tests/models/transformers/test_models_transformer_qwenimage21.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)
31 changes: 31 additions & 0 deletions tests/models/transformers/test_models_transformer_z_image.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Loading