From b8e65f2c712421ee1f8d12164302a3d03f9609bc Mon Sep 17 00:00:00 2001 From: Igor Shovkun Date: Mon, 21 Sep 2026 21:31:54 -0700 Subject: [PATCH 1/7] [None][fix] Compute the Cosmos3 rotary table with a broadcast multiply, not a K=1 matmul The 3D mRoPE angle table was built as inv_freq @ position_ids, an fp32 batched matmul with inner dimension 1. That is an outer product, but as a GEMM it runs in TF32 wherever fp32 GEMMs are allowed to (NGC PyTorch images set TORCH_ALLOW_TF32_CUBLAS_OVERRIDE=1), and TF32 cannot represent positions above 2048. Cosmos3 positions reach ~10k, so on Blackwell the rotary phase of late tokens was off by up to 5 rad and late video frames smeared. Use an elementwise multiply instead. Same arithmetic, no GEMM, immune to matmul precision settings; output is bit-identical on stacks that were already correct. Signed-off-by: Igor Shovkun --- .../models/cosmos3/transformer_cosmos3.py | 6 +- .../test_lists/test-db/l0_b200.yml | 1 + .../visual_gen/test_cosmos3_rope_precision.py | 88 +++++++++++++++++++ 3 files changed, 93 insertions(+), 2 deletions(-) create mode 100644 tests/unittest/_torch/visual_gen/test_cosmos3_rope_precision.py diff --git a/tensorrt_llm/_torch/visual_gen/models/cosmos3/transformer_cosmos3.py b/tensorrt_llm/_torch/visual_gen/models/cosmos3/transformer_cosmos3.py index 7e5d24ae7a5f..6891b0212ebc 100644 --- a/tensorrt_llm/_torch/visual_gen/models/cosmos3/transformer_cosmos3.py +++ b/tensorrt_llm/_torch/visual_gen/models/cosmos3/transformer_cosmos3.py @@ -996,9 +996,11 @@ def forward(self, x, position_ids): .expand(3, position_ids.shape[1], -1, 1) .to(x.device) ) - position_ids_expanded = position_ids[:, :, None, :].float() # shape (3, bs, 1, positions) + position_ids_expanded = position_ids[:, :, None, :] # shape (3, bs, 1, positions) - freqs = (inv_freq_expanded.float() @ position_ids_expanded.float()).transpose(2, 3) + # Elementwise, not a K=1 matmul: a GEMM under TF32 (default in NGC PyTorch + # images) rounds positions above 2048 and skews the phase of late tokens. + freqs = (inv_freq_expanded * position_ids_expanded).transpose(2, 3) freqs = self.apply_interleaved_mrope(freqs, self.mrope_section) emb = torch.cat((freqs, freqs), dim=-1) cos = emb.cos() * self.attention_scaling diff --git a/tests/integration/test_lists/test-db/l0_b200.yml b/tests/integration/test_lists/test-db/l0_b200.yml index 6d5edcb10f32..ab86b19a81a4 100644 --- a/tests/integration/test_lists/test-db/l0_b200.yml +++ b/tests/integration/test_lists/test-db/l0_b200.yml @@ -339,6 +339,7 @@ l0_b200: - unittest/_torch/visual_gen/test_wan_transformer.py - unittest/_torch/visual_gen/test_cosmos3_action.py - unittest/_torch/visual_gen/test_cosmos3_transformer.py + - unittest/_torch/visual_gen/test_cosmos3_rope_precision.py - unittest/_torch/visual_gen/test_cosmos3_pipeline.py - unittest/_torch/visual_gen/test_cosmos3_distilled.py - unittest/_torch/visual_gen/test_cosmos3_transfer.py diff --git a/tests/unittest/_torch/visual_gen/test_cosmos3_rope_precision.py b/tests/unittest/_torch/visual_gen/test_cosmos3_rope_precision.py new file mode 100644 index 000000000000..8760974ad458 --- /dev/null +++ b/tests/unittest/_torch/visual_gen/test_cosmos3_rope_precision.py @@ -0,0 +1,88 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +"""The Cosmos3 rotary table must stay fp32-accurate when fp32 GEMMs run in TF32. + +NGC PyTorch images set TORCH_ALLOW_TF32_CUBLAS_OVERRIDE=1. Positions above 2048 +do not fit TF32's 10-bit mantissa, so a rotary table built as a K=1 matmul is +off by radians for late tokens on GPUs where cuBLAS picks a TF32 kernel for +that shape (Blackwell). The table is an outer product and must not go through +a GEMM. +""" + +from types import SimpleNamespace + +import pytest +import torch + +from tensorrt_llm._torch.visual_gen.models.cosmos3.transformer_cosmos3 import ( + Qwen3VLTextRotaryEmbedding, +) + +pytestmark = pytest.mark.cosmos3 + +HEAD_DIM = 128 +ROPE_AXES = [24, 20, 20] # sums to HEAD_DIM // 2, the Cosmos3-Nano layout +# Text tokens (up to 4096) + margin + fps-scaled vision positions land well +# above 2048, the largest integer TF32 represents exactly. cuBLAS picks a TF32 +# kernel for the K=1 GEMM only at some problem sizes, so sweep the sizes a +# Cosmos3 request actually produces. +SEQUENCE_LENGTHS = [4096, 6240, 8192, 10336, 16384] + + +def _rotary() -> Qwen3VLTextRotaryEmbedding: + pretrained = SimpleNamespace( + rope_theta=1_000_000.0, + head_dim=HEAD_DIM, + max_position_embeddings=262_144, + rope_axes_dim=ROPE_AXES, + rope_scaling=None, + ) + return Qwen3VLTextRotaryEmbedding(SimpleNamespace(pretrained_config=pretrained)) + + +def _reference_cos_sin(rotary: Qwen3VLTextRotaryEmbedding, position_ids: torch.Tensor): + """Same math as the module, in float64 and without any matmul.""" + inv = rotary.inv_freq.double()[None, None, :, None] # [1, 1, D/2, 1] + pos = position_ids.double()[:, :, None, :] # [3, B, 1, N] + freqs = (inv * pos).transpose(2, 3) # [3, B, N, D/2] + freqs = rotary.apply_interleaved_mrope(freqs.clone(), rotary.mrope_section) + emb = torch.cat((freqs, freqs), dim=-1) + return emb.cos(), emb.sin() + + +@pytest.mark.skipif(not torch.cuda.is_available(), reason="needs a GPU") +@pytest.mark.parametrize("allow_tf32", [False, True]) +@pytest.mark.parametrize("seq_len", SEQUENCE_LENGTHS) +def test_rotary_table_matches_fp64_under_tf32(allow_tf32: bool, seq_len: int): + device = torch.device("cuda") + saved = torch.backends.cuda.matmul.allow_tf32 + torch.backends.cuda.matmul.allow_tf32 = allow_tf32 + try: + rotary = _rotary().to(device) + base = torch.arange(seq_len, dtype=torch.float32, device=device) + # 3D mRoPE ids: temporal axis fps-scaled (fractional), spatial axes integer + position_ids = torch.stack([base * 24.0 / 10.0, base, base * 0.5], dim=0)[:, None, :] + probe = torch.empty(0, dtype=torch.float32, device=device) + + cos, sin = rotary(probe, position_ids) + ref_cos, ref_sin = _reference_cos_sin(rotary, position_ids) + + # fp32 evaluates angles of ~4e4 rad with ~6e-8 relative precision, so a + # few 1e-3 absolute is the honest fp32 floor; TF32 breakage is O(1). + tol = 5e-3 + assert (cos.double() - ref_cos).abs().max().item() < tol + assert (sin.double() - ref_sin).abs().max().item() < tol + finally: + torch.backends.cuda.matmul.allow_tf32 = saved From 3d315f55352164727f3907f4f3c3e86fdc390eb4 Mon Sep 17 00:00:00 2001 From: Igor Shovkun Date: Tue, 22 Sep 2026 10:32:46 -0700 Subject: [PATCH 2/7] Move the rotary TF32 test into test_cosmos3_transformer.py Keeps the check next to the other Cosmos3 transformer tests instead of a new file, so no test-list change is needed. Signed-off-by: Igor Shovkun --- .../test_lists/test-db/l0_b200.yml | 1 - .../visual_gen/test_cosmos3_rope_precision.py | 88 ------------------- .../visual_gen/test_cosmos3_transformer.py | 66 ++++++++++++++ 3 files changed, 66 insertions(+), 89 deletions(-) delete mode 100644 tests/unittest/_torch/visual_gen/test_cosmos3_rope_precision.py diff --git a/tests/integration/test_lists/test-db/l0_b200.yml b/tests/integration/test_lists/test-db/l0_b200.yml index ab86b19a81a4..6d5edcb10f32 100644 --- a/tests/integration/test_lists/test-db/l0_b200.yml +++ b/tests/integration/test_lists/test-db/l0_b200.yml @@ -339,7 +339,6 @@ l0_b200: - unittest/_torch/visual_gen/test_wan_transformer.py - unittest/_torch/visual_gen/test_cosmos3_action.py - unittest/_torch/visual_gen/test_cosmos3_transformer.py - - unittest/_torch/visual_gen/test_cosmos3_rope_precision.py - unittest/_torch/visual_gen/test_cosmos3_pipeline.py - unittest/_torch/visual_gen/test_cosmos3_distilled.py - unittest/_torch/visual_gen/test_cosmos3_transfer.py diff --git a/tests/unittest/_torch/visual_gen/test_cosmos3_rope_precision.py b/tests/unittest/_torch/visual_gen/test_cosmos3_rope_precision.py deleted file mode 100644 index 8760974ad458..000000000000 --- a/tests/unittest/_torch/visual_gen/test_cosmos3_rope_precision.py +++ /dev/null @@ -1,88 +0,0 @@ -# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. -# SPDX-License-Identifier: Apache-2.0 -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. -"""The Cosmos3 rotary table must stay fp32-accurate when fp32 GEMMs run in TF32. - -NGC PyTorch images set TORCH_ALLOW_TF32_CUBLAS_OVERRIDE=1. Positions above 2048 -do not fit TF32's 10-bit mantissa, so a rotary table built as a K=1 matmul is -off by radians for late tokens on GPUs where cuBLAS picks a TF32 kernel for -that shape (Blackwell). The table is an outer product and must not go through -a GEMM. -""" - -from types import SimpleNamespace - -import pytest -import torch - -from tensorrt_llm._torch.visual_gen.models.cosmos3.transformer_cosmos3 import ( - Qwen3VLTextRotaryEmbedding, -) - -pytestmark = pytest.mark.cosmos3 - -HEAD_DIM = 128 -ROPE_AXES = [24, 20, 20] # sums to HEAD_DIM // 2, the Cosmos3-Nano layout -# Text tokens (up to 4096) + margin + fps-scaled vision positions land well -# above 2048, the largest integer TF32 represents exactly. cuBLAS picks a TF32 -# kernel for the K=1 GEMM only at some problem sizes, so sweep the sizes a -# Cosmos3 request actually produces. -SEQUENCE_LENGTHS = [4096, 6240, 8192, 10336, 16384] - - -def _rotary() -> Qwen3VLTextRotaryEmbedding: - pretrained = SimpleNamespace( - rope_theta=1_000_000.0, - head_dim=HEAD_DIM, - max_position_embeddings=262_144, - rope_axes_dim=ROPE_AXES, - rope_scaling=None, - ) - return Qwen3VLTextRotaryEmbedding(SimpleNamespace(pretrained_config=pretrained)) - - -def _reference_cos_sin(rotary: Qwen3VLTextRotaryEmbedding, position_ids: torch.Tensor): - """Same math as the module, in float64 and without any matmul.""" - inv = rotary.inv_freq.double()[None, None, :, None] # [1, 1, D/2, 1] - pos = position_ids.double()[:, :, None, :] # [3, B, 1, N] - freqs = (inv * pos).transpose(2, 3) # [3, B, N, D/2] - freqs = rotary.apply_interleaved_mrope(freqs.clone(), rotary.mrope_section) - emb = torch.cat((freqs, freqs), dim=-1) - return emb.cos(), emb.sin() - - -@pytest.mark.skipif(not torch.cuda.is_available(), reason="needs a GPU") -@pytest.mark.parametrize("allow_tf32", [False, True]) -@pytest.mark.parametrize("seq_len", SEQUENCE_LENGTHS) -def test_rotary_table_matches_fp64_under_tf32(allow_tf32: bool, seq_len: int): - device = torch.device("cuda") - saved = torch.backends.cuda.matmul.allow_tf32 - torch.backends.cuda.matmul.allow_tf32 = allow_tf32 - try: - rotary = _rotary().to(device) - base = torch.arange(seq_len, dtype=torch.float32, device=device) - # 3D mRoPE ids: temporal axis fps-scaled (fractional), spatial axes integer - position_ids = torch.stack([base * 24.0 / 10.0, base, base * 0.5], dim=0)[:, None, :] - probe = torch.empty(0, dtype=torch.float32, device=device) - - cos, sin = rotary(probe, position_ids) - ref_cos, ref_sin = _reference_cos_sin(rotary, position_ids) - - # fp32 evaluates angles of ~4e4 rad with ~6e-8 relative precision, so a - # few 1e-3 absolute is the honest fp32 floor; TF32 breakage is O(1). - tol = 5e-3 - assert (cos.double() - ref_cos).abs().max().item() < tol - assert (sin.double() - ref_sin).abs().max().item() < tol - finally: - torch.backends.cuda.matmul.allow_tf32 = saved diff --git a/tests/unittest/_torch/visual_gen/test_cosmos3_transformer.py b/tests/unittest/_torch/visual_gen/test_cosmos3_transformer.py index 9a8105b14d57..36456efd1eb2 100644 --- a/tests/unittest/_torch/visual_gen/test_cosmos3_transformer.py +++ b/tests/unittest/_torch/visual_gen/test_cosmos3_transformer.py @@ -28,6 +28,7 @@ from tensorrt_llm._torch.visual_gen.models.cosmos3.transformer_cosmos3 import ( PRETRAINED_CONFIG_COMPAT_DEFAULTS, Cosmos3VFMTransformer, + Qwen3VLTextRotaryEmbedding, apply_pretrained_config_compat_defaults, ) from tensorrt_llm._torch.visual_gen.pipeline_loader import PipelineComponent, PipelineLoader @@ -854,3 +855,68 @@ def test_constructs_without_audio_or_action_towers(self): assert not hasattr(model, "audio_modality_embed") assert model.base_fps == 16 assert len(model.gen_layers) == 2 + + +class TestRotaryTablePrecision: + """The mRoPE angle table must stay fp32-accurate when fp32 GEMMs run in TF32. + + NGC PyTorch images set TORCH_ALLOW_TF32_CUBLAS_OVERRIDE=1. Positions above + 2048 do not fit TF32's 10-bit mantissa, so a table built as a K=1 matmul is + off by radians for late tokens wherever cuBLAS picks a TF32 kernel for that + shape (Blackwell). The table is an outer product and must not go through a + GEMM. Checkpoint-free; Nano head_dim and rope axes.""" + + HEAD_DIM = 128 + ROPE_AXES = [24, 20, 20] + # Text tokens (up to 4096) + margin + fps-scaled vision positions land well + # above 2048. cuBLAS picks a TF32 kernel for the K=1 GEMM only at some + # problem sizes, so sweep the sizes a Cosmos3 request actually produces. + SEQUENCE_LENGTHS = [4096, 6240, 8192, 10336, 16384] + + @pytest.fixture(autouse=True) + def _require_cuda(self): + if not torch.cuda.is_available(): + pytest.skip("CUDA not available") + + def _rotary(self) -> Qwen3VLTextRotaryEmbedding: + pretrained = SimpleNamespace( + rope_theta=1_000_000.0, + head_dim=self.HEAD_DIM, + max_position_embeddings=262_144, + rope_axes_dim=self.ROPE_AXES, + rope_scaling=None, + ) + return Qwen3VLTextRotaryEmbedding(SimpleNamespace(pretrained_config=pretrained)) + + @staticmethod + def _reference_cos_sin(rotary: Qwen3VLTextRotaryEmbedding, position_ids: torch.Tensor): + """Same math as the module, in float64 and without any matmul.""" + inv = rotary.inv_freq.double()[None, None, :, None] # [1, 1, D/2, 1] + pos = position_ids.double()[:, :, None, :] # [3, B, 1, N] + freqs = (inv * pos).transpose(2, 3) # [3, B, N, D/2] + freqs = rotary.apply_interleaved_mrope(freqs.clone(), rotary.mrope_section) + emb = torch.cat((freqs, freqs), dim=-1) + return emb.cos(), emb.sin() + + @pytest.mark.parametrize("allow_tf32", [False, True]) + @pytest.mark.parametrize("seq_len", SEQUENCE_LENGTHS) + def test_rotary_table_matches_fp64_under_tf32(self, allow_tf32: bool, seq_len: int): + saved = torch.backends.cuda.matmul.allow_tf32 + torch.backends.cuda.matmul.allow_tf32 = allow_tf32 + try: + rotary = self._rotary().to(DEVICE) + base = torch.arange(seq_len, dtype=torch.float32, device=DEVICE) + # 3D mRoPE ids: temporal axis fps-scaled (fractional), spatial axes integer + position_ids = torch.stack([base * 24.0 / 10.0, base, base * 0.5], dim=0)[:, None, :] + probe = torch.empty(0, dtype=torch.float32, device=DEVICE) + + cos, sin = rotary(probe, position_ids) + ref_cos, ref_sin = self._reference_cos_sin(rotary, position_ids) + + # fp32 evaluates angles of ~4e4 rad with ~6e-8 relative precision, so + # a few 1e-3 absolute is the honest fp32 floor; TF32 breakage is O(1). + tol = 5e-3 + assert (cos.double() - ref_cos).abs().max().item() < tol + assert (sin.double() - ref_sin).abs().max().item() < tol + finally: + torch.backends.cuda.matmul.allow_tf32 = saved From ce900bd5efaea245ca66f7d44bb922a7f6eff2c7 Mon Sep 17 00:00:00 2001 From: Igor Shovkun Date: Tue, 22 Sep 2026 10:36:54 -0700 Subject: [PATCH 3/7] Drop container-specific wording from the rotary table comment Signed-off-by: Igor Shovkun --- .../_torch/visual_gen/models/cosmos3/transformer_cosmos3.py | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/tensorrt_llm/_torch/visual_gen/models/cosmos3/transformer_cosmos3.py b/tensorrt_llm/_torch/visual_gen/models/cosmos3/transformer_cosmos3.py index 6891b0212ebc..0e6a35c53492 100644 --- a/tensorrt_llm/_torch/visual_gen/models/cosmos3/transformer_cosmos3.py +++ b/tensorrt_llm/_torch/visual_gen/models/cosmos3/transformer_cosmos3.py @@ -998,8 +998,8 @@ def forward(self, x, position_ids): ) position_ids_expanded = position_ids[:, :, None, :] # shape (3, bs, 1, positions) - # Elementwise, not a K=1 matmul: a GEMM under TF32 (default in NGC PyTorch - # images) rounds positions above 2048 and skews the phase of late tokens. + # Elementwise, not a K=1 matmul: a GEMM may run in TF32, which rounds + # positions above 2048 and skews the phase of late tokens. freqs = (inv_freq_expanded * position_ids_expanded).transpose(2, 3) freqs = self.apply_interleaved_mrope(freqs, self.mrope_section) emb = torch.cat((freqs, freqs), dim=-1) From b9254f8053ed5999bb76fe543fa8bc5ba22c9f55 Mon Sep 17 00:00:00 2001 From: Igor Shovkun Date: Tue, 22 Sep 2026 10:39:03 -0700 Subject: [PATCH 4/7] Explain why the rotary table departs from the diffusers matmul Signed-off-by: Igor Shovkun --- .../visual_gen/models/cosmos3/transformer_cosmos3.py | 7 +++++-- 1 file changed, 5 insertions(+), 2 deletions(-) diff --git a/tensorrt_llm/_torch/visual_gen/models/cosmos3/transformer_cosmos3.py b/tensorrt_llm/_torch/visual_gen/models/cosmos3/transformer_cosmos3.py index 0e6a35c53492..95a21311ce44 100644 --- a/tensorrt_llm/_torch/visual_gen/models/cosmos3/transformer_cosmos3.py +++ b/tensorrt_llm/_torch/visual_gen/models/cosmos3/transformer_cosmos3.py @@ -998,8 +998,11 @@ def forward(self, x, position_ids): ) position_ids_expanded = position_ids[:, :, None, :] # shape (3, bs, 1, positions) - # Elementwise, not a K=1 matmul: a GEMM may run in TF32, which rounds - # positions above 2048 and skews the phase of late tokens. + # The diffusers/transformers reference writes this outer product as + # `inv_freq @ position_ids`. That is a K=1 GEMM, and a GEMM may run in + # TF32, which cannot represent positions above 2048 and skews the + # rotary phase of late tokens. The broadcast multiply is the same + # arithmetic with no GEMM dispatch, so no precision mode applies. freqs = (inv_freq_expanded * position_ids_expanded).transpose(2, 3) freqs = self.apply_interleaved_mrope(freqs, self.mrope_section) emb = torch.cat((freqs, freqs), dim=-1) From 293c769a3e67fc2504e2b4c49d487b36595f0b68 Mon Sep 17 00:00:00 2001 From: Igor Shovkun Date: Tue, 22 Sep 2026 11:12:11 -0700 Subject: [PATCH 5/7] Address review: gate the rotary TF32 test on SM80+ and use integer spatial ids TF32 only exists from Ampere on, so allow_tf32=True passed trivially on older GPUs and looked like coverage it did not provide. The third mRoPE axis also used half-integral positions, which no Cosmos3 pathway produces: text ids share one integer ramp across all three axes and only the vision temporal axis is fps-scaled. Signed-off-by: Igor Shovkun --- .../_torch/visual_gen/test_cosmos3_transformer.py | 11 ++++++++--- 1 file changed, 8 insertions(+), 3 deletions(-) diff --git a/tests/unittest/_torch/visual_gen/test_cosmos3_transformer.py b/tests/unittest/_torch/visual_gen/test_cosmos3_transformer.py index 36456efd1eb2..3e84a5fd4589 100644 --- a/tests/unittest/_torch/visual_gen/test_cosmos3_transformer.py +++ b/tests/unittest/_torch/visual_gen/test_cosmos3_transformer.py @@ -874,9 +874,13 @@ class TestRotaryTablePrecision: SEQUENCE_LENGTHS = [4096, 6240, 8192, 10336, 16384] @pytest.fixture(autouse=True) - def _require_cuda(self): + def _require_tf32_capable_cuda(self): if not torch.cuda.is_available(): pytest.skip("CUDA not available") + # TF32 exists from Ampere on. Below that, allow_tf32=True changes + # nothing and the case would pass without covering the regression. + if torch.cuda.get_device_capability() < (8, 0): + pytest.skip("TF32-capable GPU (SM80+) required") def _rotary(self) -> Qwen3VLTextRotaryEmbedding: pretrained = SimpleNamespace( @@ -906,8 +910,9 @@ def test_rotary_table_matches_fp64_under_tf32(self, allow_tf32: bool, seq_len: i try: rotary = self._rotary().to(DEVICE) base = torch.arange(seq_len, dtype=torch.float32, device=DEVICE) - # 3D mRoPE ids: temporal axis fps-scaled (fractional), spatial axes integer - position_ids = torch.stack([base * 24.0 / 10.0, base, base * 0.5], dim=0)[:, None, :] + # Text ids share one integer ramp across all three axes; vision + # scales the temporal axis by fps, so that axis can be fractional. + position_ids = torch.stack([base * 24.0 / 10.0, base, base], dim=0)[:, None, :] probe = torch.empty(0, dtype=torch.float32, device=DEVICE) cos, sin = rotary(probe, position_ids) From dded22964a2ea172e01c63133f9d0c6553f434aa Mon Sep 17 00:00:00 2001 From: Igor Shovkun Date: Wed, 23 Sep 2026 21:23:20 -0700 Subject: [PATCH 6/7] Cover int64 position ids in the rotary precision test Every Cosmos3 request calls the rotary forward with two dtypes: the text tower passes an int64 ramp, and vision, audio and action pass fp32 when fps modulation is on and int64 when it is off. The test only built fp32 ids, so the int64 path was unexercised. Add it as a separate parametrization, since torch.stack needs one common dtype, and assert the returned dtype follows the probe rather than the position ids. Signed-off-by: Igor Shovkun --- .../visual_gen/test_cosmos3_transformer.py | 27 +++++++++++++++---- 1 file changed, 22 insertions(+), 5 deletions(-) diff --git a/tests/unittest/_torch/visual_gen/test_cosmos3_transformer.py b/tests/unittest/_torch/visual_gen/test_cosmos3_transformer.py index 3e84a5fd4589..c0c2cd039421 100644 --- a/tests/unittest/_torch/visual_gen/test_cosmos3_transformer.py +++ b/tests/unittest/_torch/visual_gen/test_cosmos3_transformer.py @@ -872,6 +872,10 @@ class TestRotaryTablePrecision: # above 2048. cuBLAS picks a TF32 kernel for the K=1 GEMM only at some # problem sizes, so sweep the sizes a Cosmos3 request actually produces. SEQUENCE_LENGTHS = [4096, 6240, 8192, 10336, 16384] + # Each request feeds this two dtypes: the text tower passes one int64 ramp + # shared by all three axes, fps-modulated vision a fractional fp32 + # temporal axis. Vision is int64 too when fps modulation is off. + POSITION_MODES = ["text_int64", "vision_fps_fp32"] @pytest.fixture(autouse=True) def _require_tf32_capable_cuda(self): @@ -902,22 +906,35 @@ def _reference_cos_sin(rotary: Qwen3VLTextRotaryEmbedding, position_ids: torch.T emb = torch.cat((freqs, freqs), dim=-1) return emb.cos(), emb.sin() + @staticmethod + def _position_ids(mode: str, seq_len: int) -> torch.Tensor: + """``[3, 1, seq_len]`` mRoPE ids in one of the two production shapes.""" + if mode == "text_int64": + ramp = torch.arange(seq_len, dtype=torch.long, device=DEVICE) + ids = ramp.unsqueeze(0).expand(3, -1).contiguous() + else: + base = torch.arange(seq_len, dtype=torch.float32, device=DEVICE) + ids = torch.stack([base * 24.0 / 10.0, base, base], dim=0) + return ids[:, None, :] + @pytest.mark.parametrize("allow_tf32", [False, True]) + @pytest.mark.parametrize("mode", POSITION_MODES) @pytest.mark.parametrize("seq_len", SEQUENCE_LENGTHS) - def test_rotary_table_matches_fp64_under_tf32(self, allow_tf32: bool, seq_len: int): + def test_rotary_table_matches_fp64_under_tf32(self, allow_tf32: bool, mode: str, seq_len: int): saved = torch.backends.cuda.matmul.allow_tf32 torch.backends.cuda.matmul.allow_tf32 = allow_tf32 try: rotary = self._rotary().to(DEVICE) - base = torch.arange(seq_len, dtype=torch.float32, device=DEVICE) - # Text ids share one integer ramp across all three axes; vision - # scales the temporal axis by fps, so that axis can be fractional. - position_ids = torch.stack([base * 24.0 / 10.0, base, base], dim=0)[:, None, :] + position_ids = self._position_ids(mode, seq_len) probe = torch.empty(0, dtype=torch.float32, device=DEVICE) cos, sin = rotary(probe, position_ids) ref_cos, ref_sin = self._reference_cos_sin(rotary, position_ids) + # The returned dtype follows the probe, not the position ids. + assert cos.dtype == probe.dtype + assert sin.dtype == probe.dtype + # fp32 evaluates angles of ~4e4 rad with ~6e-8 relative precision, so # a few 1e-3 absolute is the honest fp32 floor; TF32 breakage is O(1). tol = 5e-3 From 6b24db546170f6f967de3067df0a59d112370a2e Mon Sep 17 00:00:00 2001 From: Igor Shovkun Date: Wed, 23 Sep 2026 21:28:49 -0700 Subject: [PATCH 7/7] Parametrize the rotary precision test directly over position dtype Signed-off-by: Igor Shovkun --- .../visual_gen/test_cosmos3_transformer.py | 26 ++++++------------- 1 file changed, 8 insertions(+), 18 deletions(-) diff --git a/tests/unittest/_torch/visual_gen/test_cosmos3_transformer.py b/tests/unittest/_torch/visual_gen/test_cosmos3_transformer.py index c0c2cd039421..c167094dc84b 100644 --- a/tests/unittest/_torch/visual_gen/test_cosmos3_transformer.py +++ b/tests/unittest/_torch/visual_gen/test_cosmos3_transformer.py @@ -872,10 +872,8 @@ class TestRotaryTablePrecision: # above 2048. cuBLAS picks a TF32 kernel for the K=1 GEMM only at some # problem sizes, so sweep the sizes a Cosmos3 request actually produces. SEQUENCE_LENGTHS = [4096, 6240, 8192, 10336, 16384] - # Each request feeds this two dtypes: the text tower passes one int64 ramp - # shared by all three axes, fps-modulated vision a fractional fp32 - # temporal axis. Vision is int64 too when fps modulation is off. - POSITION_MODES = ["text_int64", "vision_fps_fp32"] + # Text ids are int64; fps-modulated vision, audio and action ids are fp32. + POSITION_DTYPES = [torch.int64, torch.float32] @pytest.fixture(autouse=True) def _require_tf32_capable_cuda(self): @@ -906,26 +904,18 @@ def _reference_cos_sin(rotary: Qwen3VLTextRotaryEmbedding, position_ids: torch.T emb = torch.cat((freqs, freqs), dim=-1) return emb.cos(), emb.sin() - @staticmethod - def _position_ids(mode: str, seq_len: int) -> torch.Tensor: - """``[3, 1, seq_len]`` mRoPE ids in one of the two production shapes.""" - if mode == "text_int64": - ramp = torch.arange(seq_len, dtype=torch.long, device=DEVICE) - ids = ramp.unsqueeze(0).expand(3, -1).contiguous() - else: - base = torch.arange(seq_len, dtype=torch.float32, device=DEVICE) - ids = torch.stack([base * 24.0 / 10.0, base, base], dim=0) - return ids[:, None, :] - @pytest.mark.parametrize("allow_tf32", [False, True]) - @pytest.mark.parametrize("mode", POSITION_MODES) + @pytest.mark.parametrize("pos_dtype", POSITION_DTYPES, ids=str) @pytest.mark.parametrize("seq_len", SEQUENCE_LENGTHS) - def test_rotary_table_matches_fp64_under_tf32(self, allow_tf32: bool, mode: str, seq_len: int): + def test_rotary_table_matches_fp64_under_tf32( + self, allow_tf32: bool, pos_dtype: torch.dtype, seq_len: int + ): saved = torch.backends.cuda.matmul.allow_tf32 torch.backends.cuda.matmul.allow_tf32 = allow_tf32 try: rotary = self._rotary().to(DEVICE) - position_ids = self._position_ids(mode, seq_len) + ramp = torch.arange(seq_len, dtype=pos_dtype, device=DEVICE) + position_ids = ramp.expand(3, -1)[:, None, :] probe = torch.empty(0, dtype=torch.float32, device=DEVICE) cos, sin = rotary(probe, position_ids)