differentiating through initialization
- 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
Assessment
This issue has not been assessed yet.