Skip to main content

MiniML

Run tests PyPI

MiniML (pronounced "minimal") is a tiny machine-learning framework which uses Jax as its core engine, but mixes a PyTorch inspired approach to building model with Scikit-learn's interface (using the .fit and .predict methods), and is powered by SciPy's optimization algorithms. It's meant for simple prototyping of small ML architectures that allows more flexibility than Scikit's built-in models without sacrificing too much on performance.

Training a linear model in MiniML for example looks as simple as this:

class LinearModel(MiniMLModel):
    A: MiniMLParam
    b: MiniMLParam

    def __init__(self, n_in: int, n_out: int):
        self.A = MiniMLParam((n_in,n_out))
        self.b = MiniMLParam((n_out,))
        super().__init__()

    def _predict_kernel(self, X, buffer):
        return X@self.A(buffer)+self.b(buffer)

lin_model = LinearModel(X.shape[1], y.shape[1])
lin_model.randomize()
lin_model.fit(X, y)
y_hat = lin_model.predict(X)

Note that calling a parameter with the buffer as an argument returns the value of that parameter.

Installation

Simply install this package from PyPi:

pip install miniml-jax

or to use the CUDA-enabled version of Jax:

pip install miniml-jax[cuda]

Usage

The two core types are MiniMLParam and MiniMLModel. There are also MiniMLParamList and MiniMLModelList containers to store multiple of either inside.

To define a model in MiniML, subclass MiniMLModel and define your parameters as MiniMLParam attributes in the __init__ method. Remember to make sure that:

  • every parameter or child model is stored either directly as a class member, or inside a corresponding List class;
  • the super().__init__() constructor is called at the end.

Then, implement the internal _predict_kernel method, which takes an input array as well as a memory buffer containing the parameters and returns the model's prediction. After instantiating your model, call bind() to initialize parameter buffers, or use directly randomize() to initialize parameter values. You can then use methods like fit, save, and load.

Example: Linear Model

import jax.numpy as jnp
from miniml.param import MiniMLParam
from miniml.model import MiniMLModel

class LinearModel(MiniMLModel):
    def __init__(self):
        self.a = MiniMLParam((1,))
        self.b = MiniMLParam((1,))
        super().__init__()

    def _predict_kernel(self, X, buffer):
        return X@self.A(buffer)+self.b(buffer)

# Create and bind the model
model = LinearModel()
model.bind()
model.randomize()

# Fit to data (e.g., y = 2x + 1)
X = jnp.linspace(0, 10, 20)
y = 2 * X + 1
model.fit(X, y)

# Save and load
model.save('model.npz')
model.load('model.npz')

Nested Models

You can compose models by including other MiniMLModel instances as attributes. In this case, remember to always call child models via their own _predict_kernel methods too. For example:

class ConstantModel(MiniMLModel):

    def __init__(self):
        self._c = MiniMLParam((1,))
        super().__init__()

    def _predict_kernel(self, X, buffer):
        return self._c(buffer)

class LinearWithConstant(MiniMLModel):
    def __init__(self):
        self._b = MiniMLParam((5,))
        self._M = MiniMLParam((5, 5))
        self._c = ConstantModel()
        super().__init__()

    def _predict_kernel(self, X, buffer):
        return self._M(buffer) @ X + self._b(buffer)[:, None] + self._c._predict_kernel(X, buffer)

See the full documentation.

Download files

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

Source Distribution

miniml_jax-0.10.0.tar.gz (31.1 kB view details)

Uploaded Source

Built Distribution

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

miniml_jax-0.10.0-py3-none-any.whl (42.2 kB view details)

Uploaded Python 3

File details

Details for the file miniml_jax-0.10.0.tar.gz.

File metadata

  • Download URL: miniml_jax-0.10.0.tar.gz
  • Upload date:
  • Size: 31.1 kB
  • Tags: Source
  • Uploaded using Trusted Publishing? Yes
  • Uploaded via: twine/7.0.0 CPython/3.13.14

File hashes

Hashes for miniml_jax-0.10.0.tar.gz
Algorithm Hash digest
SHA256 e577de95f90f2a6e760671062eef64193303f67523aee2c151777f1c5db9c693
MD5 85f8874b1871a098a0d915c2acf5afae
BLAKE2b-256 09d3ca293b5746c5b44a31b0b460d9291cfab9219aea8a80a903b90374953605

See more details on using hashes here.

Provenance

The following attestation bundles were made for miniml_jax-0.10.0.tar.gz:

Publisher: python-publish.yml on stur86/miniml

Attestations: Values shown here reflect the state when the release was signed and may no longer be current.

File details

Details for the file miniml_jax-0.10.0-py3-none-any.whl.

File metadata

  • Download URL: miniml_jax-0.10.0-py3-none-any.whl
  • Upload date:
  • Size: 42.2 kB
  • Tags: Python 3
  • Uploaded using Trusted Publishing? Yes
  • Uploaded via: twine/7.0.0 CPython/3.13.14

File hashes

Hashes for miniml_jax-0.10.0-py3-none-any.whl
Algorithm Hash digest
SHA256 d1dc9f930cd3fe996a65106ffcc3aa633ec427d2f871e2abdd2d74ec02d86ae8
MD5 3cfb837733a3b1a4b58fdf3ded1d2b9f
BLAKE2b-256 8d7ac3eef6644b53962b79b78d0c1913d97eaece3c00f2ec423c32870ec16859

See more details on using hashes here.

Provenance

The following attestation bundles were made for miniml_jax-0.10.0-py3-none-any.whl:

Publisher: python-publish.yml on stur86/miniml

Attestations: Values shown here reflect the state when the release was signed and may no longer be current.

Release history Release notifications | RSS feed

This release

0.10.0 This release

2 files

0.9.0

2 files

0.8.0

2 files

0.7.0

2 files

0.6.0

2 files

0.5.0

2 files

0.4.0

2 files

0.3.0

2 files

0.2.0

2 files

0.1.1

2 files

0.1.0

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