diff --git a/torax/_src/fvm/newton_raphson_solve_block.py b/torax/_src/fvm/newton_raphson_solve_block.py index a4baf3ace..458121163 100644 --- a/torax/_src/fvm/newton_raphson_solve_block.py +++ b/torax/_src/fvm/newton_raphson_solve_block.py @@ -242,9 +242,7 @@ def newton_raphson_solve_block( tau_min=tau_min, log_iterations=log_iterations, ) - root_finder = jax_utils.xla_metadata_call( - jax.jit(root_finder), compilation_unit='newton_raphson_root_finder' - ) + root_finder = jax.jit(root_finder, inline=jax.Inline.XLA_LATE) x_root, metadata = root_finder(x0=init_x_new_vec) diff --git a/torax/_src/solver/jax_root_finding.py b/torax/_src/solver/jax_root_finding.py index ed9d1234d..2965bc805 100644 --- a/torax/_src/solver/jax_root_finding.py +++ b/torax/_src/solver/jax_root_finding.py @@ -85,17 +85,11 @@ def root_newton_raphson( def _newton_raphson(f, x, jacobian_fun=None): init_x_new_vec = x - f = jax.jit(f) - - residual_fun = jax_utils.xla_metadata_call( - f, compilation_unit='residual_fun_block' - ) + residual_fun = f if jacobian_fun is None: jacobian_fun = jax.jacfwd(f) - jacobian_fun = jax_utils.xla_metadata_call( - jax.jit(jacobian_fun), compilation_unit='jacobian_fun_block' - ) + jacobian_fun = jax.jit(jacobian_fun, inline=jax.Inline.XLA_LATE) # initialize state dict being passed around Newton-Raphson iterations residual_vec_init_x_new = residual_fun(init_x_new_vec)