google / google/jaxopt

differentiating through initialization

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

Description

Hi, all
I found that the parameters of NN model are not updated when using gradient descent solver right before SGD.
```python
import haiku as hk
import jax
import jax.numpy as jnp
import numpy as np
import haiku as hk
import optax as optix

from jax import nn
from jaxopt import LBFGS, GradientDescent
import optax as optix
import numpy as np

data_use = jnp.array([[1., 1.]])
OUTPUT_SIZE = 2

def square_function(x, data):
return jnp.mean((x - data)**2)

def inner_loss(x, lower_bounds = 0, upper_bounds = 1.5):
return jnp.mean(nn.relu(x-upper_bounds)**2 + nn.relu(-x+lower_bounds)**2)

def Gradient_solver(x_init):
gd = GradientDescent(fun = inner_loss, maxiter = 10)
x_opt = gd.run(x_init).params
return x_opt
```

```python
# define my neural network
class MySolver(hk.Module):
def __init__(self, output_size = OUTPUT_SIZE):
super().__init__()
self.output_size = output_size
self.MLP = hk.nets.MLP([self.output_size, self.output_size])
def __call__(self, data):
x_init = self.MLP(data)
return x_init

FuncSolverNet = hk.transform(lambda x : MySolver()(x))
```
```python
# training process!
def train_model(data) -> hk.Params:
rng = jax.random.PRNGKey(428)
opt = optix.sgd(1e-3)

@jax.jit
def outer_loss(params, data):
x_opt = Gradient_solver(FuncSolverNet.apply(params, None, data))
return square_function(x_opt, data)

@jax.jit
def update(params, opt_state, data):
l, grads = jax.value_and_grad(outer_loss)(params, data) # return loss and grad
grads, opt_state = opt.update(grads, opt_state)
params = optix.apply_updates(params, grads)
return l, params, opt_state

# Initialize state.
params = FuncSolverNet.init(rng, np.zeros_like(data))
opt_state = opt.init(params)
for epoch in range(5):
train_loss, params, opt_state = update(params, opt_state, data)
print("epoch {}: train loss total {:.10e}".format(epoch, train_loss))

return params

trained_params = train_model(data_use)
```
```
epoch 0: train loss total 6.2500000000e-01
epoch 1: train loss total 6.2500000000e-01
epoch 2: train loss total 6.2500000000e-01
epoch 3: train loss total 6.2500000000e-01
epoch 4: train loss total 6.2500000000e-01
```
The purpose of the gradient descent solver is solving $x_{opt} = \arg\min_x (max(-x+lowerbound,0))^2+(max(x-upperbound,0))^2$.

As you can check, the parameters of `FuncSolverNet ` are not updated. However, the parameters are updated after I delete the gradient descent solver and start the training process.

```python
# training without gradient descent solver
def train_model(data) -> hk.Params:
rng = jax.random.PRNGKey(428)
opt = optix.sgd(1e-3)

@jax.jit
def outer_loss(params, data): # <--outer_loss without gradient descent solver
x_opt = FuncSolverNet.apply(params, None, data)
return square_function(x_opt, data)

@jax.jit
def update(params, opt_state, data):
l, grads = jax.value_and_grad(outer_loss)(params, data) # return loss and grad
grads, opt_state = opt.update(grads, opt_state)
params = optix.apply_updates(params, grads)
return l, params, opt_state

# Initialize state.
params = FuncSolverNet.init(rng, np.zeros_like(data))
opt_state = opt.init(params)
for epoch in range(5):
train_loss, params, opt_state = update(params, opt_state, data)
print("epoch {}: train loss total {:.10e}".format(epoch, train_loss))

return params

trained_params = train_model(data_use)
```
```
epoch 0: train loss total 4.7954492569e+00
epoch 1: train loss total 4.6865110397e+00
epoch 2: train loss total 4.5809593201e+00
epoch 3: train loss total 4.4786567688e+00
epoch 4: train loss total 4.3794698715e+00
```

I'm not sure why this happens. Could anyone provide a hint to resolve this issue?
Thanks for reading.
MK.

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.