Skip to main content

ebmkit

CI PyPI coverage Python License: MIT

Energy-based models in PyTorch — samplers, training losses, and honest evaluation, as composable objects with tested defaults.

An EBM is an unnormalized density p(x) ∝ exp(-E(x)) defined by a network E: (B, *shape) -> (B,). torch is the only runtime dependency.

Not the Explainable Boosting Machines that also go by "EBM" — this is the deep-learning kind (LeCun et al. 2006; Du & Mordatch 2019; Song & Kingma 2021).

Install

pip install ebmkit          # runtime dependency is just torch>=2.0
pip install "ebmkit[viz]"   # + matplotlib for the plotting helpers

The import name is ebm:

import torch, ebm

energy = ebm.nets.MLPEnergy(dim=2, hidden=(128, 128))
sampler = ebm.LangevinDynamics(step_size=1e-2, steps=60)
loss_fn = ebm.ContrastiveDivergence(sampler, buffer=ebm.ReplayBuffer(8192, (2,)))

trainer = ebm.Trainer(energy, loss_fn, lr=1e-3)
trainer.fit(ebm.datasets.two_moons(8192), steps=3000, batch_size=256)

samples = sampler.sample(energy, torch.randn(2000, 2), steps=500)

The Trainer is optional sugar — the loop underneath is plain PyTorch (each loss returns LossOutput(loss, metrics, x_neg); call out.loss.backward()).

What's in the box

Piece Contents
Energies any callable (B, *shape) -> (B,); nets.MLPEnergy / ConvEnergy / ResNetEnergy / ConvClassifier (SiLU, optional spectral norm, no batch norm), nets.RBM (Bernoulli RBM with exact log_z), nets.IsingEnergy / PottsEnergy (discrete lattices), nets.FunnelEnergy / GaussianMixtureEnergy / BananaEnergy (closed-form targets), noise-conditional variants for NCSN; EnergyModel, ebm.score
Samplers LangevinDynamics (ULA/SGLD), MALA, HMC, UnderdampedLangevin (SGHMC), PreconditionedLangevin, ParallelTempering (replica exchange), TemperedTransitions, SVGD (Stein variational), GibbsSampler (block Gibbs), GibbsWithGradients + CategoricalGibbsWithGradients, AnnealedLangevinDynamics, ProbabilityFlowODE / PredictorCorrector (score-SDE)
Losses ContrastiveDivergence (CD-k / persistent CD), DiffusionRecoveryLikelihood + drl_sample, DenoisingScoreMatching / MultiSigmaDenoisingScoreMatching (NCSN), SlicedScoreMatching, ExactScoreMatching, EnergyDiscrepancy (MCMC-free), PseudoLikelihood / RatioMatching / ConcreteScoreMatching (MCMC-free, discrete), NoiseContrastiveEstimation, JEMLoss
Composition SumEnergy (product of experts), MixtureEnergy, TemperedEnergy — energies compose like densities and nest
Training thin Trainer (device, EMA, supervised batches, save/load checkpointing), ReplayBuffer, EMA
Eval ais_log_z / reverse_ais_log_z (bracket log Z), bits_per_dim, frechet_distance (FID), mmd, precision_recall, inception_score, kernel_stein_discrepancy / classifier_two_sample_test / fisher_divergence (goodness-of-fit), mutual_information (MINE), ood_auroc, effective_sample_size / split_rhat (MCMC diagnostics)
Data & viz 2D toys (two_moons, eight_gaussians, checkerboard, rings, spirals) and torchvision-free image loaders (mnist, fashion_mnist, cifar10, cifar100); viz.energy_contour / plot_samples / energy_histogram / show_images

Examples

Runnable scripts in examples/ (python examples/<name>.py):

  • train_two_moons.py — the canonical 2D contrastive-divergence smoke test
  • train_jem.py / train_mnist_jem.py — classify, generate, and detect OOD with one network
  • train_mnist.py — the image-scale IGEBM short-run recipe
  • train_composition.py — product of experts / mixture / tempering, without retraining
  • train_ising.py / train_potts.py — discrete lattices via (categorical) Gibbs-with-Gradients
  • train_ncsn.py — score-based generation: multi-sigma denoising + annealed Langevin
  • deterministic_sampling.py — annealed Langevin vs the (deterministic) probability-flow ODE vs predictor-corrector
  • train_cifar_ood.py — energy-based OOD at color scale (CIFAR-10 vs CIFAR-100)
  • train_rbm.py — a Bernoulli RBM on binary bars via CD-1, with the exact log Z
  • train_ising_pseudolikelihood.py — recover an Ising coupling with no MCMC in the loop
  • train_potts_concrete.py — recover a categorical density with concrete score matching (no MCMC)
  • train_energy_discrepancy.py — two-moons trained MCMC-free (energy discrepancy)
  • sampling_hard_targets.py — parallel tempering escapes a trapped mode; ESS / R̂ diagnostics
  • goodness_of_fit.py — KSD for model selection; classifier two-sample test
  • mine_mutual_information.py — estimate mutual information with MINE vs the Gaussian closed form
  • benchmark_samplers.py — rank samplers on the banana against exact draws (ESS, R̂, MMD)
  • checkpoint_resume.py — save a run and resume it in a fresh process

Conventions

  • Sign: p ∝ exp(-E) — low energy is high probability. Samplers descend the energy gradient; training pushes data energy down. Never flip this.
  • Stop-gradients: MCMC negatives are detached and the energy's parameters are frozen during sampling; score-matching losses instead keep the graph (create_graph=True).
  • The CD loss value is not a convergence signal — it hovers near zero at equilibrium; watch metrics["energy_gap"] and energy histograms.

See CONTRIBUTING.md for the rest.

Development

uv run pytest        # tests (CPU-only, seeded)
uv run ruff check .  # lint
uv run mypy          # type-check

Citation

If you use ebmkit in your research, please cite it — see CITATION.cff.

License

MIT — see LICENSE.

Metadata

Release files for ebmkit 0.13.0

For a detailed explanation of source distributions (sdists) and built distributions (wheels), please see the package formats documentation.

Source distribution (sdist)

Source distribution for ebmkit 0.13.0
File Size Uploaded
ebmkit-0.13.0.tar.gz 3.0 MB Details

Built distribution (wheel)

Table of built distributions (wheels) for ebmkit 0.13.0
File Interpreter ABI Platform
ebmkit-0.13.0-py3-none-any.whl Python 3 none any Details

Total release size: 3.1 MB

Release files / ebmkit-0.13.0.tar.gz

Download URL ebmkit-0.13.0.tar.gz
Size 3.0 MB
Tags Source
SHA-256 checksum
How to use checksums
08a22ff0da9158c70fbcc98a36c278b10f6812840d276a349e5dc43c35707c01
BLAKE2b-256 checksum
How to use checksums
91d04b2d8435812fb7e42abe758f4e2d33f795a9681b41c6292dbf3d1519abde
Upload date
Uploaded using Trusted Publishing?
What is trusted publishing?
Yes
Uploaded via twine/7.0.0 CPython/3.13.14

Provenance

Provenance describes where a file came from. On PyPI, provenance is shared via attestations, which provide a verifiable record of the build or publishing details. View details, limitations and caveats.

PyPI Publish Attestation

PyPI verified that this artifact, at this checksum, originated from the publisher listed below.

Signed by GitHub Actions, verified by PyPI on Aug 17, 2026.

Transparency log

Release files / ebmkit-0.13.0-py3-none-any.whl

Download URL ebmkit-0.13.0-py3-none-any.whl
Size 66.0 kB
Tags Python 3
SHA-256 checksum
How to use checksums
c1607a65268f85419e08bb747d2666c37574c5f32d93e32ff089149e4bf34c95
BLAKE2b-256 checksum
How to use checksums
5d5da123031fff723602f8cb0902b46de49189d055edb41f66bbe280a6359a48
Upload date
Uploaded using Trusted Publishing?
What is trusted publishing?
Yes
Uploaded via twine/7.0.0 CPython/3.13.14

Provenance

Provenance describes where a file came from. On PyPI, provenance is shared via attestations, which provide a verifiable record of the build or publishing details. View details, limitations and caveats.

PyPI Publish Attestation

PyPI verified that this artifact, at this checksum, originated from the publisher listed below.

Signed by GitHub Actions, verified by PyPI on Aug 17, 2026.

Transparency log

Release history Release notifications | RSS feed

0.16.0

2 release files

0.15.0

2 release files

0.14.0

2 release files

This release

0.13.0 This release

2 release files

0.12.0

2 release 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