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), GuidedEnergy (classifier-free guidance), nets.IsingEnergy / PottsEnergy (discrete lattices), nets.FunnelEnergy / GaussianMixtureEnergy / BananaEnergy (closed-form targets), nets.AffineCouplingFlow (RealNVP) / nets.NeuralSplineCouplingFlow (rational-quadratic spline) / nets.ContinuousNormalizingFlow (FFJORD — trainable exact-likelihood flows / self-normalized energies), noise-conditional variants for NCSN; EnergyModel, ebm.score
Samplers LangevinDynamics (ULA/SGLD), MALA, AdaptiveMALA (dual-averaging step-size warmup + diagonal metric), HMC, UnderdampedLangevin (SGHMC), PreconditionedLangevin, ParallelTempering (replica exchange), TemperedTransitions, SVGD (Stein variational), GibbsSampler (block Gibbs), GibbsWithGradients + CategoricalGibbsWithGradients, AnnealedLangevinDynamics, ProbabilityFlowODE / PredictorCorrector (score-SDE), DDPMAncestralSampler (VP diffusion)
Losses ContrastiveDivergence (CD-k / persistent CD), DiffusionRecoveryLikelihood + drl_sample, DenoisingScoreMatching / MultiSigmaDenoisingScoreMatching (NCSN), VPDenoisingScoreMatching (DDPM), SlicedScoreMatching, ExactScoreMatching, EnergyDiscrepancy (MCMC-free), PseudoLikelihood / RatioMatching / ConcreteScoreMatching (MCMC-free, discrete), NoiseContrastiveEstimation, JEMLoss
Composition SumEnergy (product of experts), MixtureEnergy, EnsembleEnergy (deep-ensemble mean energy + member disagreement), 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), expected_calibration_error / reliability_curve / temperature_scale (calibration), ood_auroc, ensemble_disagreement (epistemic OOD), effective_sample_size / split_rhat / autocorrelation (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 / autocorrelation_plot / rank_plot / trace_plot

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
  • jem_guidance.py — classifier-free guidance sharpening class-conditional samples
  • 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_diffusion.py — a variance-preserving (DDPM) diffusion, energy-parameterized
  • 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_spline_flow.py — a rational-quadratic spline flow: sharper fit than affine at equal depth
  • train_cnf.py — a continuous normalizing flow (FFJORD): exact likelihood by ODE
  • 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
  • adaptive_mala.py — a self-tuning MALA: dual-averaging step size + a learned diagonal metric
  • goodness_of_fit.py — KSD for model selection; classifier two-sample test
  • ensemble_ood.py — a deep-ensemble EBM whose member disagreement flags OOD
  • 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.15.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.15.0
File Size Uploaded
ebmkit-0.15.0.tar.gz 4.3 MB Details

Built distribution (wheel)

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

Total release size: 4.4 MB

Release files / ebmkit-0.15.0.tar.gz

Download URL ebmkit-0.15.0.tar.gz
Size 4.3 MB
Tags Source
SHA-256 checksum
How to use checksums
b7a22671106ae2640bb4b376c92021689aa25d6d50024e691bbdb39d38556a90
BLAKE2b-256 checksum
How to use checksums
b5e91748e2e04a26a6f0749d5e70eab19d7455943cee1ede6d343b3591554a69
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 20, 2026.

Transparency log

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

Download URL ebmkit-0.15.0-py3-none-any.whl
Size 89.9 kB
Tags Python 3
SHA-256 checksum
How to use checksums
224a9cc6feefeb48694f3c9d69833d8e55568dab42c6825780f0cd3c79aacb5a
BLAKE2b-256 checksum
How to use checksums
2c90550dbd3d1087e20d7baad0cb0922b6a7a1a556b44d6a88744a2d443215a7
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 20, 2026.

Transparency log

Release history Release notifications | RSS feed

0.16.0

2 release files

This release

0.15.0 This release

2 release files

0.14.0

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