Skip to main content

🌲pytreeclass🌲

Write pytorch-like layers with keras-like visualizations in JAX.

Installation |Description |Examples

Tests pyver codestyle Open In Colab

🛠️ 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() and plot_model visualizations for pytrees wrapped with tree.
  • 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()

image

Visualize
>>> print(tree_viz.summary(model))
┌──────┬───────┬─────────┬───────────────────┐
Type  Param #│Size     │Config             │
├──────┼───────┼─────────┼───────────────────┤
Linear256    1.000 KB bias=f32[1,128]    
                      weight=f32[1,128]  
├──────┼───────┼─────────┼───────────────────┤
Linear16,512 64.500 KBbias=f32[1,128]    
                      weight=f32[128,128]
├──────┼───────┼─────────┼───────────────────┤
Linear129    516.000 Bbias=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)

Uploaded Source

Built Distribution

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

pytreeclass-0.0.4-py3-none-any.whl (13.9 kB view details)

Uploaded Python 3

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

Hashes for pytreeclass-0.0.4.tar.gz
Algorithm Hash digest
SHA256 584d62756fd12634d47112b15d54fbdc66ac5888337f3ef0508c732d510271ec
MD5 58bef61f9a7b0d7046ec6eef7fb8b8df
BLAKE2b-256 612330b71928e74183ab3ce3257e45b2d3b01553049d50da6e530d8b2c25fb7e

See more details on using hashes here.

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

Hashes for pytreeclass-0.0.4-py3-none-any.whl
Algorithm Hash digest
SHA256 616d82ed24d1df85b65be0ff72a0160d1614d9b57646ed3a27c723a4232ff25d
MD5 13a33d7382bd8594bc3625ff29859b11
BLAKE2b-256 8c5a718cd8d07ad657abfc75aa50ecc99ef6f9148960a195ebe6e488e6d03c6a

See more details on using hashes here.

Release history Release notifications | RSS feed

0.11.1

2 files

0.11.0

2 files

0.9.2

2 files

0.9.1

2 files

0.9.0

2 files

0.8.0

2 files

0.7.0

2 files

0.6.0.post0

2 files

0.6.0

2 files

0.5.0.post0

2 files

0.5.0

2 files

0.4.0

2 files

0.3.8

2 files

0.3.7

2 files

0.3.6

2 files

0.3.4

2 files

0.3.3

2 files

0.3.2

2 files

0.3.1

2 files

0.3.0

2 files

0.2.8

2 files

0.2.7

2 files

0.2.6

2 files

0.2.5

2 files

0.2.4

2 files

0.2.3

2 files

0.2.2

2 files

0.2.1

2 files

0.2.0

2 files

0.1.13

2 files

0.1.12

2 files

0.1.11

2 files

0.1.10

2 files

0.1.9

2 files

0.1.8

2 files

0.1.7

2 files

0.1.6

2 files

0.1.5

2 files

0.1.4

2 files

0.1.3

2 files

0.1.2

2 files

0.1.1

2 files

0.1.0

2 files

0.0.11

2 files

0.0.9.post1

2 files

0.0.9

2 files

0.0.8

2 files

0.0.7

2 files

0.0.6.post2

2 files

0.0.6.post1

2 files

0.0.6.post0

2 files

0.0.6

2 files

0.0.5

2 files

This release

0.0.4 This release

2 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