def train(model, key, n_burnin, n_steps, B=BATCH):
opt = optax.adam(LR)
params, static = eqx.partition(model, eqx.is_inexact_array)
opt_state = opt.init(params)
w, y, z = eq_state(B)
def burn(carry, k): # simulate, no training
w, y, z = carry
_, (w, y, z) = residual_loss(
eqx.combine(params, static), w, y, z, *draw_shocks(k, B)
)
return (w, y, z), None
(w, y, z), _ = lax.scan(burn, (w, y, z), jr.split(key, n_burnin))
gfn = eqx.filter_value_and_grad(loss_fn, has_aux=True)
def step(carry, k): # one grad step + one sim step
params, opt_state, w, y, z = carry
(loss, (w, y, z)), grads = gfn(
eqx.combine(params, static), w, y, z, draw_shocks(k, B)
)
gp = eqx.filter(grads, eqx.is_inexact_array)
upd, opt_state = opt.update(gp, opt_state, params)
params = eqx.apply_updates(params, upd)
return (params, opt_state, w, y, z), loss
carry, losses = lax.scan(
step, (params, opt_state, w, y, z), jr.split(jr.fold_in(key, 1), n_steps)
)
return eqx.combine(carry[0], static), losses