test_index.py 1.39 KB
Newer Older
1
2
3
4
import dgl
import dgl.ndarray as nd
from dgl.utils import toindex
import numpy as np
5
import backend as F
VoVAllen's avatar
VoVAllen committed
6
import unittest
7

VoVAllen's avatar
VoVAllen committed
8
@unittest.skipIf(dgl.backend.backend_name == "tensorflow", reason="TF doesn't support inplace update")
9
10
11
12
13
14
15
16
def test_dlpack():
    # test dlpack conversion.
    def nd2th():
        ans = np.array([[1., 1., 1., 1.],
                        [0., 0., 0., 0.],
                        [0., 0., 0., 0.]])
        x = nd.array(np.zeros((3, 4), dtype=np.float32))
        dl = x.to_dlpack()
17
        y = F.zerocopy_from_dlpack(dl)
18
        y[0] = 1
19
20
        print(x)
        print(y)
21
22
23
24
25
26
        assert np.allclose(x.asnumpy(), ans)

    def th2nd():
        ans = np.array([[1., 1., 1., 1.],
                        [0., 0., 0., 0.],
                        [0., 0., 0., 0.]])
27
28
        x = F.zeros((3, 4))
        dl = F.zerocopy_to_dlpack(x)
29
30
        y = nd.from_dlpack(dl)
        x[0] = 1
31
32
        print(x)
        print(y)
33
34
        assert np.allclose(y.asnumpy(), ans)

35
    def th2nd_incontiguous():
36
        x = F.astype(F.tensor([[0, 1], [2, 3]]), F.int64)
37
38
39
40
41
42
        ans = np.array([0, 2])
        y = x[:2, 0]
        # Uncomment this line and comment the one below to observe error
        #dl = dlpack.to_dlpack(y)
        dl = F.zerocopy_to_dlpack(y)
        z = nd.from_dlpack(dl)
43
44
        print(x)
        print(z)
45
46
        assert np.allclose(z.asnumpy(), ans)

47
48
    nd2th()
    th2nd()
49
    th2nd_incontiguous()