MCPcopy Create free account
hub / github.com/PersiaML/PERSIA / assert_ndarray_base_data

Function assert_ndarray_base_data

test/test_ctx.py:28–38  ·  view source on GitHub ↗
(
    ndarray_base_data_list: List[np.ndarray],
    tensors: List[torch.Tensor],
    use_cuda: bool,
)

Source from the content-addressed store, hash-verified

26
27
28def assert_ndarray_base_data(
29 ndarray_base_data_list: List[np.ndarray],
30 tensors: List[torch.Tensor],
31 use_cuda: bool,
32):
33 assert len(ndarray_base_data_list) == len(tensors)
34 for ndarray_base_data, tensor in zip(ndarray_base_data_list, tensors):
35 if use_cuda:
36 tensor = tensor.cpu()
37
38 np.testing.assert_equal(ndarray_base_data, tensor.numpy())
39
40
41def assert_id_type_feature_data(tensors: List[torch.Tensor], config: dict):

Callers 1

test_data_ctxFunction · 0.85

Calls 1

numpyMethod · 0.80

Tested by

no test coverage detected