google-deepmind/dm-haiku
JAX-based neural network library
What it solves
JAX 是一个强大的数值计算库,但它要求函数必须是纯函数才能使用其转换(如 jit 和 grad)。Haiku 通过允许开发者使用熟悉的面向对象编程模型来定义神经网络,并自动将这些“非纯”面向对象的定义转换为 JAX 能处理的纯函数,从而解决了这个问题。
How it works
Haiku 提供两个主要工具,弥合面向对象设计与函数纯度之间的鸿沟:
hk.Module:用于定义网络层和组件的 Python 对象。这些模块持有参数和方法的引用,允许用户编写看起来像标准神经网络库的代码。hk.transform:一个函数转换器,将使用hk.Module的函数转换为一对纯函数:init(收集初始参数值)和apply(将这些参数注入函数进行计算)。
对于需要内部可变状态(如批量归一化)的模型,Haiku 提供 hk.transform_with_state,可以分别管理参数和状态。
Who it’s for
需要面向对象 API 提高构建神经网络生产力,同时又想完整利用 JAX 的函数转换和硬件加速的研究人员和开发者。
Highlights
- DeepMind Scale:已由 DeepMind 研究人员在大规模图像、语言和强化学习实验中测试。
- Library, Not Framework:设计为轻量库,专注于参数和状态管理,不强加自定义优化器或检查点格式。
hk.next_rng_key():通过提供确定性的键序列,简化 JAX 中的随机数生成。- JAX Compatibility:完全兼容
jax.pmap,可在多个加速器上进行分布式训练。