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.jitjax.gradjax.pmap)的研究人员和开发人员。

亮点

  • JAX 集成:与 JAX 的自动微分和 GPU/TPU 支持完全兼容。
  • 类 Sonnet API:设计上与 Sonnet 2 API 近乎匹配,使从 TensorFlow/Sonnet 的迁移变得容易。
  • 简化的 RNG 管理:提供 hk.next_rng_key(),以便在转换后的函数中确定性地处理随机数生成。
  • 可扩展性:由 DeepMind 研究人员在图像、语言和强化学习任务中进行了大规模测试。
  • 是库,而非框架:严格专注于参数和状态管理,将优化器和检查点保存留给其他库。

相关

  • 项目
  • 项目
  • 项目
  • 项目
  • 项目