Haiku is a JAX-based neural network library that provides an object-oriented programming model for building neural networks while maintaining access to JAX's pure function transformations.
Haiku solves the tension between writing neural networks in a familiar object-oriented style and leveraging JAX's functional programming paradigm. It accomplishes this through two core abstractions: `hk.Module`, which lets developers write classes that hold parameters and methods, and `hk.transform`, a function transformation that converts these stateful modules into pure functions compatible with JAX's transformations like `jax.jit`, `jax.grad`, and `jax.pmap`. This approach allows researchers to use intuitive object-oriented patterns without sacrificing JAX's performance and composability benefits.
Haiku suits researchers and practitioners who want to build neural networks with familiar class-based abstractions while retaining full access to JAX's ecosystem. The library has been validated at scale through DeepMind's internal use across image processing, language models, generative models, and reinforcement learning. However, the README contains an important notice: Google DeepMind recommends that new projects adopt Flax instead, which offers a superset of Haiku's features, more extensive documentation, and a larger development community. Haiku is positioned as a library rather than a framework, giving developers more control over their training loops and architecture choices.
The project has entered maintenance mode, with development efforts focused on bug fixes and compatibility with new JAX releases rather than new feature development. Haiku continues to receive updates to maintain compatibility with newer Python and JAX versions. The tool remains supported indefinitely for internal use at Google DeepMind despite the shift away from recommending it for new external projects.