test_eye.py 1.63 KB
Newer Older
rusty1s's avatar
rusty1s committed
1
from itertools import product
rusty1s's avatar
rusty1s committed
2

rusty1s's avatar
rusty1s committed
3
4
import pytest
from torch_sparse.tensor import SparseTensor
rusty1s's avatar
rusty1s committed
5

rusty1s's avatar
rusty1s committed
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
from .utils import dtypes, devices


@pytest.mark.parametrize('dtype,device', product(dtypes, devices))
def test_eye(dtype, device):
    mat = SparseTensor.eye(3, dtype=dtype, device=device)
    assert mat.storage.index.tolist() == [[0, 1, 2], [0, 1, 2]]
    assert mat.storage.value.tolist() == [1, 1, 1]
    assert len(mat.cached_keys()) == 0

    mat = SparseTensor.eye(3, dtype=dtype, device=device, no_value=True)
    assert mat.storage.index.tolist() == [[0, 1, 2], [0, 1, 2]]
    assert mat.storage.value is None
    assert len(mat.cached_keys()) == 0

    mat = SparseTensor.eye(3, 4, dtype=dtype, device=device, fill_cache=True)
    assert mat.storage.index.tolist() == [[0, 1, 2], [0, 1, 2]]
    assert len(mat.cached_keys()) == 6
    assert mat.storage.rowcount.tolist() == [1, 1, 1]
    assert mat.storage.rowptr.tolist() == [0, 1, 2, 3]
    assert mat.storage.colcount.tolist() == [1, 1, 1, 0]
    assert mat.storage.colptr.tolist() == [0, 1, 2, 3, 3]
    assert mat.storage.csr2csc.tolist() == [0, 1, 2]
    assert mat.storage.csc2csr.tolist() == [0, 1, 2]

    mat = SparseTensor.eye(4, 3, dtype=dtype, device=device, fill_cache=True)
    assert mat.storage.index.tolist() == [[0, 1, 2], [0, 1, 2]]
    assert len(mat.cached_keys()) == 6
    assert mat.storage.rowcount.tolist() == [1, 1, 1, 0]
    assert mat.storage.rowptr.tolist() == [0, 1, 2, 3, 3]
    assert mat.storage.colcount.tolist() == [1, 1, 1]
    assert mat.storage.colptr.tolist() == [0, 1, 2, 3]
    assert mat.storage.csr2csc.tolist() == [0, 1, 2]
    assert mat.storage.csc2csr.tolist() == [0, 1, 2]