Skip to content

Add a real-valued RoPE fallback for backends without complex64 kernels (Qwen-Image / Qwen-Image 2.1 / Z-Image) - #14905

Open
cyzlmh wants to merge 1 commit into
huggingface:mainfrom
cyzlmh:rope-complex-less-backends
Open

cyzlmh wants to merge 1 commit into
huggingface:mainfrom
cyzlmh:rope-complex-less-backends

Conversation

@cyzlmh

@cyzlmh cyzlmh commented Sep 29, 2026 •

Copy link
Copy Markdown

Title: Extend the real-valued RoPE fallback to complex-less backends (Qwen-Image 2.1, Z-Image; register NPU in Qwen-Image)


What does this PR do?

The Qwen-Image-lineage RoPE (torch.polar frequency table → complex64 index/mul
in apply_rotary_emb_qwen(..., use_real=False)) only works on backends with
complex64 kernels. On backends without them it is a hard failure, not a slowdown:

#13718 already solved exactly this for AWS Neuron: a numerically equivalent
real-valued path (apply_rotary_emb_qwen_neuron), a per-device dispatch
(ROPE_PER_DEVICE), and angle extraction at the device boundary
(_get_device_freqs, torch.angle while the table is still on CPU). However,
the two pipelines added since never received the dispatch:

  • transformer_qwenimage21.py still hardcodes torch.polar + on-device complex
    indexing + use_real=False,
  • transformer_z_image.py likewise (torch.polar, view_as_complex).

This PR extends the established Neuron pattern instead of introducing a new one:

  1. Qwen-Image 2.1 (transformer_qwenimage21.py): port ROPE_PER_DEVICE
    and reuse the real-valued apply. The per-request frequency selection in
    QwenImage21Rope.forward keeps the complex table on CPU, indexes there, and
    crosses the device boundary as reals — see the parity section below.
  2. Z-Image (transformer_z_image.py): same port — the attention
    processor's RoPE helper moves to module level as apply_rotary_emb_zimage
    (dtype-driven dual path), and RopeEmbedder keeps its tables on CPU,
    indexes there, and returns view_as_real pairs. The # Copied from blocks
    in controlnet_z_image.py pick this up via make fix-copies, importing the
    shared helper/constant from transformer_z_image (same pattern as
    controlnet_qwenimage.py).
  3. Qwen-Image (transformer_qwenimage.py): register "npu" in the existing
    ROPE_PER_DEVICE, reusing the Neuron real path unchanged.

Scope note, deliberately: this does not flip the CUDA default to the
real-valued path, so #12668's torch.compile aspect on CUDA is mitigated for
complex-less backends but not closed. Making the real path the universal default
(numerically equivalent, inductor-friendly) is a reasonable follow-up; kept out
here to keep the diff minimal and behavior-neutral on the reference path.

Related: #12668 (does not fully fix — see above), apple/coremltools#2258.

Numerical parity (measured, CPU, torch 2.14.0)

Two real-valued formulations are available at the device boundary:

  • Neuron's existing one: torch.angle(table) on CPU → cos/sin on device.
    Measured to add up to 1.19e-07 error from the atan2 round-trip.
  • This PR, for the per-request index path in Qwen-Image 2.1: index the
    complex table on CPU, then torch.view_as_real. Since torch.polar(1, θ)
    stores exactly (cos θ, sin θ), the extracted operands are bit-identical
    to the complex path's — the selected frequency rows compare at 0 ulp.

The rotation arithmetic itself expands identically in both paths:

(xr + i·xi)(cos + i·sin) = (xr·cos − xi·sin) + i(xr·sin + xi·cos)

but torch's complex64 multiply kernel contracts FMAs differently from explicit
real multiply-add, so outputs differ by ≤ 0.82 ulp of the fp32 compute domain
(both paths upcast with x.float()). Measured: max 4.77e-07 across
fp64/fp32/bf16 inputs at shapes up to [2, 4099, 8, 128]. For reference,
upstream's own use_real=True path shows the same ≤ 0.82 ulp relationship to
the complex path — the candidate is exactly as exact as the real-valued RoPE
diffusers already ships for flux et al.

Happy to switch to torch.angle for consistency with #13718 if reviewers prefer
one formulation; the parity test covers either choice.

Hardware validation

  • Ascend 910B2 (torch 2.10.0 / torch_npu 2.10.0.post4): an equivalent of this
    patch has been run against Qwen-Image-2.1 BF16 — RoPE outputs match the
    complex reference to ≤ 1 ulp, and the pipeline produces expected outputs at
    1024×1024 / 40 steps.
  • Parity harness and full log attached below (pure CPU, imports the pinned
    upstream module as reference — no re-implementation): covers the apply math
    (fp64/fp32/bf16 × 3 shapes), the full QwenImage21Rope.forward selection
    (t2i / two-image edit / trailing text), and the torch.angle round-trip
    comparison. All rows PASS at ≤ 1 ulp; the forward selection is bit-exact.

Before submitting

  • Did you use an AI agent (Claude Code, Codex, Cursor, etc.) to help with this PR? If so:
    • Did you read the Coding with AI agents guide?
    • Did you run the self-review skill on the diff?
    • Did you share the final self-review notes in the PR description or a comment? → see "Self-review notes" below
  • Did you read the contributor guideline?
  • Did you read our philosophy doc?
  • Was this discussed/approved via a GitHub issue? → Complex numbers in Qwen/Z-Image Image pipeline incompatible with torch.compile #12668 (stale since 2026-01; this PR revives it for the complex-less-backend half)
  • Did you make sure to update the documentation with your changes? → docstrings only; behavior-neutral on existing backends
  • Did you write any new necessary tests? → three RoPE real/complex parity test classes (CPU, monkeypatching the backend list)
  • Are you the author (or part of the team) of the model/pipeline? → No

Self-review notes (final round)

  • Propagation checked: controlnet_z_image.py copy-syncs
    ZSingleStreamAttnProcessor/RopeEmbedder; propagated with
    python utils/check_copies.py --fix_and_overwrite (no hand-edited copied
    blocks) and check_copies passes.
  • Invariant preserved: upstream moves the frequency table to the indexing
    device on every call; the fallback keeps that shape (target = CPU for
    complex-less backends), so a table cached for a complex-capable device is
    moved back if a later call targets a complex-less one (offload transitions).
  • Verified: ruff check, ruff format --check, check_copies, the three new
    parity tests, and the full transformer test files for the three models —
    pass/fail sets identical to the base commit in this environment (the 32
    pre-existing failures are group-offloading/disk tests, unrelated to RoPE).
  • Intentionally not fixed: CUDA default stays on the complex path (see scope
    note); apply_rotary_emb_qwen_neuron is reused rather than renamed to a
    generic *_real — a rename touches [core] Support tensor parallelism for model inference (CUDA, Neuron) #13718's public-ish surface for cosmetic
    gain, left to reviewer preference.
  • Intentionally not fixed: "mps" dispatch entry. PyTorch MPS also lacks full
    complex64, so it is a natural one-line addition once someone with the
    hardware confirms it.
  • Risk called out: QwenImage21Rope.forward's CPU-side indexing mirrors the
    current upstream logic 1:1 (including the Python-loop position bookkeeping);
    the diff only changes where it runs and the dtype crossing the boundary.
  • The dispatch lives at the single application site in
    _qwenimage21_prepare_qkv, which all KV-cache modes also route through.

Who can review?

@sayakpaul (engaged on #12668), @dg845 (pipelines/models; prior Qwen-Image RoPE
work), @JingyaHuang (author of #13718 — this PR extends your Neuron pattern).


Appendix: parity harness and log

upstream_pr_rope_parity.py (pure CPU; imports the pinned upstream module as reference)
"""RoPE complex-vs-real parity evidence for the upstream diffusers PR (transient, do not commit).

Reference: upstream diffusers @ 6256aa7666cedd47443adc8f82da9a10e110b09c
(src/diffusers/models/transformers/transformer_qwenimage21.py), imported from the
pinned source tree — not re-implemented here. Candidate: this repo's
_rope_patch.py, the exact code the PR proposes to upstream.

Criterion: both paths cast to fp32 before computing (upstream does `x.float()`),
so bitwise equality is NOT the bar — torch's complex64 multiply kernel contracts
FMAs differently from explicit real arithmetic. The bar is <= 1 ulp of the fp32
compute domain (bf16 outputs: <= 1 ulp of bf16), the same relationship upstream's
own `use_real=True` path has to the complex path (shown in T1b).

Run (pure CPU, no NPU/GPU needed):

    uv run --with torch --with /tmp/diffusers-6256aa7.tar.gz \
      python upstream_pr_rope_parity.py
"""
import os
import sys

import torch

sys.path.insert(0, os.path.dirname(os.path.abspath(__file__)))

from diffusers.models.transformers.transformer_qwenimage21 import (  # noqa: E402
    QwenImage21Rope,
    apply_rotary_emb_qwen,
)

import _rope_patch  # noqa: E402  candidate: repo's real-domain implementation

THETA = 10000
AXES_DIM = [16, 56, 56]  # upstream default axes_dims_rope
HEAD_DIM = sum(AXES_DIM)  # 128
FP32_EPS = 2**-23


def ulp_report(name, ref, cand):
    """PASS iff max|diff| <= 1 ulp at the largest magnitude of ref (compute domain)."""
    md = (ref.double() - cand.double()).abs().max().item()
    peak = ref.double().abs().max().item()
    eps = FP32_EPS if ref.dtype != torch.bfloat16 else 2**-8
    ulps = md / (eps * max(peak, 1e-30))
    ok = ulps <= 1.0
    print(f"  {name:<40} dtype={str(ref.dtype):<15} max|diff|={md:.3e} = {ulps:.2f} ulp  {'PASS' if ok else 'FAIL'}")
    return ok


def make_freqs(S, D, seed):
    g = torch.Generator().manual_seed(seed)
    angles = (torch.rand(S, D // 2, generator=g, dtype=torch.float64) * 2 - 1) * torch.pi
    return torch.polar(torch.ones_like(angles), angles).to(torch.complex64), angles


def t1_apply_parity():
    print("T1a complex apply (upstream) vs candidate real apply (_rope_patch.py)")
    ok = True
    for dtype in (torch.float64, torch.float32, torch.bfloat16):
        for shape in [(2, 1101, 8, HEAD_DIM), (1, 1, 1, HEAD_DIM), (1, 4099, 8, HEAD_DIM)]:
            g = torch.Generator().manual_seed(42)
            x = torch.randn(*shape, generator=g, dtype=torch.float64).to(dtype)
            freqs, _ = make_freqs(shape[1], shape[3], 43)
            ref = apply_rotary_emb_qwen(x, freqs, use_real=False)
            cand = _rope_patch._apply_rotary_real(x, torch.view_as_real(freqs), use_real=False)
            ok &= ulp_report(f"x{list(shape)}", ref, cand)
    return ok


def t1b_upstream_real_path():
    """The same 1-ulp relationship holds for upstream's own use_real=True path."""
    print("T1b complex apply vs upstream's OWN real path (use_real=True, flux convention)")
    ok = True
    for dtype in (torch.float32, torch.bfloat16):
        g = torch.Generator().manual_seed(44)
        x = torch.randn(2, 1101, 8, HEAD_DIM, generator=g, dtype=torch.float64).to(dtype)
        freqs, angles = make_freqs(1101, HEAD_DIM, 45)
        ref = apply_rotary_emb_qwen(x, freqs, use_real=False)
        cos = torch.cos(angles).repeat_interleave(2, dim=-1).float()  # [S, D], pair-shared angles
        sin = torch.sin(angles).repeat_interleave(2, dim=-1).float()
        # upstream's real path expects the flux layout [B, H, S, D] (cos gets [None, None])
        cand = apply_rotary_emb_qwen(x.transpose(1, 2), (cos, sin), use_real=True, use_real_unbind_dim=-1).transpose(1, 2)
        ok &= ulp_report(f"dtype={dtype}", ref, cand)
    return ok


def t2_operand_identity():
    print("T2  view_as_real(polar(1, th)) == (cos th, sin th) — operands bit-exact")
    freqs, angles = make_freqs(1101, 64, 46)
    pair = torch.view_as_real(freqs)
    exact = torch.equal(pair[..., 0], freqs.real) and torch.equal(pair[..., 1], freqs.imag)
    print(f"  bit-identical: {exact}  {'PASS' if exact else 'FAIL'}")
    return exact


def build_mask(img_shapes, leading_text, trailing_text):
    n_img = sum(h * w for _, h, w in img_shapes)
    total = leading_text + n_img + trailing_text
    mask = torch.zeros(total, dtype=torch.bool)
    cursor = leading_text
    for _, h, w in img_shapes:
        mask[cursor : cursor + h * w] = True
        cursor += h * w
    return mask  # 1D [S], as upstream forward expects


def t3_forward_parity():
    print("T3  QwenImage21Rope.forward + apply end-to-end (t2i / edit / trailing text)")
    ok = True
    cases = {
        "t2i": ([(1, 32, 32)], 77, 0),
        "edit, 2 images": ([(1, 32, 32), (1, 16, 48)], 77, 0),
        "trailing text": ([(1, 32, 32)], 77, 52),
    }
    for name, (img_shapes, lead, trail) in cases.items():
        mask = build_mask(img_shapes, lead, trail)
        ref = QwenImage21Rope(theta=THETA, axes_dim=AXES_DIM).forward(img_shapes, mask, torch.device("cpu"))
        cand = _rope_patch._rope_forward_real(
            QwenImage21Rope(theta=THETA, axes_dim=AXES_DIM), img_shapes, mask, torch.device("cpu")
        )
        assert ref.dtype == torch.complex64
        ok &= ulp_report(f"freqs {name}", torch.view_as_real(ref).flatten(1), cand.flatten(1))
        for dtype in (torch.float32, torch.bfloat16):
            g = torch.Generator().manual_seed(47)
            x = torch.randn(1, mask.shape[0], 8, HEAD_DIM, generator=g, dtype=torch.float64).to(dtype)
            r = apply_rotary_emb_qwen(x, ref, use_real=False)
            c = _rope_patch._apply_rotary_real(x, cand, use_real=False)
            ok &= ulp_report(f"apply {name} {dtype}", r, c)
    return ok


def t4_angle_roundtrip():
    print("T4  torch.angle round-trip vs view_as_real (design note: why not Neuron's formulation)")
    freqs, _ = make_freqs(1101, 64, 48)
    ang = torch.angle(freqs)
    rec = torch.polar(torch.ones_like(ang), ang)
    md = (torch.view_as_real(freqs) - torch.view_as_real(rec)).abs().max().item()
    print(f"  angle->cos/sin round-trip adds up to {md:.2e} error; view_as_real adds none (T2)")
    return True  # informational


def main():
    import diffusers

    print(f"torch {torch.__version__} | diffusers {diffusers.__version__} | {sys.platform}")
    print("reference: diffusers @ 6256aa7666cedd47443adc8f82da9a10e110b09c (transformer_qwenimage21.py)")
    print("candidate: this repo's deployments/qwen-image-2.1/_rope_patch.py")
    print("criterion: <= 1 ulp of the fp32 compute domain (upstream casts with x.float())\n")
    results = [t2_operand_identity(), t1_apply_parity(), t1b_upstream_real_path(), t3_forward_parity(), t4_angle_roundtrip()]
    print("\n" + ("ALL PASS" if all(results) else "FAILURES PRESENT"))
    return 0 if all(results) else 1


if __name__ == "__main__":
    sys.exit(main())
Full log (all PASS)
torch 2.14.0 | diffusers 0.41.0.dev0 | darwin
reference: diffusers @ 6256aa7666cedd47443adc8f82da9a10e110b09c (transformer_qwenimage21.py)
candidate: this repo's deployments/qwen-image-2.1/_rope_patch.py
criterion: <= 1 ulp of the fp32 compute domain (upstream casts with x.float())

T2  view_as_real(polar(1, th)) == (cos th, sin th) — operands bit-exact
  bit-identical: True  PASS
T1a complex apply (upstream) vs candidate real apply (_rope_patch.py)
  x[2, 1101, 8, 128]                       dtype=torch.float64   max|diff|=4.768e-07 = 0.76 ulp  PASS
  x[1, 1, 1, 128]                          dtype=torch.float64   max|diff|=1.192e-07 = 0.36 ulp  PASS
  x[1, 4099, 8, 128]                       dtype=torch.float64   max|diff|=4.768e-07 = 0.76 ulp  PASS
  x[2, 1101, 8, 128]                       dtype=torch.float32   max|diff|=4.768e-07 = 0.76 ulp  PASS
  x[1, 1, 1, 128]                          dtype=torch.float32   max|diff|=1.192e-07 = 0.36 ulp  PASS
  x[1, 4099, 8, 128]                       dtype=torch.float32   max|diff|=4.768e-07 = 0.76 ulp  PASS
  x[2, 1101, 8, 128]                       dtype=torch.bfloat16  max|diff|=1.562e-02 = 0.76 ulp  PASS
  x[1, 1, 1, 128]                          dtype=torch.bfloat16  max|diff|=0.000e+00 = 0.00 ulp  PASS
  x[1, 4099, 8, 128]                       dtype=torch.bfloat16  max|diff|=1.562e-02 = 0.76 ulp  PASS
T1b complex apply vs upstream's OWN real path (use_real=True, flux convention)
  dtype=torch.float32                      dtype=torch.float32   max|diff|=4.768e-07 = 0.82 ulp  PASS
  dtype=torch.bfloat16                     dtype=torch.bfloat16  max|diff|=7.812e-03 = 0.41 ulp  PASS
T3  QwenImage21Rope.forward + apply end-to-end (t2i / edit / trailing text)
  freqs t2i                                dtype=torch.float32   max|diff|=0.000e+00 = 0.00 ulp  PASS
  apply t2i torch.float32                  dtype=torch.float32   max|diff|=4.768e-07 = 0.82 ulp  PASS
  apply t2i torch.bfloat16                 dtype=torch.bfloat16  max|diff|=3.906e-03 = 0.21 ulp  PASS
  freqs edit, 2 images                     dtype=torch.float32   max|diff|=0.000e+00 = 0.00 ulp  PASS
  apply edit, 2 images torch.float32       dtype=torch.float32   max|diff|=4.768e-07 = 0.80 ulp  PASS
  apply edit, 2 images torch.bfloat16      dtype=torch.bfloat16  max|diff|=7.812e-03 = 0.40 ulp  PASS
  freqs trailing text                      dtype=torch.float32   max|diff|=0.000e+00 = 0.00 ulp  PASS
  apply trailing text torch.float32        dtype=torch.float32   max|diff|=4.768e-07 = 0.82 ulp  PASS
  apply trailing text torch.bfloat16       dtype=torch.bfloat16  max|diff|=3.906e-03 = 0.21 ulp  PASS
T4  torch.angle round-trip vs view_as_real (design note: why not Neuron's formulation)
  angle->cos/sin round-trip adds up to 1.19e-07 error; view_as_real adds none (T2)

ALL PASS

Ascend NPU (like AWS Neuron) cannot execute torch's complex64 ops, so
Qwen-Image, Qwen-Image 2.1 and Z-Image crash on it (huggingface#12668).

- Qwen-Image: register "npu" in the existing ROPE_PER_DEVICE dispatch
  next to Neuron (huggingface#13718); factor the backend list into
  COMPLEX_LESS_ROPE_BACKENDS.
- Qwen-Image 2.1: add apply_rotary_emb_qwen_real, picked by carrier
  dtype; QwenImage21Rope keeps its complex64 table on CPU, runs the
  per-request selection there, and returns view_as_real pairs whose
  operands are bit-exact to the complex path's.
- Z-Image: hoist the processor RoPE helper to module level with the
  same dtype-driven dual path; RopeEmbedder uses the same CPU-selection
  scheme. controlnet_z_image re-imports the shared pieces (fix-copies).

The complex path is untouched and stays the default everywhere. New
tests pin the real fallback against it on CPU: frequency selection is
bit-exact, the rotation matches within float32 rounding.
@github-actions github-actions Bot added models tests size/L PR with diff > 200 LOC labels Sep 29, 2026

This branch has not been deployed

No deployments
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

models size/L PR with diff > 200 LOC tests

Projects

None yet

Development

Successfully merging this pull request may close these issues.

1 participant