Skip to content
zichunhaoPublic

About

RINO: Renormalization Group Invariance with No Labels

Resources

Stars

1 star

Watchers

0 watching

Forks

Latest commit

 

History

9 Commits

Folders and files

NameName
Last commit message
Last commit date
 
 
 
 
 
 
 
 
 
 
 
 
 
 

Repository files navigation

RINO: Renormalization Group Distillation with No Labels

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).

Repository structure

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/.

Setup

conda env create -f environment.yml
conda activate rino

environment.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.

Data

See configs/data-README.md to download JetClass, JetNet and Top Quark Tagging and to build the $k_\mathrm{T}$-clustered pretraining set.

Pretraining RINO

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.yaml

Checkpoints 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 $k_\mathrm{T}$ subjets, and the student-only views have N = 2, 3, 4 and the unclustered jet. The paper's ablation variants start from the same config with the teacher schedule before tuning (teacher temperature warmed up from 0.04 to 0.07 over epochs 0–19, final teacher EMA momentum 0.999); studies/evaluation/README.md lists the change for each variant.

Frozen-backbone probes

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.yaml

Finetuning and OOD evaluation

Finetuning 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.0

The 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 10

The paper uses 10 seeds for JetNet and 6 for Top Quark Tagging.

Baselines

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.ckpt

MPMv1 (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.pt

MPMv2 (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.pt

OmniJet-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.pt

JetCLR (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.pt

JetCLR-scale (NT-Xent between $k_\mathrm{T}$ views of different scales, plus masked reconstruction). It is trained by the RINO code and needs no conversion:

accelerate launch --multi_gpu --num_processes=2 \
  dino/jetclr_train.py -c configs/jetclr/3tier-recon.yaml

Frozen-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>

Paper studies

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.

About

RINO: Renormalization Group Invariance with No Labels

Resources

Stars

1 star

Watchers

0 watching

Forks

Releases

Packages

Contributors

Languages