From 4bee4d54ac9e2feada735382b9fd2ed6ac357a79 Mon Sep 17 00:00:00 2001 From: Dashener2 Date: Tue, 29 Sep 2026 21:06:55 +0800 Subject: [PATCH] test: add GPU-free unit tests for host-side utility modules Importing deep_ep runs check_nccl_so() and init_jit(), which need the compiled extension and a NCCL install, so the pure-Python helpers under deep_ep/utils/ had no tests that could run on a machine without a GPU. Load those modules directly by path and cover math, semantic and testing.parse_num_bytes. hash_tensor() assumed a contiguous tensor with 4-byte elements and raised for bool, int64, float16 and non-contiguous inputs; normalize through a byte view for other element sizes. --- deep_ep/utils/math.py | 8 ++- tests/utils/test_math.py | 101 +++++++++++++++++++++++++++++++++++ tests/utils/test_semantic.py | 66 +++++++++++++++++++++++ tests/utils/test_testing.py | 45 ++++++++++++++++ 4 files changed, 219 insertions(+), 1 deletion(-) create mode 100644 tests/utils/test_math.py create mode 100644 tests/utils/test_semantic.py create mode 100644 tests/utils/test_testing.py diff --git a/deep_ep/utils/math.py b/deep_ep/utils/math.py index 199418968..4d31a1496 100644 --- a/deep_ep/utils/math.py +++ b/deep_ep/utils/math.py @@ -80,7 +80,13 @@ def create_grouped_scores(scores: torch.Tensor, group_idx: torch.Tensor, num_gro def hash_tensor(t: torch.Tensor) -> int: - return t.view(torch.int).sum().item() + # `view(torch.int)` needs a contiguous tensor whose elements are 4 bytes wide, + # so it raises for non-contiguous tensors and for dtypes such as `bool`, + # `int64` or `float16`. Normalize with a byte view to stay dtype-agnostic. + t = t.contiguous() + if t.element_size() == 4: + return t.view(torch.int).sum().item() + return t.view(torch.uint8).sum().item() def hash_tensors(*tensors) -> int: diff --git a/tests/utils/test_math.py b/tests/utils/test_math.py new file mode 100644 index 000000000..6ca68b728 --- /dev/null +++ b/tests/utils/test_math.py @@ -0,0 +1,101 @@ +import importlib.util +import sys +from pathlib import Path + +import torch + +REPO_ROOT = Path(__file__).resolve().parents[2] + + +def _load_module(name: str, relpath: str): + """Load a host-side module by path. + + Importing `deep_ep` runs `check_nccl_so()` and `init_jit()`, which need the + compiled extension and a NCCL install. These helpers are pure host logic, so + they are loaded directly to keep the tests runnable without a GPU build. + """ + spec = importlib.util.spec_from_file_location(name, REPO_ROOT / relpath) + module = importlib.util.module_from_spec(spec) + sys.modules[name] = module + spec.loader.exec_module(module) + return module + + +math = _load_module('deep_ep_utils_math', 'deep_ep/utils/math.py') + + +def test_ceil_div_and_align(): + assert math.ceil_div(0, 128) == 0 + assert math.ceil_div(1, 128) == 1 + assert math.ceil_div(128, 128) == 1 + assert math.ceil_div(129, 128) == 2 + assert math.align(0, 128) == 0 + assert math.align(1, 128) == 128 + assert math.align(128, 128) == 128 + assert math.align(129, 128) == 256 + + +def test_safe_div(): + assert math.safe_div(6, 3) == 2 + assert math.safe_div(0, 5) == 0 + assert math.safe_div(0, 0) == 0 + try: + math.safe_div(1, 0) + except ZeroDivisionError: + pass + else: + raise AssertionError('A non-zero numerator over zero must re-raise') + + +def test_calc_diff(): + x = torch.randn(4, 8) + assert math.calc_diff(x, x.clone()) == 0 + + +def test_inplace_unique(): + x = torch.tensor([[0, 1, 0, 2], [3, 3, -1, -1]]) + math.inplace_unique(x, num_slots=4) + expected = [{0, 1, 2}, {3}] + for row, kept_expected in zip(x.tolist(), expected): + kept = [value for value in row if value >= 0] + assert len(kept) == len(set(kept)), 'kept slots must be unique' + assert set(kept) == kept_expected + + +def test_hash_tensor_accepts_any_dtype_and_layout(): + values = { + 'bool': torch.tensor([True, False, True, False]), + 'int64': torch.arange(8), + 'float16': torch.arange(8).half(), + 'float32': torch.arange(8), + 'float64': torch.arange(8).double(), + } + for name, t in values.items(): + assert isinstance(math.hash_tensor(t), int) + # Stable for equal content. + assert math.hash_tensor(t) == math.hash_tensor(t.clone()) + # Non-contiguous and empty tensors must not raise. + math.hash_tensor(torch.arange(16).reshape(4, 4).t()) + math.hash_tensor(torch.tensor([])) + + +def test_hash_tensors_nested_and_none(): + a, b = torch.arange(4), torch.zeros(4) + expected = math.hash_tensor(a) ^ math.hash_tensor(b) + assert math.hash_tensors(a, b, None) == expected + assert math.hash_tensors(a, [b], None) == expected + + +def test_count_bytes(): + a = torch.zeros(3, dtype=torch.float32) + b = torch.zeros(2, dtype=torch.int64) + assert math.count_bytes(a, b) == 3 * 4 + 2 * 8 + assert math.count_bytes([a, b], None) == 3 * 4 + 2 * 8 + assert math.count_bytes(None) == 0 + + +if __name__ == '__main__': + for _name, _fn in sorted(globals().items()): + if _name.startswith('test_') and callable(_fn): + _fn() + print('All math tests passed') diff --git a/tests/utils/test_semantic.py b/tests/utils/test_semantic.py new file mode 100644 index 000000000..7904e72a7 --- /dev/null +++ b/tests/utils/test_semantic.py @@ -0,0 +1,66 @@ +import importlib.util +import sys +from pathlib import Path + +REPO_ROOT = Path(__file__).resolve().parents[2] + + +def _load_module(name: str, relpath: str): + """Load a host-side module by path, avoiding `deep_ep`'s GPU initialization.""" + spec = importlib.util.spec_from_file_location(name, REPO_ROOT / relpath) + module = importlib.util.module_from_spec(spec) + sys.modules[name] = module + spec.loader.exec_module(module) + return module + + +semantic = _load_module('deep_ep_utils_semantic', 'deep_ep/utils/semantic.py') + + +def test_value_or(): + assert semantic.value_or(None, 5) == 5 + assert semantic.value_or(0, 5) == 0 + assert semantic.value_or(False, True) is False + + +def test_weak_lru_caches_per_instance(): + calls = [] + + class Obj: + + @semantic.weak_lru(maxsize=None) + def compute(self, key): + calls.append((id(self), key)) + return key * 2 + + a, b = Obj(), Obj() + assert a.compute(3) == 6 + assert a.compute(3) == 6 + assert len(calls) == 1, 'repeated calls on the same instance must hit the cache' + assert b.compute(3) == 6 + assert len(calls) == 2, 'a different instance must get its own cache entry' + + +def test_weak_lru_releases_referents(): + class Obj: + + @semantic.weak_lru(maxsize=None) + def compute(self, key): + return key + + obj = Obj() + assert obj.compute(1) == 1 + # A weak reference must not keep the instance alive. + import gc + import weakref + ref = weakref.ref(obj) + del obj + gc.collect() + assert ref() is None + + +if __name__ == '__main__': + for _name, _fn in sorted(globals().items()): + if _name.startswith('test_') and callable(_fn): + _fn() + print('All semantic tests passed') diff --git a/tests/utils/test_testing.py b/tests/utils/test_testing.py new file mode 100644 index 000000000..dd08c8166 --- /dev/null +++ b/tests/utils/test_testing.py @@ -0,0 +1,45 @@ +import importlib.util +import sys +from pathlib import Path + +REPO_ROOT = Path(__file__).resolve().parents[2] + + +def _load_module(name: str, relpath: str): + """Load a host-side module by path, avoiding `deep_ep`'s GPU initialization.""" + spec = importlib.util.spec_from_file_location(name, REPO_ROOT / relpath) + module = importlib.util.module_from_spec(spec) + sys.modules[name] = module + spec.loader.exec_module(module) + return module + + +testing = _load_module('deep_ep_utils_testing', 'deep_ep/utils/testing.py') + + +def test_parse_num_bytes_accepts_binary_suffixes(): + assert testing.parse_num_bytes('1') == 1 + assert testing.parse_num_bytes('1B') == 1 + assert testing.parse_num_bytes('1k') == 1 << 10 + assert testing.parse_num_bytes('1K') == 1 << 10 + assert testing.parse_num_bytes('64M') == 1 << 26 + assert testing.parse_num_bytes('2G') == 1 << 31 + assert testing.parse_num_bytes('1GiB') == 1 << 30 + assert testing.parse_num_bytes(' 4 K ') == 1 << 12 + assert testing.parse_num_bytes('1.5M') == int(1.5 * (1 << 20)) + + +def test_parse_num_bytes_rejects_bad_input(): + for bad in ('', 'abc', '-1', '0', '1e3', '1.2.3', '1G1', 'nan'): + try: + testing.parse_num_bytes(bad) + except ValueError: + continue + raise AssertionError(f'{bad!r} should have been rejected') + + +if __name__ == '__main__': + for _name, _fn in sorted(globals().items()): + if _name.startswith('test_') and callable(_fn): + _fn() + print('All testing-utils tests passed')