Skip to content

Commit c69c9bb

Browse files
committed
update tests
1 parent 8aedcbe commit c69c9bb

4 files changed

Lines changed: 86 additions & 51 deletions

File tree

‎benchmark/README.md‎

Lines changed: 5 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -11,6 +11,11 @@ Available implementations:
1111
- Triton: adapter from Whisper, CUDA only.
1212
- PyTorch C++ extension: functions from `torchdtw`.
1313

14+
Run with:
15+
```bash
16+
python -m dtw_benchmark
17+
```
18+
1419
Computation time for DTW on array of shape (n, n), on one A40 GPU:
1520

1621
```

‎tests/conftest.py‎

Lines changed: 39 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,39 @@
1+
"""pytest configuration."""
2+
3+
import pytest
4+
import torch
5+
from hypothesis import settings
6+
from hypothesis import strategies as st
7+
8+
settings.register_profile("default", deadline=None)
9+
settings.load_profile("default")
10+
11+
DIM, BATCH = st.integers(1, 1280), st.integers(1, 3)
12+
LOW, HIGH_MINUS_LOW = st.floats(-100, 100), st.floats(0.1, 100)
13+
14+
15+
def make_tensor(shape: tuple[int, ...], *, dtype: torch.dtype, low: float, high: float) -> torch.Tensor:
16+
"""Build a tensor for testing."""
17+
if low == high and dtype == torch.long:
18+
return torch.ones(shape, dtype=torch.long, device="cpu")
19+
return torch.testing.make_tensor(shape, dtype=dtype, device="cpu", low=low, high=high)
20+
21+
22+
def assert_equal(actual: torch.Tensor, expected: torch.Tensor) -> None:
23+
"""Assert tensors equal."""
24+
torch.testing.assert_close(actual, expected, rtol=0, atol=0)
25+
26+
27+
def pytest_configure(config: pytest.Config) -> None:
28+
"""Add 'requires_gpu' marker."""
29+
config.addinivalue_line("markers", "requires_gpu: skip test if no GPU is available")
30+
31+
32+
def pytest_collection_modifyitems(session: pytest.Session, config: pytest.Config, items: list[pytest.Item]) -> None: # noqa: ARG001
33+
"""Skip tests marked with 'requires_gpu' if CUDA not available."""
34+
if torch.cuda.is_available():
35+
return
36+
skip_gpu = pytest.mark.skip(reason="CUDA not available")
37+
for item in items:
38+
if "requires_gpu" in item.keywords:
39+
item.add_marker(skip_gpu)

‎tests/test_dtw.py‎

Lines changed: 8 additions & 30 deletions
Original file line numberDiff line numberDiff line change
@@ -2,62 +2,40 @@
22

33
import pytest
44
import torch
5-
from hypothesis import given, settings
6-
from hypothesis import strategies as st
5+
from hypothesis import given
76

87
from torchdtw import dtw, dtw_batch
98

10-
rtol, atol = 0, 1e-9
11-
skipifnogpu = pytest.mark.skipif(not torch.cuda.is_available(), reason="No GPU available")
9+
from .conftest import BATCH, DIM, HIGH_MINUS_LOW, LOW, assert_equal, make_tensor
1210

13-
DIM, BATCH = st.integers(1, 1280), st.integers(1, 3)
14-
LOW, HIGH_MINUS_LOW = st.floats(-100, 100), st.floats(0.1, 100)
1511

16-
17-
def make_tensor(shape: tuple[int, ...], *, dtype: torch.dtype, low: float, high: float) -> torch.Tensor:
18-
"""Build a tensor for testing."""
19-
if low == high and dtype == torch.long:
20-
return torch.ones(shape, dtype=torch.long, device="cpu")
21-
return torch.testing.make_tensor(shape, dtype=dtype, device="cpu", low=low, high=high)
22-
23-
24-
@skipifnogpu
12+
@pytest.mark.requires_gpu
2513
@given(x=DIM, y=DIM, low=LOW, high_minus_low=HIGH_MINUS_LOW)
26-
@settings(deadline=None)
2714
def test_dtw(x: int, y: int, low: float, high_minus_low: float) -> None:
2815
"""Compare the output of dtw between CPU and GPU implementations."""
2916
d = make_tensor((x, y), dtype=torch.float32, low=low, high=high_minus_low + low)
30-
torch.testing.assert_close(dtw(d), dtw(d.cuda()).cpu(), rtol=rtol, atol=atol)
17+
assert_equal(dtw(d), dtw(d.cuda()).cpu())
3118

3219

33-
@skipifnogpu
20+
@pytest.mark.requires_gpu
3421
@given(n=BATCH, x=DIM, low=LOW, high_minus_low=HIGH_MINUS_LOW)
35-
@settings(deadline=None)
3622
def test_dtw_batch_symmetric(n: int, x: int, low: float, high_minus_low: float) -> None:
3723
"""Compare the output of dtw_batch between CPU and GPU implementations, symmetric case."""
3824
d = make_tensor((n, n, x, x), dtype=torch.float32, low=low, high=high_minus_low + low)
3925
sx = make_tensor((n,), dtype=torch.long, low=1, high=x + 1)
4026
i, j = torch.triu_indices(n, n)
4127
d[i, j] = d[j, i]
42-
torch.testing.assert_close(
43-
dtw_batch(d, sx, sx, symmetric=True),
44-
dtw_batch(d.cuda(), sx.cuda(), sx.cuda(), symmetric=True).cpu(),
45-
rtol=rtol,
46-
atol=atol,
47-
)
28+
assert_equal(dtw_batch(d, sx, sx, symmetric=True), dtw_batch(d.cuda(), sx.cuda(), sx.cuda(), symmetric=True).cpu())
4829

4930

50-
@skipifnogpu
31+
@pytest.mark.requires_gpu
5132
@given(n=BATCH, m=BATCH, x=DIM, y=DIM, low=LOW, high_minus_low=HIGH_MINUS_LOW)
52-
@settings(deadline=None)
5333
def test_dtw_batch_not_symmetric(n: int, m: int, x: int, y: int, low: float, high_minus_low: float) -> None:
5434
"""Compare the output of dtw_batch between CPU and GPU implementations, non symmetric case."""
5535
d = make_tensor((n, m, x, y), dtype=torch.float32, low=low, high=high_minus_low + low)
5636
sx = make_tensor((n,), dtype=torch.long, low=1, high=x + 1)
5737
sy = make_tensor((m,), dtype=torch.long, low=1, high=y + 1)
58-
torch.testing.assert_close(
38+
assert_equal(
5939
dtw_batch(d, sx, sy, symmetric=False),
6040
dtw_batch(d.cuda(), sx.cuda(), sy.cuda(), symmetric=False).cpu(),
61-
rtol=rtol,
62-
atol=atol,
6341
)

‎tests/test_opcheck.py‎

Lines changed: 34 additions & 21 deletions
Original file line numberDiff line numberDiff line change
@@ -1,54 +1,67 @@
11
"""Check for compatibility with torch.compile."""
22

3+
import pytest
34
import torch
4-
from hypothesis import given, settings
5-
from hypothesis import strategies as st
5+
from hypothesis import given
66
from torch.library import opcheck
77

88
import torchdtw # noqa: F401 # Need to import it to register dtw operation
99

10-
DIM, BATCH = st.integers(1, 1280), st.integers(1, 3)
11-
LOW, HIGH_MINUS_LOW = st.floats(-100, 100), st.floats(0.1, 100)
12-
CUDA_AVAILABLE = torch.cuda.is_available()
13-
14-
15-
def make_tensor(shape: tuple[int, ...], *, dtype: torch.dtype, low: float, high: float) -> torch.Tensor:
16-
"""Build a tensor for testing."""
17-
if low == high and dtype == torch.long:
18-
return torch.ones(shape, dtype=torch.long, device="cpu")
19-
return torch.testing.make_tensor(shape, dtype=dtype, device="cpu", low=low, high=high)
10+
from .conftest import BATCH, DIM, HIGH_MINUS_LOW, LOW, make_tensor
2011

2112

2213
@given(x=DIM, y=DIM, low=LOW, high_minus_low=HIGH_MINUS_LOW)
23-
@settings(deadline=None)
2414
def test_opcheck_dtw(x: int, y: int, low: float, high_minus_low: float) -> None:
2515
"""Verify that dtw can be torch compiled."""
2616
sample = make_tensor((x, y), dtype=torch.float32, low=low, high=high_minus_low + low)
2717
opcheck(torch.ops.torchdtw.dtw.default, (sample,))
28-
if CUDA_AVAILABLE:
29-
opcheck(torch.ops.torchdtw.dtw.default, (sample.cuda(),))
18+
19+
20+
@pytest.mark.requires_gpu
21+
@given(x=DIM, y=DIM, low=LOW, high_minus_low=HIGH_MINUS_LOW)
22+
def test_opcheck_dtw_cuda(x: int, y: int, low: float, high_minus_low: float) -> None:
23+
"""Verify that dtw can be torch compiled on CUDA."""
24+
sample = make_tensor((x, y), dtype=torch.float32, low=low, high=high_minus_low + low)
25+
opcheck(torch.ops.torchdtw.dtw.default, (sample.cuda(),))
3026

3127

3228
@given(n=BATCH, x=DIM, low=LOW, high_minus_low=HIGH_MINUS_LOW)
33-
@settings(deadline=None)
3429
def test_opcheck_dtw_batch_symmetric(n: int, x: int, low: float, high_minus_low: float) -> None:
3530
"""Verify that dtw_batch can be torch compiled, with symmetric input."""
3631
sample = make_tensor((n, n, x, x), dtype=torch.float32, low=low, high=high_minus_low + low)
3732
sx = make_tensor((n,), dtype=torch.long, low=1, high=x)
3833
i, j = torch.triu_indices(n, n)
3934
sample[i, j] = sample[j, i]
4035
opcheck(torch.ops.torchdtw.dtw_batch.default, (sample, sx, sx), {"symmetric": True})
41-
if CUDA_AVAILABLE:
42-
opcheck(torch.ops.torchdtw.dtw_batch.default, (sample.cuda(), sx.cuda(), sx.cuda()), {"symmetric": True})
36+
37+
38+
@pytest.mark.requires_gpu
39+
@given(n=BATCH, x=DIM, low=LOW, high_minus_low=HIGH_MINUS_LOW)
40+
def test_opcheck_dtw_batch_symmetric_cuda(n: int, x: int, low: float, high_minus_low: float) -> None:
41+
"""Verify that dtw_batch can be torch compiled on CUDA, with symmetric input."""
42+
sample = make_tensor((n, n, x, x), dtype=torch.float32, low=low, high=high_minus_low + low)
43+
sx = make_tensor((n,), dtype=torch.long, low=1, high=x)
44+
i, j = torch.triu_indices(n, n)
45+
sample[i, j] = sample[j, i]
46+
opcheck(torch.ops.torchdtw.dtw_batch.default, (sample.cuda(), sx.cuda(), sx.cuda()), {"symmetric": True})
4347

4448

4549
@given(n=BATCH, m=BATCH, x=DIM, y=DIM, low=LOW, high_minus_low=HIGH_MINUS_LOW)
46-
@settings(deadline=None)
4750
def test_opcheck_dtw_batch_not_symmetric(n: int, m: int, x: int, y: int, low: float, high_minus_low: float) -> None:
4851
"""Verify that dtw_batch can be torch compiled, with symmetric input."""
4952
sample = make_tensor((n, m, x, y), dtype=torch.float32, low=low, high=high_minus_low + low)
5053
sx = make_tensor((n,), dtype=torch.long, low=1, high=x)
5154
sy = make_tensor((m,), dtype=torch.long, low=1, high=y)
5255
opcheck(torch.ops.torchdtw.dtw_batch.default, (sample, sx, sy), {"symmetric": False})
53-
if CUDA_AVAILABLE:
54-
opcheck(torch.ops.torchdtw.dtw_batch.default, (sample.cuda(), sx.cuda(), sy.cuda()), {"symmetric": False})
56+
57+
58+
@pytest.mark.requires_gpu
59+
@given(n=BATCH, m=BATCH, x=DIM, y=DIM, low=LOW, high_minus_low=HIGH_MINUS_LOW)
60+
def test_opcheck_dtw_batch_not_symmetric_cuda(
61+
n: int, m: int, x: int, y: int, low: float, high_minus_low: float
62+
) -> None:
63+
"""Verify that dtw_batch can be torch compiled on CUDA, with symmetric input."""
64+
sample = make_tensor((n, m, x, y), dtype=torch.float32, low=low, high=high_minus_low + low)
65+
sx = make_tensor((n,), dtype=torch.long, low=1, high=x)
66+
sy = make_tensor((m,), dtype=torch.long, low=1, high=y)
67+
opcheck(torch.ops.torchdtw.dtw_batch.default, (sample.cuda(), sx.cuda(), sy.cuda()), {"symmetric": False})

0 commit comments

Comments
 (0)