(trace_mode)
| 96 | |
| 97 | @pytest.mark.parametrize("trace_mode", [False, True]) |
| 98 | def test_tensor_detach(trace_mode): |
| 99 | @trace(symbolic=True) |
| 100 | def f(x): |
| 101 | y = x.detach() ** 2 |
| 102 | z = y.detach() + 1 |
| 103 | return z.detach() |
| 104 | |
| 105 | x = tensor([1, 2, 3, 4]) |
| 106 | for _ in range(3): |
| 107 | f(x).numpy() |
| 108 | |
| 109 | |
| 110 | @pytest.mark.parametrize("trace_mode", [False, True]) |