๐๏ธ Mixed precision training#
The following example uses deepmind jmp library to train for mixed precision.
[4]:
!pip install git+https://github.com/ASEM000/serket --quiet
!pip install git+https://github.com/deepmind/jmp --quiet
[5]:
import serket as sk
import jax
import jax.numpy as jnp
import jmp # https://github.com/deepmind/jmp mixed precision library
import matplotlib.pyplot as plt
half = jnp.float16 # On TPU this should be jnp.bfloat16.
full = jnp.float32
k1, k2 = jax.random.split(jax.random.key(0), 2)
mp_policy = jmp.Policy(compute_dtype=half, param_dtype=full, output_dtype=half)
net = sk.Sequential(
sk.nn.Linear(
in_features=1,
out_features=50,
dtype=mp_policy.param_dtype, # +
weight_init="he_normal",
key=k1,
),
jax.numpy.tanh,
sk.nn.Linear(
in_features=50,
out_features=1,
dtype=mp_policy.param_dtype, # +
weight_init="he_normal",
key=k2,
),
mp_policy.cast_to_output, # +
)
net = sk.tree_mask(net)
x = jnp.linspace(-1, 1, 100)[..., None]
y = jnp.sin(x * 3.14)
print(sk.tree_summary(net))
@jax.jit
def train_step(net: sk.Sequential, x: jax.Array, y: jax.Array) -> sk.Sequential:
def loss_func(net, x, y):
net = sk.tree_unmask(net)
net, x = mp_policy.cast_to_compute((net, x)) # +
ypred = jax.vmap(net)(x)
loss = jnp.mean((ypred - y) ** 2)
return loss
grad = jax.grad(loss_func)(net, x, y)
net = jax.tree_map(lambda p, g: p - 1e-3 * g, net, grad)
return net
for i in range(10_000):
net = train_step(net, x, y)
net = sk.tree_unmask(net)
plt.plot(x, y, label="true")
plt.plot(x, jax.vmap(net)(x), label="pred")
โโโโโโโโโโโโโโโโโโโฌโโโโโโโโโโโฌโโโโโโฌโโโโโโโโ
โName โType โCountโSize โ
โโโโโโโโโโโโโโโโโโโผโโโโโโโโโโโผโโโโโโผโโโโโโโโค
โ.layers[0].weightโf32[1,50] โ50 โ200.00Bโ
โโโโโโโโโโโโโโโโโโโผโโโโโโโโโโโผโโโโโโผโโโโโโโโค
โ.layers[0].bias โf32[50] โ50 โ200.00Bโ
โโโโโโโโโโโโโโโโโโโผโโโโโโโโโโโผโโโโโโผโโโโโโโโค
โ.layers[2].weightโf32[50,1] โ50 โ200.00Bโ
โโโโโโโโโโโโโโโโโโโผโโโโโโโโโโโผโโโโโโผโโโโโโโโค
โ.layers[2].bias โf32[1] โ1 โ4.00B โ
โโโโโโโโโโโโโโโโโโโผโโโโโโโโโโโผโโโโโโผโโโโโโโโค
โฮฃ โSequentialโ151 โ604.00Bโ
โโโโโโโโโโโโโโโโโโโดโโโโโโโโโโโดโโโโโโดโโโโโโโโ
[5]:
[<matplotlib.lines.Line2D at 0x14acf01d0>]