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.
Describe the bug
QwenImage21TemporalTimestepsstores its sinusoid frequencies as a non-persistentfreqsbuffer and multiplies them by1000 * 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 latermodel.to(dtype=torch.bfloat16)roundsfreqsto 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.pyloads the transformer withtorch_dtype=weight_dtypeand then callstransformer.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()inforwarddoesn't help, because the buffer itself has already been rounded.Reproduction
CPU, tiny config, same load/cast sequence as the training script:
Output on
main(8b33bfc):Suggested fix
Compute the frequencies in fp32 inside
forward, as the genericTimesteps/get_timestep_embeddingalready 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):
from_pretrained, the training sequence above, andfrom_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.pyandtests/pipelines/qwenimage21/test_qwenimage21.pygive 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.