|
1 | 1 | """Check for compatibility with torch.compile.""" |
2 | 2 |
|
| 3 | +import pytest |
3 | 4 | import torch |
4 | | -from hypothesis import given, settings |
5 | | -from hypothesis import strategies as st |
| 5 | +from hypothesis import given |
6 | 6 | from torch.library import opcheck |
7 | 7 |
|
8 | 8 | import torchdtw # noqa: F401 # Need to import it to register dtw operation |
9 | 9 |
|
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 |
20 | 11 |
|
21 | 12 |
|
22 | 13 | @given(x=DIM, y=DIM, low=LOW, high_minus_low=HIGH_MINUS_LOW) |
23 | | -@settings(deadline=None) |
24 | 14 | def test_opcheck_dtw(x: int, y: int, low: float, high_minus_low: float) -> None: |
25 | 15 | """Verify that dtw can be torch compiled.""" |
26 | 16 | sample = make_tensor((x, y), dtype=torch.float32, low=low, high=high_minus_low + low) |
27 | 17 | 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(),)) |
30 | 26 |
|
31 | 27 |
|
32 | 28 | @given(n=BATCH, x=DIM, low=LOW, high_minus_low=HIGH_MINUS_LOW) |
33 | | -@settings(deadline=None) |
34 | 29 | def test_opcheck_dtw_batch_symmetric(n: int, x: int, low: float, high_minus_low: float) -> None: |
35 | 30 | """Verify that dtw_batch can be torch compiled, with symmetric input.""" |
36 | 31 | sample = make_tensor((n, n, x, x), dtype=torch.float32, low=low, high=high_minus_low + low) |
37 | 32 | sx = make_tensor((n,), dtype=torch.long, low=1, high=x) |
38 | 33 | i, j = torch.triu_indices(n, n) |
39 | 34 | sample[i, j] = sample[j, i] |
40 | 35 | 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}) |
43 | 47 |
|
44 | 48 |
|
45 | 49 | @given(n=BATCH, m=BATCH, x=DIM, y=DIM, low=LOW, high_minus_low=HIGH_MINUS_LOW) |
46 | | -@settings(deadline=None) |
47 | 50 | def test_opcheck_dtw_batch_not_symmetric(n: int, m: int, x: int, y: int, low: float, high_minus_low: float) -> None: |
48 | 51 | """Verify that dtw_batch can be torch compiled, with symmetric input.""" |
49 | 52 | sample = make_tensor((n, m, x, y), dtype=torch.float32, low=low, high=high_minus_low + low) |
50 | 53 | sx = make_tensor((n,), dtype=torch.long, low=1, high=x) |
51 | 54 | sy = make_tensor((m,), dtype=torch.long, low=1, high=y) |
52 | 55 | 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