Skip to content
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
78 changes: 78 additions & 0 deletions tests/unittest/_torch/visual_gen/test_cosmos3_transformer.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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():
Comment thread
coderabbitai[bot] marked this conversation as resolved.
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)
Comment thread
ishovkun marked this conversation as resolved.
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
Loading