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, NUTS (No-U-Turn Sampler, self-tuning trajectory length), 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, LatentEBM (joint E(x, z) with a prior + decoder, block-Gibbs sampled) — 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
  • nuts_sampling.py — the No-U-Turn Sampler on Neal's funnel: trajectory length adapts per draw
  • goodness_of_fit.py — KSD for model selection; classifier two-sample test
  • ensemble_ood.py — a deep-ensemble EBM whose member disagreement flags OOD
  • latent_ebm.py — a latent-variable EBM: block-Gibbs on a joint E(x, z) matches ancestral sampling
  • 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.16.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.16.0
File Size Uploaded
ebmkit-0.16.0.tar.gz 4.5 MB Details

Built distribution (wheel)

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

Total release size: 4.6 MB

Release files / ebmkit-0.16.0.tar.gz

Download URL ebmkit-0.16.0.tar.gz
Size 4.5 MB
Tags Source
SHA-256 checksum
How to use checksums
6f2af7309f6461a9aa241cb1234b8f718ea701ad0bff52b37980cbe580c61a06
BLAKE2b-256 checksum
How to use checksums
2f46db94590b3c0fcd64c699c9dd54ae7afbaec7a5aa82b4e9ca506fd2ff4a96
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 21, 2026.

Transparency log

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

Download URL ebmkit-0.16.0-py3-none-any.whl
Size 96.1 kB
Tags Python 3
SHA-256 checksum
How to use checksums
5ece460e0e465287a9d3b78341dcf1f79810f3eef56f7e90c562bd24cf5650f5
BLAKE2b-256 checksum
How to use checksums
f00ec22dab010bbe683c66336f045af628217e0274a93cf4c289cbe76b03609e
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 21, 2026.

Transparency log

Release history Release notifications | RSS feed

This release

0.16.0 This release

2 release files

0.15.0

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