Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
8 changes: 7 additions & 1 deletion deep_ep/utils/math.py
Original file line number Diff line number Diff line change
Expand Up @@ -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:
Expand Down
101 changes: 101 additions & 0 deletions tests/utils/test_math.py
Original file line number Diff line number Diff line change
@@ -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')
66 changes: 66 additions & 0 deletions tests/utils/test_semantic.py
Original file line number Diff line number Diff line change
@@ -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')
45 changes: 45 additions & 0 deletions tests/utils/test_testing.py
Original file line number Diff line number Diff line change
@@ -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')