[Kernel] Remove unused implementations and stale registry entries (#32636)

This commit is contained in:
Xiaoyu Zhang
2026-07-28 18:12:10 +08:00
committed by GitHub
parent dde03d7c4a
commit 9cffc2ba52
27 changed files with 21 additions and 4674 deletions
@@ -1,292 +0,0 @@
from __future__ import annotations
import sys
from typing import Any
import torch
import triton
from sglang.kernels.jit.benchmark.utils import get_benchmark_range, run_benchmark
from sglang.kernels.ops.quantization.mxfp8 import (
es_sm100_mxfp8_blockscaled_grouped_quant,
es_sm100_mxfp8_blockscaled_moe_grouped_gemm,
)
from sglang.test.ci.ci_register import register_cuda_ci
register_cuda_ci(
est_time=5, stage="base-b-kernel-benchmark", runner_config="1-gpu-large"
)
def is_sm100_supported(device=None) -> bool:
if not torch.cuda.is_available():
return False
return (torch.cuda.get_device_capability(device)[0] == 10) and (
torch.version.cuda >= "12.8"
)
_SM100_SUPPORTED = is_sm100_supported()
def _probe_sgl_kernel_group_mm() -> tuple[bool, str]:
if not _SM100_SUPPORTED:
return False, "MXFP8 MoE benchmark requires sm100+ with CUDA 12.8+."
try:
import sgl_kernel # noqa: F401
except Exception as e:
return False, f"import sgl_kernel failed: {e}"
if not hasattr(sgl_kernel, "es_sm100_mxfp8_blockscaled_grouped_mm"):
return False, "sgl_kernel.es_sm100_mxfp8_blockscaled_grouped_mm is missing."
try:
pass
# We assume if it's imported, it works
except Exception as e:
return False, f"calling sgl-kernel grouped_mm op failed: {e}"
return True, ""
_SGL_KERNEL_AVAILABLE, _SGL_KERNEL_REASON = _probe_sgl_kernel_group_mm()
def align(val: int, alignment: int = 128) -> int:
return int((val + alignment - 1) // alignment * alignment)
def _prepare_case(
total_tokens: int, n_g: int, k_g: int, num_experts: int, dtype: torch.dtype
) -> dict[str, Any]:
device = torch.device("cuda")
base = total_tokens // num_experts
rem = total_tokens % num_experts
m_per_expert = [base + (1 if i < rem else 0) for i in range(num_experts)]
expert_offset = 0
expert_offsets = []
aux_expert_offset = 0
aux_expert_offsets = []
a_blockscale_offset = 0
a_blockscale_offsets = []
b_blockscale_offset = 0
b_blockscale_offsets = []
tokens_per_expert_list = []
expert_ranges = []
problem_sizes = []
a_list = []
b_list = []
for g in range(num_experts):
m_g = m_per_expert[g]
tokens_per_expert_list.append(m_g)
expert_ranges.append((expert_offset, expert_offset + m_g))
expert_offsets.append(expert_offset)
expert_offset += m_g
aux_expert_offsets.append(aux_expert_offset)
aux_expert_offset += n_g
a_blockscale_offsets.append(a_blockscale_offset)
a_blockscale_offset += align(m_g, 128)
b_blockscale_offsets.append(b_blockscale_offset)
b_blockscale_offset += n_g # n_g already align to 128 in practice
problem_sizes.append([m_g, n_g, k_g])
a = torch.randn((m_g, k_g), device=device, dtype=dtype) * 0.1
b = torch.randn((n_g, k_g), device=device, dtype=dtype) * 0.1
a_list.append(a)
b_list.append(b)
a = torch.concat(a_list, dim=0)
b = torch.concat(b_list, dim=0)
_expert_offsets = torch.tensor(expert_offsets).to(device=device, dtype=torch.int32)
_aux_expert_offsets = torch.tensor(aux_expert_offsets).to(
device=device, dtype=torch.int32
)
_a_blockscale_offsets = torch.tensor(a_blockscale_offsets).to(
device=device, dtype=torch.int32
)
_b_blockscale_offsets = torch.tensor(b_blockscale_offsets).to(
device=device, dtype=torch.int32
)
_tokens_per_expert = torch.tensor(tokens_per_expert_list).to(
device=device, dtype=torch.int32
)
_problem_sizes = torch.tensor(problem_sizes).to(device=device, dtype=torch.int32)
a_quant = torch.zeros_like(a, dtype=torch.float8_e4m3fn, device=device)
a_scale_factor = torch.zeros(
(a_blockscale_offset, k_g // 32), dtype=torch.uint8, device=device
)
b_quant = torch.zeros_like(b, dtype=torch.float8_e4m3fn, device=device)
b_scale_factor = torch.zeros(
(num_experts * n_g, k_g // 32), dtype=torch.uint8, device=device
)
# Use a global workspace to avoid allocating 1GB every time
workspace = torch.empty((1024, 1024, 1024), dtype=torch.uint8, device=device)
es_sm100_mxfp8_blockscaled_grouped_quant(
a,
_tokens_per_expert,
_expert_offsets,
_a_blockscale_offsets,
a_quant,
a_scale_factor,
)
es_sm100_mxfp8_blockscaled_grouped_quant(
b,
torch.ones_like(_tokens_per_expert) * n_g,
_aux_expert_offsets,
_b_blockscale_offsets,
b_quant,
b_scale_factor,
)
b_quant = b_quant.view(num_experts, n_g, k_g)
b_scale_factor = b_scale_factor.view(num_experts, n_g, k_g // 32)
sgl_b_quant = b_quant.transpose(1, 2)
sgl_b_scale_factor = b_scale_factor.transpose(1, 2)
return {
"a": a,
"b": b.view(num_experts, n_g, k_g),
"b_quant": b_quant,
"a_quant": a_quant,
"b_scale_factor": b_scale_factor,
"a_scale_factor": a_scale_factor,
"expert_offsets": _expert_offsets,
"a_blockscale_offsets": _a_blockscale_offsets,
"tokens_per_expert": _tokens_per_expert,
"problem_sizes": _problem_sizes,
"sgl_b_quant": sgl_b_quant,
"sgl_b_scale_factor": sgl_b_scale_factor,
"workspace": workspace,
"expert_ranges": expert_ranges,
"dtype": dtype,
}
def _sgl_kernel_group_mm(case: dict[str, Any]) -> torch.Tensor:
from sgl_kernel import es_sm100_mxfp8_blockscaled_grouped_mm
a_quant = case["a_quant"]
sgl_b_quant = case["sgl_b_quant"]
a_scale_factor = case["a_scale_factor"]
sgl_b_scale_factor = case["sgl_b_scale_factor"]
problem_sizes = case["problem_sizes"]
expert_offsets = case["expert_offsets"]
a_blockscale_offsets = case["a_blockscale_offsets"]
dtype = case["dtype"]
total_tokens = a_quant.shape[0]
n_g = sgl_b_quant.shape[2]
# sgl-kernel takes output pre-allocated
d = torch.empty((total_tokens, n_g), device=a_quant.device, dtype=dtype)
es_sm100_mxfp8_blockscaled_grouped_mm(
d,
a_quant,
sgl_b_quant,
a_scale_factor,
sgl_b_scale_factor,
problem_sizes,
expert_offsets,
a_blockscale_offsets,
)
return d
shape_range = get_benchmark_range(
full_range=[
# (total_tokens, n_g, k_g, num_experts)
(1024, 4096, 4096, 64),
(2048, 4096, 4096, 64),
(4096, 4096, 4096, 64),
]
+ [
(total_tokens, n_g, k_g, num_experts)
for total_tokens in [32 * (2**i) for i in range(9)] # 32 to 8192
for n_g, k_g, num_experts in [
# DeepSeek-V3/R1, gateup, TP = 1, EP = 8
(4096, 7168, 32),
# DeepSeek-V3/R1, down, TP = 1, EP = 8
(7168, 2048, 32),
]
],
ci_range=[(1024, 2048, 2048, 8)],
)
line_vals = ["jit"]
line_names = ["JIT MXFP8 MoE GroupMM"]
styles = [("green", "-")]
if _SGL_KERNEL_AVAILABLE:
line_vals.append("sgl_kernel")
line_names.append("sgl-kernel MXFP8 MoE GroupMM")
styles.append(("orange", "-"))
@triton.testing.perf_report(
triton.testing.Benchmark(
x_names=["total_tokens", "n_g", "k_g", "num_experts"],
x_vals=shape_range,
x_log=False,
line_arg="provider",
line_vals=line_vals,
line_names=line_names,
styles=styles,
ylabel="us",
plot_name="mxfp8-moe-groupmm-performance",
args={},
)
)
def benchmark(total_tokens, n_g, k_g, num_experts, provider):
case = _prepare_case(total_tokens, n_g, k_g, num_experts, torch.bfloat16)
if provider == "jit":
fn = lambda: es_sm100_mxfp8_blockscaled_moe_grouped_gemm(
case["b_quant"],
case["a_quant"],
case["b_scale_factor"],
case["a_scale_factor"],
case["expert_offsets"],
case["a_blockscale_offsets"],
case["tokens_per_expert"],
case["workspace"],
case["dtype"],
)
elif provider == "sgl_kernel":
fn = lambda: _sgl_kernel_group_mm(case)
else:
raise ValueError(f"Unknown provider: {provider}")
# Warm up
fn()
# Profile
if provider == "jit":
torch.cuda.nvtx.range_push("jit")
fn()
torch.cuda.nvtx.range_pop()
elif provider == "sgl_kernel":
torch.cuda.nvtx.range_push("sgl_kernel")
fn()
torch.cuda.nvtx.range_pop()
return run_benchmark(fn)
if __name__ == "__main__":
if not _SM100_SUPPORTED:
print("[skip] MXFP8 MoE GroupMM benchmark requires sm100+ with CUDA 12.8+.")
sys.exit(0)
if not _SGL_KERNEL_AVAILABLE:
print(f"[info] sgl-kernel baseline unavailable: {_SGL_KERNEL_REASON}")
benchmark.run(print_data=True)
@@ -1,77 +0,0 @@
import itertools
import torch
import triton
import triton.testing
from sglang.kernels.jit.benchmark.utils import (
DEFAULT_DEVICE,
get_benchmark_range,
run_benchmark,
)
from sglang.kernels.ops.speculative.resolve_future_token_ids import (
resolve_future_token_ids_cuda,
)
from sglang.srt.utils import get_compiler_backend
from sglang.test.ci.ci_register import register_amd_ci, register_cuda_ci
register_cuda_ci(
est_time=10, stage="base-b-kernel-benchmark", runner_config="1-gpu-large"
)
register_amd_ci(est_time=10, stage="jit-kernel-benchmark", runner_config="amd")
SIZE_LIST = get_benchmark_range(
full_range=[2**n for n in range(4, 16)], # 16 … 32K elements
ci_range=[256, 4096],
)
configs = list(itertools.product(SIZE_LIST))
def _torch_resolve(input_ids, future_map):
input_ids[:] = torch.where(
input_ids < 0,
future_map[torch.clamp(-input_ids, min=0)],
input_ids,
)
_compiled_resolve = torch.compile(
_torch_resolve, dynamic=True, backend=get_compiler_backend()
)
@triton.testing.perf_report(
triton.testing.Benchmark(
x_names=["size"],
x_vals=configs,
line_arg="provider",
line_vals=["jit", "torch_compile", "torch"],
line_names=["SGL JIT Kernel", "torch.compile", "PyTorch"],
styles=[("blue", "-"), ("green", "-."), ("red", "--")],
ylabel="us",
plot_name="resolve-future-token-ids-performance",
args={},
)
)
def benchmark(size: int, provider: str):
map_size = 8192
future_map = torch.randint(
0, 50000, (map_size,), dtype=torch.int64, device=DEFAULT_DEVICE
)
input_ids = torch.randint(
-map_size + 1, 50000, (size,), dtype=torch.int64, device=DEFAULT_DEVICE
)
if provider == "jit":
fn = lambda: resolve_future_token_ids_cuda(input_ids.clone(), future_map)
elif provider == "torch_compile":
fn = lambda: _compiled_resolve(input_ids.clone(), future_map)
else:
fn = lambda: _torch_resolve(input_ids.clone(), future_map)
return run_benchmark(fn)
if __name__ == "__main__":
benchmark.run(print_data=True)
@@ -1,6 +1,7 @@
"""GPU-free import / registry / selector tests for ``sglang.kernels`` (RFC #29630)."""
import importlib
import importlib.util
import subprocess
import sys
@@ -84,6 +85,20 @@ def test_specs_well_formed():
assert sep == ":" and mod and attr, spec.target
def test_internal_registry_target_modules_exist():
for spec in K.registry.all_specs():
module, _, _ = spec.target.partition(":")
if module.startswith("sglang.kernels."):
assert importlib.util.find_spec(module) is not None, spec.target
def test_sparse_linear_attention_registry_targets_forward_kernel():
spec = K.registry.get_backend(
"diffusion.sparse_linear_attn_fwd", KernelBackend.TRITON
)
assert spec.target.endswith(":_attn_fwd")
def test_single_backend_resolves_without_backend():
assert K.select_kernel("gemm.fp8_scaled_mm").backend is KernelBackend.AOT
@@ -1,136 +0,0 @@
"""Test for the kpool top-k transform JIT kernel.
Ported from the former AOT sgl-kernel test (sgl-kernel/tests/test_topk.py).
The kernel selects pool groups at pool granularity, expands each selected group
to ``pool_size`` token indices, and optionally transforms those token indices
through a page table or a ragged offset.
"""
from __future__ import annotations
import sys
from typing import Optional
import pytest
import torch
from sglang.kernels.ops.moe.kpool_topk_transform import fast_kpool_topk_transform_fused
from sglang.test.ci.ci_register import register_cuda_ci
register_cuda_ci(est_time=60, stage="base-b-kernel-unit", runner_config="1-gpu-large")
def _ref_torch_kpool_transform_impl(
score: torch.Tensor,
lengths: torch.Tensor,
pool_size: int,
topk: int,
page_table: Optional[torch.Tensor] = None,
topk_indices_offset: Optional[torch.Tensor] = None,
seq_lens: Optional[torch.Tensor] = None,
) -> torch.Tensor:
rows = score.shape[0]
group_topk = topk // pool_size
offsets = torch.arange(pool_size, dtype=torch.int32, device=score.device)
out_cols = topk + (pool_size - 1 if seq_lens is not None else 0)
out = torch.full((rows, out_cols), -1, dtype=torch.int32, device=score.device)
for i in range(rows):
length = int(lengths[i].item())
valid_count = min(length, group_topk)
write_pos = 0
if valid_count == 0:
token_ids = torch.empty((0,), dtype=torch.int32, device=score.device)
elif length <= group_topk:
selected = torch.arange(length, dtype=torch.int32, device=score.device)
token_ids = (selected.unsqueeze(1) * pool_size + offsets).reshape(-1)
else:
selected = torch.topk(
score[i, :length], group_topk, dim=-1, sorted=False
).indices.to(torch.int32)
token_ids = (selected.unsqueeze(1) * pool_size + offsets).reshape(-1)
if token_ids.numel() > 0:
if page_table is not None:
token_ids = page_table[i, token_ids.long()].to(torch.int32)
elif topk_indices_offset is not None:
token_ids = token_ids + topk_indices_offset[i].to(torch.int32)
write_pos = valid_count * pool_size
out[i, :write_pos] = token_ids[:write_pos]
if seq_lens is not None:
tail_count = int(seq_lens[i].item()) % pool_size
if tail_count > 0:
raw_tail = length * pool_size + torch.arange(
tail_count, dtype=torch.int32, device=score.device
)
if page_table is not None:
tail = page_table[i, raw_tail.long()].to(torch.int32)
elif topk_indices_offset is not None:
tail = raw_tail + topk_indices_offset[i].to(torch.int32)
else:
tail = raw_tail
out[i, write_pos : write_pos + tail_count] = tail
return out
@pytest.mark.parametrize(
"pool_size,group_topk",
[(16, 128), (16, 160), (16, 192), (16, 224), (8, 256), (4, 512)],
)
@pytest.mark.parametrize("mode", ["raw", "paged", "ragged"])
@pytest.mark.parametrize("append_tail", [False, True])
@torch.inference_mode()
def test_kpool_topk_transform_kernel(
pool_size: int, group_topk: int, mode: str, append_tail: bool
) -> None:
torch.manual_seed(42)
bs = 17
topk = pool_size * group_topk
num_groups = 4096
score = torch.randn(bs, num_groups, dtype=torch.float32, device="cuda")
lengths = torch.randint(
group_topk + 1, num_groups + 1, (bs,), dtype=torch.int32, device="cuda"
)
page_table = None
topk_indices_offset = None
seq_lens = None
tail_counts = torch.randint(0, pool_size, (bs,), dtype=torch.int32, device="cuda")
if append_tail:
seq_lens = lengths * pool_size + tail_counts
if mode == "paged":
page_table = torch.arange(
bs * (num_groups * pool_size + pool_size),
dtype=torch.int32,
device="cuda",
).view(bs, num_groups * pool_size + pool_size)
elif mode == "ragged":
topk_indices_offset = torch.randint(
0, 2048, (bs,), dtype=torch.int32, device="cuda"
)
out_ref = _ref_torch_kpool_transform_impl(
score,
lengths,
pool_size,
topk,
page_table=page_table,
topk_indices_offset=topk_indices_offset,
seq_lens=seq_lens,
)
out_our = fast_kpool_topk_transform_fused(
score,
lengths,
pool_size,
topk,
page_table=page_table,
topk_indices_offset=topk_indices_offset,
seq_lens=seq_lens,
)
torch.cuda.synchronize()
out_ref = torch.sort(out_ref, dim=-1).values
out_our = torch.sort(out_our, dim=-1).values
torch.testing.assert_close(out_our, out_ref, atol=0, rtol=0)
if __name__ == "__main__":
sys.exit(pytest.main([__file__, "-v", "-s"]))
@@ -1,152 +0,0 @@
import random
import sys
import pytest
import torch
from sglang.kernels.ops.quantization.mxfp8 import (
es_sm100_mxfp8_blockscaled_grouped_quant,
es_sm100_mxfp8_blockscaled_moe_grouped_gemm,
)
from sglang.test.ci.ci_register import register_cuda_ci
register_cuda_ci(est_time=5, stage="base-b-kernel-unit", runner_config="1-gpu-large")
def align(val: int, alignment: int = 128) -> int:
return int((val + alignment - 1) // alignment * alignment)
# Copy from: https://github.com/deepseek-ai/DeepGEMM/blob/main/deep_gemm/utils.py
def calc_diff(x, y):
x, y = x.double(), y.double()
denominator = (x * x + y * y).sum()
sim = 2 * (x * y).sum() / denominator
return 1 - sim
def is_sm100_supported(device=None) -> bool:
return (torch.cuda.get_device_capability(device)[0] == 10) and (
torch.version.cuda >= "12.8"
)
@pytest.mark.skipif(
not is_sm100_supported(),
reason="test_mxfp8_moe at jit kernen is only supported on sm100",
)
@pytest.mark.parametrize("num_experts", [8, 16, 32, 64])
@pytest.mark.parametrize("out_dtype", [torch.half, torch.bfloat16])
def test_es_sm100_mxfp8_blockscaled_grouped_mm(num_experts, out_dtype):
device = "cuda"
alignment = 128
n_g = random.randint(1, 64) * alignment
k_g = random.randint(1, 64) * alignment
expert_offset = 0
expert_offsets = []
aux_expert_offset = 0
aux_expert_offsets = []
a_blockscale_offset = 0
a_blockscale_offsets = []
b_blockscale_offset = 0
b_blockscale_offsets = []
a_list = []
b_list = []
ref_d_list = []
tokens_per_expert = []
for g in range(num_experts):
m_g = random.randint(1, 512)
tokens_per_expert.append(m_g)
expert_offsets.append(expert_offset)
expert_offset += m_g
aux_expert_offsets.append(aux_expert_offset)
aux_expert_offset += n_g
a_blockscale_offsets.append(a_blockscale_offset)
a_blockscale_offset += align(m_g, 128)
b_blockscale_offsets.append(b_blockscale_offset)
b_blockscale_offset += n_g # n_g already align to 128
a = torch.normal(
0.0, std=1.0, size=(m_g, k_g), device=device, dtype=out_dtype
) # (M, K):(K, 1)
b = torch.normal(
0.0, std=1.0, size=(n_g, k_g), device=device, dtype=out_dtype
) # (N, K):(K, 1)
a_list.append(a)
b_list.append(b)
ref_d = a @ b.T
ref_d_list.append(ref_d)
a = torch.concat(a_list, dim=0)
b = torch.concat(b_list, dim=0)
_expert_offsets = torch.tensor(expert_offsets).to(device=device, dtype=torch.int32)
_aux_expert_offsets = torch.tensor(aux_expert_offsets).to(
device=device, dtype=torch.int32
)
_a_blockscale_offsets = torch.tensor(a_blockscale_offsets).to(
device=device, dtype=torch.int32
)
_b_blockscale_offsets = torch.tensor(b_blockscale_offsets).to(
device=device, dtype=torch.int32
)
a_quant = torch.zeros_like(a, dtype=torch.float8_e4m3fn, device=device)
a_scale_factor = torch.zeros(
(a_blockscale_offset, k_g // 32), dtype=torch.uint8, device=device
)
b_quant = torch.zeros_like(b, dtype=torch.float8_e4m3fn, device=device)
b_scale_factor = torch.zeros(
(num_experts * n_g, k_g // 32), dtype=torch.uint8, device=device
)
tokens_per_expert = torch.tensor(tokens_per_expert).to(
device=device, dtype=torch.int32
)
workspace = torch.empty((1024, 1024, 1024), dtype=torch.uint8, device=device)
es_sm100_mxfp8_blockscaled_grouped_quant(
a,
tokens_per_expert,
_expert_offsets,
_a_blockscale_offsets,
a_quant,
a_scale_factor,
)
es_sm100_mxfp8_blockscaled_grouped_quant(
b,
torch.ones_like(tokens_per_expert) * n_g,
_aux_expert_offsets,
_b_blockscale_offsets,
b_quant,
b_scale_factor,
)
b_quant = b_quant.view(num_experts, n_g, k_g)
b_scale_factor = b_scale_factor.view(num_experts, n_g, k_g // 32)
d = es_sm100_mxfp8_blockscaled_moe_grouped_gemm(
b_quant,
a_quant,
b_scale_factor,
a_scale_factor,
_expert_offsets,
_a_blockscale_offsets,
tokens_per_expert,
workspace,
a.dtype,
)
for g in range(num_experts):
baseline = ref_d_list[g]
actual = d[expert_offsets[g] : (expert_offsets[g] + tokens_per_expert[g])]
diff = calc_diff(actual, baseline)
assert diff < 0.001
print(
f"m_g={baseline.shape[0]} n_g={n_g} k_g={k_g} num_experts={num_experts}, out_dtype={out_dtype}, diff={diff:.5f}: OK"
)
if __name__ == "__main__":
sys.exit(pytest.main([__file__]))
@@ -1,71 +0,0 @@
import sys
import pytest
import torch
from sglang.kernels.ops.speculative.resolve_future_token_ids import (
resolve_future_token_ids_cuda,
)
from sglang.test.ci.ci_register import register_amd_ci, register_cuda_ci
register_cuda_ci(est_time=9, stage="base-b-kernel-unit", runner_config="1-gpu-large")
register_amd_ci(est_time=9, stage="jit-kernel-unit", runner_config="amd")
def _reference_resolve(input_ids, future_map):
"""Reference implementation using plain torch."""
result = input_ids.clone()
result[:] = torch.where(
result < 0,
future_map[torch.clamp(-result, min=0)],
result,
)
return result
@pytest.mark.parametrize("size", [1, 2, 127, 128, 255, 256, 1024, 4097])
@pytest.mark.parametrize("dtype", [torch.int32, torch.int64])
class TestResolveFutureTokenIds:
def test_all_negative(self, size: int, dtype: torch.dtype) -> None:
map_size = 8192
future_map = torch.randint(0, 50000, (map_size,), dtype=dtype, device="cuda")
# Negative indices in range [-map_size+1, -1]
input_ids = -torch.randint(1, map_size, (size,), dtype=dtype, device="cuda")
expected = _reference_resolve(input_ids, future_map)
resolve_future_token_ids_cuda(input_ids, future_map)
assert torch.equal(input_ids, expected)
def test_all_non_negative(self, size: int, dtype: torch.dtype) -> None:
map_size = 16
future_map = torch.randint(0, 50000, (map_size,), dtype=dtype, device="cuda")
input_ids = torch.randint(0, 50000, (size,), dtype=dtype, device="cuda")
expected = input_ids.clone()
resolve_future_token_ids_cuda(input_ids, future_map)
assert torch.equal(input_ids, expected)
def test_mixed(self, size: int, dtype: torch.dtype) -> None:
map_size = 8192
future_map = torch.randint(0, 50000, (map_size,), dtype=dtype, device="cuda")
# Mix of negative and non-negative
input_ids = torch.randint(
-map_size + 1, 50000, (size,), dtype=dtype, device="cuda"
)
expected = _reference_resolve(input_ids, future_map)
resolve_future_token_ids_cuda(input_ids, future_map)
assert torch.equal(input_ids, expected)
def test_zeros(self, size: int, dtype: torch.dtype) -> None:
map_size = 16
future_map = torch.randint(0, 50000, (map_size,), dtype=dtype, device="cuda")
input_ids = torch.zeros(size, dtype=dtype, device="cuda")
expected = input_ids.clone()
resolve_future_token_ids_cuda(input_ids, future_map)
assert torch.equal(input_ids, expected)
if __name__ == "__main__":
sys.exit(pytest.main([__file__, "-v", "-s"]))