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..95a21311ce44 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,14 @@ 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) - - freqs = (inv_freq_expanded.float() @ position_ids_expanded.float()).transpose(2, 3) + position_ids_expanded = position_ids[:, :, None, :] # shape (3, bs, 1, positions) + + # 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) cos = emb.cos() * self.attention_scaling diff --git a/tests/unittest/_torch/visual_gen/test_cosmos3_transformer.py b/tests/unittest/_torch/visual_gen/test_cosmos3_transformer.py index 9a8105b14d57..c167094dc84b 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,80 @@ 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] + # 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): + 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( + 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("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, 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) + 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) + 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 + 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