Skip to content

test: add GPU-free unit tests for host-side utility modules - #768

Open
Dashener2 wants to merge 1 commit into
deepseek-ai:mainfrom
Dashener2:test/host-utils-gpu-free
Open

Dashener2 wants to merge 1 commit into
deepseek-ai:mainfrom
Dashener2:test/host-utils-gpu-free

Conversation

@Dashener2

Copy link
Copy Markdown

Problem

Importing deep_ep runs check_nccl_so() and init_jit(), which require the compiled
extension and a NCCL install, so the pure-Python helpers under deep_ep/utils/ could not
be tested on a machine without a GPU. While adding those tests, hash_tensor() turned out
to raise for common inputs.

Root cause

deep_ep/utils/math.py:

def hash_tensor(t: torch.Tensor) -> int:
    return t.view(torch.int).sum().item()

view(torch.int) requires a contiguous tensor with a 4-byte element size. It raises for
bool, int64, float16/bfloat16, and for non-contiguous tensors. hash_tensors is a
public helper under deep_ep.utils.math, so any caller hashing non-float32 test data hits
this.

Fix

Load the host modules directly by path so the tests do not import deep_ep, and make
hash_tensor dtype/layout agnostic:

def hash_tensor(t: torch.Tensor) -> int:
    t = t.contiguous()
    if t.element_size() == 4:
        return t.view(torch.int).sum().item()
    return t.view(torch.uint8).sum().item()

Testing

No GPU, no CUDA and no compiled extension required:

python tests/utils/test_math.py
python tests/utils/test_semantic.py
python tests/utils/test_testing.py

Verified with torch 2.8.0+cpu and numpy 1.26. The added hash test fails against the old
implementation with RuntimeError: self.stride(-1) must be 1 to view Long as Int.

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.
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

1 participant