diff --git a/python/sglang/kernels/ops/mamba/mamba_state_scatter_triton.py b/python/sglang/kernels/ops/mamba/mamba_state_scatter_triton.py index 29be4ccc5..1824fcd96 100644 --- a/python/sglang/kernels/ops/mamba/mamba_state_scatter_triton.py +++ b/python/sglang/kernels/ops/mamba/mamba_state_scatter_triton.py @@ -10,6 +10,8 @@ import torch import triton import triton.language as tl +from sglang.kernels.ops.memory.ptr_table import make_ptr_table + def _require_entry_contiguous_dst( 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) - 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 diff --git a/python/sglang/kernels/ops/memory/ptr_table.py b/python/sglang/kernels/ops/memory/ptr_table.py new file mode 100644 index 000000000..44b9cecf1 --- /dev/null +++ b/python/sglang/kernels/ops/memory/ptr_table.py @@ -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) diff --git a/python/sglang/srt/mem_cache/mamba_slot_fused.py b/python/sglang/srt/mem_cache/mamba_slot_fused.py index ccce228ef..5689d0505 100644 --- a/python/sglang/srt/mem_cache/mamba_slot_fused.py +++ b/python/sglang/srt/mem_cache/mamba_slot_fused.py @@ -24,11 +24,13 @@ import torch import triton import triton.language as tl +from sglang.kernels.ops.memory.ptr_table import make_ptr_table + _BLOCK = 1024 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) layer_stride: torch.Tensor # [T] int64 element stride between layers 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()) to_i64 = lambda xs: torch.tensor(xs, dtype=torch.int64, device=device) return ConvSlotDescriptor( - ptr=to_i64(ptr), + ptr=make_ptr_table(ptr, device=device), feat=to_i64(feat), layer_stride=to_i64(layer_stride), slot_stride=to_i64(slot_stride), diff --git a/python/sglang/srt/mem_cache/memory_pool.py b/python/sglang/srt/mem_cache/memory_pool.py index 5133a2276..0b23b9156 100644 --- a/python/sglang/srt/mem_cache/memory_pool.py +++ b/python/sglang/srt/mem_cache/memory_pool.py @@ -489,7 +489,7 @@ class MambaPool: *physical_conv_shape, ), dtype=conv_dtype, - device="cuda", + device=self.device, ) physical_conv_strides = phys.stride()[2:] window_stride = physical_conv_strides[window_axis] @@ -771,7 +771,7 @@ class MambaPool: temporal_state_shape[2], ), dtype=ssm_dtype, - device="cuda", + device=device, ) # Cache intermediate conv windows (last K-1 inputs) per draft token # during target verify. @@ -839,7 +839,7 @@ class MambaPool: conv_shape[1], ), dtype=conv_dtype, - device="cuda", + device=device, ) for conv_shape in dense_conv_shapes ] diff --git a/scripts/lint/check_registered_tests.py b/scripts/lint/check_registered_tests.py index 8699b67c3..507609bb1 100755 --- a/scripts/lint/check_registered_tests.py +++ b/scripts/lint/check_registered_tests.py @@ -41,6 +41,10 @@ _LEGACY_CUDA_PREFIXES = ("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: """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("/") 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: return [ f"{path}: registered tests must live under " diff --git a/test/registered/layers/mamba/test_mamba2_mixer.py b/test/registered/layers/mamba/test_mamba2_mixer.py index 69b38472b..063b84f37 100644 --- a/test/registered/layers/mamba/test_mamba2_mixer.py +++ b/test/registered/layers/mamba/test_mamba2_mixer.py @@ -10,14 +10,16 @@ from sglang.srt.distributed.device_communicators.custom_all_reduce_utils import update_environment_variables, ) from sglang.srt.distributed.parallel_state import ( + get_default_distributed_backend, init_distributed_environment, initialize_model_parallel, ) from sglang.srt.runtime_context import get_parallel 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_xpu_ci(est_time=60, suite="nightly-xpu-2-gpu", nightly=True) 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( - 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) diff --git a/test/registered/layers/mamba/test_mamba_slot_fused.py b/test/registered/layers/mamba/test_mamba_slot_fused.py index e7fc04afe..484518be4 100644 --- a/test/registered/layers/mamba/test_mamba_slot_fused.py +++ b/test/registered/layers/mamba/test_mamba_slot_fused.py @@ -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 / partial / full index sets, int32 indices, and the strided per-slot-envelope 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 @@ -16,10 +19,29 @@ from sglang.srt.mem_cache.mamba_slot_fused import ( fused_clear_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 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 # 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] -@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): def test_clear_matches_reference(self): - dev = "cuda" + dev = DEVICE for dims, num_layers, pool in CONFIGS: 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): @@ -93,7 +118,7 @@ class TestMambaSlotFused(CustomTestCase): got = [t.clone() for t in base] _ref_clear(ref, 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): self.assertTrue(torch.equal(r, g)) # Cleared slots are exactly zero; the rest is untouched. @@ -104,7 +129,7 @@ class TestMambaSlotFused(CustomTestCase): self.assertTrue(torch.equal(g[:, keep], b[:, keep])) def test_copy_matches_reference(self): - dev = "cuda" + dev = DEVICE for dims, num_layers, pool in CONFIGS: with self.subTest(dims=dims, num_layers=num_layers, pool=pool): base = _make_convs(dims, num_layers, pool, dev, seed=1) @@ -116,7 +141,7 @@ class TestMambaSlotFused(CustomTestCase): got = [t.clone() for t in base] _ref_copy(ref, 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): 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 # buffer (including the other streams' bytes in each slot) must be # bit-exact vs the reference, proving no cross-stream clobber. - dev = "cuda" + dev = DEVICE num_layers, pool = 2, 48 dims = [128, 256, 6144] envelope = sum(CONV_LEN * d for d in dims) @@ -148,7 +173,7 @@ class TestMambaSlotFused(CustomTestCase): ), idx, ) - torch.cuda.synchronize() + get_device_module().synchronize() self.assertTrue(torch.equal(ref_buf, got_buf)) # copy on the same strided layout @@ -164,31 +189,31 @@ class TestMambaSlotFused(CustomTestCase): src, dst, ) - torch.cuda.synchronize() + get_device_module().synchronize() self.assertTrue(torch.equal(ref_buf, got_buf)) def test_empty_indices_is_noop(self): - dev = "cuda" + dev = DEVICE base = _make_convs(HETERO_DIMS, 1, 16, dev, seed=2) got = [t.clone() for t in base] empty = torch.empty(0, dtype=torch.int64, device=dev) desc = build_conv_slot_descriptor(got) fused_clear_conv_slots(desc, empty) fused_copy_conv_slots(desc, empty, empty) - torch.cuda.synchronize() + get_device_module().synchronize() for b, g in zip(base, got): self.assertTrue(torch.equal(b, g)) def test_int32_indices_accepted(self): # 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) idx = torch.tensor([1, 5, 9], dtype=torch.int32, device=dev) ref = [t.clone() for t in base] got = [t.clone() for t in base] _ref_clear(ref, idx.long()) fused_clear_conv_slots(build_conv_slot_descriptor(got), idx) - torch.cuda.synchronize() + get_device_module().synchronize() for r, g in zip(ref, got): self.assertTrue(torch.equal(r, g)) diff --git a/test/registered/unit/layers/test_mamba_state_scatter_triton.py b/test/registered/unit/layers/test_mamba_state_scatter_triton.py index aa0897e10..145f20fe7 100644 --- a/test/registered/unit/layers/test_mamba_state_scatter_triton.py +++ b/test/registered/unit/layers/test_mamba_state_scatter_triton.py @@ -2,10 +2,12 @@ from sglang.test.ci.ci_register import ( register_amd_ci, register_cpu_ci, register_cuda_ci, + register_xpu_ci, ) 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_xpu_ci(est_time=20, suite="stage-b-test-1-gpu-xpu") # The dst layout-contract tests run on CPU (no kernel launch). register_cpu_ci(est_time=6, suite="base-a-test-cpu") @@ -16,6 +18,7 @@ import torch try: from sglang.kernels.ops.mamba.mamba_state_scatter_triton import ( _require_entry_contiguous_dst, + fused_conv_window_scatter_multi, fused_conv_window_scatter_with_mask, fused_mamba_state_scatter_with_mask, ) @@ -23,6 +26,7 @@ try: _FUSED_IMPORT_ERROR = None except Exception as e: # pragma: no cover _require_entry_contiguous_dst = None + fused_conv_window_scatter_multi = None fused_conv_window_scatter_with_mask = None fused_mamba_state_scatter_with_mask = None _FUSED_IMPORT_ERROR = e @@ -31,6 +35,25 @@ from sglang.srt.mem_cache.layout.page_major import ( build_page_major_mamba_views, 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): @@ -365,5 +388,78 @@ class TestMambaStateScatterEnvelopeDst(unittest.TestCase): 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 unittest.main() diff --git a/test/registered/unit/mem_cache/test_mamba_unittest.py b/test/registered/unit/mem_cache/test_mamba_unittest.py index 7bfa12dd4..54b8dd21e 100755 --- a/test/registered/unit/mem_cache/test_mamba_unittest.py +++ b/test/registered/unit/mem_cache/test_mamba_unittest.py @@ -26,10 +26,15 @@ from sglang.srt.mem_cache.radix_cache import RadixKey from sglang.srt.sampling.sampling_params import SamplingParams from sglang.srt.server_args import ServerArgs, set_global_server_args_for_scheduler 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_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): @@ -168,6 +173,8 @@ class TestMamba(unittest.TestCase): conv_dim = 5 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( conv_shape=(window_size, conv_dim), num_mamba_layers=num_mamba_layers, diff --git a/test/registered/unit/memory/test_ptr_table.py b/test/registered/unit/memory/test_ptr_table.py new file mode 100644 index 000000000..0f5d09061 --- /dev/null +++ b/test/registered/unit/memory/test_ptr_table.py @@ -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() diff --git a/test/registered/xpu/test_ptr_table.py b/test/registered/xpu/test_ptr_table.py new file mode 100644 index 000000000..3d22cfbcb --- /dev/null +++ b/test/registered/xpu/test_ptr_table.py @@ -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()