Skip to main content

The ✨Magical✨ JAX NN Library.

*Serket is the goddess of magic in Egyptian mythology

Installation |Description |Quick Example |Freezing/Fine tuning |Filtering

Tests pyver codestyle Downloads codecov DOI PyPI

🛠️ Installation

pip install serket

Install development version

pip install git+https://github.com/ASEM000/serket

📖 Description

  • serket aims to be the most intuitive and easy-to-use physics-based Neural network library in JAX.

  • serket is built on top of pytreeclass

  • serket currently implements

🧠 Neural network package: serket.nn 🧠

Group Layers
Linear Linear, Bilinear,Identity
Densely connected FNN (Fully connected network), PFNN (Parallel fully connected network)
Convolution Conv1D, Conv2D, Conv3D, Conv1DTranspose , Conv2DTranspose, Conv3DTranspose, DepthwiseConv1D, DepthwiseConv2D, DepthwiseConv3D, SeparableConv1D, SeparableConv2D, SeparableConv3D, Conv1DLocal, Conv2DLocal, Conv3DLocal, Conv1DSemiLocal, Conv2DSemiLocal, Conv3DSemiLocal*
Containers Sequential, Lambda
Pooling MaxPool1D, MaxPool2D, MaxPool3D, AvgPool1D, AvgPool2D, AvgPool3D GlobalMaxPool1D, GlobalMaxPool2D, GlobalMaxPool3D, GlobalAvgPool1D, GlobalAvgPool2D, GlobalAvgPool3D (kernex backend)
Reshaping Flatten, Unflatten, FlipLeftRight2D, FlipUpDown2D, Repeat1D, Repeat2D, Repeat3D, Resize1D, Resize2D, Resize3D, Upsample1D, Upsample2D, Upsample3D, Pad1D, Pad2D, Pad3D
Crop Crop1D, Crop2D,
Normalization LayerNorm, InstanceNorm, GroupNorm
Blurring AvgBlur2D, GaussianBlur2D
Dropout Dropout, Dropout1D, Dropout2D, Dropout3D,
Random transforms RandomCrop1D, RandomCrop2D, RandomApply, RandomCutout1D, RandomCutout2D, RandomZoom2D, RandomContrast2D
Preprocessing HistogramEqualization2D, AdjustContrast2D
Activations AdaptiveLeakyReLU,AdaptiveReLU,AdaptiveSigmoid,AdaptiveTanh,
CeLU,ELU,GELU,GLU
,HardSILU,HardShrink,HardSigmoid,HardSwish,HardTanh,
LeakyReLU,LogSigmoid,LogSoftmax,Mish,PReLU,
ReLU,ReLU6,SILU,SeLU,Sigmoid,SoftPlus,SoftShrink,
SoftSign,Swish,Tanh,TanhShrink, ThresholdedReLU
Blocks VGG16Block, VGG19Block, UNetBlock

* Apply set of different shared kernel weights to each spatial group, where spatial groups<= Total patches of the input.

➖➕Finite difference package: serket.fd➕➖

Group Function/Layer
Finite difference layer Difference: apply finite difference to input array to any derivative order and accuracy
Finite difference functions - difference: finite difference of array with any accuracy and derivative order
-generate_finitediff_coeffs : generate coeffs using sample points and derivative order
- fgrad: differentiate functions (similar to jax.grad) with custom accuracy and derivative order
Vector operator layers Curl, Divergence, Gradient, Laplacian
Vector operator function curl, divergence, gradient, laplacian

⏩ Quick Example:

Lazy initialization

In cases where in_features needs to be inferred from input, use None instead of in_features to infer the value at runtime. However, since the lazy module initialize it's state after the first call (i.e. mutate it's state) jax transformation ex: vmap, grad ... is not allowed before initialization. Using any jax transformation before initialization will throw a ValueError.

import serket as sk 
import jax
import jax.numpy as jnp 

model = sk.nn.Sequential(
    [
        sk.nn.Conv2D(None, 128, 3),
        sk.nn.ReLU(),
        sk.nn.MaxPool2D(2, 2),
        sk.nn.Conv2D(128, 64, 3),
        sk.nn.ReLU(),
        sk.nn.MaxPool2D(2, 2),
        sk.nn.Flatten(),
        sk.nn.Linear(None, 128),
        sk.nn.ReLU(),
        sk.nn.Linear(128, 1),
    ]
)

# print the first `Conv2D` layer before initialization
print(model[0].__repr__())
# Conv2D(
#   weight=None,
#   bias=None,
#   *in_features=None,
#   *out_features=None,
#   *kernel_size=None,
#   *strides=None,
#   *padding=None,
#   *input_dilation=None,
#   *kernel_dilation=None,
#   weight_init_func=None,
#   bias_init_func=None,
#   *groups=None
# )

try :
    jax.vmap(model)(jnp.ones((10, 1,28, 28)))
except ValueError:
    print("***** Not initialized *****")
# ***** Not initialized *****

# dry run to initialize the model
model(jnp.empty([3,128,128]))

print(model[0].__repr__())
# Conv2D(
#   weight=f32[128,3,3,3],
#   bias=f32[128,1,1],
#   *in_features=3,
#   *out_features=128,
#   *kernel_size=(3,3),
#   *strides=(1,1),
#   *padding=((1,1),(1,1)),
#   *input_dilation=(1,1),
#   *kernel_dilation=(1,1),
#   weight_init_func=Partial(glorot_uniform(key,shape,dtype)),
#   bias_init_func=Partial(zeros(key,shape,dtype)),
#   *groups=1
# )
Train MNIST

We will use tensorflow datasets for dataloading. for more on interface of jax/tensorflow dataset see here

# imports
import tensorflow as tf
# Ensure TF does not see GPU and grab all GPU memory.
tf.config.set_visible_devices([], device_type="GPU")
import tensorflow_datasets as tfds
import tensorflow.experimental.numpy as tnp
import jax
import jax.numpy as jnp
import jax.random as jr 
import optax  # for gradient optimization
import serket as sk
import matplotlib.pyplot as plt
import functools as ft
# Construct a tf.data.Dataset
batch_size = 128

# convert the samples from integers to floating-point numbers
# and channel first format
def preprocess_data(x):
    # convert to channel first format
    image = tnp.moveaxis(x["image"], -1, 0)
    # normalize to [0, 1]
    image = tf.cast(image, tf.float32) / 255.0

    # one-hot encode the labels
    label = tf.one_hot(x["label"], 10) / 1.0
    return {"image": image, "label": label}


ds_train, ds_test = tfds.load("mnist", split=["train", "test"], shuffle_files=True)
# (batches, batch_size, 1, 28, 28)
ds_train = ds_train.shuffle(1024).map(preprocess_data).batch(batch_size).prefetch(tf.data.AUTOTUNE)

# (batches, 1, 28, 28)
ds_test = ds_test.map(preprocess_data).prefetch(tf.data.AUTOTUNE)

🏗️ Model definition

We will use jax.vmap(model) to apply model on batches.

@sk.treeclass
class CNN:
    def __init__(self):
        self.conv1 = sk.nn.Conv2D(1, 32, (3, 3), padding="valid")
        self.relu1 = sk.nn.ReLU()
        self.pool1 = sk.nn.MaxPool2D((2, 2), strides=(2, 2))
        self.conv2 = sk.nn.Conv2D(32, 64, (3, 3), padding="valid")
        self.relu2 = sk.nn.ReLU()
        self.pool2 = sk.nn.MaxPool2D((2, 2), strides=(2, 2))
        self.flatten = sk.nn.Flatten(start_dim=0)
        self.dropout = sk.nn.Dropout(0.5)
        self.linear = sk.nn.Linear(5*5*64, 10)

    def __call__(self, x):
        x = self.conv1(x)
        x = self.relu1(x)
        x = self.pool1(x)
        x = self.conv2(x)
        x = self.relu2(x)
        x = self.pool2(x)
        x = self.flatten(x)
        x = self.dropout(x)
        x = self.linear(x)
        return x

model = CNN()

🎨 Visualize model

Model summary
print(model.summary(show_config=False, array=jnp.empty((1, 28, 28))))  
┌───────┬─────────┬─────────┬──────────────┬─────────────┬─────────────┐
Name   Type     Param #  │Size          │Input        │Output       │
├───────┼─────────┼─────────┼──────────────┼─────────────┼─────────────┤
conv1  Conv2D   320(0)   1.25KB(0.00B) f32[1,28,28] f32[32,26,26]
├───────┼─────────┼─────────┼──────────────┼─────────────┼─────────────┤
relu1  ReLU     0(0)     0.00B(0.00B)  f32[32,26,26]f32[32,26,26]
├───────┼─────────┼─────────┼──────────────┼─────────────┼─────────────┤
pool1  MaxPool2D0(0)     0.00B(0.00B)  f32[32,26,26]f32[32,13,13]
├───────┼─────────┼─────────┼──────────────┼─────────────┼─────────────┤
conv2  Conv2D   18,496(0)72.25KB(0.00B)f32[32,13,13]f32[64,11,11]
├───────┼─────────┼─────────┼──────────────┼─────────────┼─────────────┤
relu2  ReLU     0(0)     0.00B(0.00B)  f32[64,11,11]f32[64,11,11]
├───────┼─────────┼─────────┼──────────────┼─────────────┼─────────────┤
pool2  MaxPool2D0(0)     0.00B(0.00B)  f32[64,11,11]f32[64,5,5]  
├───────┼─────────┼─────────┼──────────────┼─────────────┼─────────────┤
flattenFlatten  0(0)     0.00B(0.00B)  f32[64,5,5]  f32[1600]    
├───────┼─────────┼─────────┼──────────────┼─────────────┼─────────────┤
dropoutDropout  0(0)     0.00B(0.00B)  f32[1600]    f32[1600]    
├───────┼─────────┼─────────┼──────────────┼─────────────┼─────────────┤
linear Linear   16,010(0)62.54KB(0.00B)f32[1600]    f32[10]      
└───────┴─────────┴─────────┴──────────────┴─────────────┴─────────────┘
Total count :	34,826(0)
Dynamic count :	34,826(0)
Frozen count :	0(0)
------------------------------------------------------------------------
Total size :	136.04KB(0.00B)
Dynamic size :	136.04KB(0.00B)
Frozen size :	0.00B(0.00B)
========================================================================
tree diagram
print(model.tree_diagram())
CNN
    ├── conv1=Conv2D
       ├── weight=f32[32,1,3,3]
       ├── bias=f32[32,1,1]
       * in_features=1
       * out_features=32
       * kernel_size=(3, 3)
       * strides=(1, 1)
       * padding=((0, 0), (0, 0))
       * input_dilation=(1, 1)
       * kernel_dilation=(1, 1)
       ├── weight_init_func=Partial(init(key,shape,dtype))
       ├── bias_init_func=Partial(zeros(key,shape,dtype))
       * groups=1    
    ├── relu1=ReLU  
    * pool1=MaxPool2D
       * kernel_size=(2, 2)
       * strides=(2, 2)
       * padding='valid' 
    ├── conv2=Conv2D
       ├── weight=f32[64,32,3,3]
       ├── bias=f32[64,1,1]
       * in_features=32
       * out_features=64
       * kernel_size=(3, 3)
       * strides=(1, 1)
       * padding=((0, 0), (0, 0))
       * input_dilation=(1, 1)
       * kernel_dilation=(1, 1)
       ├── weight_init_func=Partial(init(key,shape,dtype))
       ├── bias_init_func=Partial(zeros(key,shape,dtype))
       * groups=1    
    ├── relu2=ReLU  
    * pool2=MaxPool2D
       * kernel_size=(2, 2)
       * strides=(2, 2)
       * padding='valid' 
    * flatten=Flatten
       * start_dim=0
       * end_dim=-1  
    ├── dropout=Dropout
       * p=0.5
       └── eval=None   
    └── linear=Linear
        ├── weight=f32[1600,10]
        ├── bias=f32[10]
        * in_features=1600
        * out_features=10  
    
Plot sample predictions before training
 
# set all dropout off
test_model = model.at[model == "eval"].set(True, is_leaf=lambda x: x is None)

def show_images_with_predictions(model, images, one_hot_labels):
    logits = jax.vmap(model)(images)
    predictions = jnp.argmax(logits, axis=-1)
    fig, axes = plt.subplots(5, 5, figsize=(10, 10))
    for i, ax in enumerate(axes.flat):
        ax.imshow(images[i].reshape(28, 28), cmap="binary")
        ax.set(title=f"Prediction: {predictions[i]}\nLabel: {jnp.argmax(one_hot_labels[i], axis=-1)}")
        ax.set_xticks([])
        ax.set_yticks([])
    plt.show()

example = ds_test.take(25).as_numpy_iterator()
example = list(example)
sample_test_images = jnp.stack([x["image"] for x in example])
sample_test_labels = jnp.stack([x["label"] for x in example])

show_images_with_predictions(test_model, sample_test_images, sample_test_labels)

image

🏃 Train the model

@ft.partial(jax.value_and_grad, has_aux=True)
def loss_func(model, batched_images, batched_one_hot_labels):
    logits = jax.vmap(model)(batched_images)
    loss = jnp.mean(optax.softmax_cross_entropy(logits=logits, labels=batched_one_hot_labels))
    return loss, logits


# using optax for gradient updates
optim = optax.adam(1e-3)
optim_state = optim.init(model)


@jax.jit
def batch_step(model, batched_images, batched_one_hot_labels, optim_state):
    (loss, logits), grads = loss_func(model, batched_images, batched_one_hot_labels)
    updates, optim_state = optim.update(grads, optim_state)
    model = optax.apply_updates(model, updates)
    accuracy = jnp.mean(jnp.argmax(logits, axis=-1) == jnp.argmax(batched_one_hot_labels, axis=-1))
    return model, optim_state, loss, accuracy


epochs = 5

for i in range(epochs):
    epoch_accuracy = []
    epoch_loss = []

    for example in ds_train.as_numpy_iterator():
        image, label = example["image"], example["label"]
        model, optim_state, loss, accuracy = batch_step(model, image, label, optim_state)
        epoch_accuracy.append(accuracy)
        epoch_loss.append(loss)

    epoch_loss = jnp.mean(jnp.array(epoch_loss))
    epoch_accuracy = jnp.mean(jnp.array(epoch_accuracy))

    print(f"epoch:{i+1:00d}\tloss:{epoch_loss:.4f}\taccuracy:{epoch_accuracy:.4f}")
    
# epoch:1	loss:0.2706	accuracy:0.9268
# epoch:2	loss:0.0725	accuracy:0.9784
# epoch:3	loss:0.0533	accuracy:0.9836
# epoch:4	loss:0.0442	accuracy:0.9868
# epoch:5	loss:0.0368	accuracy:0.9889

🎨 Visualize After training

test_model = model.at[model == "eval"].set(True, is_leaf=lambda x: x is None)
show_images_with_predictions(test_model, sample_test_images, sample_test_labels)

image

🥶 Freezing parameters /Fine tuning

See here for more about freezing

🔘 Filtering by masking

See here for more about filterning

Download files

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

Source Distribution

serket-0.0.7.tar.gz (61.4 kB view details)

Uploaded Source

Built Distribution

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

serket-0.0.7-py3-none-any.whl (72.5 kB view details)

Uploaded Python 3

File details

Details for the file serket-0.0.7.tar.gz.

File metadata

  • Download URL: serket-0.0.7.tar.gz
  • Upload date:
  • Size: 61.4 kB
  • Tags: Source
  • Uploaded using Trusted Publishing? No
  • Uploaded via: twine/4.0.1 CPython/3.10.7

File hashes

Hashes for serket-0.0.7.tar.gz
Algorithm Hash digest
SHA256 1ea034a8b13cd0b9ea5fd8d9b29fd5be908fee9524ef06cc3ca334b6979f0eae
MD5 561ee09759f4cda9c995d115a3178889
BLAKE2b-256 93556d6fe7770b433d8dc344ee312819e17c7711d6f46cb0eb08d208ff45bc79

See more details on using hashes here.

File details

Details for the file serket-0.0.7-py3-none-any.whl.

File metadata

  • Download URL: serket-0.0.7-py3-none-any.whl
  • Upload date:
  • Size: 72.5 kB
  • Tags: Python 3
  • Uploaded using Trusted Publishing? No
  • Uploaded via: twine/4.0.1 CPython/3.10.7

File hashes

Hashes for serket-0.0.7-py3-none-any.whl
Algorithm Hash digest
SHA256 3bb881c044d41bd6340eba53ce6166efae5ef78a6a260d9a06e02562c4c33556
MD5 574e79b8da960835a08c7be49d322f2c
BLAKE2b-256 64b279ab9baaf111496f352346bdf847c5623e520496c7bcb26fe4c31d2f1617

See more details on using hashes here.

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