(is_varnode)
| 476 | |
| 477 | @pytest.mark.parametrize("is_varnode", [True, False]) |
| 478 | def test_squeeze(is_varnode): |
| 479 | if is_varnode: |
| 480 | network = Network() |
| 481 | saved_symbolic_shape = set_symbolic_shape(False) |
| 482 | else: |
| 483 | network = None |
| 484 | |
| 485 | x = Tensor(np.array([1, 2], dtype=np.int32).reshape(1, 1, 2, 1)) |
| 486 | y = F.squeeze(x, -1) |
| 487 | np.testing.assert_equal(y.numpy(), np.array([[[1, 2]]]).astype(np.int32)) |
| 488 | |
| 489 | x = np.arange(6, dtype="float32").reshape(1, 2, 3, 1) |
| 490 | xx = make_tensor(x, network) |
| 491 | |
| 492 | for axis in [None, 3, -4, (3, -4)]: |
| 493 | y = np.squeeze(x, axis) |
| 494 | yy = F.squeeze(xx, axis) |
| 495 | np.testing.assert_equal(y, yy.numpy()) |
| 496 | |
| 497 | if is_varnode: |
| 498 | set_symbolic_shape(saved_symbolic_shape) |
| 499 | |
| 500 | |
| 501 | @pytest.mark.parametrize("is_varnode", [True, False]) |
nothing calls this directly
no test coverage detected