Skip to main content

Like Karpathy's micrograd but following JAX's functional style

Project description

microjax

A tiny autograd engine following the spirit of Karpathy's micrograd.

Like micrograd, it implements backpropagation (reverse-mode autodiff) over a dynamically built DAG of scalar values and a small neural networks library on top of it.

Unlike micrograd, which is implemented following PyTorch's OOP API, microjax replicates JAX's functional API. In particular, it exposes the transformation microjax.engine.grad. If you have a Python function f that evaluates the mathematical function $f$, then grad(f) is a Python function that evaluates the mathematical function $\nabla f$. That means that grad(f)(x1, ..., xn) represents the value $\nabla f(x_1, \ldots, x_n)$. For univariate functions, grad can be applied to its own output to compute higher order derivatives. For example given the mathematical function $g(x)$ and its Python representation g(x), grad(grad(g))(x) represents the value $g''(x)$.

In combination with micrograd, microjax could be useful to illustrate the differences between the OOP and functional paradigms. The functional paradigm is characterized by the use of pure functions acting on immutable state, and higher order functions (transformations) that act on pure functions to return new pure functions. These are all apparent in the implementation of microjax, e.g. f -> grad(f).

In micrograd, one composes differentiable functions as a succession of operations acting on instances of Value, which is micrograd's object to represent nodes in the DAG. Each new operation produces an output that is a new instance of Value, aware of its parents and the operation that created it, thus building the computational DAG. Finally, one can call .backward() on the output Value of the function to compute its gradient with respect to all nodes in the computational DAG. The gradient with respect to a Value of name, e.g., x is accessed as x.grad.

In microjax, one composes differentiable functions as a succession of operations defined as instances of Primitive, which is microjax's object to represent primitive operations that are differentiable and traceable. When grad(f) is called, it wraps the function's arguments as instances of Tracer, which is microjax's object to represent nodes in the DAG. Then, it evaluates f on the traced inputs, generating the computational DAG in the process, and computes the gradient with respect to all nodes in the computational DAG. Finally, it returns the gradient with respect to the input arguments.

Installation

pip install microjax

Example usage

Below is a slightly contrived example showing a number of supported operations. It is a replica of micrograd's example, for comparison.

from microjax.engine import grad, relu

def g_fn(a, b):
    c = a + b
    d = a * b + b**3
    c += c + 1
    c += 1 + c + (-a)
    d += d * 2 + relu(b + a)
    d += 3 * d + relu(b - a)
    e = c - d
    f = e**2
    g = f / 2.0
    g += 10.0 / f

    return g

a = -4.0
b = 2.0

g = g_fn(a, b)
dgda, dgdb = grad(g_fn)(a, b)

print(f'{g:.4f}') # prints 24.7041, the outcome of this forward pass
print(f'{dgda:.4f}') # prints 138.8338, i.e. the numerical value of dg/da
print(f'{dgdb:.4f}') # prints 645.5773, i.e. the numerical value of dg/db

Training a neural net

The notebook demo.ipynb provides a full demo of training a 2-layer neural network (MLP) binary classifier. This is achieved by initializing a neural net from microjax.nn module, implementing a simple svm "max-margin" binary classification loss and using GD for optimization. As shown in the notebook, using a 2-layer neural net with two 16-node hidden layers we achieve the following decision boundary on the moon dataset:

2d neuron

Again, this is a replica of micrograd's demo, for comparison.

The demo.ipynb uses additional libraries for visualization and training examples. This project uses Hatch for environment managing and testing.

If you don't already have hatch installed:

pip install hatch

Then select .venv.default as the kernel when opening demo.ipynb.

Running tests

Tests use PyTorch as a reference for verifying the correctness of the calculated gradients.

If you have installed hatch, from the root of your repository:

hatch run test

This will:

  • Automatically create a virtual environment (if needed),
  • Install all development and testing dependencies (including PyTorch),
  • Run the test suite using pytest.

License

MIT

Project details


Download files

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

Source Distribution

microjax-0.1.0.tar.gz (86.3 kB view details)

Uploaded Source

Built Distribution

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

microjax-0.1.0-py3-none-any.whl (7.0 kB view details)

Uploaded Python 3

File details

Details for the file microjax-0.1.0.tar.gz.

File metadata

  • Download URL: microjax-0.1.0.tar.gz
  • Upload date:
  • Size: 86.3 kB
  • Tags: Source
  • Uploaded using Trusted Publishing? No
  • Uploaded via: twine/6.1.0 CPython/3.13.5

File hashes

Hashes for microjax-0.1.0.tar.gz
Algorithm Hash digest
SHA256 0d4edbbf0d7b72a1dd76c9870f7d0b34a209812b97af9809b81f866db616d72a
MD5 ed09e40206ef1781382957147fb4b66d
BLAKE2b-256 b2eb810cf9e01492a885e2cb5216321d131a4861a7167a3d8a39660d251d31bd

See more details on using hashes here.

File details

Details for the file microjax-0.1.0-py3-none-any.whl.

File metadata

  • Download URL: microjax-0.1.0-py3-none-any.whl
  • Upload date:
  • Size: 7.0 kB
  • Tags: Python 3
  • Uploaded using Trusted Publishing? No
  • Uploaded via: twine/6.1.0 CPython/3.13.5

File hashes

Hashes for microjax-0.1.0-py3-none-any.whl
Algorithm Hash digest
SHA256 01cdd491108a28bca66e8f8e89f3653a13db236d8eb4e2e3c981dcf77876a0e9
MD5 aed324c81cd02d0ce280bc9374525d5f
BLAKE2b-256 aec368ec253bd05b95012bfe3c772fa1368d6d5ce5945381791968922ed31c60

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