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), nets.AffineCouplingFlow (RealNVP — exact-likelihood flow / self-normalized energy), 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), pf_ode_log_likelihood (exact likelihood via the probability-flow ODE), 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
  • exact_likelihood_ode.py — exact bits/dim for a score model via the probability-flow ODE
  • train_flow.py — a RealNVP normalizing flow with exact density and sampling
  • 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.14.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.14.0
File Size Uploaded
ebmkit-0.14.0.tar.gz 3.2 MB Details

Built distribution (wheel)

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

Total release size: 3.3 MB

Release files / ebmkit-0.14.0.tar.gz

Download URL ebmkit-0.14.0.tar.gz
Size 3.2 MB
Tags Source
SHA-256 checksum
How to use checksums
16771d450e9c55aeec72b2949c9f4baa0ef36f1150c8ea90a479968f373f388f
BLAKE2b-256 checksum
How to use checksums
db0c54f06f73321163af72456b5287a139bc055a4ce225f8df4789fd61abdfe6
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.14.0-py3-none-any.whl

Download URL ebmkit-0.14.0-py3-none-any.whl
Size 68.2 kB
Tags Python 3
SHA-256 checksum
How to use checksums
9ad977b6ef94a3bb5cf80d49b7049061e909890ae5c4aa9c7a5b61e60f7b6113
BLAKE2b-256 checksum
How to use checksums
d72f6220c6ca6b2c45004cfe121144610d7e61d6b828bfc58aee0c0cdceb861e
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

This release

0.14.0 This release

2 release files

0.13.0

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