🌲pytreeclass🌲
Write pytorch-like layers with keras-like visualizations in JAX.
Installation |Description |Examples
🛠️ Installation
pip install pytreeclass
📖 Description
A JAX compatible dataclass like datastructure with the following functionalities
- Create PyTorch like NN classes like equinox and Treex
- Provides Keras-like
model.summary()andplot_modelvisualizations for pytrees wrapped withtree. - Apply math/numpy operations like tree-math
- Registering user-defined reduce operations on each class.
- Some fancy indexing syntax functionalities like
x[x>0]on pytrees
🔢 Examples
Write PyTorch like NN classes
# construct a Pytorch like NN classes with JAX
import jax
from jax import numpy as jnp
from pytreeclass import treeclass,static_field,tree_viz
@treeclass
class Linear :
weight : jnp.ndarray
bias : jnp.ndarray
def __init__(self,key,in_dim,out_dim):
self.weight = jax.random.normal(key,shape=(in_dim, out_dim)) * jnp.sqrt(2/in_dim)
self.bias = jnp.ones((1,out_dim))
def __call__(self,x):
return x @ self.weight + self.bias
@treeclass
class StackedLinear:
l1 : Linear
l2 : Linear
l3 : Linear
def __init__(self,key,in_dim,out_dim):
keys= jax.random.split(key,3)
self.l1 = Linear(key=keys[0],in_dim=in_dim,out_dim=128)
self.l2 = Linear(key=keys[1],in_dim=128,out_dim=128)
self.l3 = Linear(key=keys[2],in_dim=128,out_dim=out_dim)
def __call__(self,x):
x = self.l1(x)
x = jax.nn.tanh(x)
x = self.l2(x)
x = jax.nn.tanh(x)
x = self.l3(x)
return x
x = jnp.linspace(0,1,100)[:,None]
y = x**3 + jax.random.uniform(jax.random.PRNGKey(0),(100,1))*0.01
model = StackedLinear(in_dim=1,out_dim=1,key=jax.random.PRNGKey(0))
def loss_func(model,x,y):
return jnp.mean((model(x)-y)**2 )
@jax.jit
def update(model,x,y):
value,grads = jax.value_and_grad(loss_func)(model,x,y)
# no need to use `jax.tree_map` to update the model
# as it model is wrapped by @treeclass
return value , model-1e-3*grads
for _ in range(1,2001):
value,model = update(model,x,y)
plt.scatter(x,model(x),color='r',label = 'Prediction')
plt.scatter(x,y,color='k',label='True')
plt.legend()
Visualize
>>> print(tree_viz.summary(model))
┌──────┬───────┬─────────┬───────────────────┐
│Type │Param #│Size │Config │
├──────┼───────┼─────────┼───────────────────┤
│Linear│256 │1.000 KB │bias=f32[1,128] │
│ │ │ │weight=f32[1,128] │
├──────┼───────┼─────────┼───────────────────┤
│Linear│16,512 │64.500 KB│bias=f32[1,128] │
│ │ │ │weight=f32[128,128]│
├──────┼───────┼─────────┼───────────────────┤
│Linear│129 │516.000 B│bias=f32[1,1] │
│ │ │ │weight=f32[128,1] │
└──────┴───────┴─────────┴───────────────────┘
Total params : 16,897
Inexact params: 16,897
Other params: 0
----------------------------------------------
Total size : 66.004 KB
Inexact size: 66.004 KB
Other size: 0.000 B
==============================================
>>> print(tree_viz.tree_box(model,array=x))
# using jax.eval_shape (no-flops operation)
┌──────────────────────────────────────┐
│StackedLinear(Parent) │
├──────────────────────────────────────┤
│┌────────────┬────────┬──────────────┐│
││ │ Input │ f32[100,1] ││
││ Linear(l1) │────────┼──────────────┤│
││ │ Output │ f32[100,128] ││
│└────────────┴────────┴──────────────┘│
│┌────────────┬────────┬──────────────┐│
││ │ Input │ f32[100,128] ││
││ Linear(l2) │────────┼──────────────┤│
││ │ Output │ f32[100,128] ││
│└────────────┴────────┴──────────────┘│
│┌────────────┬────────┬──────────────┐│
││ │ Input │ f32[100,128] ││
││ Linear(l3) │────────┼──────────────┤│
││ │ Output │ f32[100,1] ││
│└────────────┴────────┴──────────────┘│
└──────────────────────────────────────┘
>>> print(tree_viz.tree_diagram(model))
StackedLinear
├── l1=Linear
│ ├── weight=f32[1,128]
│ └── bias=f32[1,128]
├── l2=Linear
│ ├── weight=f32[128,128]
│ └── bias=f32[1,128]
└──l3=Linear
├── weight=f32[128,1]
└── bias=f32[1,1]
Perform Math operations on JAX pytrees
@treeclass
class Test :
a : float
b : float
c : float
name : str = static_field() # ignore from jax computations
# basic operations
A = Test(10,20,30,'A')
assert (A + A) == Test(20,40,60,'A')
assert (A - A) == Test(0,0,0,'A')
assert (A*A).reduce_mean() == 1400
assert (A + 1) == Test(11,21,31,'A')
# selective operations
# only add 1 to field `a`
# all other fields are set to None and returns the same class
assert (A['a'] + 1) == Test(11,None,None,'A')
# use `|` to merge classes by performing ( left_node or right_node )
Aa = A['a'] + 10 # Test(a=20,b=None,c=None,name=A)
Ab = A['b'] + 10 # Test(a=None,b=30,c=None,name=A)
assert (Aa | Ab | A ) == Test(20,30,30,'A')
# indexing by class
assert A[A>10] == Test(a=None,b=20,c=30,name='A')
# Register custom operations
B = Test([10,10],20,30,'B')
B.register_op( func=lambda node:node+1,name='plus_one')
assert B.plus_one() == Test(a=[11, 11],b=21,c=31,name='B')
# Register custom reduce operations ( similar to functools.reduce)
C = Test(jnp.array([10,10]),20,30,'C')
C.register_op(
func=jnp.prod, # function applied on each node
name='product', # name of the function
reduce_op=lambda x,y:x*y, # function applied between nodes (accumulated * current node)
init_val=1 # initializer for the reduce function
)
# product applies only on each node
# and returns an instance of the same class
assert C.product() == Test(a=100,b=20,c=30,name='C')
# `reduce_` + name of the registered function (`product`)
# reduces the class and returns a value
assert C.reduce_product() == 60000
Download files
Download the file for your platform. If you're not sure which to choose, learn more about installing packages.
Source Distribution
pytreeclass-0.0.4.tar.gz
(15.3 kB
view details)
Built Distribution
Filter files by name, interpreter, ABI, and platform.
If you're not sure about the file name format, learn more about wheel file names.
Copy a direct link to the current filters
File details
Details for the file pytreeclass-0.0.4.tar.gz.
File metadata
- Download URL: pytreeclass-0.0.4.tar.gz
- Upload date:
- Size: 15.3 kB
- Tags: Source
- Uploaded using Trusted Publishing? No
- Uploaded via:
twine/4.0.1 CPython/3.10.5
File hashes
| Algorithm | Hash digest | |
|---|---|---|
| SHA256 |
584d62756fd12634d47112b15d54fbdc66ac5888337f3ef0508c732d510271ec
|
|
| MD5 |
58bef61f9a7b0d7046ec6eef7fb8b8df
|
|
| BLAKE2b-256 |
612330b71928e74183ab3ce3257e45b2d3b01553049d50da6e530d8b2c25fb7e
|
File details
Details for the file pytreeclass-0.0.4-py3-none-any.whl.
File metadata
- Download URL: pytreeclass-0.0.4-py3-none-any.whl
- Upload date:
- Size: 13.9 kB
- Tags: Python 3
- Uploaded using Trusted Publishing? No
- Uploaded via:
twine/4.0.1 CPython/3.10.5
File hashes
| Algorithm | Hash digest | |
|---|---|---|
| SHA256 |
616d82ed24d1df85b65be0ff72a0160d1614d9b57646ed3a27c723a4232ff25d
|
|
| MD5 |
13a33d7382bd8594bc3625ff29859b11
|
|
| BLAKE2b-256 |
8c5a718cd8d07ad657abfc75aa50ecc99ef6f9148960a195ebe6e488e6d03c6a
|