[AMD] Fix stale imports in test_fused_fp8_kv_write.py (#31515)
This commit is contained in:
@@ -31,10 +31,10 @@ def _fp32_divide_write(k, v, k_cache, v_cache, loc, k_scale, v_scale, fp8_dtype)
|
|||||||
@unittest.skipUnless(_HAS_CUDA, "Triton kernels require a GPU")
|
@unittest.skipUnless(_HAS_CUDA, "Triton kernels require a GPU")
|
||||||
class TestFusedFp8KvWrite(unittest.TestCase):
|
class TestFusedFp8KvWrite(unittest.TestCase):
|
||||||
def _run(self, num_tokens, num_heads, head_dim, total_slots=None, seed=0xC0FFEE):
|
def _run(self, num_tokens, num_heads, head_dim, total_slots=None, seed=0xC0FFEE):
|
||||||
from sglang.srt.layers.attention.utils import (
|
from sglang.kernels.ops.attention.utils import (
|
||||||
launch_reshape_and_cache_flash,
|
launch_reshape_and_cache_flash,
|
||||||
)
|
)
|
||||||
from sglang.srt.layers.quantization.fp8_kernel import fp8_dtype
|
from sglang.kernels.ops.quantization.fp8_kernel import fp8_dtype
|
||||||
|
|
||||||
torch.manual_seed(seed)
|
torch.manual_seed(seed)
|
||||||
dev = "cuda"
|
dev = "cuda"
|
||||||
@@ -232,7 +232,7 @@ class TestAiterFp8KvDispatch(unittest.TestCase):
|
|||||||
v_scale falls back to self.v_scale (not self.k_scale)."""
|
v_scale falls back to self.v_scale (not self.k_scale)."""
|
||||||
from unittest import mock
|
from unittest import mock
|
||||||
|
|
||||||
from sglang.srt.layers.quantization.fp8_kernel import fp8_dtype
|
from sglang.kernels.ops.quantization.fp8_kernel import fp8_dtype
|
||||||
|
|
||||||
dev = "cuda"
|
dev = "cuda"
|
||||||
heads, dim, n = 2, 64, 4
|
heads, dim, n = 2, 64, 4
|
||||||
@@ -280,7 +280,7 @@ class TestAiterFp8KvDispatch(unittest.TestCase):
|
|||||||
K's head_dim for V)."""
|
K's head_dim for V)."""
|
||||||
from unittest import mock
|
from unittest import mock
|
||||||
|
|
||||||
from sglang.srt.layers.quantization.fp8_kernel import fp8_dtype
|
from sglang.kernels.ops.quantization.fp8_kernel import fp8_dtype
|
||||||
|
|
||||||
dev = "cuda"
|
dev = "cuda"
|
||||||
heads, qk_dim, v_dim, n = 2, 64, 32, 4 # qk_head_dim != v_head_dim
|
heads, qk_dim, v_dim, n = 2, 64, 32, 4 # qk_head_dim != v_head_dim
|
||||||
@@ -337,7 +337,7 @@ class TestAiterFp8KvDispatch(unittest.TestCase):
|
|||||||
without scale args (its signature takes no scales)."""
|
without scale args (its signature takes no scales)."""
|
||||||
from unittest import mock
|
from unittest import mock
|
||||||
|
|
||||||
from sglang.srt.layers.quantization.fp8_kernel import fp8_dtype
|
from sglang.kernels.ops.quantization.fp8_kernel import fp8_dtype
|
||||||
|
|
||||||
dev = "cuda"
|
dev = "cuda"
|
||||||
heads, dim, n = 2, 64, 4
|
heads, dim, n = 2, 64, 4
|
||||||
@@ -389,8 +389,8 @@ class TestUseFusedFp8KvWritePredicate(unittest.TestCase):
|
|||||||
"""Unit test for the shared _use_fused_fp8_kv_write predicate."""
|
"""Unit test for the shared _use_fused_fp8_kv_write predicate."""
|
||||||
|
|
||||||
def _backend(self, **overrides):
|
def _backend(self, **overrides):
|
||||||
|
from sglang.kernels.ops.quantization.fp8_kernel import fp8_dtype
|
||||||
from sglang.srt.layers.attention.aiter_backend import AiterAttnBackend
|
from sglang.srt.layers.attention.aiter_backend import AiterAttnBackend
|
||||||
from sglang.srt.layers.quantization.fp8_kernel import fp8_dtype
|
|
||||||
|
|
||||||
be = AiterAttnBackend.__new__(AiterAttnBackend)
|
be = AiterAttnBackend.__new__(AiterAttnBackend)
|
||||||
be.kv_cache_dtype = overrides.get("kv_cache_dtype", fp8_dtype)
|
be.kv_cache_dtype = overrides.get("kv_cache_dtype", fp8_dtype)
|
||||||
|
|||||||
Reference in New Issue
Block a user