[AMD][CI] Remove obsolete split-dim check from Kimi-K3 prefill test (#38288)
This commit is contained in:
@@ -3,9 +3,6 @@ import unittest
|
|||||||
import torch
|
import torch
|
||||||
|
|
||||||
from sglang.kernels.ops.attention.extend_attention import extend_attention_fwd
|
from sglang.kernels.ops.attention.extend_attention import extend_attention_fwd
|
||||||
from sglang.kernels.ops.attention.extend_attention_split_dim import (
|
|
||||||
can_use_split_dim_absorbed_extend,
|
|
||||||
)
|
|
||||||
from sglang.srt.environ import envs
|
from sglang.srt.environ import envs
|
||||||
from sglang.srt.utils import get_device, is_gfx95_supported, is_hip
|
from sglang.srt.utils import get_device, is_gfx95_supported, is_hip
|
||||||
from sglang.test.ci.ci_register import register_amd_ci
|
from sglang.test.ci.ci_register import register_amd_ci
|
||||||
@@ -98,82 +95,6 @@ class TestKimiK3TritonPrefill(unittest.TestCase):
|
|||||||
|
|
||||||
torch.testing.assert_close(output.float(), reference, rtol=1e-2, atol=1e-2)
|
torch.testing.assert_close(output.float(), reference, rtol=1e-2, atol=1e-2)
|
||||||
|
|
||||||
def test_split_dim_dispatch_gates(self):
|
|
||||||
device = get_device()
|
|
||||||
q = torch.empty(1, 12, 576, dtype=torch.bfloat16, device=device)
|
|
||||||
k = torch.empty(1, 1, 576, dtype=torch.bfloat16, device=device)
|
|
||||||
v = torch.empty(1, 1, 512, dtype=torch.bfloat16, device=device)
|
|
||||||
o = torch.empty(1, 12, 512, dtype=torch.bfloat16, device=device)
|
|
||||||
k_buffer = torch.empty(1, 1, 576, dtype=torch.bfloat16, device=device)
|
|
||||||
v_buffer = torch.empty(1, 1, 512, dtype=torch.bfloat16, device=device)
|
|
||||||
kwargs = dict(
|
|
||||||
lse=None,
|
|
||||||
sinks=None,
|
|
||||||
k_scale=1.0,
|
|
||||||
v_scale=1.0,
|
|
||||||
custom_mask=None,
|
|
||||||
is_causal=True,
|
|
||||||
sliding_window_size=-1,
|
|
||||||
logit_cap=0.0,
|
|
||||||
xai_temperature_len=-1,
|
|
||||||
skip_prefix=False,
|
|
||||||
skip_extend=False,
|
|
||||||
page_size=1,
|
|
||||||
score_mod=None,
|
|
||||||
aux_tensors=None,
|
|
||||||
)
|
|
||||||
self.assertTrue(
|
|
||||||
can_use_split_dim_absorbed_extend(q, k, v, o, k_buffer, v_buffer, **kwargs)
|
|
||||||
)
|
|
||||||
|
|
||||||
fp8_k_buffer = k_buffer.to(torch.float8_e4m3fn)
|
|
||||||
fp8_v_buffer = v_buffer.to(torch.float8_e4m3fn)
|
|
||||||
with envs.SGLANG_TRITON_FP8_PREFILL_ATTN.override(False):
|
|
||||||
self.assertFalse(
|
|
||||||
can_use_split_dim_absorbed_extend(
|
|
||||||
q, k, v, o, fp8_k_buffer, fp8_v_buffer, **kwargs
|
|
||||||
)
|
|
||||||
)
|
|
||||||
with envs.SGLANG_TRITON_FP8_PREFILL_ATTN.override(True):
|
|
||||||
self.assertTrue(
|
|
||||||
can_use_split_dim_absorbed_extend(
|
|
||||||
q, k, v, o, fp8_k_buffer, fp8_v_buffer, **kwargs
|
|
||||||
)
|
|
||||||
)
|
|
||||||
self.assertTrue(
|
|
||||||
can_use_split_dim_absorbed_extend(
|
|
||||||
q,
|
|
||||||
k,
|
|
||||||
v,
|
|
||||||
o,
|
|
||||||
fp8_k_buffer,
|
|
||||||
fp8_v_buffer,
|
|
||||||
**{**kwargs, "k_scale": 0.5, "v_scale": 0.25},
|
|
||||||
)
|
|
||||||
)
|
|
||||||
|
|
||||||
for override in (
|
|
||||||
{"page_size": 2},
|
|
||||||
{"logit_cap": 1.0},
|
|
||||||
{"sliding_window_size": 128},
|
|
||||||
{"skip_prefix": True},
|
|
||||||
{"is_causal": False},
|
|
||||||
{"lse": torch.empty(1, 12, dtype=torch.float32, device=device)},
|
|
||||||
{"sinks": torch.empty(12, dtype=torch.float32, device=device)},
|
|
||||||
{"k_scale": 0.5},
|
|
||||||
):
|
|
||||||
self.assertFalse(
|
|
||||||
can_use_split_dim_absorbed_extend(
|
|
||||||
q,
|
|
||||||
k,
|
|
||||||
v,
|
|
||||||
o,
|
|
||||||
k_buffer,
|
|
||||||
v_buffer,
|
|
||||||
**{**kwargs, **override},
|
|
||||||
)
|
|
||||||
)
|
|
||||||
|
|
||||||
def test_zero_prefix_fp8_flag(self):
|
def test_zero_prefix_fp8_flag(self):
|
||||||
device = get_device()
|
device = get_device()
|
||||||
tokens, heads, d_qk, d_v = 128, 12, 192, 128
|
tokens, heads, d_qk, d_v = 128, 12, 192, 128
|
||||||
|
|||||||
Reference in New Issue
Block a user