Self-supervised representation learning for jets. RINO builds views of each jet at different energy scales by clustering it into different numbers of exclusive-$k_\mathrm{T}$ subjets, and aligns teacher and student representations across these scales with DINO + iBOT self-distillation and a KoLeo regularizer. It is pretrained on JetClass QCD jets only, finetuned on JetNet (or Top Quark Tagging) and evaluated out of distribution on JetClass (top vs. QCD).
dino/ # RINO: pretraining, finetuning, inference, models, losses
preprocess/ # dataset download and preprocessing scripts
configs/
dino/ # RINO pretraining config
jetclr/ # JetCLR-scale pretraining config
dataloaders/ # dataloader configs (data paths, features)
gen_le_configs.py # generator for all finetuning configs
data-README.md # data preparation
baselines/ # MPMv1, MPMv2, OmniJet-alpha, JetCLR, shared VQ-VAE,
# checkpoint conversion and probes
studies/ # scripts for the paper's tables and figures
Config files use the placeholders PROJECT_ROOT (repository root), JOBNAME (the config's name) and EPOCHNUM (checkpoint epoch); they are resolved at run time. Outputs go to experiments/.
conda env create -f environment.yml
conda activate rinoenvironment.yml was exported on macOS (osx-arm64) with a PyTorch nightly build. On other platforms (e.g. Linux with CUDA), create a Python 3.12 environment, install PyTorch >= 2.6 for your CUDA version, and then:
pip install accelerate omegaconf pyyaml numpy scipy scikit-learn pandas matplotlib \
tqdm prettytable einops h5py awkward uproot vector fastjet wandb
pip install jetnet datasets zenodo-get # data download (see configs/data-README.md)The baselines have their own dependencies (Hydra, Lightning, ...): see baselines/mpmv1/requirements.txt, baselines/mpmv2/requirements.txt and baselines/omnijet_alpha/docker/requirements.txt. MPMv1, OmniJet-alpha and the shared VQ-VAE also need the bundled vqtorch (requires CUDA, cupy and torchpq): pip install baselines/mpmv1/vqtorch.
W&B logging is optional: add --use-wandb to the training and inference scripts.
See configs/data-README.md to download JetClass, JetNet and Top Quark Tagging and to build the
The paper's RINO model (teacher temperature and EMA momentum tuned on the OOD top-vs-QCD accuracy) is trained with configs/dino/rino.yaml on two GPUs:
accelerate launch --multi_gpu --num_processes=2 \
dino/dino_train.py -c configs/dino/rino.yamlCheckpoints are written to experiments/<name>/checkpoints/ (model_checkpoint_best.pt is used downstream). With training.auto_resume: true (set in rino.yaml) a restarted run continues from its latest checkpoint; otherwise set training.load_epoch. Single-GPU runs use python dino/dino_train.py -c ....
In rino.yaml, the teacher views (seen by teacher and student) have N = 6, 8 and 16 studies/evaluation/README.md lists the change for each variant.
dino_inference.py extracts representations from the pretrained backbone and runs the k-NN (k = 20) and linear probes on JetClass top vs. QCD configured under inference.acc_tests:
python dino/dino_inference.py -c configs/dino/rino.yamlFinetuning configs (RINO and baselines, JetNet or Top Quark Tagging, any label fraction) are generated by configs/gen_le_configs.py and written to configs/finetune/<subdir>/<name>.yaml. The models are the keys of its MODELS dictionary; each entry points to the pretrained checkpoint it finetunes. The methods of the main comparison are:
| Method | --models key |
|---|---|
| RINO | rino |
| Supervised | sup |
| MPMv1 / MPMv2 | mpmv1 / mpmv2 |
| JetCLR / JetCLR-scale | jetclr-orig / jclr-3tier-recon |
| OmniJet-alpha | omnijet-mean |
# Full JetNet training set (Table 1)
python configs/gen_le_configs.py --task jn --fractions 1.0 \
--models rino sup mpmv1 mpmv2 jetclr-orig jclr-3tier-recon omnijet-mean
# Label efficiency (0.1%, 1%, 10%, 50% of JetNet)
python configs/gen_le_configs.py --task jn --fractions 0.001 0.01 0.1 0.5
# Top Quark Tagging
python configs/gen_le_configs.py --task tt --fractions 1.0The generated configs use the MLP head (hidden layers of 256 and 128); for a linear head, set hidden_dims, activations, batch_norms and dropouts under models.head.params to [].
Each config is finetuned and evaluated with several seeds; --run-index N sets the seed and writes to a run-N/ subdirectory. Inference evaluates the JetNet (or Top Quark Tagging) validation and test sets and the JetClass test set (top vs. QCD), and eval_per_class.py aggregates the runs (mean and standard deviation, plus the ensemble):
CFG=configs/finetune/finetune-jetnet-final/finetune-jn-rino.yaml
for i in $(seq 1 10); do
python dino/classification_train.py -c $CFG --run-index $i
python dino/dino_inference.py -c $CFG --run-index $i
done
python dino/eval_per_class.py \
--exp-dir experiments/finetune-jetnet-final/finetune-jn-rino --num-runs 10The paper uses 10 seeds for JetNet and 6 for Top Quark Tagging.
All baselines use the same 8-layer JetTransformerEncoder backbone (OmniJet-alpha uses the matching causal decoder), the same QCD-only JetClass data and the same finetuning protocol. Each pretrained model is converted to a backbone checkpoint that the finetuning configs above load; the --output of each conversion must be the backbone_weight_path of the corresponding model in configs/gen_le_configs.py. Commands are run from the repository root unless stated otherwise; the data preparation is in configs/data-README.md.
Supervised. No pretraining: model sup in gen_le_configs.py.
Shared VQ-VAE tokenizer (used by MPMv1 and OmniJet-alpha):
python baselines/vqvae/train.py -c baselines/vqvae/configs/vqvae-shared.yaml
# -> experiments/vqvae/vqvae-shared/checkpoints/last.ckptMPMv1 (masked prediction of VQ-VAE tokens):
python baselines/mpmv1/scripts/precompute_tokens.py # optional *_tokens.h5 next to the data
cd baselines/mpmv1
python scripts/train.py experiment=train_mpm_rino network_name=mpmv1-rino
# -> experiments/mpmv1/mpmv1-rino/mpmv1-rino/checkpoints/best_<epoch>.ckpt
cd ../..
python baselines/scripts/convert_checkpoint.py --model mpmv1 \
--input experiments/mpmv1/mpmv1-rino/mpmv1-rino/checkpoints/best_<epoch>.ckpt \
--output experiments/mpmv1/mpmv1-rino/backbone.ptMPMv2 (masked K-means classification, L1 regression and a conditional flow):
cd baselines/mpmv2
python scripts/train.py experiment=pretrain_rino
cd ../..
python baselines/scripts/convert_checkpoint.py --model mpmv2 \
--input <path/to/mpmv2/run>/checkpoints/<checkpoint>.ckpt \
--output experiments/mpmv2/pretrain-rino-a100/backbone.ptOmniJet-alpha (autoregressive next-token prediction on VQ-VAE tokens):
cd baselines/omnijet_alpha
python gabbro/train.py experiment=rino_generative \
trainer.devices=<n_gpus> +trainer.precision=bf16-mixed
# -> experiments/omnijet_alpha/omnijet-rino/runs/<run>/checkpoints/best.ckpt
cd ../..
python baselines/scripts/convert_checkpoint.py --model omnijet \
--input experiments/omnijet_alpha/omnijet-rino/runs/<run>/checkpoints/best.ckpt \
--output experiments/omnijet_alpha/omnijet-rino/backbone.ptJetCLR (NT-Xent with the original JetCLR augmentations):
accelerate launch --multi_gpu --num_processes=<n_gpus> \
baselines/jetclr/JetCLR/scripts/jetclr_train.py \
--config baselines/jetclr/configs/jetclr-rinomodel.yaml
python baselines/scripts/convert_checkpoint.py --model jetclr \
--input experiments/jetclr/jetclr-rinomodel/<checkpoint>.pt \
--output experiments/jetclr/jetclr-rinomodel/backbone.ptJetCLR-scale (NT-Xent between
accelerate launch --multi_gpu --num_processes=2 \
dino/jetclr_train.py -c configs/jetclr/3tier-recon.yamlFrozen-backbone probes for a converted baseline backbone (k-NN and linear probe on JetClass top vs. QCD):
python baselines/scripts/backbone_probe.py \
--backbone-path <path/to/backbone.pt> \
--dataloader-config configs/dataloaders/jetclass-raw/kinematics.yaml \
--pooling mean --output <path/to/probe_results.json>The scripts behind the paper's tables and figures (result tables, data efficiency, domain shift, scale structure, scale-factor proxy and the physics studies) are in studies/; see studies/README.md.