Add FP32 dtype support for RoPE - Part2 (#13328)

This commit is contained in:
iLeGend
2025-11-19 21:19:53 -08:00
committed by GitHub
parent bc42c8c415
commit 10e0b83a4c
3 changed files with 9 additions and 7 deletions
+2 -6
View File
@@ -113,7 +113,7 @@ class RotaryEmbedding(CustomOp):
if not _is_cuda:
cache = cache.to(dtype)
if dtype == torch.float32 or (
if (
(not (_is_cuda or _is_npu) or self.head_size not in [64, 128, 256, 512])
and not (_is_cpu and _is_cpu_amx_available)
and not (_is_xpu)
@@ -273,11 +273,7 @@ class RotaryEmbedding(CustomOp):
offsets: Optional[torch.Tensor] = None,
fused_set_kv_buffer_arg: Optional[FusedSetKVBufferArg] = None,
) -> Tuple[torch.Tensor, torch.Tensor]:
if (
_is_cuda
and (self.head_size in [64, 128, 256, 512])
and self.dtype != torch.float32
):
if _is_cuda and (self.head_size in [64, 128, 256, 512]):
apply_rope_with_cos_sin_cache_inplace(
positions=positions,
query=query,
+6
View File
@@ -146,6 +146,12 @@ class TestROPE(CustomTestCase):
(128, 128, 2048, 10000, False, torch.bfloat16, "cpu", 2, 512, 32, 8),
(128, 128, 2048, 10000, False, torch.bfloat16, "cpu", 2, 512, 16, 4),
(512, 128, 311, 10000, False, torch.bfloat16, "cpu", 3, 39, 4, 2),
(64, 64, 32, 8000, True, torch.float32, "cpu", 32, 32, 1, 1),
(256, 128, 4096, 10000, True, torch.float32, "cpu", 2, 512, 32, 8),
(512, 128, 311, 10000, True, torch.float32, "cpu", 3, 39, 4, 2),
(128, 128, 2048, 10000, False, torch.float32, "cpu", 2, 512, 32, 8),
(128, 128, 2048, 10000, False, torch.float32, "cpu", 2, 512, 16, 4),
(512, 128, 311, 10000, False, torch.float32, "cpu", 3, 39, 4, 2),
]
for (
+1 -1
View File
@@ -76,7 +76,7 @@ num_tokens_list = [11, 8192]
],
)
@pytest.mark.parametrize("tp_size", [1, 2])
@pytest.mark.parametrize("dtype", [torch.bfloat16])
@pytest.mark.parametrize("dtype", [torch.bfloat16, torch.float32])
@pytest.mark.parametrize("num_tokens", num_tokens_list)
def test_mrope(
model_name: str,