Skip to main content

jax_dataclasses

build mypy lint codecov

Overview

jax_dataclasses provides a simple wrapper around dataclasses.dataclass for use in JAX, which enables automatic support for:

  • Pytree registration. This allows dataclasses to be used at API boundaries in JAX.
  • Serialization via flax.serialization.

Distinguishing features include:

  • An annotation-based interface for marking static fields.
  • Improved ergonomics for "model surgery" in nested structures.

Installation

In Python >=3.7:

pip install jax_dataclasses

We can then import:

import jax_dataclasses as jdc

Core interface

jax_dataclasses is meant to provide a drop-in replacement for dataclasses.dataclass: jdc.pytree_dataclass has the same interface as dataclasses.dataclass, but also registers the target class as a pytree node.

We also provide several aliases: jdc.[field, asdict, astuples, is_dataclass, replace] are identical to their counterparts in the standard dataclasses library.

Static fields

To mark a field as static (in this context: constant at compile-time), we can wrap its type with jdc.Static[]:

@jdc.pytree_dataclass
class A:
    a: jax.Array
    b: jdc.Static[bool]

In a pytree node, static fields will be treated as part of the treedef instead of as a child of the node; all fields that are not explicitly marked static should contain arrays or child nodes.

Bonus: if you like jdc.Static[], we also introduce jdc.jit(). This enables use in function signatures, for example:

@jdc.jit
def f(a: jax.Array, b: jdc.Static[bool]) -> jax.Array:
  ...

Mutations

All dataclasses are automatically marked as frozen and thus immutable (even when no frozen= parameter is passed in). To make changes to nested structures easier, jdc.copy_and_mutate (a) makes a copy of a pytree and (b) returns a context in which any of that copy's contained dataclasses are temporarily mutable:

import jax
from jax import numpy as jnp
import jax_dataclasses as jdc

@jdc.pytree_dataclass
class Node:
  child: jax.Array

obj = Node(child=jnp.zeros(3))

with jdc.copy_and_mutate(obj) as obj_updated:
  # Make mutations to the dataclass. This is primarily useful for nested
  # dataclasses.
  #
  # Does input validation by default: if the treedef, leaf shapes, or dtypes
  # of `obj` and `obj_updated` don't match, an AssertionError will be raised.
  # This can be disabled with a `validate=False` argument.
  obj_updated.child = jnp.ones(3)

print(obj)
print(obj_updated)

Alternatives

A few other solutions exist for automatically integrating dataclass-style objects into pytree structures. Great ones include: chex.dataclass, flax.struct, and tjax.dataclass. These all influenced this library.

The main differentiators of jax_dataclasses are:

  • Static analysis support. tjax has a custom mypy plugin to enable type checking, but isn't supported by other tools. flax.struct implements the dataclass_transform spec proposed by pyright, but isn't supported by other tools. Because @jdc.pytree_dataclass has the same API as @dataclasses.dataclass, it can include pytree registration behavior at runtime while being treated as the standard decorator during static analysis. This means that all static checkers, language servers, and autocomplete engines that support the standard dataclasses library should work out of the box with jax_dataclasses.

  • Nested dataclasses. Making replacements/modifications in deeply nested dataclasses can be really frustrating. The three alternatives all introduce a .replace(self, ...) method to dataclasses that's a bit more convenient than the traditional dataclasses.replace(obj, ...) API for shallow changes, but still becomes really cumbersome to use when dataclasses are nested. jdc.copy_and_mutate() is introduced to address this.

  • Static field support. Parameters that should not be traced in JAX should be marked as static. This is supported in flax, tjax, and jax_dataclasses, but not chex.

  • Serialization. When working with flax, being able to serialize dataclasses is really handy. This is supported in flax.struct (naturally) and jax_dataclasses, but not chex or tjax.

You can also eschew the dataclass-style interface entirely; see how brax registers pytrees. This is a reasonable thing to prefer: it requires some floating strings and breaks things that I care about but you may not (like immutability and __post_init__), but gives more flexibility with custom __init__ methods.

Misc

jax_dataclasses was originally written for and factored out of jaxfg, where Nick Heppert provided valuable feedback.

Metadata

Release files for jax-dataclasses 1.6.3

For a detailed explanation of source distributions (sdists) and built distributions (wheels), please see the package formats documentation.

Source distribution (sdist)

Source distribution for jax-dataclasses 1.6.3
File Size Uploaded
jax_dataclasses-1.6.3.tar.gz 18.6 kB Details

Built distribution (wheel)

Table of built distributions (wheels) for jax-dataclasses 1.6.3
File Interpreter ABI Platform
jax_dataclasses-1.6.3-py3-none-any.whl Python 3 none any Details

Total release size: 33.2 kB

Release files / jax_dataclasses-1.6.3.tar.gz

Download URL jax_dataclasses-1.6.3.tar.gz
Size 18.6 kB
Tags Source
SHA-256 checksum
How to use checksums
e8bbd61221083d618a71217b2ee35ee26d438df1ac0a67eb25292f5a6ffa7f73
BLAKE2b-256 checksum
How to use checksums
b39b92ad5eb43efad3d0794e2ebc9c8047fefa113b7cb758578277ca5c663a5e
Upload date
Uploaded using Trusted Publishing?
What is trusted publishing?
No
Uploaded via twine/6.2.0 CPython/3.14.2

Release files / jax_dataclasses-1.6.3-py3-none-any.whl

Download URL jax_dataclasses-1.6.3-py3-none-any.whl
Size 14.6 kB
Tags Python 3
SHA-256 checksum
How to use checksums
d1c56fd6b34b522772bfde2a9224d9d678a990cad11db8e1a460eaf36fa80e6c
BLAKE2b-256 checksum
How to use checksums
542d49b9c938d1dfa4861aacadd1ba7eb2d421f27198135686010c51b171961a
Upload date
Uploaded using Trusted Publishing?
What is trusted publishing?
No
Uploaded via twine/6.2.0 CPython/3.14.2

Release history Release notifications | RSS feed

This release

1.6.3 This release

2 release files

1.6.2

2 release files

1.6.1

2 release files

1.6.0

2 release files

1.5.1

2 release files

1.5.0

2 release files

1.4.4

2 release files

1.4.3

2 release files

1.4.2

2 release files

1.4.1

2 release files

1.4.0

2 release files

1.3.0

2 release files

1.2.2

2 release files

1.2.1

2 release files

1.2.0

2 release files

1.1.0

2 release files

1.0.2

2 release files

1.0.1

2 release files

1.0.0

2 release files

0.0.3

2 release files

0.0.2

2 release files

0.0.1

2 release files

0.0

2 release files

Anthropic, PBC Visionary sponsor Bloomberg Visionary sponsor Hudson River Trading Visionary sponsor Meta Visionary sponsor NVIDIA Visionary sponsor Microsoft Sustainability sponsor Depot Continuous Integration AWS Cloud computing and Security Sponsor Datadog Monitoring Fastly CDN Google Download Analytics Sentry Error logging StatusPage Status page