Simple Pytree
A dead simple Python package for creating custom JAX pytree objects.
- Strives to be minimal, the implementation is just ~100 lines of code
- Has no dependencies other than JAX
- Its compatible with both
dataclassesand regular classes - It has no intention of supporting Neural Network use cases (e.g. partitioning)
Installation
pip install simple-pytree
Usage
import jax
from simple_pytree import Pytree
class Foo(Pytree):
def __init__(self, x, y):
self.x = x
self.y = y
foo = Foo(1, 2)
foo = jax.tree_map(lambda x: -x, foo)
assert foo.x == -1 and foo.y == -2
Static fields
You can mark fields as static by assigning static_field() to a class attribute with the same name
as the instance attribute:
import jax
from simple_pytree import Pytree, static_field
class Foo(Pytree):
y = static_field()
def __init__(self, x, y):
self.x = x
self.y = y
foo = Foo(1, 2)
foo = jax.tree_map(lambda x: -x, foo) # y is not modified
assert foo.x == -1 and foo.y == 2
Static fields are not included in the pytree leaves, they are passed as pytree metadata instead.
Dataclasses
simple_pytree provides a dataclass decorator you can use with classes
that contain static_fields:
import jax
from simple_pytree import Pytree, dataclass, static_field
@dataclass
class Foo(Pytree):
x: int
y: int = static_field(default=2)
foo = Foo(1)
foo = jax.tree_map(lambda x: -x, foo) # y is not modified
assert foo.x == -1 and foo.y == 2
simple_pytree.dataclass is just a wrapper around dataclasses.dataclass but
when used static analysis tools and IDEs will understand that static_field is a
field specifier just like dataclasses.field.
Mutability
Pytree objects are immutable by default after __init__:
from simple_pytree import Pytree, static_field
class Foo(Pytree):
y = static_field()
def __init__(self, x, y):
self.x = x
self.y = y
foo = Foo(1, 2)
foo.x = 3 # AttributeError
If you want to make them mutable, you can use the mutable argument in class definition:
from simple_pytree import Pytree, static_field
class Foo(Pytree, mutable=True):
y = static_field()
def __init__(self, x, y):
self.x = x
self.y = y
foo = Foo(1, 2)
foo.x = 3 # OK
Replacing fields
If you want to make a copy of a Pytree object with some fields modified, you can use the .replace() method:
from simple_pytree import Pytree, static_field
class Foo(Pytree):
y = static_field()
def __init__(self, x, y):
self.x = x
self.y = y
foo = Foo(1, 2)
foo = foo.replace(x=10)
assert foo.x == 10 and foo.y == 2
replace works for both mutable and immutable Pytree objects. If the class
is a dataclass, replace internally use dataclasses.replace.
Release files for simple-pytree 0.2.2
For a detailed explanation of source distributions (sdists) and built distributions (wheels), please see the package formats documentation.
Source distribution (sdist)
| File | Size | Uploaded | |
|---|---|---|---|
| simple_pytree-0.2.2.tar.gz | 5.3 kB | Details |
Built distribution (wheel)
| File | Interpreter | ABI | Platform | Reset |
|---|---|---|---|---|
| simple_pytree-0.2.2-py3-none-any.whl | Python 3 | none | any | Details |
Total release size: 11.6 kB
Release files / simple_pytree-0.2.2.tar.gz
| Download URL | simple_pytree-0.2.2.tar.gz |
|---|---|
| Size | 5.3 kB |
| Tags | Source |
|
SHA-256 checksum How to use checksums |
b61eddb5b5de558209dfb6464a041c5faf06af2c3cab582f4d69543a773dab0e
|
|
BLAKE2b-256 checksum How to use checksums |
8eb0b2e7ea15dfb26bf014cfb6243a9bb20b9477ee2f12d754257514f508639a
|
| Upload date | |
|
Uploaded using Trusted Publishing? What is trusted publishing? |
No |
| Uploaded via |
poetry/1.4.0 CPython/3.8.17 Linux/5.15.0-1041-azure
|
Release files / simple_pytree-0.2.2-py3-none-any.whl
| Download URL | simple_pytree-0.2.2-py3-none-any.whl |
|---|---|
| Size | 6.2 kB |
| Tags | Python 3 |
|
SHA-256 checksum How to use checksums |
3a7a2f66883194ab14875dd01c4306c4f102f31d90e7e2d421366f12a2c49bab
|
|
BLAKE2b-256 checksum How to use checksums |
7c160272467306ef489512a843222567c9939b9aff7003f15474a1ef90168c8f
|
| Upload date | |
|
Uploaded using Trusted Publishing? What is trusted publishing? |
No |
| Uploaded via |
poetry/1.4.0 CPython/3.8.17 Linux/5.15.0-1041-azure
|