registering a hook with pytorch does not change the gradient

Viewed 99

I'm trying to remove nans/infs from a gradient calculation in

the problematic bit is in pytorch's multivariate normal logp. If the mahalanobis distance is too large, then it throws nans/infs. Specifically, it is the triangular_solve which throws bad results and then the .pow(2) will raise the error when the detect_anomaly is turned on:

torch.autograd.set_detect_anomaly(True)
unsquared = torch.triangular_solve(flat_x_swap, flat_L, upper=False)[0]
M_swap = unsquared.pow(2).sum(-2)  # shape = b x c

I was under the impression that I could just introduce a hook like below, that could turn infs/nans into just large (but computable) gradients.

unsquared = torch.triangular_solve(flat_x_swap, flat_L, upper=False)[0]
unsquared.register_hook(lambda grad: torch.where(torch.isfinite(grad), grad, torch.scalar_tensor(-1e+9, dtype=flat_L.dtype)))
M_swap = unsquared.pow(2).sum(-2)  # shape = b x c

However, the error is still raised at that point: RuntimeError: Function 'PowBackward0' returned nan values in its 0th output.

Even if replace -1e+9 with a smaller number say 1, it still fails

What am I doing wrong? All I need to do is cap the gradient at this point!

0 Answers
Related