Deconvolve
Unbinned, full-phase-space unfolding. An adversarial neural network learns per-event weights that correct simulated (Monte Carlo) distributions so they match observed data. Built with Keras 3 on the JAX backend.
Deconvolution is the inverse problem of recovering a true signal from detector-smeared observations — the information was never lost, only transformed. That is the detector, and the weights this package learns perform the deconvolution.
Motivation
In particle physics, Monte Carlo (MC) simulations are used to model detector responses and physical processes. These simulations never perfectly reproduce real data. There are always residual mismodelling effects. Traditional reweighting uses hand-tuned correction factors binned in one or two variables, which scales poorly to high-dimensional feature spaces.
Deconvolve replaces this with a learned reweighting: a generator network predicts a continuous per-event weight from particle-level (truth) features, while an adversarial discriminator tries to distinguish the reweighted simulation from real data. At convergence the discriminator can no longer tell them apart, and the generator's weights constitute an optimal correction.
Model
The system is a two-player adversarial game over event weights:
| Component | Input | Output | Role |
|---|---|---|---|
| Generator $g(z)$ | Particle-level feature $z$ | Per-event weight (logplus = $\log(1+\exp)$ ) |
Predict weights that make MC look like nature |
| Discriminator $d(x)$ | Detector-level feature $x$ | Data vs MC probability (sigmoid) |
Distinguish real data from reweighted MC |
Training loop
- Discriminator step: freeze $g$, update $d$ to maximize weighted binary cross-entropy (classify Data vs reweighted Simulation)
- Generator step: freeze $d$, update $g$ to minimize the same loss (fool the discriminator).
- Repeat with 5:1 D:G update ratio
Weight normalization ensures the total MC yield is preserved:
$$w_i = \frac{g(z_i)}{\text{mean}(g(z))}$$
In equilibrium, both losses converge to $\log(2)$ and the reweighted MC matches data.
Installation
Requires Python >= 3.13. Uses uv for dependency management. One way to install it is with pip install uv; for alternatives see the uv documentation.
git clone https://github.com/krishdesai7/deconvolve.git
cd deconvolve
uv sync
This installs the deconvolve console script into .venv/bin. Commands below are
written as deconvolve ...; from a checkout without an activated virtualenv, prefix
them with uv run (uv run deconvolve train --config params/1d_default.yaml).
Tab completion for subcommands, flags and enum values is available with:
deconvolve --install-completion
GPU Support
The JAX dependency is platform-resolved:
Linux, x86_64
Built against jax[cuda13] on x86_64 Linux, compiled against CUDA version 13.0. For NVIDIA GPUs, the CUDA 13 runtime libraries are available as pypi wheels that the JAX binary is built against, so only a compatible NVIDIA driver is needed.
macOS, arm64 (Apple Silicon)
The official macOS arm64 wheels for JAX do not provide GPU acceleration. Therefore JAX and consequentially Deconvolve only offer CPU support on Apple Silicon;
Experimental alternatives, such as jax-mps or IREE-based workflows, may enable Metal acceleration, but these configurations are not tested or supported by Deconvolve. Users should independently validate their correctness and performance.
Usage
Gaussian Datasets
Gaussian datasets are configured via YAML files. Examples are provided in params/:
# 1D uncorrelated Gaussian
deconvolve train --config params/1d_default.yaml
# 2D with correlated covariance
deconvolve train --config params/2d_correlated.yaml
# 4D and 6D correlated
deconvolve train --config params/4d_correlated.yaml
deconvolve train --config params/6d_correlated.yaml
# Customize network and training
deconvolve train --config params/1d_default.yaml --hidden-units 128 --n-layers 3 --n-epochs 200
YAML config format (see params/ for examples):
mu_gen: [0.5]
mu_true: [0.0]
sigma_gen: 0.9 # scalar, vector, or full covariance matrix
sigma_true: 1.0
sigma_detector: 0.5
Sigma values are promoted to covariance matrices:
- scalar $\to \sigma^2 I$
- vector $\to \text{diag}(\sigma^2)$
- matrix $\to$ as-is
Jet Substructure
# All 6 jet variables
deconvolve train --dataset jets
# Specific variables
deconvolve train --dataset jets --var m --var w
Other Options
# Reload an existing run (regenerate plots/metrics)
deconvolve train --load-run runs/2026-03-14T061023Z
# Enable debug logging for any command
deconvolve --log-level DEBUG train --config params/1d_default.yaml
# SLURM submission
sbatch scripts/submit.sh --config params/2d_correlated.yaml
sbatch scripts/submit.sh --dataset jets
| Flag | Default | Description |
|---|---|---|
--config |
None |
Path to Gaussian YAML config |
--dataset |
gaussian |
Dataset type: gaussian or jets |
--n-samples |
500_000 |
Number of events per class (data + MC) |
--batch-size |
1024 |
Training batch size |
--hidden-units |
64 |
Units per hidden layer |
--n-layers |
2 |
Number of hidden layers |
--n-epochs |
100 |
Epochs to train (best checkpoint is always restored) |
--var |
all 6 | Repeat once for each jet substructure variable to use |
--load-run |
None |
Path to an existing run directory to reload |
--seed |
system entropy | Weight-initialization seed (see Seeding) |
--data-seed |
42 |
Data generation, shuffle, split and batch order |
--n-samples, --batch-size and --var also accept the short forms -n, -b
and -v; the global --log-level accepts -l.
The pipeline will:
- Generate (or load from cache) the dataset
- Split into train / validation / test sets (70 / 10 / 20%)
- Train the Deconvolve with early stopping
- Save models, training history, and plots to
runs/<UTC-timestamp>/ - Compute distance metrics on the test set
Evaluation
Distance metrics can be computed independently on existing runs:
# Evaluate all runs
deconvolve evaluate
# Evaluate a single run
deconvolve evaluate --run-dir runs/2026-03-14T061023Z
# Recompute even if metrics.json exists
deconvolve evaluate --force
This computes per-dimension 1D Wasserstein distances, Jensen-Shannon divergences, and triangular discriminator (Vincze-LeCam divergence) [$\times10^3$] at both detector and particle level, before and after reweighting. Results are saved to metrics.json in each run directory.
Reports
One PDF dossier per run — configuration, timing, both metrics tables and every figure — built from the JSON a run already writes:
# Compile runs/<timestamp>/report.pdf
deconvolve report runs/2026-03-14T061023Z
# Rebuild one that already exists
deconvolve report runs/2026-03-14T061023Z --force
# Emit artifacts/report.tex alone, without a TeX installation
deconvolve report runs/2026-03-14T061023Z --no-compile
report.tex is written into artifacts/; report.pdf lands at the run root
beside config.json. Compilation needs pdflatex on PATH. A run missing its
baseline, its timings or even its metrics still reports: the affected cells
degrade to dashes or a labelled row rather than failing. scripts/submit.sh
ends with deconvolve report, so the report sees the IBU overlay, the redrawn figures
and the recomputed metrics.
Baseline Comparisons
Run IBU (Iterative Bayesian Unfolding) on the same datasets for head-to-head comparison:
# IBU — single run
deconvolve baseline ibu --run-dir runs/2026-03-14T061023Z
# IBU — all runs
deconvolve baseline ibu
Results are saved to metrics_ibu.json in each run directory using the same metric format as Deconvolve.
Leakage Verification
A core correctness requirement is that the generator $g(z)$ never receives $z_\text{true}$, the particle-level values of measured data events, which are unknowable in a real experiment. The leakage-check command verifies this empirically via a data poisoning test:
# Clean run — z_true drawn from N(0, 1) as normal
deconvolve leakage-check --clean
# Poisoned run — z_true overwritten with -999 after x_data is generated
deconvolve leakage-check --poison
The poisoned run corrupts every data particle-level value to a nonsense sentinel (-999) while leaving $x_\text{data}$ (the reco-level observations the discriminator actually sees) unchanged. If $g$ had any access to $z_\text{true}$, the poisoned run would produce degraded weights. Both runs should report statistically identical Wasserstein and triangular discriminator improvements. Matching results confirm that no leakage path exists.
Both arms must share --seed, or initialization variance swamps the effect and
the arms differ even with no leakage. With it fixed, detector-level results are
bit-identical between the clean and poisoned arms.
Backend
JAX is the only backend in the build; TensorFlow is not a dependency, direct or transitive.
src/deconvolve/__init__.py sets KERAS_BACKEND=jax and JAX_ENABLE_X64=0. Keras 3 still defaults to TensorFlow when that variable is unset, so the pin is what makes import keras work here at all. The backend is fixed at the first keras import, so the pin has to land before it — which is why it lives in the package __init__, and why any deconvolve.* import must come before import keras. src/deconvolve/train.py raises a clear error if the backend has been initialized to something else.
Precision
The project runs in float32 end to end. The pin is a single constant, EVENT_DTYPE in src/deconvolve/coretypes/constants.py, with the annotation alias EventArray alongside it; JAX_ENABLE_X64=0 and the dtype= arguments in src/deconvolve/models.py follow from it.
This is a measured choice, not a default. Every jet observable is float32-clean — mass and mult survive a float32 round trip bit-exactly, and the other four lose exactly half a ULP, the least a cast can cost. Across 20 paired seeds, float32 and float64 agree on unfolding improvement to within ±0.5 percentage points (equivalence test p=0.015), while the seed-to-seed spread within either precision is larger than the gap between them. benchmarks/precision.py reproduces the comparison and benchmarks/compare_precision.py runs the statistics.
deconvolve.data.download computes jet observables in float64, because the ε protecting degenerate jets is below the smallest float32 denormal.
src/deconvolve/train.py is a hand-rolled loop, since the two-optimizer min-max game does not fit a standard keras.Model.fit. It does, however, follow the standard Keras 3 + JAX pattern:
- Model state lives in JAX pytrees (
TrainState) for the duration of training - Updates are applied through
stateless_call/stateless_apply - Each step is a single jitted function.
- Values are written back into the Keras models at the end, so the returned objects are ordinary saveable
keras.Models. - Loss math is plain
jnp.stateless_call/stateless_applyare the only Keras calls inside the trace;lax.scan,lax.while_loopandjax.randomare all native JAX, so backend-agnostickeras.opsbought nothing this module could still use.
One unexpected behaviour is worth flagging, because it is the reason the reduction is written the way it is:
keras.ops.meanis not float64-safe.- For float64 input, it selects a float32 compute dtype internally and returns a float64 result carrying ~1e-8 relative error.
src/deconvolve/train.pyno longer toucheskeras.ops, but it still reduces withjnp.sum(...) / nrather than a mean, andtests/test_train.pyguards the accuracy either way.- Anything that reaches for
keras.opsagain needs to know.ops.sumis unaffected.
Seeding
Two independent randomness axes, deliberately kept separate:
| Seed | Controls |
|---|---|
--data-seed |
Generation, shuffle, train/val/test split, batch order |
--seed |
Weight initialization only |
--seed defaults to a draw from system entropy, and the value used is recorded in config.json, so a run stays reproducible after the fact.
Configs predating this default used data_seed=42.
To estimate model uncertainty, ensemble, i.e. rerun on the same inputs with fresh initializations and take the variance as the model uncertainty, is a loop over --seed at fixed --data-seed.
Because the networks are Dense-only (no dropout or batch norm) and Adam is deterministic, the two seeds together fully determine a run, up to non-deterministic GPU reductions.
Force bitwise reproducibility with XLA_FLAGS=--xla_gpu_deterministic_ops=true. This costs throughput and is not needed for variance estimates.
Project Structure
Deconvolve/
├── src/deconvolve/ Python package
│ ├── __init__.py Pins KERAS_BACKEND=jax and JAX_ENABLE_X64=1
│ ├── __main__.py Fallback entry point (python -m deconvolve)
│ ├── cli.py Unified Typer command tree; target of the `deconvolve` script
│ ├── workflow.py Training and reload workflow
│ ├── logging_config.py Structured application logging
│ ├── leakage.py Data-poisoning leakage check
│ ├── py.typed PEP 561 typing marker
│ ├── coretypes/
│ │ ├── events.py Split, Events, ZXY, Populations, DatasetSplits
│ │ ├── configs.py GaussianConfig, RunConfig
│ │ ├── results.py UnfoldingPopulations, VariableOutcome, IBUResult
│ │ ├── constants.py Zenodo record, cache layout, jet plot metadata
│ │ ├── enums.py CLI choice enums
│ │ └── types.py TypedDicts and array aliases
│ ├── data/
│ │ ├── config.py YAML config parsing, sigma promotion
│ │ ├── datasets.py DatasetSplits, DeconvolveDataset, caching
│ │ ├── jets.py Jet substructure loading and standardization
│ │ ├── device.py Device-resident training form (TrainSplit/EvalSplit)
│ │ └── download.py One-time Zenodo data download
│ ├── baselines/
│ │ ├── _shared.py Run config and populations a baseline needs, minus the unfolder
│ │ └── ibu.py IBU (Iterative Bayesian Unfolding) baseline
│ ├── models.py Generator and discriminator architectures
│ ├── train.py JAX adversarial training loop with early stopping
│ ├── plotting.py Detector-level, particle-level, and loss curve plots
│ └── evaluate.py Post-hoc distance metrics (Wasserstein, JS, triangular)
├── params/ Gaussian config YAML files
│ ├── 1d_default.yaml
│ ├── 2d_correlated.yaml
│ ├── 4d_correlated.yaml
│ └── 6d_correlated.yaml
├── scripts/
│ ├── submit.sh Training and baseline SLURM submission script
├── tests/ pytest tests
├── .github/workflows/ci.yml Lint, format, types, complexity, tests, audit
├── Justfile Development recipes (just validate, just lint-fix, ...)
├── pyproject.toml Project metadata and dependencies
├── runs/ Output directory (timestamped subdirectories)
└── .cache/ Cached datasets
src/deconvolve/coretypes/, src/deconvolve/data/ and src/deconvolve/baselines/ carry their own README.md with module-level detail.
Datasets
Gaussian (Synthetic)
Configurable multivariate Gaussian distributions with correlated covariance matrices. Supports arbitrary dimensionality and correlation structure via YAML config files. Both truth and MC samples are smeared by additive Gaussian noise to simulate detector resolution, producing paired particle-level ($z$) and detector-level ($x$) features.
Jet Substructure (Physics)
Herwig (data) vs Pythia26 (MC) $Z+$ jets at high $p_T$ (200 GeV), with Delphes detector simulation. Automatically downloaded from from Zenodo record 3548091 if not already present in .cache/.
| Variable | Symbol | Description |
|---|---|---|
m |
$m/\text{GeV}$ | Jet mass |
M |
$M$ | Jet constituent multiplicity |
w |
$w$ | Jet width |
tau21 |
$\tau_{21}$ | N-subjettiness ratio |
zg |
$z_g$ | Groomed jet momentum fraction |
sdm |
$\ln\rho$ | Log soft-drop jet mass |
All variables are z-score standardized using MC gen-level statistics only (no information leakage).
Output
Each run produces a timestamped directory under runs/. The root holds only
what a person opens by hand -- config.json (run configuration, for
reproducibility) and, later, report.pdf. Everything else is supporting
material and lives one level down, flat, in artifacts/:
runs/<timestamp>/
├── report.pdf
├── config.json
└── artifacts/ figures, metrics/timings JSON, checkpoints, arrays
generator.keras/discriminator.keras-- Saved model checkpointshistory.npz-- Training loss historydetector_level.pdf-- Histogram comparing data, MC, and reweighted MC at detector level with ratio panelparticle_level.pdf-- Same comparison at particle levellosses.pdf-- Training curves with log(2) equilibrium targetselection.pdf-- Per-epoch MMD curves and the epoch model selection restoredmetrics.json-- Wasserstein, JS divergence, and triangular discriminator (before/after)metrics_ibu.json-- Same metrics from IBU baseline (if run)timings.json-- Per-phase wall clock, when the run was made underDECONVOLVE_TIMING=1report.tex-- The LaTeX sourcedeconvolve reportcompiles into the run root'sreport.pdf
Training Hyperparameters
These are internal training defaults in src/deconvolve/train.py; the CLI-exposed
training options are listed above.
| Parameter | Default | Description |
|---|---|---|
n_epochs |
100 | Training epochs — a fixed scan trip count |
n_disc_steps |
5 | Discriminator updates per generator update |
lr_g |
3e-5 | Generator learning rate (Adam) |
lr_d |
1e-4 | Discriminator learning rate (Adam) |
lambda_dispersion |
0.015 | Penalty on the variance of g's weights |
hidden_units |
64 | Units per hidden layer |
n_layers |
2 | Number of hidden layers |
lr_g and lambda_dispersion are both measured rather than chosen, and they
act on the same axis: the dispersion of g's normalized MC weights. See "What
tuning actually found" and "The dispersion penalty: the trade made explicit" in
benchmarks/README.md. Because the penalty is on by default, a run left at
these defaults already carries it — which is the configuration any comparison
should be made against, not a variant of it.
n_epochs is not a maximum in the early-stopping sense. scan needs a fixed
trip count, so every run executes all of them; the best epoch is then restored
on the host by the detector-level MMD argmin.
Development
just validate # all local, read-only validation (format, lint, types, complexity, tests)
just lint-fix # apply safe lint fixes, then format
just test # pytest, forwards extra args
just type-check # pyrefly
just ci # the full CI suite
just # list every recipe
GitHub Actions runs the same suite on push. Lint and format are
ruff, type checking is
pyrefly at --min-severity info, and
complexipy enforces a maximum
cognitive complexity of 10.
Dependencies
Release files for deconvolve 0.3.0
For a detailed explanation of source distributions (sdists) and built distributions (wheels), please see the package formats documentation.
Source distribution (sdist)
| File | Size | Uploaded | |
|---|---|---|---|
| deconvolve-0.3.0.tar.gz | 153.7 kB | Details |
Built distribution (wheel)
| File | Interpreter | ABI | Platform | Reset |
|---|---|---|---|---|
| deconvolve-0.3.0-py3-none-any.whl | Python 3 | none | any | Details |
Total release size: 327.9 kB
Release files / deconvolve-0.3.0.tar.gz
| Download URL | deconvolve-0.3.0.tar.gz |
|---|---|
| Size | 153.7 kB |
| Tags | Source |
|
SHA-256 checksum How to use checksums |
2f7d01ca27f44c2b385c253bb4c386d6adbd14ffe8091e719532ec460c7faa0f
|
|
BLAKE2b-256 checksum How to use checksums |
3e910e4e58feaba3375745ae02f8fe2582bd6bcd08291e4f1de99725a012e4fe
|
| Upload date | |
|
Uploaded using Trusted Publishing? What is trusted publishing? |
Yes |
| Uploaded via |
uv/0.12.17 {"installer":{"name":"uv","version":"0.12.17","subcommand":["publish"]},"python":null,"implementation":{"name":null,"version":null},"distro":{"name":"Ubuntu","version":"24.04","id":"noble","libc":null},"system":{"name":null,"release":null},"cpu":null,"openssl_version":null,"setuptools_version":null,"rustc_version":null,"ci":true}
|
Release files / deconvolve-0.3.0-py3-none-any.whl
| Download URL | deconvolve-0.3.0-py3-none-any.whl |
|---|---|
| Size | 174.1 kB |
| Tags | Python 3 |
|
SHA-256 checksum How to use checksums |
f311ac4900a0c0d89216688c1fc077588a104b49072283aace5586337a455421
|
|
BLAKE2b-256 checksum How to use checksums |
4d856f5b6499d1ff2d765d4a77d60411b31622853d1d4273dcf843d0064410b5
|
| Upload date | |
|
Uploaded using Trusted Publishing? What is trusted publishing? |
Yes |
| Uploaded via |
uv/0.12.17 {"installer":{"name":"uv","version":"0.12.17","subcommand":["publish"]},"python":null,"implementation":{"name":null,"version":null},"distro":{"name":"Ubuntu","version":"24.04","id":"noble","libc":null},"system":{"name":null,"release":null},"cpu":null,"openssl_version":null,"setuptools_version":null,"rustc_version":null,"ci":true}
|