| 273 | |
| 274 | def test_tensor_gpu_squeeze(): |
| 275 | def check_squeeze(shape, dim, in_layout, expected_out_layout): |
| 276 | arr = cp.random.rand(*shape) |
| 277 | t = TensorGPU(arr, in_layout) |
| 278 | is_squeezed = t.squeeze(dim) |
| 279 | should_squeeze = len(expected_out_layout) < len(in_layout) |
| 280 | arr_squeeze = arr.squeeze(dim) |
| 281 | t_shape = tuple(t.shape()) |
| 282 | assert t_shape == arr_squeeze.shape, f"{t_shape} != {arr_squeeze.shape}" |
| 283 | assert t.layout() == expected_out_layout, f"{t.layout()} != {expected_out_layout}" |
| 284 | assert cp.allclose(arr_squeeze, cp.asanyarray(t)) |
| 285 | assert is_squeezed == should_squeeze, f"{is_squeezed} != {should_squeeze}" |
| 286 | |
| 287 | for dim, shape, in_layout, expected_out_layout in [ |
| 288 | (None, (3, 5, 6), "ABC", "ABC"), |