ParticleViT-XL

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 XL rung (396.8M parameters, pretraining compute 3.4e20 FLOPs) of the scaling ladder from the paper Predict before you train: scaling laws for particle physics foundation models (Uslu, Nachman, Re).

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-XL --local-dir ParticleViT-XL

Usage

import torch
from modeling_particlevit import ParticleViT
from preprocessing import normalize, build_attn_mask

model = ParticleViT.from_pretrained("jaluus/ParticleViT-XL").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
14 1536 24 64 396.8M

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

peak LR global batch GPUs (A100) batch/GPU steps jets seen passes pretraining compute
5e-4 16384 64 256 183,104 3.0B ~2.8 3.4e20 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 (396.8M 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 8x832 67M 618 / 2903 43.2 / 109.8
ParticleViT-L 10x1024 126M 631 / 3042 43.4 / 112.2
ParticleViT-XL (this model) 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
Safetensors
Model size
0.4B params
Tensor type
F32
·
Inference Providers NEW
This model isn't deployed by any Inference Provider. 🙋 Ask for provider support

Collection including jaluus/ParticleViT-XL

Paper for jaluus/ParticleViT-XL