google / google/jaxopt

second-order derivatives with implicit diff

Open
#474 2 comments 6 reactions 0 assignees View on GitHub
enhancement
Dominant language
Python
Stars
1.1k
Forks
76
Avg merge
2d 21h
Merged PRs (30d)
1

Description

I was trying to use the `custom_root` decorator to differentiate through a solver. When I try to take the gradients, it works well. However, if I try to use `jax.hessian`, I get the error that "cannot use forward-mode autodiff with a custom_vjp function". When searching the JAX documents, it shows that we can use both modes of differentiation if and only if we use `custom_jvp` instead of `custom_vjp`.

I saw that internally, `custom_root` implements only a `custom_vjp` rule. Is there any way to to choose the `custom_jvp` rule instead ?

A minimal example is as follows:

```
import jax
import jax.numpy as jnp
import numpy as onp
from jaxopt import implicit_diff

def f(x, theta, X_train, y_train): # Objective function
residual = jnp.dot(X_train, x) - y_train
return (jnp.sum(residual ** 2) + theta * jnp.sum(x ** 2)) / 2

F = jax.grad(f, argnums=0)

@implicit_diff.custom_root(F)
def ridge_solver(init_x, theta, X_train, y_train):
del init_x # Initialization not used in this solver
XX = jnp.dot(X_train.T, X_train)
Xy = jnp.dot(X_train.T, y_train)
I = jnp.eye(X_train.shape[1]) # Identity matrix
# Finds the ridge reg solution by solving a linear system
return jnp.linalg.solve(XX + theta * I, Xy)

init_x = None
# Create some data.
onp.random.seed(0)
X_train = onp.random.randn(100, 10)
y_train = onp.random.randn(100)

print(jax.hessian(ridge_solver, argnums=1)(init_x, 10.0, X_train, y_train))
```

with the error
```
TypeError: can't apply forward-mode autodiff (jvp) to a custom_vjp function.
```

Contributor guide

Open the contributing guide

Assessment

This issue has not been assessed yet.

Get new issues in your inbox

A short digest of beginner-friendly GitHub issues.