pytorch
33 строки · 922.0 Байт
1
2
3
4
5
6from hypothesis import given
7import numpy as np
8
9from caffe2.python import core
10import caffe2.python.hypothesis_test_util as hu
11
12
13class TestFlatten(hu.HypothesisTestCase):
14@given(X=hu.tensor(min_dim=2, max_dim=4),
15**hu.gcs)
16def test_flatten(self, X, gc, dc):
17for axis in range(X.ndim + 1):
18op = core.CreateOperator(
19"Flatten",
20["X"],
21["Y"],
22axis=axis)
23
24def flatten_ref(X):
25shape = X.shape
26outer = np.prod(shape[:axis]).astype(int)
27inner = np.prod(shape[axis:]).astype(int)
28return np.copy(X).reshape(outer, inner),
29
30self.assertReferenceChecks(gc, op, [X], flatten_ref)
31
32# Check over multiple devices
33self.assertDeviceChecks(dc, op, [X], [0])
34
35
36if __name__ == "__main__":
37import unittest
38unittest.main()
39