From 227f7068f3b62467f0d1d142519506ba9ee0f737 Mon Sep 17 00:00:00 2001 From: GokayAI <60583610+gokay-ai@users.noreply.github.com> Date: Sun, 4 Oct 2026 09:46:51 +0300 Subject: [PATCH] Keep Stable Audio 3 RoPE angles in fp32 bf16 and fp16 loads cast the persistent inv_freq buffer and built positions in that dtype, so neighbouring latents shared a rotary angle. Compute the outer product in fp32 and keep rotary_pos_emb in fp32 on from_pretrained. Signed-off-by: GokayAI <60583610+gokay-ai@users.noreply.github.com> --- .../transformers/transformer_stable_audio3.py | 9 +++++++-- .../test_models_transformer_stable_audio3.py | 18 ++++++++++++++++++ 2 files changed, 25 insertions(+), 2 deletions(-) diff --git a/src/diffusers/models/transformers/transformer_stable_audio3.py b/src/diffusers/models/transformers/transformer_stable_audio3.py index aab80b744d8f..d80654458b92 100644 --- a/src/diffusers/models/transformers/transformer_stable_audio3.py +++ b/src/diffusers/models/transformers/transformer_stable_audio3.py @@ -103,8 +103,10 @@ def __init__(self, dim: int, base: int = 10000): self.register_buffer("inv_freq", inv_freq, persistent=True) def forward(self, seq_len: int, device: torch.device) -> torch.Tensor: - t = torch.arange(seq_len, device=device, dtype=self.inv_freq.dtype) - freqs = torch.outer(t, self.inv_freq) + # bf16/fp16 cannot represent large integer positions exactly, so build the + # angles in fp32 even when `inv_freq` was cast with the rest of the model. + t = torch.arange(seq_len, device=device, dtype=torch.float32) + freqs = torch.outer(t, self.inv_freq.float()) return torch.cat((freqs, freqs), dim=-1) # (seq_len, rot_dim) @@ -434,6 +436,9 @@ class StableAudio3DiTModel(ModelMixin, ConfigMixin, AttentionMixin): _supports_gradient_checkpointing = True _no_split_modules = ["StableAudio3DiTBlock"] _repeated_blocks = ["StableAudio3DiTBlock"] + # `inv_freq` is a persistent buffer. Keep it fp32 so `from_pretrained(torch_dtype=...)` + # does not round the rotary frequencies. + _keep_in_fp32_modules = ["rotary_pos_emb"] @register_to_config def __init__( diff --git a/tests/models/transformers/test_models_transformer_stable_audio3.py b/tests/models/transformers/test_models_transformer_stable_audio3.py index 4b770d95361c..9eab35d5e000 100644 --- a/tests/models/transformers/test_models_transformer_stable_audio3.py +++ b/tests/models/transformers/test_models_transformer_stable_audio3.py @@ -215,6 +215,24 @@ def test_memory_tokens_present(self): self.assertEqual(self.model.memory_tokens.shape[0], TINY_CFG["num_memory_tokens"]) self.assertEqual(self.model.memory_tokens.shape[1], TINY_CFG["embed_dim"]) + def test_low_precision_rope_keeps_distinct_positions(self): + """Loading with bf16/fp16 must not collapse RoPE positions (#14934).""" + import math + import tempfile + + with tempfile.TemporaryDirectory() as tmp: + self.model.save_pretrained(tmp) + ref = StableAudio3DiTModel.from_pretrained(tmp).rotary_pos_emb + seq_len = 64 + math.ceil(120 * 44100 / 4096) + ref_freqs = ref(seq_len, "cpu") + for dtype in (torch.float16, torch.bfloat16): + rope = StableAudio3DiTModel.from_pretrained(tmp, torch_dtype=dtype).rotary_pos_emb + self.assertEqual(rope.inv_freq.dtype, torch.float32) + freqs = rope(seq_len, "cpu") + self.assertEqual(freqs.unique(dim=0).shape[0], seq_len) + err = (freqs.double() - ref_freqs.double() + math.pi).remainder(2 * math.pi) - math.pi + self.assertLess(err.abs().max().item(), 1e-4) + # ────────────────────────────────────────────────────────────────────────────── # Structural parity with the released SA3 Medium checkpoint