本文拷贝自我的知乎文章 👇

前言


本文适用对象:任何接触过 TensorFlowPytorchKeras 并且已经开始了解或尝鲜 Jax 的人群。如果是没有接触过任何深度学习框架的人群,这篇文章可能不适合你。在开始学习之前,你应该对 PyTorch 或 TensorFlow 有一定的了解。Jax 可能是一个比较难学的库,但值得一学。为什么使用 Jax 的理由这里就不多赘述了。就我个人而言,tf 或者 torch 自定义损失函数训练的速度实在是不太满意,即使这过程中换用了 numba 仍然差强人意,加之我是 tf 的旧党,所以不难预料的投入了 Jax 的怀抱。

正文开始


如果是经常用 Tensorflow Keras 或者 Pytorch Lightning 的炼丹师,一定会喜欢 fit 这个方法。所以本文以实现一个简单且常用的 fit 方法来快速上手 Jax,而且实现的这个 fit 方法基本上可以复用在很多项目中。另外再次强调,这篇文章可能不适合入门,但是很适合快速上工(从删库到跑路)。

本文实现的 fit 方法需要安装如下依赖,如果你已经使用过 Jax ,基本以下依赖库想必都已经了解了。

本文也是用的这个经典组合:Jax + Flax + Optax + Orbax,硬件加速 + 网络结构 + 损失函数 + 保存储存点

快速上手


先看看一个训练模型的模板,但只需要修改脚本中的三个关键代码部分。

import jax, flax, optax, orbax
from fit import lr_schedule, TrainState

# 准备你自己的数据集
train_ds, test_ds = your_dataset()
# 学习率
lr_fn = lr_schedule(
    base_lr=1e-3,
    steps_per_epoch=len(train_ds),
    epochs=100,
    warmup_epochs=5,
)

# key 1: 你的模型
model = YourModel()

# 初始化 key 和你的模型
key = jax.random.PRNGKey(0)
x = jnp.ones((1, 28, 28, 1)) # MNIST 示例输入大小
# 注意这里 train=True, 区别模型的训练和评价模式
var = model.init(key, x, train=True)
# 固定模板,直接复制就能用
state = TrainState.create(
    apply_fn=model.apply,
    params=var['params'],
    batch_stats=var['batch_stats'],
    tx=optax.inject_hyperparams(optax.adam)(lr_fn),
)

# 你的训练函数,详情参考下个章节
@jax.jit
def train_step():
    # key 2: 你的损失函数
    def loss_fn():
        ...
    return state, loss_dict, opt_state

# 你的评价函数
@jax.jit
def eval_step():
    # key 3: 你的评价函数
    ...
    return acc

# 一些必要的参数,epoches 之类
fit(state, train_ds, test_ds,
    train_step=train_step,
    eval_step=eval_step,
    eval_freq=1,
    num_epochs=10,
    log_name='mnist',
)