项目文件夹

文件
wehub-resource-sync 59a0a3844c
PR Test AMD / cancel-on-close (push) Has been skipped
PR Test NVIDIA ARM / scan (push) Has been skipped
PR Test NVIDIA / cancel-on-close (push) Has been skipped
PR Test AMD / scan (push) Has been skipped
PR Test NVIDIA ARM / cancel-on-close (push) Has been skipped
PR Test NVIDIA / scan (push) Has been skipped
Release Docker Images / build (cu129-torch-2.11.0) (push) Has been skipped
Release Docker Images / build (cu130-torch-2.11.0) (push) Has been skipped
Release PyPI / publish (push) Has been skipped
Scheduler Python Test / test (push) Successful in 27m19s
Docs / build (push) Successful in 28m8s
Scheduler C++ Test / test (push) Successful in 28m19s
Scheduler C++ Test / test-flat (push) Successful in 28m18s
Docs / deploy (push) Has been cancelled
PR Test AMD / finish (push) Has been cancelled
PR Test NVIDIA / finish (push) Has been cancelled
PR Test NVIDIA ARM / finish (push) Has been cancelled
PR Test NVIDIA ARM / ${{ matrix.name }} (${{ matrix.runner }}) (push) Has been cancelled
PR Test AMD / ${{ matrix.name }} (${{ matrix.runner }}) (push) Has been cancelled
PR Test NVIDIA / ${{ matrix.name }} (${{ matrix.runner }}) (push) Has been cancelled
chore: import upstream snapshot with attribution
2026-07-13 12:32:31 +08:00

531 行
18 KiB
Python

# Copyright (c) 2026 LightSeek Foundation
#
# Permission is hereby granted, free of charge, to any person obtaining a copy
# of this software and associated documentation files (the "Software"), to deal
# in the Software without restriction, including without limitation the rights
# to use, copy, modify, merge, publish, distribute, sublicense, and/or sell
# copies of the Software, and to permit persons to whom the Software is
# furnished to do so, subject to the following conditions:
#
# The above copyright notice and this permission notice shall be included in
# all copies or substantial portions of the Software.
#
# THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY,
# FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE
# AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER
# LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM,
# OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE
# SOFTWARE.
from __future__ import annotations
import pytest
import torch
from tokenspeed_kernel.platform import ArchVersion, CapabilityRequirement
from tokenspeed_kernel.registry import (
KernelRegistry,
KernelSpec,
describe_kernel,
register_kernel,
)
from tokenspeed_kernel.signature import (
ScaleFormat,
dense_tensor_format,
format_signature,
format_signatures,
tensor_format,
)
from utils import dummy_impl, register_all_samples
pytestmark = pytest.mark.usefixtures("fresh_registry")
class TestKernelSpec:
def test_frozen_dataclass(self):
spec = KernelSpec(name="k1", family="attention", mode="decode")
with pytest.raises(AttributeError):
spec.name = "k2"
def test_default_values(self):
spec = KernelSpec(name="k1", family="attention", mode="decode")
assert spec.features == frozenset()
assert spec.solution == ""
assert spec.priority == 10
assert spec.tags == frozenset()
assert spec.format_signatures == frozenset()
def test_hashable_without_dict_traits(self):
spec = KernelSpec(name="k1", family="attention", mode="decode", traits={})
with pytest.raises(TypeError):
hash(spec)
def test_format_signature_bundles_scale_metadata(self):
scale = ScaleFormat(
storage_dtype=torch.float32,
granularity="block",
block_shape=(32,),
)
mixed = format_signature(
a=dense_tensor_format(torch.bfloat16),
b=tensor_format("mxfp4", torch.uint8, scale=scale),
)
dense = format_signature(
a=dense_tensor_format(torch.bfloat16),
b=dense_tensor_format(torch.uint8),
)
assert mixed != dense
assert mixed.format_for("b").scale == scale
def test_block_scale_requires_shape_or_dynamic_marker(self):
with pytest.raises(ValueError, match="requires block_shape"):
ScaleFormat(storage_dtype=torch.float32, granularity="block")
dynamic = ScaleFormat(
storage_dtype=torch.float32,
granularity="block",
dynamic_block_shape=True,
)
assert dynamic.block_shape is None
assert str(dynamic) == "scale(block, storage=torch.float32, block=dynamic)"
with pytest.raises(ValueError, match="mutually exclusive"):
ScaleFormat(
storage_dtype=torch.float32,
granularity="block",
block_shape=(16,),
dynamic_block_shape=True,
)
with pytest.raises(ValueError, match="only valid for block"):
ScaleFormat(
storage_dtype=torch.float32,
granularity="tensor",
block_shape=(16,),
)
def test_fp8_tensor_format_names_are_unambiguous(self):
dense = dense_tensor_format(torch.float8_e4m3fn)
assert dense.format == "dense"
assert dense.scale is None
scale = ScaleFormat(storage_dtype=torch.float32, granularity="tensor")
scaled = tensor_format("scaled-fp8", torch.float8_e4m3fn, scale=scale)
assert scaled.format == "scaled-fp8"
assert scaled.scale == scale
with pytest.raises(ValueError, match="ambiguous"):
tensor_format("fp8", torch.float8_e4m3fn)
with pytest.raises(ValueError, match="requires scale"):
tensor_format("scaled-fp8", torch.float8_e4m3fn)
def test_format_signature_storage_dtype_for_role(self):
signature = format_signature(
a=dense_tensor_format(torch.bfloat16),
b=dense_tensor_format(torch.float16),
)
assert signature.storage_dtype_for("a") == torch.bfloat16
assert signature.storage_dtype_for("b") == torch.float16
assert signature.storage_dtype_for("missing") is None
def test_format_signatures_for_storage_dtype_uses_explicit_roles(self):
mxfp4_scale = ScaleFormat(
storage_dtype=torch.uint8,
granularity="block",
block_shape=(32,),
)
nvfp4_scale = ScaleFormat(
storage_dtype=torch.float32,
granularity="block",
dynamic_block_shape=True,
)
bf16_mxfp4 = format_signature(
x=dense_tensor_format(torch.bfloat16),
weight=tensor_format("mxfp4", torch.uint8, scale=mxfp4_scale),
)
bf16_nvfp4 = format_signature(
x=dense_tensor_format(torch.bfloat16),
weight=tensor_format("nvfp4", torch.uint8, scale=nvfp4_scale),
)
fp16_mxfp4 = format_signature(
x=dense_tensor_format(torch.float16),
weight=tensor_format("mxfp4", torch.uint8, scale=mxfp4_scale),
)
spec = KernelSpec(
name="moe_fused",
family="moe",
mode="fused",
format_signatures=frozenset({bf16_mxfp4, bf16_nvfp4, fp16_mxfp4}),
)
matches = spec.format_signatures_for_storage_dtype(torch.bfloat16, "x")
assert set(matches) == {bf16_mxfp4, bf16_nvfp4}
assert spec.storage_dtypes_for_role("x") == {torch.bfloat16, torch.float16}
assert spec.format_signature_for_storage_dtype(torch.float16, "x") == fp16_mxfp4
with pytest.raises(ValueError, match="multiple format signatures"):
spec.format_signature_for_storage_dtype(torch.bfloat16, "x")
def test_equality(self):
spec1 = KernelSpec(name="k1", family="attention", mode="decode")
spec2 = KernelSpec(name="k1", family="attention", mode="decode")
assert spec1 == spec2
class TestRegistrySingleton:
def test_get_returns_same_instance(self):
r1 = KernelRegistry.get()
r2 = KernelRegistry.get()
assert r1 is r2
def test_reset_creates_new_instance(self):
r1 = KernelRegistry.get()
KernelRegistry.reset()
r2 = KernelRegistry.get()
assert r1 is not r2
class TestRegistryRegister:
def test_register_and_retrieve(self):
reg = KernelRegistry.get()
spec = KernelSpec(name="test_k", family="attention", mode="decode")
impl = dummy_impl("test_k")
reg.register(spec, impl)
assert reg.get_by_name("test_k") is spec
assert reg.get_impl("test_k") is impl
def test_register_multiple_kernels(self, sample_specs):
reg = KernelRegistry.get()
register_all_samples(reg, sample_specs)
assert reg.get_by_name("flashinfer_decode") is not None
assert reg.get_by_name("triton_decode") is not None
assert reg.get_by_name("cutlass_prefill") is not None
assert reg.get_by_name("nonexistent") is None
def test_reregister_replaces_old(self):
reg = KernelRegistry.get()
spec1 = KernelSpec(name="k", family="attention", mode="decode", priority=5)
spec2 = KernelSpec(name="k", family="attention", mode="decode", priority=15)
impl1 = dummy_impl("old")
impl2 = dummy_impl("new")
reg.register(spec1, impl1)
reg.register(spec2, impl2)
assert reg.get_by_name("k") is spec2
assert reg.get_impl("k") is impl2
assert len(reg.get_for_operator("attention", "decode")) == 1
def test_sorted_by_priority_descending(self, sample_specs):
reg = KernelRegistry.get()
register_all_samples(reg, sample_specs)
decode_specs = reg.get_for_operator("attention", "decode")
priorities = [s.priority for s in decode_specs]
assert priorities == sorted(priorities, reverse=True)
class TestRegistryQueries:
def test_get_for_operator_basic(self, sample_specs):
reg = KernelRegistry.get()
register_all_samples(reg, sample_specs)
decode = reg.get_for_operator("attention", "decode")
assert len(decode) >= 3
for s in decode:
assert s.family == "attention"
assert s.mode == "decode"
def test_get_for_operator_empty(self):
reg = KernelRegistry.get()
assert reg.get_for_operator("nonexistent", "op") == []
def test_filter_by_features(self, sample_specs):
reg = KernelRegistry.get()
register_all_samples(reg, sample_specs)
paged = reg.get_for_operator(
"attention", "decode", features=frozenset({"paged"})
)
for s in paged:
assert "paged" in s.features
def test_filter_by_platform(
self, sample_specs, h100_platform, mi300_platform, mi350_platform
):
reg = KernelRegistry.get()
register_all_samples(reg, sample_specs)
nvidia_kernels = reg.get_for_operator(
"attention", "decode", platform=h100_platform
)
nvidia_names = {s.name for s in nvidia_kernels}
assert "aiter_decode" not in nvidia_names
assert "flashinfer_decode" in nvidia_names
amd_kernels = reg.get_for_operator(
"attention", "decode", platform=mi300_platform
)
amd_names = {s.name for s in amd_kernels}
assert "flashinfer_decode" not in amd_names
assert "aiter_decode" in amd_names
mi350_kernels = reg.get_for_operator(
"attention", "decode", platform=mi350_platform
)
mi350_names = {s.name for s in mi350_kernels}
assert "flashinfer_decode" not in mi350_names
assert "aiter_decode" in mi350_names
assert "triton_decode" in mi350_names
def test_filter_by_signature(self, sample_specs):
reg = KernelRegistry.get()
register_all_samples(reg, sample_specs)
signature = next(
iter(
format_signatures(("q", "k_cache", "v_cache"), "dense", {torch.float32})
)
)
fp32 = reg.get_for_operator("attention", "decode", format_signature=signature)
names = {s.name for s in fp32}
assert "reference_decode" in names
assert "flashinfer_decode" not in names
def test_filter_by_tags(self, sample_specs):
reg = KernelRegistry.get()
register_all_samples(reg, sample_specs)
latency = reg.get_for_operator("attention", "decode", tags={"latency"})
for s in latency:
assert "latency" in s.tags
def test_filter_by_solution(self, sample_specs):
reg = KernelRegistry.get()
register_all_samples(reg, sample_specs)
triton = reg.get_for_operator("attention", "decode", solution="triton")
assert all(s.solution == "triton" for s in triton)
assert len(triton) == 1
def test_list_operators(self, sample_specs):
reg = KernelRegistry.get()
register_all_samples(reg, sample_specs)
ops = reg.list_operators()
assert ("attention", "decode") in ops
assert ("attention", "prefill") in ops
assert ("gemm", "mm") in ops
def test_list_kernels_all(self, sample_specs):
reg = KernelRegistry.get()
register_all_samples(reg, sample_specs)
all_kernels = reg.list_kernels()
assert len(all_kernels) == len(sample_specs)
def test_list_kernels_by_family(self, sample_specs):
reg = KernelRegistry.get()
register_all_samples(reg, sample_specs)
attn = reg.list_kernels(family="attention")
assert all(s.family == "attention" for s in attn)
def test_list_kernels_by_mode(self, sample_specs):
reg = KernelRegistry.get()
register_all_samples(reg, sample_specs)
decode = reg.list_kernels(mode="decode")
assert all(s.mode == "decode" for s in decode)
def test_list_kernels_by_family_and_mode(self, sample_specs):
reg = KernelRegistry.get()
register_all_samples(reg, sample_specs)
decode = reg.list_kernels(family="attention", mode="decode")
assert all(s.family == "attention" and s.mode == "decode" for s in decode)
def test_list_solutions(self, sample_specs):
reg = KernelRegistry.get()
register_all_samples(reg, sample_specs)
solutions = reg.list_solutions("attention", "decode")
assert "flashinfer" in solutions
assert "triton" in solutions
assert "reference" in solutions
class TestRegistryCache:
def test_cache_put_and_get(self):
reg = KernelRegistry.get()
key = ("attention", "decode", torch.bfloat16, "sm_90")
impl = dummy_impl("cached")
assert reg.cache_get(key) is None
reg.cache_put(key, impl)
assert reg.cache_get(key) is impl
def test_clear_cache(self):
reg = KernelRegistry.get()
key = ("attention", "decode", torch.bfloat16, "sm_90")
reg.cache_put(key, dummy_impl("cached"))
reg.clear_cache()
assert reg.cache_get(key) is None
def test_invalidate_cache_on_register(self):
reg = KernelRegistry.get()
key = ("attention", "decode", torch.bfloat16, "sm_90")
reg.cache_put(key, dummy_impl("cached"))
spec = KernelSpec(name="new_k", family="attention", mode="decode")
reg.register(spec, dummy_impl("new_k"))
assert reg.cache_get(key) is None
def test_invalidate_preserves_other_ops(self):
reg = KernelRegistry.get()
attn_key = ("attention", "decode", torch.bfloat16, "sm_90")
gemm_key = ("gemm", "mm", torch.bfloat16, "sm_90")
reg.cache_put(attn_key, dummy_impl("attn"))
reg.cache_put(gemm_key, dummy_impl("gemm"))
spec = KernelSpec(name="new_attn", family="attention", mode="decode")
reg.register(spec, dummy_impl("new_attn"))
assert reg.cache_get(attn_key) is None
assert reg.cache_get(gemm_key) is not None
class TestRegisterKernelDecorator:
def test_basic_decorator(self):
@register_kernel(
"gemm",
"mm",
solution="reference",
signatures=format_signatures(("a", "b"), "dense", {torch.bfloat16}),
priority=12,
)
def my_torch_gemm(a, b):
return a @ b
reg = KernelRegistry.get()
spec = reg.get_by_name("reference_gemm_mm")
assert spec is not None
assert spec.solution == "reference"
assert spec.priority == 12
assert (
next(iter(format_signatures(("a", "b"), "dense", {torch.bfloat16})))
in spec.format_signatures
)
impl = reg.get_impl("reference_gemm_mm")
assert impl is my_torch_gemm
def test_custom_name(self):
@register_kernel(
"attention",
"decode",
name="my_custom_kernel",
solution="custom",
signatures=format_signatures(
("q", "k_cache", "v_cache"), "dense", {torch.float16}
),
)
def some_func():
pass
reg = KernelRegistry.get()
assert reg.get_by_name("my_custom_kernel") is not None
def test_decorator_with_features_and_tags(self):
@register_kernel(
"attention",
"decode",
features={"paged", "rope"},
solution="triton",
capability=CapabilityRequirement(
min_arch_version=ArchVersion(8, 0),
),
signatures=format_signatures(
("q", "k_cache", "v_cache"), "dense", {torch.float16, torch.bfloat16}
),
tags={"determinism", "latency"},
)
def decorated_kernel():
pass
reg = KernelRegistry.get()
spec = reg.get_by_name("triton_attention_decode")
assert spec is not None
assert spec.features == frozenset({"paged", "rope"})
assert spec.tags == frozenset({"determinism", "latency"})
assert spec.capability.min_arch_version == ArchVersion(8, 0)
def test_decorator_returns_original_function(self):
@register_kernel(
"gemm",
"mm",
solution="test",
signatures=format_signatures(("a", "b"), "dense", {torch.float16}),
)
def original(x):
return x * 2
assert original(5) == 10
class TestDescribeKernel:
def test_describe_existing(self, sample_specs):
reg = KernelRegistry.get()
register_all_samples(reg, sample_specs)
desc = describe_kernel("flashinfer_decode")
assert "flashinfer_decode" in desc
assert "attention" in desc
assert "flashinfer" in desc
def test_describe_includes_weight_preprocessor_link(self):
def moe_weights(**_):
return None
reg = KernelRegistry.get()
spec = KernelSpec(
name="moe_apply",
family="moe",
mode="apply",
weight_preprocessor=moe_weights,
)
reg.register(spec, dummy_impl("moe_apply"))
desc = describe_kernel("moe_apply")
assert "Weight preprocessor: moe_weights" in desc
def test_describe_not_found(self):
desc = describe_kernel("nonexistent_kernel")
assert "not found" in desc.lower()
class TestUnregister:
def test_unregister_removes_from_all_lookups(self, sample_specs):
reg = KernelRegistry.get()
register_all_samples(reg, sample_specs)
assert reg.get_by_name("triton_decode") is not None
reg._unregister("triton_decode")
assert reg.get_by_name("triton_decode") is None
assert reg.get_impl("triton_decode") is None
names = {s.name for s in reg.get_for_operator("attention", "decode")}
assert "triton_decode" not in names
def test_unregister_nonexistent_is_noop(self):
reg = KernelRegistry.get()
reg._unregister("does_not_exist")