def train(net, data_iter, lr, num_epochs):
optimizer = nnx.Optimizer(net, optax.adam(lr), wrt=nnx.Param)
animator = d2l.Animator(xlabel='epoch', ylabel='loss',
xlim=[1, num_epochs])
@nnx.jit
def train_step(net, optimizer, center, context_negative, mask, label):
def compute_loss(model):
pred = skip_gram(center, context_negative, model[0], model[1])
l = (loss(pred.reshape(label.shape), label, mask)
/ mask.sum(axis=1) * mask.shape[1])
return l.sum(), l.size
(loss_val, l_size), grads = nnx.value_and_grad(
compute_loss, has_aux=True)(net)
optimizer.update(net, grads)
return loss_val, l_size
for epoch in range(num_epochs):
timer, num_batches = d2l.Timer(), len(data_iter)
# Accumulate on device to avoid per-batch host syncs
loss_sum, count = jnp.array(0.0), jnp.array(0, dtype=jnp.int32)
for i, batch in enumerate(data_iter):
center, context_negative, mask, label = batch
loss_val, l_size = train_step(
net, optimizer, center, context_negative, mask, label)
loss_sum = loss_sum + loss_val
count = count + l_size
if (i + 1) % (num_batches // 5) == 0 or i == num_batches - 1:
animator.add(epoch + (i + 1) / num_batches,
(float(loss_sum / count),))
total_loss = float(loss_sum)
total_count = int(count)
print(f'loss {total_loss / total_count:.3f}, '
f'{total_count / timer.stop():.1f} tokens/sec')
return net