Repository navigation
Fix NaN gradients of the CARMA kernel - #284
Open
raashish1601 wants to merge 1 commit into
Open
raashish1601 wants to merge 1 commit into
raashish1601 wants to merge 1 commit into
Conversation
This branch has not been deployed
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
Sign up for free
to join this conversation on GitHub.
Already have an account?
Sign in to comment
Add this suggestion to a batch that can be applied as a single commit.This suggestion is invalid because no changes were made to the code.Suggestions cannot be applied while the pull request is closed.Suggestions cannot be applied while viewing a subset of changes.Only one suggestion per line can be applied in a batch.Add this suggestion to a batch that can be applied as a single commit.Applying suggestions on deleted lines is not supported.You must change the existing code in this line in order to create a valid suggestion.Outdated suggestions cannot be applied.This suggestion has been applied or marked resolved.Suggestions cannot be applied from pending reviews.Suggestions cannot be applied on multi-line comments.Suggestions cannot be applied while the pull request is queued to merge.Suggestion cannot be applied right now. Please check back later.
Fixes #228.
CARMA.__init__builds the observation model from square roots that are computed for every root, real or complex, and then picks the right one withjnp.where. For a real root the complex terms aresqrt(0), and for some complex cases (e.g. CARMA(2,0)) theh1radicand is exactly 0 as well. The derivative ofjnp.sqrtat 0 is infinite, andjnp.wheremultiplies it by 0 for the unused branch, which gives NaN. Sojax.gradof a log likelihood with a CARMA kernel was NaN for CARMA(1,0), CARMA(2,0), CARMA(2,1) and CARMA(3,1) alike, which is what the issue hit in NumPyro.The fix adds a small
_safe_sqrtthat returns 0 with a zero derivative for non-positive input, and feeds it 0 for the terms thatjnp.wherediscards. The values that are used do not change. As a side effect, a radicand that comes out as a tiny negative number from rounding now gives 0 instead of a NaN in the forward pass.With the example from the issue, the gradient now matches central finite differences:
Tests:
test_carma_gradintests/test_kernels/test_quasisep.pychecks that the gradient is finite and passesjax.test_util.check_gradsfor a real root, two complex cases and a CARMA(3,1) with both kinds of roots. All four fail onmain.tests/test_kernels/test_quasisep.pypasses locally, black and ruff at the pre-commit versions are clean, and there is a news fragment innews/228.bugfix.