Skip to content

v0.1.0 — Multi-Backend TP-GMM

Latest

Choose a tag to compare

@RobinU434 RobinU434 released this 19 Feb 12:22

First stable release of tpgmm with full support for three computational backends.

Highlights

  • Multi-backend architecture — unified API across NumPy, PyTorch, and JAX with shared abstract base classes in tpgmm._core
  • JIT-compiled JAX backend — entire EM iteration compiled via @jax.jit, making JAX the fastest backend on CPU (~10 ms per fit)
  • Vectorized Gaussian PDF — einsum-based implementation with no Python loops in the hot path across all backends
  • Model selection — BIC, AIC, silhouette score, and log-likelihood scoring built into every backend

Core

  • TPGMM — Task Parameterized Gaussian Mixture Model with K-Means initialization and Expectation Maximization
  • GaussianMixtureRegression — conditional regression from a fitted TP-GMM, with from_tpgmm() convenience constructor
  • Abstract base classes (BaseTPGMM, BaseGMR, LearningModule) ensuring consistent interfaces and documentation

Backends

Backend TPGMM GMR Notes
NumPy yes yes Pure NumPy + SciPy, no GPU required
PyTorch yes yes GPU-ready, eager execution
JAX yes yes IT-compiled EM loop, GPU/TPU-ready

Extras

  • Three example notebooks (examples/example_numpy.ipynb, example_torch.ipynb, example_jax.ipynb) with identical structure: data loading → model fitting → model selection → GMR → visualization
  • Comprehensive test suite (46 tests) with pytest-benchmark runtime profiling
  • Optional dependency groups: torch, jax, examples, dev, all

Installation

Currently, only supported via install from the repository.

pip install tpgmm              # NumPy backend
pip install tpgmm[torch]       # + PyTorch
pip install tpgmm[jax]         # + JAX
pip install tpgmm[all]         # Everything

Requirements

  • Python >= 3.12
  • NumPy, SciPy, scikit-learn (base)
  • PyTorch >= 2.0 (optional)
  • JAX >= 0.4 (optional)