google-deepmind/dm-haiku
JAX-based neural network library
解决的问题
Haiku 是一个用于 JAX 的神经网络库,它允许开发人员使用熟悉的面向对象编程模型(如 Sonnet 或 TensorFlow 中使用的模型),同时保持对 JAX 纯函数转换的完全访问。它简化了模型参数和内部状态的管理,而无需用户为初始化编写大量的样板代码。
工作原理
Haiku 提供了两种主要工具来弥合面向对象设计与函数纯粹性之间的鸿っ:
hk.Module: 用于定义神经网络层和组件的 Python 对象。这些模块持有参数和方法的引用,但在定义过程中被视为函数上的“不纯”对象。hk.transform: 一种函数转换,可将这些基于模块的不纯函数转换为一对纯函数:init(收集初始参数值)和apply(将这些参数重新注入函数进行计算)。
对于需要可变状态(如 Batch Normalization 的移动平均值)的模型,Haiku 提供了 hk.transform_with_state,它可以分别管理参数和状态。
适用对象
想要利用面向对象 API 的生产力来构建神经网络,但又需要 JAX 的性能和转换能力(如 jax.jit、jax.grad 和 jax.pmap)的研究人员和开发人员。
亮点
- JAX 集成:与 JAX 的自动微分和 GPU/TPU 支持完全兼容。
- 类 Sonnet API:设计上与 Sonnet 2 API 近乎匹配,使从 TensorFlow/Sonnet 的迁移变得容易。
- 简化的 RNG 管理:提供
hk.next_rng_key(),以便在转换后的函数中确定性地处理随机数生成。 - 可扩展性:由 DeepMind 研究人员在图像、语言和强化学习任务中进行了大规模测试。
- 是库,而非框架:严格专注于参数和状态管理,将优化器和检查点保存留给其他库。
相关
- 项目
- 项目
- 项目
- 项目
- 项目