Metadata-Version: 2.4
Name: jaxtree-extensions
Version: 0.1.0
Summary: Additional pytree utilities for JAX
Project-URL: Homepage, https://paulgekeler.github.io/jaxtree_extensions/
Project-URL: Repository, https://github.com/paulgekeler/jaxtree_extensions
Author-email: Paul Gekeler <pgekeler@gmail.com>
License: MIT
License-File: LICENSE
Keywords: jax,pytree,utilities
Requires-Python: >=3.10
Requires-Dist: jax
Description-Content-Type: text/markdown

# Jaxtree Extensions

`jaxtree-extensions` is a Python library that extends JAX's native `jax.tree_util` with additional tree utilities. It is designed to provide clean, reusable helper functions for mapping, filtering, and manipulating PyTree structures in common JAX workflows.

## Installation

You can install the package locally using `uv`:

```bash
uv pip install -e .
```

## Quick Start

```python
import jax.numpy as jnp
import jaxtree_extensions as jte

# Example tree structure
tree = {
    "a": jnp.ones((3, 3)),
    "b": jnp.zeros((3,)),
}
```

## Running Tests

To run the test suite:

```bash
uv run pytest
```
