ParticleViT-B
ParticleViT is a generic transformer over the constituents of a collider jet, with no physics-specific inductive bias. Each particle is one token; a single prepended class token is read out. This is the B rung (66.7M parameters, pretraining compute 5.7e19 FLOPs) of the scaling ladder from the paper Predict before you train: scaling laws for particle physics foundation models (Uslu, Nachman, Re).
- Code: https://github.com/Jaluus/ParticleViT
- Paper corpus: OmniLearned (~1.06B jets)
- All five models: ParticleViT collection (S, M, B, L, XL)
Installation
pip install "torch>=2.4" safetensors huggingface_hub
That is all you need: the model code (modeling_particlevit.py) and the input
normalization (preprocessing.py) ship inside this repository and depend only on
PyTorch. To grab the files explicitly:
huggingface-cli download jaluus/ParticleViT-B --local-dir ParticleViT-B
Usage
import torch
from modeling_particlevit import ParticleViT
from preprocessing import normalize, build_attn_mask
model = ParticleViT.from_pretrained("jaluus/ParticleViT-B").eval()
# X_raw: (B, 150, 9) raw OmniLearned features per particle:
# 0:4 kinematics (delta eta, delta phi, log pT, log E), 4 PID id, 5:9 vertex.
# Padded particle slots are all-zero.
X_raw = torch.zeros(1, 150, 9)
# ... fill X_raw[0, :n_particles] with your jet ...
mask = build_attn_mask(X_raw) # real-particle mask
X = normalize(X_raw, "omnilearned_parametric_normalization.json", mask)
with torch.no_grad():
logits = model(X, attn_mask=mask) # (B, 210) pretraining logits
For downstream tagging, replace model.head with a fresh linear head and
fine-tune the full network end to end (see the paper's fine-tuning protocol and
the code repository).
The parametric input normalization is required. Feeding raw features without
applying omnilearned_parametric_normalization.json produces meaningless
predictions.
Attention backend
This repository ships the padded SDPA inference path
(F.scaled_dot_product_attention with a key-padding mask): pure PyTorch, no
flash-attn or custom kernels, running on CPU or any GPU in any precision.
Training instead used a mathematically-equivalent variable-length packed
attention kernel (torch.nn.attention.varlen, BF16/FP16-only) purely for
throughput — it strips padding and is ~70% faster. The two paths share the
exact same parameters (the weights load strict=True into either), and the
class-token readout matches up to floating-point precision, so predictions are
the same.
If you want the bleeding-edge packed/varlen model and the full training / fine-tuning code, use the GitHub repository: https://github.com/Jaluus/ParticleViT.
Model architecture
A generic set-transformer over particle tokens with no physics inductive bias (no Lorentz equivariance, no pairwise interaction features). Design choices follow modern open language models:
- Reordered RMSNorm kept outside the residual stream (double-norm blocks).
- Query-key normalization (QK-Norm) for attention stability.
- SwiGLU feedforward with the 8/3 width convention.
- A single prepended class token for readout; no positional encoding (a jet's constituents form a set, not a sequence).
- Truncated-normal (OLMo-style) initialization; head dimension 64.
| depth | width | heads | head dim | params |
|---|---|---|---|---|
| 8 | 832 | 13 | 64 | 66.7M |
Pretraining recipe
Objective. Softmax cross-entropy over the 210-class OmniLearned label space, read from the prepended class token, with no label smoothing. An output z-loss (weight 1e-5) keeps the logits bounded; it is excluded from the reported loss.
Optimization (shared across the ladder).
- Optimizer: AdamW (beta1 0.9, beta2 0.95), weight decay 0.1 (no decay on embeddings, norm gains, or any 1-D parameter).
- Schedule: linear warmup (2000 steps) then cosine decay to 10% of the peak LR.
- No gradient clipping. BF16 mixed precision with FP32 master weights.
- Global batch size 16384 jets.
- Variable-length attention with sequence packing (removes padding, ~70% faster).
- Frozen parametric Gaussian transform on the four kinematic input features.
This model (ParticleViT-B).
| peak LR | global batch | GPUs (A100) | batch/GPU | steps | jets seen | passes | pretraining compute |
|---|---|---|---|---|---|---|---|
| 1e-3 | 16384 | 16 | 1024 | 183,105 | 3.0B | ~2.8 | 5.7e19 FLOPs |
Trained on the Perlmutter supercomputer (NERSC) with PyTorch distributed data parallelism. Compute is accounted as 6 FLOPs per parameter per token at the measured mean occupancy of the 150 particle slots.
Pretraining data
ParticleViT is pretrained on the OmniLearned bundle (Bhimji, Harris, Mikuni,
Nachman; arXiv:2510.24066, Phys. Rev. D 113,
032020), a union of seven simulated and real jet datasets totaling
~1.06 billion training jets (~102M validation, ~68M test). Upstream code and
hosting: ViniciusMikuni/OmniLearned,
served from https://portal.nersc.gov/cfs/dasrepo/omnilearned/.
| Subset | Train jets | Collider / physics | Production chain |
|---|---|---|---|
jetclass |
100M | pp to jets (10 flavors) | MadGraph5 to Pythia 8 to Delphes (CMS card), anti-kt R=0.8 |
jetclass2 |
200M | pp to jets (188 labels) | Same chain, fine-grained parton labels |
aspen |
125M | pp, CMS 2016 open data | Real data + matched simulation |
atlas |
178M | pp to top vs QCD | Pythia 8 + ATLAS Geant4 full sim, anti-kt R=1.0 + Soft Drop |
h1 |
42.2M | ep DIS, 27.6 x 920 GeV | Rapgap / Djangoh + Geant3, kt R=1.0 |
cms_qcd |
239M | pp to QCD | CMS 2016 simulation framework |
cms_bsm |
173.5M | pp to BSM signals | VLQ, charged Higgs, graviton, radion, SUSY, Z', X to YY |
Clustering radius and detector handling deliberately differ across subsets (R = 0.4 / 0.8 / 1.0; Delphes vs Geant4 vs Geant3 vs real data), so the model must generalize across regimes.
Per-particle features (9 channels). Each jet is up to 150 particles:
| idx | feature | notes |
|---|---|---|
| 0 | delta eta | particle eta minus jet axis |
| 1 | delta phi | particle phi minus jet axis (wrapped to [-pi, pi]) |
| 2 | log pT | padding sentinel: == 0 marks a padded slot |
| 3 | log E | clamped >= log pT |
| 4 | PID | categorical particle-ID code (integer, embedded) |
| 5 | tanh(D0) | transverse impact parameter |
| 6 | D0 err | |
| 7 | tanh(Dz) | longitudinal impact parameter |
| 8 | Dz err |
Features 0-3 are the minimal kinematic set; 4-8 are optional and zeroed where a
subset does not provide them. The four kinematic features are mapped to an
approximately standard-normal distribution by the frozen parametric transform in
omnilearned_parametric_normalization.json (required at inference).
Labels. A flat 210-class label space with disjoint per-dataset ranges
(e.g. jetclass 2-11, jetclass2 12-199, single buckets for aspen, cms_qcd,
cms_bsm, ...). The pretraining objective is classification over these 210 classes.
Files
model.safetensors- FP32 weights (66.7M parameters).config.json- architecture configuration.modeling_particlevit.py- self-contained model (PyTorch only).preprocessing.py- self-contained input normalization (PyTorch only).omnilearned_parametric_normalization.json- frozen normalization constants.
Downstream benchmark performance
Background rejection R50 / R30 after fine-tuning (higher is better), from the paper's Table I. R50 and R30 are 1/eps_B at 50% and 30% signal efficiency:
| Model | depth x width | params | Top tagging R50 / R30 | Quark/gluon R50 / R30 |
|---|---|---|---|---|
| ParticleViT-S | 5x512 | 16M | 616 / 2707 | 42.8 / 110.3 |
| ParticleViT-M | 7x640 | 35M | 606 / 2928 | 43.0 / 108.9 |
| ParticleViT-B (this model) | 8x832 | 67M | 618 / 2903 | 43.2 / 109.8 |
| ParticleViT-L | 10x1024 | 126M | 631 / 3042 | 43.4 / 112.2 |
| ParticleViT-XL | 14x1536 | 397M | 651 / 3008 | 43.5 / 110.2 |
License
Released under the MIT license.
Citation
@article{uslu2026predict,
title = {Predict before you train: scaling laws for particle physics foundation models},
author = {Uslu, Jan-Lucas and Nachman, Benjamin and R\'e, Christopher},
year = {2026}
}
@article{bhimji2025omnilearned,
title = {OmniLearned: A Foundation Model Framework for All Tasks Involving Jet Physics},
author = {Bhimji, Wahid and Harris, Chris and Mikuni, Vinicius and Nachman, Benjamin},
journal = {Phys. Rev. D},
volume = {113}, number = {3}, pages = {032020}, year = {2026},
eprint = {2510.24066}, archivePrefix = {arXiv}
}
- Downloads last month
- 9