Skip to main content

xmmutablemap

JAX-compatible Immutable Mapping

JAX prefers immutable objects but neither Python nor JAX provide an immutable dictionary. 😢
This repository defines a light-weight immutable map (lower-level than a dict) that JAX understands as a PyTree. 🎉 🕶️

Installation

PyPI platforms PyPI version

pip install xmmutablemap
using uv
uv add xmmutablemap
from source, using pip
pip install git+https://github.com/GalacticDynamics/xmmutablemap.git
building from source
cd /path/to/parent
git clone https://github.com/GalacticDynamics/xmmutablemap.git
cd xmmutablemap
pip install -e .  # editable mode

Documentation

xmutablemap provides the class ImmutableMap, which is a full implementation of Python's Mapping ABC. If you've used a dict then you already know how to use ImmutableMap! The things ImmutableMap adds is 1) immutability (and related benefits like hashability) and 2) compatibility with JAX.

from xmmutablemap import ImmutableMap

print(ImmutableMap(a=1, b=2, c=3))
# ImmutableMap({'a': 1, 'b': 2, 'c': 3})

print(ImmutableMap({"a": 1, "b": 2.0, "c": "3"}))
# ImmutableMap({'a': 1, 'b': 2.0, 'c': '3'})

JAX Integration

One of the key benefits of ImmutableMap is its compatibility with JAX. Since it's immutable and hashable, it can be used in places where JAX would normally complain about mutable objects like regular dictionaries.

Using ImmutableMap as a Default in JAX Dataclasses

Here's an example showing how ImmutableMap can be used as a default value in a dataclass, which is particularly useful with JAX:

import functools
import jax
import jax.numpy as jnp
from dataclasses import dataclass
from xmmutablemap import ImmutableMap


@functools.partial(
    jax.tree_util.register_dataclass, data_fields=["params"], meta_fields=["batch_size"]
)
@dataclass(frozen=True)
class Config:
    """Configuration with immutable default parameters."""

    # This works! ImmutableMap is immutable and hashable
    params: ImmutableMap[str, float] = ImmutableMap(
        learning_rate=0.001, momentum=0.9, weight_decay=1e-4
    )
    batch_size: int = 32


# JAX can safely transform functions using this dataclass
@jax.jit
def train_step(config: Config, data: jnp.ndarray) -> jnp.ndarray:
    """Example training step that uses config parameters."""
    lr = config.params["learning_rate"]
    return data * lr


# This works perfectly
config = Config()
data = jnp.array([1.0, 2.0, 3.0])
result = train_step(config, data)
print(f"Result: {result}")
# Result: [0.001 0.002 0.003]

Key Benefits for JAX

  • Immutability: Once created, ImmutableMap cannot be modified, preventing accidental mutations that could break JAX's functional programming model
  • Hashability: JAX can safely cache and memoize functions that use ImmutableMap instances
  • PyTree Support: ImmutableMap is registered as a JAX PyTree, so it works seamlessly with JAX transformations like jit, grad, vmap, etc.
  • Safe Defaults: Can be used as default values in dataclasses without the typical pitfalls of mutable defaults

Development

Actions Status

We welcome contributions!

Metadata

Release files for xmmutablemap 0.2.2

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

Source distribution (sdist)

Source distribution for xmmutablemap 0.2.2
File Size Uploaded
xmmutablemap-0.2.2.tar.gz 116.0 kB Details

Built distribution (wheel)

Table of built distributions (wheels) for xmmutablemap 0.2.2
File Interpreter ABI Platform
xmmutablemap-0.2.2-py3-none-any.whl Python 3 none any Details

Total release size: 123.7 kB

Release files / xmmutablemap-0.2.2.tar.gz

Download URL xmmutablemap-0.2.2.tar.gz
Size 116.0 kB
Tags Source
SHA-256 checksum
How to use checksums
c47d46fd2c6e622861125c5c521a0c5e03e992b6cd462a76f2796741cb11e774
BLAKE2b-256 checksum
How to use checksums
c5c545a697fb03521d8caefe1d1c9e32153e461e1d4f334636bbafc56cff66b0
Upload date
Uploaded using Trusted Publishing?
What is trusted publishing?
Yes
Uploaded via twine/7.0.0 CPython/3.13.14

Provenance

Provenance describes where a file came from. On PyPI, provenance is shared via attestations, which provide a verifiable record of the build or publishing details. View details, limitations and caveats.

PyPI Publish Attestation

PyPI verified that this artifact, at this checksum, originated from the publisher listed below.

Signed by GitHub Actions, verified by PyPI on Sep 7, 2026.

Transparency log

Release files / xmmutablemap-0.2.2-py3-none-any.whl

Download URL xmmutablemap-0.2.2-py3-none-any.whl
Size 7.7 kB
Tags Python 3
SHA-256 checksum
How to use checksums
fd5ecef589efd3038db2ccaca67dfc26df19cce478eb4968a0e058c6f24b1c17
BLAKE2b-256 checksum
How to use checksums
9782591c092e6b350ced7245f5b34e625b69f57cb076ff16845f7e6d6f396202
Upload date
Uploaded using Trusted Publishing?
What is trusted publishing?
Yes
Uploaded via twine/7.0.0 CPython/3.13.14

Provenance

Provenance describes where a file came from. On PyPI, provenance is shared via attestations, which provide a verifiable record of the build or publishing details. View details, limitations and caveats.

PyPI Publish Attestation

PyPI verified that this artifact, at this checksum, originated from the publisher listed below.

Signed by GitHub Actions, verified by PyPI on Sep 7, 2026.

Transparency log

Release history Release notifications | RSS feed

This release

0.2.2 This release

2 release files

0.2.1

2 release files

0.2.0

2 release files

0.1

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