Skip to content

QwenImage21: model.to(bf16) (as in train_dreambooth_lora_qwenimage21.py) rounds the timestep-embedding freqs buffer, so LoRA training sees a different embedding than inference #14935

Description

@tritsystem

Describe the bug

QwenImage21TemporalTimesteps stores its sinusoid frequencies as a non-persistent freqs buffer and multiplies them by 1000 * t, so the arguments reach about 1000 rad. from_pretrained(torch_dtype=torch.bfloat16) keeps the buffer in fp32, so inference is exact. But a later model.to(dtype=torch.bfloat16) rounds freqs to bf16 and shifts the high-frequency channels of the timestep embedding by up to about 1.4 rad.

The official LoRA script does exactly that. examples/dreambooth/train_dreambooth_lora_qwenimage21.py loads the transformer with torch_dtype=weight_dtype and then calls transformer.to(device=accelerator.device, dtype=weight_dtype). So with --mixed_precision bf16, LoRAs are trained against a different timestep embedding than the one the inference pipeline uses.

The timestep.float() in forward doesn't help, because the buffer itself has already been rounded.

Reproduction

CPU, tiny config, same load/cast sequence as the training script:

import math, torch
from diffusers import QwenImage21Transformer2DModel

torch.manual_seed(0)
QwenImage21Transformer2DModel(
    num_layers=1, attention_head_dim=16, num_attention_heads=2, context_in_dim=32,
    in_channels=8, out_channels=8, axes_dims_rope=(4, 6, 6),
).save_pretrained("tiny-qwenimage21")

t = torch.linspace(0, 1, 1001)
half = 128
freqs = torch.exp(-math.log(10000) * torch.arange(half, dtype=torch.float32) / half)
args = (1000 * t)[:, None] * freqs[None]
reference = torch.cat([torch.cos(args), torch.sin(args)], dim=-1)  # exact fp32 embedding

inference = QwenImage21Transformer2DModel.from_pretrained("tiny-qwenimage21", torch_dtype=torch.bfloat16)
training = QwenImage21Transformer2DModel.from_pretrained("tiny-qwenimage21", torch_dtype=torch.bfloat16)
training.to(dtype=torch.bfloat16)  # what train_dreambooth_lora_qwenimage21.py does

for name, model in (("inference (from_pretrained)", inference), ("training (.to(dtype=bf16))", training)):
    proj = model.time_text_embed.time_proj
    emb = proj(t)
    angle = torch.atan2(emb[:, half:].double(), emb[:, :half].double())
    ref_angle = torch.atan2(reference[:, half:].double(), reference[:, :half].double())
    err = ((angle - ref_angle + math.pi) % (2 * math.pi) - math.pi).abs()
    print(f"{name:28s} freqs={proj.freqs.dtype}, max angle error {err.max():.3f} rad, "
          f"{(err > 0.1).float().mean() * 100:.1f}% of angles off by > 0.1 rad, "
          f"max |embedding diff| {(emb.float() - reference).abs().max():.3f}")

Output on main (8b33bfc):

inference (from_pretrained)  freqs=torch.float32, max angle error 0.000 rad, 0.0% of angles off by > 0.1 rad, max |embedding diff| 0.000
training (.to(dtype=bf16))   freqs=torch.bfloat16, max angle error 1.388 rad, 13.9% of angles off by > 0.1 rad, max |embedding diff| 1.272

Suggested fix

Compute the frequencies in fp32 inside forward, as the generic Timesteps / get_timestep_embedding already does, instead of keeping them in a buffer. The buffer is non-persistent, so no checkpoint key changes:

     def __init__(self, timestep_dim: int, max_period: int = 10000, time_factor: float = 1000.0):
         super().__init__()
         self.timestep_dim = timestep_dim
         self.time_factor = time_factor
-        half = timestep_dim // 2
-        freqs = torch.exp(-math.log(max_period) * torch.arange(start=0, end=half, dtype=torch.float32) / half)
-        self.register_buffer("freqs", freqs, persistent=False)
+        self.max_period = max_period

     def forward(self, timestep: torch.Tensor) -> torch.Tensor:
         timestep = self.time_factor * timestep.float()
-        args = timestep[:, None] * self.freqs[None].to(timestep.device)
+        half = self.timestep_dim // 2
+        exponent = torch.arange(start=0, end=half, dtype=torch.float32, device=timestep.device) / half
+        freqs = torch.exp(-math.log(self.max_period) * exponent)
+        args = timestep[:, None] * freqs[None]

What I checked with it (CPU):

  • For fp32 from_pretrained, the training sequence above, and from_pretrained().to(torch.bfloat16), the timestep embedding is bit-identical to the exact fp32 reference. Existing fp32 and inference outputs are unchanged.
  • tests/models/transformers/test_models_transformer_qwenimage21.py and tests/pipelines/qwenimage21/test_qwenimage21.py give 85 passed / 62 skipped / 1 xfailed both before and after.

I didn't measure the effect on trained LoRA quality, so I can't say how much it matters in practice. The mismatch itself is deterministic.

Happy to open a PR with this. The same buffer-cast pattern also affects StableAudio3, where it's worse because the positions themselves get rounded on the default load path (#14934).

System Info

diffusers main @ 8b33bfc, torch 2.11, CPU.

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

Metadata

Metadata

Assignees

No one assigned

    Type

    No type

    Projects

    No projects

      Milestone

      No milestone

      Relationships

      None yet

      Development

      No branches or pull requests

      Issue actions