(f, primals, tangents)
| 342 | # instead compose reverse-mode AD with reverse-mode AD: |
| 343 | |
| 344 | def hvp_revrev(f, primals, tangents): |
| 345 | _, vjp_fn = vjp(grad(f), *primals) |
| 346 | return vjp_fn(*tangents) |
| 347 | |
| 348 | result_hvp_revrev = hvp_revrev(f, (x,), (tangent,)) |
| 349 | assert torch.allclose(result, result_hvp_revrev[0]) |