def forward(params, x): x = jax.nn.relu(x @ params['w1'] + params['b1']) return x @ params['w2'] + params['b2'] def train_step(params, x, y): loss, grads = jax.value_and_grad(loss_fn)(params, x, y) new_params = update_params(params, grads) return new_params, loss # Must return the new state # You're always passing params around params, loss = train_step(params, x_batch, y_batch) params, loss = train_step(params, x_batch, y_batch) params, loss = train_step(params, x_batch, y_batch) __ __