[XPU][Fix] Pack device-pointer tables as uint64 to avoid 64-bit address overflow (#35051)
Co-authored-by: roopaksrivastav <roopak.srivastava@intel.com> Co-authored-by: Claude Opus 5 <noreply@anthropic.com>
This commit is contained in:
co-authored by
roopaksrivastav
Claude Opus 5
parent
3700c4ee26
commit
fd596a474c
@@ -10,6 +10,8 @@ import torch
|
|||||||
import triton
|
import triton
|
||||||
import triton.language as tl
|
import triton.language as tl
|
||||||
|
|
||||||
|
from sglang.kernels.ops.memory.ptr_table import make_ptr_table
|
||||||
|
|
||||||
|
|
||||||
def _require_entry_contiguous_dst(
|
def _require_entry_contiguous_dst(
|
||||||
dst: torch.Tensor, entry_start_dim: int, fn_name: str
|
dst: torch.Tensor, entry_start_dim: int, fn_name: str
|
||||||
@@ -578,7 +580,7 @@ def _conv_multi_build_meta(pairs, block_size: int):
|
|||||||
]
|
]
|
||||||
)
|
)
|
||||||
block_start += triton.cdiv(elem, block_size)
|
block_start += triton.cdiv(elem, block_size)
|
||||||
meta = torch.tensor(rows, dtype=torch.int64, device=pairs[0][0].device)
|
meta = make_ptr_table(rows, device=pairs[0][0].device)
|
||||||
return meta, block_start
|
return meta, block_start
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -0,0 +1,22 @@
|
|||||||
|
"""Device-pointer tables for kernels that address several tensors per launch."""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
from typing import Sequence, Union
|
||||||
|
|
||||||
|
import torch
|
||||||
|
|
||||||
|
|
||||||
|
def make_ptr_table(
|
||||||
|
rows: Union[Sequence[int], Sequence[Sequence[int]]],
|
||||||
|
device: Union[torch.device, str],
|
||||||
|
) -> torch.Tensor:
|
||||||
|
"""Pack ``data_ptr()`` values -- flat, or 2-D with companion columns such as
|
||||||
|
strides -- into an ``int64`` table a kernel bitcasts back to pointers.
|
||||||
|
|
||||||
|
Built unsigned because XPU USM addresses set the top bit, which
|
||||||
|
``dtype=torch.int64`` rejects while unpacking through ``long long``;
|
||||||
|
``view`` moves no bits, so kernels keep their signed element type.
|
||||||
|
Values must be in ``[0, 2**64)``.
|
||||||
|
"""
|
||||||
|
return torch.tensor(rows, dtype=torch.uint64, device=device).view(torch.int64)
|
||||||
@@ -24,11 +24,13 @@ import torch
|
|||||||
import triton
|
import triton
|
||||||
import triton.language as tl
|
import triton.language as tl
|
||||||
|
|
||||||
|
from sglang.kernels.ops.memory.ptr_table import make_ptr_table
|
||||||
|
|
||||||
_BLOCK = 1024
|
_BLOCK = 1024
|
||||||
|
|
||||||
|
|
||||||
class ConvSlotDescriptor(NamedTuple):
|
class ConvSlotDescriptor(NamedTuple):
|
||||||
ptr: torch.Tensor # [T] int64 base byte-addresses
|
ptr: torch.Tensor # [T] int64-viewed base byte-addresses (make_ptr_table)
|
||||||
feat: torch.Tensor # [T] int64 per-slot feature length (elements)
|
feat: torch.Tensor # [T] int64 per-slot feature length (elements)
|
||||||
layer_stride: torch.Tensor # [T] int64 element stride between layers
|
layer_stride: torch.Tensor # [T] int64 element stride between layers
|
||||||
slot_stride: torch.Tensor # [T] int64 element stride between slots
|
slot_stride: torch.Tensor # [T] int64 element stride between slots
|
||||||
@@ -116,7 +118,7 @@ def build_conv_slot_descriptor(tensors: List[torch.Tensor]) -> ConvSlotDescripto
|
|||||||
max_feat = max(max_feat, t[0, 0].numel())
|
max_feat = max(max_feat, t[0, 0].numel())
|
||||||
to_i64 = lambda xs: torch.tensor(xs, dtype=torch.int64, device=device)
|
to_i64 = lambda xs: torch.tensor(xs, dtype=torch.int64, device=device)
|
||||||
return ConvSlotDescriptor(
|
return ConvSlotDescriptor(
|
||||||
ptr=to_i64(ptr),
|
ptr=make_ptr_table(ptr, device=device),
|
||||||
feat=to_i64(feat),
|
feat=to_i64(feat),
|
||||||
layer_stride=to_i64(layer_stride),
|
layer_stride=to_i64(layer_stride),
|
||||||
slot_stride=to_i64(slot_stride),
|
slot_stride=to_i64(slot_stride),
|
||||||
|
|||||||
@@ -489,7 +489,7 @@ class MambaPool:
|
|||||||
*physical_conv_shape,
|
*physical_conv_shape,
|
||||||
),
|
),
|
||||||
dtype=conv_dtype,
|
dtype=conv_dtype,
|
||||||
device="cuda",
|
device=self.device,
|
||||||
)
|
)
|
||||||
physical_conv_strides = phys.stride()[2:]
|
physical_conv_strides = phys.stride()[2:]
|
||||||
window_stride = physical_conv_strides[window_axis]
|
window_stride = physical_conv_strides[window_axis]
|
||||||
@@ -771,7 +771,7 @@ class MambaPool:
|
|||||||
temporal_state_shape[2],
|
temporal_state_shape[2],
|
||||||
),
|
),
|
||||||
dtype=ssm_dtype,
|
dtype=ssm_dtype,
|
||||||
device="cuda",
|
device=device,
|
||||||
)
|
)
|
||||||
# Cache intermediate conv windows (last K-1 inputs) per draft token
|
# Cache intermediate conv windows (last K-1 inputs) per draft token
|
||||||
# during target verify.
|
# during target verify.
|
||||||
@@ -839,7 +839,7 @@ class MambaPool:
|
|||||||
conv_shape[1],
|
conv_shape[1],
|
||||||
),
|
),
|
||||||
dtype=conv_dtype,
|
dtype=conv_dtype,
|
||||||
device="cuda",
|
device=device,
|
||||||
)
|
)
|
||||||
for conv_shape in dense_conv_shapes
|
for conv_shape in dense_conv_shapes
|
||||||
]
|
]
|
||||||
|
|||||||
@@ -41,6 +41,10 @@ _LEGACY_CUDA_PREFIXES = ("stress",)
|
|||||||
|
|
||||||
_TEST_KINDS = {"unit", "kernel", "e2e", "accuracy", "perf", "stress"}
|
_TEST_KINDS = {"unit", "kernel", "e2e", "accuracy", "perf", "stress"}
|
||||||
|
|
||||||
|
# Flat vendor trees. Vendor-only coverage fits no kind above: no XPU/NPU suite
|
||||||
|
# carries the `-kernel-` infix `kernel` needs, and these launch device work.
|
||||||
|
_VENDOR_DIRS = {"amd", "mlx", "musa", "npu", "xpu"}
|
||||||
|
|
||||||
|
|
||||||
def _defines_testcase(tree: ast.AST) -> bool:
|
def _defines_testcase(tree: ast.AST) -> bool:
|
||||||
"""True if the file defines unittest classes, statically or via type()."""
|
"""True if the file defines unittest classes, statically or via type()."""
|
||||||
@@ -126,6 +130,8 @@ def taxonomy_errors(path: str, registries: list, tree: ast.AST) -> list[str]:
|
|||||||
|
|
||||||
parts = path.split("/")
|
parts = path.split("/")
|
||||||
relative_parts = parts[2:] if parts[:2] == ["test", "registered"] else []
|
relative_parts = parts[2:] if parts[:2] == ["test", "registered"] else []
|
||||||
|
if relative_parts and relative_parts[0] in _VENDOR_DIRS:
|
||||||
|
return []
|
||||||
if len(relative_parts) < 3 or relative_parts[0] not in _TEST_KINDS:
|
if len(relative_parts) < 3 or relative_parts[0] not in _TEST_KINDS:
|
||||||
return [
|
return [
|
||||||
f"{path}: registered tests must live under "
|
f"{path}: registered tests must live under "
|
||||||
|
|||||||
@@ -10,14 +10,16 @@ from sglang.srt.distributed.device_communicators.custom_all_reduce_utils import
|
|||||||
update_environment_variables,
|
update_environment_variables,
|
||||||
)
|
)
|
||||||
from sglang.srt.distributed.parallel_state import (
|
from sglang.srt.distributed.parallel_state import (
|
||||||
|
get_default_distributed_backend,
|
||||||
init_distributed_environment,
|
init_distributed_environment,
|
||||||
initialize_model_parallel,
|
initialize_model_parallel,
|
||||||
)
|
)
|
||||||
from sglang.srt.runtime_context import get_parallel
|
from sglang.srt.runtime_context import get_parallel
|
||||||
from sglang.srt.utils import get_device, get_device_count
|
from sglang.srt.utils import get_device, get_device_count
|
||||||
from sglang.test.ci.ci_register import register_cuda_ci
|
from sglang.test.ci.ci_register import register_cuda_ci, register_xpu_ci
|
||||||
|
|
||||||
register_cuda_ci(est_time=30, stage="base-b", runner_config="2-gpu-large")
|
register_cuda_ci(est_time=30, stage="base-b", runner_config="2-gpu-large")
|
||||||
|
register_xpu_ci(est_time=60, suite="nightly-xpu-2-gpu", nightly=True)
|
||||||
|
|
||||||
NUM_GPUS = 2
|
NUM_GPUS = 2
|
||||||
|
|
||||||
@@ -96,9 +98,12 @@ def mixer2_gated_norm_tensor_parallel(
|
|||||||
}
|
}
|
||||||
)
|
)
|
||||||
|
|
||||||
# initialize distributed
|
# nccl on CUDA, xccl on XPU, ...; the parameter default is always "nccl".
|
||||||
init_distributed_environment(
|
init_distributed_environment(
|
||||||
world_size=world_size, rank=local_rank, local_rank=local_rank
|
world_size=world_size,
|
||||||
|
rank=local_rank,
|
||||||
|
local_rank=local_rank,
|
||||||
|
backend=get_default_distributed_backend(device.type),
|
||||||
)
|
)
|
||||||
initialize_model_parallel(tensor_model_parallel_size=world_size)
|
initialize_model_parallel(tensor_model_parallel_size=world_size)
|
||||||
|
|
||||||
|
|||||||
@@ -5,6 +5,9 @@ reference loop that MambaPool.clear_slots / copy_from fall back to.
|
|||||||
Covers heterogeneous conv shapes, single- and multi-layer pools, single /
|
Covers heterogeneous conv shapes, single- and multi-layer pools, single /
|
||||||
partial / full index sets, int32 indices, and the strided per-slot-envelope
|
partial / full index sets, int32 indices, and the strided per-slot-envelope
|
||||||
layout used by page-major / unified pools.
|
layout used by page-major / unified pools.
|
||||||
|
|
||||||
|
Runs on whichever Triton-capable accelerator is present, not CUDA only: the
|
||||||
|
descriptor packs ``data_ptr()`` values, which only overflows off CUDA (#35047).
|
||||||
"""
|
"""
|
||||||
|
|
||||||
import unittest
|
import unittest
|
||||||
@@ -16,10 +19,29 @@ from sglang.srt.mem_cache.mamba_slot_fused import (
|
|||||||
fused_clear_conv_slots,
|
fused_clear_conv_slots,
|
||||||
fused_copy_conv_slots,
|
fused_copy_conv_slots,
|
||||||
)
|
)
|
||||||
from sglang.test.ci.ci_register import register_cuda_ci
|
from sglang.srt.utils import get_device
|
||||||
|
from sglang.srt.utils.common import get_device_module
|
||||||
|
from sglang.test.ci.ci_register import register_cuda_ci, register_xpu_ci
|
||||||
from sglang.test.test_utils import CustomTestCase
|
from sglang.test.test_utils import CustomTestCase
|
||||||
|
|
||||||
register_cuda_ci(est_time=10, stage="base-b", runner_config="1-gpu-small")
|
register_cuda_ci(est_time=10, stage="base-b", runner_config="1-gpu-small")
|
||||||
|
register_xpu_ci(est_time=20, suite="stage-b-test-1-gpu-xpu")
|
||||||
|
|
||||||
|
# Backends these kernels are verified against; extend as others gain Triton.
|
||||||
|
TRITON_DEVICES = ("cuda", "xpu")
|
||||||
|
|
||||||
|
|
||||||
|
def _triton_device():
|
||||||
|
# get_device() raises when the host has no accelerator; importing this
|
||||||
|
# module must not.
|
||||||
|
try:
|
||||||
|
device = get_device()
|
||||||
|
except RuntimeError:
|
||||||
|
return None
|
||||||
|
return device if device in TRITON_DEVICES else None
|
||||||
|
|
||||||
|
|
||||||
|
DEVICE = _triton_device()
|
||||||
|
|
||||||
CONV_LEN = 3
|
CONV_LEN = 3
|
||||||
# Representative hybrid conv-state trailing dims (a couple of KV-projection
|
# Representative hybrid conv-state trailing dims (a couple of KV-projection
|
||||||
@@ -80,10 +102,13 @@ def _ref_copy(convs, src, dst):
|
|||||||
t[:, dst] = t[:, src]
|
t[:, dst] = t[:, src]
|
||||||
|
|
||||||
|
|
||||||
@unittest.skipUnless(torch.cuda.is_available(), "fused conv-slot kernels need CUDA")
|
@unittest.skipUnless(
|
||||||
|
DEVICE is not None,
|
||||||
|
f"fused conv-slot kernels need one of {TRITON_DEVICES}",
|
||||||
|
)
|
||||||
class TestMambaSlotFused(CustomTestCase):
|
class TestMambaSlotFused(CustomTestCase):
|
||||||
def test_clear_matches_reference(self):
|
def test_clear_matches_reference(self):
|
||||||
dev = "cuda"
|
dev = DEVICE
|
||||||
for dims, num_layers, pool in CONFIGS:
|
for dims, num_layers, pool in CONFIGS:
|
||||||
for n in sorted({1, pool // 3, pool}): # single / partial / all slots
|
for n in sorted({1, pool // 3, pool}): # single / partial / all slots
|
||||||
with self.subTest(dims=dims, num_layers=num_layers, pool=pool, n=n):
|
with self.subTest(dims=dims, num_layers=num_layers, pool=pool, n=n):
|
||||||
@@ -93,7 +118,7 @@ class TestMambaSlotFused(CustomTestCase):
|
|||||||
got = [t.clone() for t in base]
|
got = [t.clone() for t in base]
|
||||||
_ref_clear(ref, idx)
|
_ref_clear(ref, idx)
|
||||||
fused_clear_conv_slots(build_conv_slot_descriptor(got), idx)
|
fused_clear_conv_slots(build_conv_slot_descriptor(got), idx)
|
||||||
torch.cuda.synchronize()
|
get_device_module().synchronize()
|
||||||
for r, g in zip(ref, got):
|
for r, g in zip(ref, got):
|
||||||
self.assertTrue(torch.equal(r, g))
|
self.assertTrue(torch.equal(r, g))
|
||||||
# Cleared slots are exactly zero; the rest is untouched.
|
# Cleared slots are exactly zero; the rest is untouched.
|
||||||
@@ -104,7 +129,7 @@ class TestMambaSlotFused(CustomTestCase):
|
|||||||
self.assertTrue(torch.equal(g[:, keep], b[:, keep]))
|
self.assertTrue(torch.equal(g[:, keep], b[:, keep]))
|
||||||
|
|
||||||
def test_copy_matches_reference(self):
|
def test_copy_matches_reference(self):
|
||||||
dev = "cuda"
|
dev = DEVICE
|
||||||
for dims, num_layers, pool in CONFIGS:
|
for dims, num_layers, pool in CONFIGS:
|
||||||
with self.subTest(dims=dims, num_layers=num_layers, pool=pool):
|
with self.subTest(dims=dims, num_layers=num_layers, pool=pool):
|
||||||
base = _make_convs(dims, num_layers, pool, dev, seed=1)
|
base = _make_convs(dims, num_layers, pool, dev, seed=1)
|
||||||
@@ -116,7 +141,7 @@ class TestMambaSlotFused(CustomTestCase):
|
|||||||
got = [t.clone() for t in base]
|
got = [t.clone() for t in base]
|
||||||
_ref_copy(ref, src, dst)
|
_ref_copy(ref, src, dst)
|
||||||
fused_copy_conv_slots(build_conv_slot_descriptor(got), src, dst)
|
fused_copy_conv_slots(build_conv_slot_descriptor(got), src, dst)
|
||||||
torch.cuda.synchronize()
|
get_device_module().synchronize()
|
||||||
for r, g in zip(ref, got):
|
for r, g in zip(ref, got):
|
||||||
self.assertTrue(torch.equal(r, g))
|
self.assertTrue(torch.equal(r, g))
|
||||||
|
|
||||||
@@ -126,7 +151,7 @@ class TestMambaSlotFused(CustomTestCase):
|
|||||||
# kernel reads real strides, so it must handle this; the whole envelope
|
# kernel reads real strides, so it must handle this; the whole envelope
|
||||||
# buffer (including the other streams' bytes in each slot) must be
|
# buffer (including the other streams' bytes in each slot) must be
|
||||||
# bit-exact vs the reference, proving no cross-stream clobber.
|
# bit-exact vs the reference, proving no cross-stream clobber.
|
||||||
dev = "cuda"
|
dev = DEVICE
|
||||||
num_layers, pool = 2, 48
|
num_layers, pool = 2, 48
|
||||||
dims = [128, 256, 6144]
|
dims = [128, 256, 6144]
|
||||||
envelope = sum(CONV_LEN * d for d in dims)
|
envelope = sum(CONV_LEN * d for d in dims)
|
||||||
@@ -148,7 +173,7 @@ class TestMambaSlotFused(CustomTestCase):
|
|||||||
),
|
),
|
||||||
idx,
|
idx,
|
||||||
)
|
)
|
||||||
torch.cuda.synchronize()
|
get_device_module().synchronize()
|
||||||
self.assertTrue(torch.equal(ref_buf, got_buf))
|
self.assertTrue(torch.equal(ref_buf, got_buf))
|
||||||
|
|
||||||
# copy on the same strided layout
|
# copy on the same strided layout
|
||||||
@@ -164,31 +189,31 @@ class TestMambaSlotFused(CustomTestCase):
|
|||||||
src,
|
src,
|
||||||
dst,
|
dst,
|
||||||
)
|
)
|
||||||
torch.cuda.synchronize()
|
get_device_module().synchronize()
|
||||||
self.assertTrue(torch.equal(ref_buf, got_buf))
|
self.assertTrue(torch.equal(ref_buf, got_buf))
|
||||||
|
|
||||||
def test_empty_indices_is_noop(self):
|
def test_empty_indices_is_noop(self):
|
||||||
dev = "cuda"
|
dev = DEVICE
|
||||||
base = _make_convs(HETERO_DIMS, 1, 16, dev, seed=2)
|
base = _make_convs(HETERO_DIMS, 1, 16, dev, seed=2)
|
||||||
got = [t.clone() for t in base]
|
got = [t.clone() for t in base]
|
||||||
empty = torch.empty(0, dtype=torch.int64, device=dev)
|
empty = torch.empty(0, dtype=torch.int64, device=dev)
|
||||||
desc = build_conv_slot_descriptor(got)
|
desc = build_conv_slot_descriptor(got)
|
||||||
fused_clear_conv_slots(desc, empty)
|
fused_clear_conv_slots(desc, empty)
|
||||||
fused_copy_conv_slots(desc, empty, empty)
|
fused_copy_conv_slots(desc, empty, empty)
|
||||||
torch.cuda.synchronize()
|
get_device_module().synchronize()
|
||||||
for b, g in zip(base, got):
|
for b, g in zip(base, got):
|
||||||
self.assertTrue(torch.equal(b, g))
|
self.assertTrue(torch.equal(b, g))
|
||||||
|
|
||||||
def test_int32_indices_accepted(self):
|
def test_int32_indices_accepted(self):
|
||||||
# deferred-clear/COW indices are staged as int32; the wrappers must upcast.
|
# deferred-clear/COW indices are staged as int32; the wrappers must upcast.
|
||||||
dev = "cuda"
|
dev = DEVICE
|
||||||
base = _make_convs(HETERO_DIMS, 1, 32, dev, seed=3)
|
base = _make_convs(HETERO_DIMS, 1, 32, dev, seed=3)
|
||||||
idx = torch.tensor([1, 5, 9], dtype=torch.int32, device=dev)
|
idx = torch.tensor([1, 5, 9], dtype=torch.int32, device=dev)
|
||||||
ref = [t.clone() for t in base]
|
ref = [t.clone() for t in base]
|
||||||
got = [t.clone() for t in base]
|
got = [t.clone() for t in base]
|
||||||
_ref_clear(ref, idx.long())
|
_ref_clear(ref, idx.long())
|
||||||
fused_clear_conv_slots(build_conv_slot_descriptor(got), idx)
|
fused_clear_conv_slots(build_conv_slot_descriptor(got), idx)
|
||||||
torch.cuda.synchronize()
|
get_device_module().synchronize()
|
||||||
for r, g in zip(ref, got):
|
for r, g in zip(ref, got):
|
||||||
self.assertTrue(torch.equal(r, g))
|
self.assertTrue(torch.equal(r, g))
|
||||||
|
|
||||||
|
|||||||
@@ -2,10 +2,12 @@ from sglang.test.ci.ci_register import (
|
|||||||
register_amd_ci,
|
register_amd_ci,
|
||||||
register_cpu_ci,
|
register_cpu_ci,
|
||||||
register_cuda_ci,
|
register_cuda_ci,
|
||||||
|
register_xpu_ci,
|
||||||
)
|
)
|
||||||
|
|
||||||
register_cuda_ci(est_time=7, stage="base-b", runner_config="1-gpu-small")
|
register_cuda_ci(est_time=7, stage="base-b", runner_config="1-gpu-small")
|
||||||
register_amd_ci(est_time=7, suite="stage-b-test-1-gpu-small-amd-mi35x")
|
register_amd_ci(est_time=7, suite="stage-b-test-1-gpu-small-amd-mi35x")
|
||||||
|
register_xpu_ci(est_time=20, suite="stage-b-test-1-gpu-xpu")
|
||||||
# The dst layout-contract tests run on CPU (no kernel launch).
|
# The dst layout-contract tests run on CPU (no kernel launch).
|
||||||
register_cpu_ci(est_time=6, suite="base-a-test-cpu")
|
register_cpu_ci(est_time=6, suite="base-a-test-cpu")
|
||||||
|
|
||||||
@@ -16,6 +18,7 @@ import torch
|
|||||||
try:
|
try:
|
||||||
from sglang.kernels.ops.mamba.mamba_state_scatter_triton import (
|
from sglang.kernels.ops.mamba.mamba_state_scatter_triton import (
|
||||||
_require_entry_contiguous_dst,
|
_require_entry_contiguous_dst,
|
||||||
|
fused_conv_window_scatter_multi,
|
||||||
fused_conv_window_scatter_with_mask,
|
fused_conv_window_scatter_with_mask,
|
||||||
fused_mamba_state_scatter_with_mask,
|
fused_mamba_state_scatter_with_mask,
|
||||||
)
|
)
|
||||||
@@ -23,6 +26,7 @@ try:
|
|||||||
_FUSED_IMPORT_ERROR = None
|
_FUSED_IMPORT_ERROR = None
|
||||||
except Exception as e: # pragma: no cover
|
except Exception as e: # pragma: no cover
|
||||||
_require_entry_contiguous_dst = None
|
_require_entry_contiguous_dst = None
|
||||||
|
fused_conv_window_scatter_multi = None
|
||||||
fused_conv_window_scatter_with_mask = None
|
fused_conv_window_scatter_with_mask = None
|
||||||
fused_mamba_state_scatter_with_mask = None
|
fused_mamba_state_scatter_with_mask = None
|
||||||
_FUSED_IMPORT_ERROR = e
|
_FUSED_IMPORT_ERROR = e
|
||||||
@@ -31,6 +35,25 @@ from sglang.srt.mem_cache.layout.page_major import (
|
|||||||
build_page_major_mamba_views,
|
build_page_major_mamba_views,
|
||||||
mamba_entry_bytes,
|
mamba_entry_bytes,
|
||||||
)
|
)
|
||||||
|
from sglang.srt.utils import get_device
|
||||||
|
from sglang.srt.utils.common import get_device_module
|
||||||
|
from sglang.test.test_utils import CustomTestCase
|
||||||
|
|
||||||
|
# Backends the multi-type scatter is verified against; extend as others gain Triton.
|
||||||
|
TRITON_DEVICES = ("cuda", "xpu")
|
||||||
|
|
||||||
|
|
||||||
|
def _triton_device():
|
||||||
|
# get_device() raises when the host has no accelerator; the layout-contract
|
||||||
|
# cases below still run on the CPU suite.
|
||||||
|
try:
|
||||||
|
device = get_device()
|
||||||
|
except RuntimeError:
|
||||||
|
return None
|
||||||
|
return device if device in TRITON_DEVICES else None
|
||||||
|
|
||||||
|
|
||||||
|
DEVICE = _triton_device()
|
||||||
|
|
||||||
|
|
||||||
def _ref_scatter(dst, src, dst_indices, src_indices, step_indices):
|
def _ref_scatter(dst, src, dst_indices, src_indices, step_indices):
|
||||||
@@ -365,5 +388,78 @@ class TestMambaStateScatterEnvelopeDst(unittest.TestCase):
|
|||||||
torch.testing.assert_close(conv_views[0], expect_conv)
|
torch.testing.assert_close(conv_views[0], expect_conv)
|
||||||
|
|
||||||
|
|
||||||
|
@unittest.skipUnless(
|
||||||
|
DEVICE is not None,
|
||||||
|
f"multi-type conv scatter needs one of {TRITON_DEVICES}",
|
||||||
|
)
|
||||||
|
class TestFusedConvWindowScatterMulti(CustomTestCase):
|
||||||
|
"""Multi-type conv scatter must place every conv type correctly; it reaches
|
||||||
|
them through a host-built device-pointer table (issue #35047)."""
|
||||||
|
|
||||||
|
def _run(self, num_types, n2):
|
||||||
|
# n2 sizes the optional second (interval-crossing track) index set.
|
||||||
|
torch.manual_seed(11)
|
||||||
|
dev = DEVICE
|
||||||
|
layers, slots, batch, steps, dim, km1 = 2, 16, 5, 3, 8, 3
|
||||||
|
pairs = []
|
||||||
|
for t in range(num_types):
|
||||||
|
elems = dim * (t + 1)
|
||||||
|
pairs.append(
|
||||||
|
(
|
||||||
|
torch.randn(
|
||||||
|
(layers, slots, elems, km1), dtype=torch.bfloat16, device=dev
|
||||||
|
),
|
||||||
|
torch.randn(
|
||||||
|
(layers, batch, steps, elems, km1),
|
||||||
|
dtype=torch.bfloat16,
|
||||||
|
device=dev,
|
||||||
|
),
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
|
g = torch.Generator(device=dev).manual_seed(5)
|
||||||
|
# One permutation split across the two sets: their dst slots must be
|
||||||
|
# disjoint, else the two writes to a shared slot race.
|
||||||
|
perm = torch.randperm(slots, device=dev, generator=g).to(torch.int64)
|
||||||
|
dst1, dst2 = perm[:batch], perm[batch : batch + n2]
|
||||||
|
|
||||||
|
def _steps(n):
|
||||||
|
step = torch.randint(
|
||||||
|
0, steps, (n,), device=dev, dtype=torch.int64, generator=g
|
||||||
|
)
|
||||||
|
step[0] = -1 # one rejected row per set must be skipped
|
||||||
|
return step
|
||||||
|
|
||||||
|
step1 = _steps(batch)
|
||||||
|
if n2:
|
||||||
|
step2 = _steps(n2)
|
||||||
|
else:
|
||||||
|
dst2, step2 = None, None
|
||||||
|
|
||||||
|
expect = [dst.clone() for dst, _ in pairs]
|
||||||
|
for exp, (_, src) in zip(expect, pairs):
|
||||||
|
for dsts, stps in ((dst1, step1), (dst2, step2)):
|
||||||
|
if stps is None:
|
||||||
|
continue
|
||||||
|
valid = stps >= 0
|
||||||
|
# src row = the row's position within its own index set
|
||||||
|
rows = torch.arange(stps.numel(), device=dev)[valid]
|
||||||
|
exp[:, dsts[valid]] = src[:, rows, stps[valid]]
|
||||||
|
|
||||||
|
fused_conv_window_scatter_multi(pairs, dst1, step1, dst2, step2)
|
||||||
|
get_device_module().synchronize()
|
||||||
|
for i, (exp, (got, _)) in enumerate(zip(expect, pairs)):
|
||||||
|
torch.testing.assert_close(got, exp, msg=f"conv type {i} mismatch")
|
||||||
|
|
||||||
|
def test_single_type(self):
|
||||||
|
self._run(num_types=1, n2=0)
|
||||||
|
|
||||||
|
def test_multi_type(self):
|
||||||
|
self._run(num_types=3, n2=0)
|
||||||
|
|
||||||
|
def test_multi_type_with_track_set(self):
|
||||||
|
self._run(num_types=2, n2=4)
|
||||||
|
|
||||||
|
|
||||||
if __name__ == "__main__": # pragma: no cover
|
if __name__ == "__main__": # pragma: no cover
|
||||||
unittest.main()
|
unittest.main()
|
||||||
|
|||||||
@@ -26,10 +26,15 @@ from sglang.srt.mem_cache.radix_cache import RadixKey
|
|||||||
from sglang.srt.sampling.sampling_params import SamplingParams
|
from sglang.srt.sampling.sampling_params import SamplingParams
|
||||||
from sglang.srt.server_args import ServerArgs, set_global_server_args_for_scheduler
|
from sglang.srt.server_args import ServerArgs, set_global_server_args_for_scheduler
|
||||||
from sglang.srt.utils import get_device
|
from sglang.srt.utils import get_device
|
||||||
from sglang.test.ci.ci_register import register_amd_ci, register_cuda_ci
|
from sglang.test.ci.ci_register import (
|
||||||
|
register_amd_ci,
|
||||||
|
register_cuda_ci,
|
||||||
|
register_xpu_ci,
|
||||||
|
)
|
||||||
|
|
||||||
register_cuda_ci(est_time=11, stage="base-b", runner_config="1-gpu-small")
|
register_cuda_ci(est_time=11, stage="base-b", runner_config="1-gpu-small")
|
||||||
register_amd_ci(est_time=9, suite="stage-b-test-1-gpu-small-amd")
|
register_amd_ci(est_time=9, suite="stage-b-test-1-gpu-small-amd")
|
||||||
|
register_xpu_ci(est_time=20, suite="stage-b-test-1-gpu-xpu")
|
||||||
|
|
||||||
|
|
||||||
def _event_hashes(events):
|
def _event_hashes(events):
|
||||||
@@ -168,6 +173,8 @@ class TestMamba(unittest.TestCase):
|
|||||||
conv_dim = 5
|
conv_dim = 5
|
||||||
|
|
||||||
pool = object.__new__(WindowFirstMambaPool)
|
pool = object.__new__(WindowFirstMambaPool)
|
||||||
|
# Bypasses __init__, so set the device the allocator reads directly.
|
||||||
|
pool.device = get_device()
|
||||||
physical, view = pool._allocate_deduplicated_conv_window(
|
physical, view = pool._allocate_deduplicated_conv_window(
|
||||||
conv_shape=(window_size, conv_dim),
|
conv_shape=(window_size, conv_dim),
|
||||||
num_mamba_layers=num_mamba_layers,
|
num_mamba_layers=num_mamba_layers,
|
||||||
|
|||||||
@@ -0,0 +1,124 @@
|
|||||||
|
"""Device-pointer tables must survive addresses with the top bit set (#35047).
|
||||||
|
|
||||||
|
Backends whose addresses stay ``< 2**47`` cannot catch this, so the cases below
|
||||||
|
spoof a high address onto CPU tensors and run anywhere. The round-trip through
|
||||||
|
real device memory lives in test/registered/xpu/test_ptr_table.py.
|
||||||
|
"""
|
||||||
|
|
||||||
|
from sglang.test.ci.ci_register import register_cpu_ci
|
||||||
|
|
||||||
|
register_cpu_ci(est_time=10, suite="base-a-test-cpu")
|
||||||
|
|
||||||
|
import unittest
|
||||||
|
|
||||||
|
import torch
|
||||||
|
|
||||||
|
from sglang.kernels.ops.memory.ptr_table import make_ptr_table
|
||||||
|
from sglang.test.test_utils import CustomTestCase
|
||||||
|
|
||||||
|
try:
|
||||||
|
from sglang.kernels.ops.mamba.mamba_state_scatter_triton import (
|
||||||
|
_conv_multi_build_meta,
|
||||||
|
)
|
||||||
|
from sglang.srt.mem_cache.mamba_slot_fused import build_conv_slot_descriptor
|
||||||
|
|
||||||
|
_IMPORT_ERROR = None
|
||||||
|
except Exception as e: # triton is not installed on every CPU runner
|
||||||
|
_conv_multi_build_meta = None
|
||||||
|
build_conv_slot_descriptor = None
|
||||||
|
_IMPORT_ERROR = e
|
||||||
|
|
||||||
|
# A base address from the issue report; the top bit is set.
|
||||||
|
_HIGH_BASE = 0xFFFF85ABD4E00000
|
||||||
|
|
||||||
|
|
||||||
|
class _SpoofedPtrTensor(torch.Tensor):
|
||||||
|
# Only data_ptr() is faked; shape, strides, dtype and device stay real.
|
||||||
|
|
||||||
|
_spoofed_ptr = None
|
||||||
|
|
||||||
|
def data_ptr(self) -> int:
|
||||||
|
if self._spoofed_ptr is None:
|
||||||
|
return super().data_ptr()
|
||||||
|
return self._spoofed_ptr
|
||||||
|
|
||||||
|
|
||||||
|
def _spoof(t: torch.Tensor, offset: int) -> torch.Tensor:
|
||||||
|
spoofed = t.as_subclass(_SpoofedPtrTensor)
|
||||||
|
spoofed._spoofed_ptr = _HIGH_BASE + offset
|
||||||
|
return spoofed
|
||||||
|
|
||||||
|
|
||||||
|
def _conv_pairs(elems):
|
||||||
|
# (dst, src) order, as the multi-type scatter takes them.
|
||||||
|
layers, slots, batch, steps, dim = 2, 8, 4, 3, 1
|
||||||
|
return [
|
||||||
|
(
|
||||||
|
torch.zeros((layers, slots, dim, e), dtype=torch.bfloat16),
|
||||||
|
torch.zeros((layers, batch, steps, dim, e), dtype=torch.bfloat16),
|
||||||
|
)
|
||||||
|
for e in elems
|
||||||
|
]
|
||||||
|
|
||||||
|
|
||||||
|
def _conv_tensors(feats):
|
||||||
|
# [layers, slots, conv_len, feat], as build_conv_slot_descriptor takes them.
|
||||||
|
return [torch.zeros((2, 16, 4, f), dtype=torch.bfloat16) for f in feats]
|
||||||
|
|
||||||
|
|
||||||
|
class TestMakePtrTable(CustomTestCase):
|
||||||
|
def test_full_range_addresses_round_trip(self):
|
||||||
|
# 2**63 is the exact value an int64 build rejects.
|
||||||
|
addrs = [0, 2**47, 2**63 - 1, 2**63, 2**64 - 1]
|
||||||
|
table = make_ptr_table(addrs, device="cpu")
|
||||||
|
self.assertEqual(table.dtype, torch.int64)
|
||||||
|
self.assertEqual(table.view(torch.uint64).tolist(), addrs)
|
||||||
|
|
||||||
|
|
||||||
|
@unittest.skipUnless(_IMPORT_ERROR is None, f"import failed: {_IMPORT_ERROR}")
|
||||||
|
class TestPtrTableCallSites(CustomTestCase):
|
||||||
|
"""Built from a spoofed top-bit-set address, each call site's address columns
|
||||||
|
must survive and its companion columns must match a real-address build."""
|
||||||
|
|
||||||
|
def test_conv_multi_meta_table(self):
|
||||||
|
elems = (128, 64)
|
||||||
|
real = _conv_pairs(elems)
|
||||||
|
spoofed = [
|
||||||
|
(_spoof(dst, 2 * i * 4096), _spoof(src, (2 * i + 1) * 4096))
|
||||||
|
for i, (dst, src) in enumerate(real)
|
||||||
|
]
|
||||||
|
|
||||||
|
real_meta, real_blocks = _conv_multi_build_meta(real, block_size=64)
|
||||||
|
meta, blocks = _conv_multi_build_meta(spoofed, block_size=64)
|
||||||
|
|
||||||
|
self.assertEqual(meta.dtype, torch.int64)
|
||||||
|
self.assertEqual(blocks, real_blocks)
|
||||||
|
# Columns 0/1 are the src/dst base addresses.
|
||||||
|
addrs = meta.view(torch.uint64)[:, :2].tolist()
|
||||||
|
self.assertEqual(
|
||||||
|
addrs, [[src.data_ptr(), dst.data_ptr()] for dst, src in spoofed]
|
||||||
|
)
|
||||||
|
self.assertTrue(torch.equal(meta[:, 2:], real_meta[:, 2:]))
|
||||||
|
|
||||||
|
def test_conv_slot_descriptor(self):
|
||||||
|
feats = (128, 6144)
|
||||||
|
real = _conv_tensors(feats)
|
||||||
|
spoofed = [_spoof(t, i * 4096) for i, t in enumerate(real)]
|
||||||
|
|
||||||
|
real_desc = build_conv_slot_descriptor(real)
|
||||||
|
desc = build_conv_slot_descriptor(spoofed)
|
||||||
|
|
||||||
|
self.assertEqual(desc.ptr.dtype, torch.int64)
|
||||||
|
self.assertEqual(
|
||||||
|
desc.ptr.view(torch.uint64).tolist(), [t.data_ptr() for t in spoofed]
|
||||||
|
)
|
||||||
|
self.assertEqual(desc.num_layers, real_desc.num_layers)
|
||||||
|
self.assertEqual(desc.max_feat_blocks, real_desc.max_feat_blocks)
|
||||||
|
for field in ("feat", "layer_stride", "slot_stride"):
|
||||||
|
self.assertTrue(
|
||||||
|
torch.equal(getattr(desc, field), getattr(real_desc, field)), field
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
unittest.main()
|
||||||
@@ -0,0 +1,34 @@
|
|||||||
|
"""Pointer tables built on real XPU memory must round-trip (#35047).
|
||||||
|
|
||||||
|
Level Zero / SYCL USM hands out addresses with the top bit set, which an int64
|
||||||
|
table cannot hold. The spoofed-address cases that run anywhere live in
|
||||||
|
test/registered/unit/memory/test_ptr_table.py.
|
||||||
|
"""
|
||||||
|
|
||||||
|
from sglang.test.ci.ci_register import register_xpu_ci
|
||||||
|
|
||||||
|
register_xpu_ci(est_time=10, suite="stage-b-test-1-gpu-xpu")
|
||||||
|
|
||||||
|
import unittest
|
||||||
|
|
||||||
|
import torch
|
||||||
|
|
||||||
|
from sglang.kernels.ops.memory.ptr_table import make_ptr_table
|
||||||
|
from sglang.test.test_utils import CustomTestCase
|
||||||
|
|
||||||
|
|
||||||
|
@unittest.skipUnless(torch.xpu.is_available(), "Intel XPU not available")
|
||||||
|
class TestPtrTableOnDeviceMemory(CustomTestCase):
|
||||||
|
def test_real_device_pointers_round_trip(self):
|
||||||
|
ptrs = [
|
||||||
|
torch.zeros(1024, device="xpu", dtype=torch.bfloat16).data_ptr()
|
||||||
|
for _ in range(2)
|
||||||
|
]
|
||||||
|
table = make_ptr_table(ptrs, device="xpu")
|
||||||
|
self.assertEqual(table.dtype, torch.int64)
|
||||||
|
self.assertEqual(table.device.type, "xpu")
|
||||||
|
self.assertEqual(table.view(torch.uint64).cpu().tolist(), ptrs)
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
unittest.main()
|
||||||
Reference in New Issue
Block a user