def loss_fn(params, X):
h = X
for W, b in params[:-1]:
h = jax.nn.gelu(h @ W + b)
W, b = params[-1]
return (h @ W + b).sum()
def train_step(params, X, lr=0.01): # Loss, gradients, AND the update
loss, grads = jax.value_and_grad(loss_fn)(params, X)
return loss, jax.tree.map(lambda p, g: p - lr * g, params, grads)
key = jax.random.PRNGKey(0)
shapes = [(1024, 1024), (1024, 1024), (1024, 1024)]
params = [(jax.random.normal(k, s) * 0.03, jnp.zeros(s[1]))
for k, s in zip(jax.random.split(key, 3), shapes)]
X = jax.random.normal(key, (512, 1024))
t0 = time.perf_counter()
compiled = jax.jit(train_step).lower(params, X).compile() # AOT: compile now
print(f'ahead-of-time compile: {time.perf_counter() - t0:.1f} s')
print(d2l.Benchmark(lambda: train_step(params, X), desc='eager'))
print(d2l.Benchmark(lambda: compiled(params, X), desc='compiled'))