๐ŸŽš๏ธ Mixed precision training

๐ŸŽš๏ธ 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>]
../_images/notebooks_%5Bguides%5D%5Bcore%5Dmixed_precision_2_2.png