Add FP32 dtype support for RoPE - Part2 (#13328)
This commit is contained in:
@@ -113,7 +113,7 @@ class RotaryEmbedding(CustomOp):
|
|||||||
if not _is_cuda:
|
if not _is_cuda:
|
||||||
cache = cache.to(dtype)
|
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])
|
(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_cpu and _is_cpu_amx_available)
|
||||||
and not (_is_xpu)
|
and not (_is_xpu)
|
||||||
@@ -273,11 +273,7 @@ class RotaryEmbedding(CustomOp):
|
|||||||
offsets: Optional[torch.Tensor] = None,
|
offsets: Optional[torch.Tensor] = None,
|
||||||
fused_set_kv_buffer_arg: Optional[FusedSetKVBufferArg] = None,
|
fused_set_kv_buffer_arg: Optional[FusedSetKVBufferArg] = None,
|
||||||
) -> Tuple[torch.Tensor, torch.Tensor]:
|
) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||||
if (
|
if _is_cuda and (self.head_size in [64, 128, 256, 512]):
|
||||||
_is_cuda
|
|
||||||
and (self.head_size in [64, 128, 256, 512])
|
|
||||||
and self.dtype != torch.float32
|
|
||||||
):
|
|
||||||
apply_rope_with_cos_sin_cache_inplace(
|
apply_rope_with_cos_sin_cache_inplace(
|
||||||
positions=positions,
|
positions=positions,
|
||||||
query=query,
|
query=query,
|
||||||
|
|||||||
@@ -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, 32, 8),
|
||||||
(128, 128, 2048, 10000, False, torch.bfloat16, "cpu", 2, 512, 16, 4),
|
(128, 128, 2048, 10000, False, torch.bfloat16, "cpu", 2, 512, 16, 4),
|
||||||
(512, 128, 311, 10000, False, torch.bfloat16, "cpu", 3, 39, 4, 2),
|
(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 (
|
for (
|
||||||
|
|||||||
@@ -76,7 +76,7 @@ num_tokens_list = [11, 8192]
|
|||||||
],
|
],
|
||||||
)
|
)
|
||||||
@pytest.mark.parametrize("tp_size", [1, 2])
|
@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)
|
@pytest.mark.parametrize("num_tokens", num_tokens_list)
|
||||||
def test_mrope(
|
def test_mrope(
|
||||||
model_name: str,
|
model_name: str,
|
||||||
|
|||||||
Reference in New Issue
Block a user