@nnx.jit
def eval_chunk(model, X, Y):
return ce_loss(model, X, Y)
def eval_loss(model, X, Y, chunk=4096):
return float(jnp.stack([eval_chunk(model, X[i:i + chunk],
Y[i:i + chunk])
for i in range(0, len(Y), chunk)]).mean())
def train_to_target(model, data, optimizer, X_eval, Y_eval, target,
max_steps, eval_every=10):
@nnx.jit
def step_fn(model, optimizer, X, Y):
loss, grads = nnx.value_and_grad(ce_loss)(model, X, Y)
optimizer.update(model, grads)
return loss
step = 0
while step < max_steps:
for X, Y in data.train_dataloader():
step_fn(model, optimizer, jnp.asarray(X), jnp.asarray(Y))
step += 1
if step % eval_every == 0 and \
eval_loss(model, X_eval, Y_eval) <= target:
return step
if step >= max_steps:
break
return float('inf')