Skip to main content

Jaxpr expression introspection and injection via tags

Project description

Taxpr

Run tests PyPI Version

Taxpr is a collection of utilities for performing manipulation of Jaxprs. This is achieved by tag-ing specific arrays at trace time, then extracting and manipulating those tags in the final Jaxpr.

⚠️ This package is still very experimental, so expect broken code and breaking changes.

The provided routines are designed to work seamlessly with jit, vmap, custom_jvp and cousins.

Example

The following example shows how you can use taxpr to emulate functions with side effects without violating Jax's pure function rules.

import itertools as it

import jax
import jax.numpy as jnp
from jax._src.core import eval_jaxpr
import taxpr as tx

_state_counter = it.count()

def get_state(shape, dtype):
    count = next(_state_counter)

    def set_state(value):
        return tx.tag(value, op="set", id=count)

    value = jax.numpy.zeros(shape, dtype=dtype)
    return tx.tag(value, op="get", id=count), set_state


def uncurry(fn, *args, **kwargs):
    jaxpr = jax.make_jaxpr(fn)(args, kwargs)
    states = {}

    # iterate through all tags in the jaxpr
    # this recurses all child Jaxprs too

    for params, shape in tx.iter_tags(jaxpr.jaxpr):
        if params["op"] == "get":
            states[params["id"]] = shape

    initial_states = jax.tree.map(
        lambda x: jax.numpy.full_like(x, 0), states
    )

    def injector(states, token, params):
        if params["op"] == "get":
            state = states[params["id"]]
            return state, states
        elif params["op"] == "set": 
            states[params["id"]] = token
            return token, states
        raise ValueError(f"Unknown tag op: {params['op']}")

    # replace the token with a function that performs the state manipulation
    # here we can pass our own context (`initial_states`)

    jaxpr = tx.inject(jaxpr, injector, initial_states)

    def wrapper(states, *args, **kwargs):
        return eval_jaxpr(jaxpr.jaxpr, jaxpr.consts, states, args, kwargs)

    return wrapper, initial_states

################################################

# Usage

def running_sum(x):
    a, set_state = get_state(x.shape, x.dtype)
    sum = set_state(a + x)
    return sum

rsum, state = uncurry(running_sum, jnp.zeros(0))

_, state = rsum(state, jnp.ones(1))
_, state = rsum(state, jnp.ones(1))
_, state = rsum(state, jnp.ones(1))

assert jnp.allclose(next(iter(state.values())), 3)

Project details


Download files

Download the file for your platform. If you're not sure which to choose, learn more about installing packages.

Source Distribution

taxpr-0.0.3.tar.gz (21.9 kB view details)

Uploaded Source

Built Distribution

If you're not sure about the file name format, learn more about wheel file names.

taxpr-0.0.3-py3-none-any.whl (13.1 kB view details)

Uploaded Python 3

File details

Details for the file taxpr-0.0.3.tar.gz.

File metadata

  • Download URL: taxpr-0.0.3.tar.gz
  • Upload date:
  • Size: 21.9 kB
  • Tags: Source
  • Uploaded using Trusted Publishing? Yes
  • Uploaded via: twine/6.1.0 CPython/3.13.7

File hashes

Hashes for taxpr-0.0.3.tar.gz
Algorithm Hash digest
SHA256 6587b0ceb1b06e19d60ab41b993fb42ff24d231d1644443cd0d08db02a7ef277
MD5 a617731a1843ed7edfc52985e68a4bca
BLAKE2b-256 4e278c695e485544a318ccb83c682ad67b5e99c4bf44bad791442e15e47592f3

See more details on using hashes here.

Provenance

The following attestation bundles were made for taxpr-0.0.3.tar.gz:

Publisher: release.yml on otoomey/taxpr

Attestations: Values shown here reflect the state when the release was signed and may no longer be current.

File details

Details for the file taxpr-0.0.3-py3-none-any.whl.

File metadata

  • Download URL: taxpr-0.0.3-py3-none-any.whl
  • Upload date:
  • Size: 13.1 kB
  • Tags: Python 3
  • Uploaded using Trusted Publishing? Yes
  • Uploaded via: twine/6.1.0 CPython/3.13.7

File hashes

Hashes for taxpr-0.0.3-py3-none-any.whl
Algorithm Hash digest
SHA256 d95f5f1deb1a6e494ca4dc816c93b817b3bbe60520c0af9f4dc96b0037e5899e
MD5 fa3afb0c63d84e82675b228208f23c30
BLAKE2b-256 94f4cc51e52d67638c4ada5ad3215001762eb58f4273c780861fe234effb869b

See more details on using hashes here.

Provenance

The following attestation bundles were made for taxpr-0.0.3-py3-none-any.whl:

Publisher: release.yml on otoomey/taxpr

Attestations: Values shown here reflect the state when the release was signed and may no longer be current.

Supported by

AWS Cloud computing and Security Sponsor Datadog Monitoring Depot Continuous Integration Fastly CDN Google Download Analytics Pingdom Monitoring Sentry Error logging StatusPage Status page