def train(X, contents_Y, styles_Y, lr, num_epochs, lr_decay_epoch):
X, styles_Y_gram = get_inits(X, lr, styles_Y)
schedule = optax.exponential_decay(
lr, transition_steps=lr_decay_epoch, decay_rate=0.8,
staircase=True)
optimizer = optax.adam(schedule)
opt_state = optimizer.init(X)
animator = d2l.Animator(xlabel='epoch', ylabel='loss',
xlim=[10, num_epochs],
legend=['content', 'style', 'TV'],
ncols=2, figsize=(7, 2.5))
@nnx.jit
def train_step(model, X, opt_state):
def loss_fn(X):
contents_Y_hat, styles_Y_hat = extract_features(
X, content_layers, style_layers, model)
contents_l, styles_l, tv_l, total = compute_loss(
X, contents_Y_hat, styles_Y_hat, contents_Y,
styles_Y_gram)
return total, (jnp.stack(contents_l), jnp.stack(styles_l), tv_l)
(_, losses), grads = jax.value_and_grad(
loss_fn, has_aux=True)(X)
updates, opt_state = optimizer.update(grads, opt_state, X)
return optax.apply_updates(X, updates), opt_state, losses
history = []
for epoch in range(num_epochs):
X, opt_state, (contents_l, styles_l, tv_l) = train_step(
net, X, opt_state)
if (epoch + 1) % 10 == 0:
animator.axes[1].imshow(postprocess(X))
animator.add(epoch + 1,
[float(jnp.sum(contents_l)),
float(jnp.sum(styles_l)),
float(tv_l)])
if (epoch + 1) % 50 == 0:
history.append((epoch + 1, float(jnp.sum(contents_l)),
float(jnp.sum(styles_l)), float(tv_l)))
for epoch, content_l, style_l, variation_l in history:
print(f'epoch {epoch}, content {content_l:.3f}, '
f'style {style_l:.3f}, TV {variation_l:.3f}, '
f'total {content_l + style_l + variation_l:.3f}')
return X