Repository navigation
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] # EverythingRequirements
- Python >= 3.12
- NumPy, SciPy, scikit-learn (base)
- PyTorch >= 2.0 (optional)
- JAX >= 0.4 (optional)