File size: 6,729 Bytes
ec60373 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 | """Run the ViT-BL dRICH particle-ID model on a raw Geant4/EDM4hep
simulation output file, end to end: DRICHHits.cellID -> resolved (x,y) ->
rasterized image -> model -> predicted species.
Requires: torch, uproot, awkward, numpy, huggingface_hub
pip install torch uproot awkward numpy huggingface_hub
Usage:
python predict_from_root.py sim_2212_30.0_2.0_3.14.root
The model architecture (RingViT) is copied inline below so this script
has no dependency beyond pip-installable packages -- it does not require
cloning the training repo.
"""
import re
import sys
import awkward as ak
import numpy as np
import torch
import torch.nn as nn
import uproot
from huggingface_hub import hf_hub_download
MODEL_REPO = "deepaksamuel-cuk/drich-vit-baseline"
DATA_REPO = "deepaksamuel-cuk/simhits"
NAMES = ["electron", "pion", "kaon", "proton"]
PID2LABEL = {11: 0, 211: 1, 321: 2, 2212: 3}
FULL_IMG = 384
FULL_WINDOW = 3540.0 # mm -- full detector extent, matches the real dRICH sensor plane
MOM_SCALE = 60.0
ETA_MID, ETA_HALF = 2.5, 1.0
# ---------------------------------------------------------------- model ---
class TransformerBlock(nn.Module):
def __init__(self, dim, heads, mlp_ratio=4.0, dropout=0.0):
super().__init__()
self.norm1 = nn.LayerNorm(dim)
self.attn = nn.MultiheadAttention(dim, heads, dropout=dropout, batch_first=True)
self.norm2 = nn.LayerNorm(dim)
hidden = int(dim * mlp_ratio)
self.mlp = nn.Sequential(
nn.Linear(dim, hidden), nn.GELU(), nn.Dropout(dropout),
nn.Linear(hidden, dim), nn.Dropout(dropout),
)
def forward(self, x):
h = self.norm1(x)
x = x + self.attn(h, h, h, need_weights=False)[0]
x = x + self.mlp(self.norm2(x))
return x
class RingViT(nn.Module):
"""Same architecture as the training repo's vit/model.py -- 8-layer,
256-dim ViT with an extra learned 'kinematics token' (p, eta, cos phi,
sin phi) alongside the cls token and 24x24=576 image patch tokens."""
def __init__(self, img_size=384, patch_size=16, in_chans=1, num_classes=4,
dim=256, depth=8, heads=8, mlp_ratio=4.0, dropout=0.1, n_kin=4):
super().__init__()
self.n_patches = (img_size // patch_size) ** 2
self.patch_embed = nn.Conv2d(in_chans, dim, kernel_size=patch_size, stride=patch_size)
self.cls_token = nn.Parameter(torch.zeros(1, 1, dim))
self.kin_embed = nn.Sequential(nn.Linear(n_kin, dim), nn.GELU(), nn.Linear(dim, dim))
self.pos_embed = nn.Parameter(torch.zeros(1, 2 + self.n_patches, dim))
self.pos_drop = nn.Dropout(dropout)
self.blocks = nn.ModuleList(
[TransformerBlock(dim, heads, mlp_ratio, dropout) for _ in range(depth)])
self.norm = nn.LayerNorm(dim)
self.head = nn.Linear(dim, num_classes)
def forward(self, img, kin):
B = img.shape[0]
x = self.patch_embed(img).flatten(2).transpose(1, 2)
cls = self.cls_token.expand(B, -1, -1)
k = self.kin_embed(kin).unsqueeze(1)
x = torch.cat([cls, k, x], dim=1) + self.pos_embed
x = self.pos_drop(x)
for blk in self.blocks:
x = blk(x)
x = self.norm(x)
return self.head(x[:, 0])
# ------------------------------------------------------------- rasterize --
def rasterize(pts, img_size=FULL_IMG, window=FULL_WINDOW):
"""pixel value = log1p(hit count) -- multiple hits in the same cell
are summed (np.add.at), never overwritten or capped. Same convention
as every image this model was trained on."""
img = np.zeros((img_size, img_size), dtype=np.float32)
if len(pts) == 0:
return img
rel = (pts - (-window / 2)) * (img_size / window)
ij = np.floor(rel).astype(np.int64)
ok = (ij[:, 0] >= 0) & (ij[:, 0] < img_size) & (ij[:, 1] >= 0) & (ij[:, 1] < img_size)
ij = ij[ok]
np.add.at(img, (ij[:, 1], ij[:, 0]), 1.0)
return np.log1p(img)
def parse_kinematics_from_filename(path):
"""Sim files follow the {PID}_{MOM}_{ETA}_{PHI}[_suffix].root naming
convention (see claude.md) -- pulls out the first four underscore-
separated numeric tokens, tolerant of an optional trailing suffix."""
stem = path.split("/")[-1]
stem = re.sub(r"\.root$", "", stem)
stem = re.sub(r"^(sim|rec|ana)_", "", stem)
m = re.match(r"^(-?\d+)_([\d.]+)_([\d.]+)_([\d.]+)", stem)
if not m:
raise ValueError(f"could not parse PID_MOM_ETA_PHI from filename: {path}")
pid_s, mom_s, eta_s, phi_s = m.groups()
return int(pid_s), float(mom_s), float(eta_s), float(phi_s)
def main():
if len(sys.argv) != 2:
print(f"usage: python {sys.argv[0]} <sim_file.root>")
sys.exit(1)
root_path = sys.argv[1]
pid, mom, eta, phi = parse_kinematics_from_filename(root_path)
print(f"from filename: PID={pid}, p={mom:.2f} GeV/c, eta={eta:.2f}, phi={phi:.2f} rad")
# -- extract raw photon hits from the sim file --
f = uproot.open(root_path)
t = f["events"]
cellid = ak.to_numpy(ak.flatten(t["DRICHHits/DRICHHits.cellID"].array()))
print(f"{len(cellid)} raw photon hits")
# -- resolve cellID -> (x, y) mm via the precomputed sensor geometry lookup --
lookup_path = hf_hub_download(DATA_REPO, "cellid_positions.npz", repo_type="dataset")
lk = np.load(lookup_path)
cid_sorted, xyz = lk["cellids"], lk["xyz"]
idx = np.searchsorted(cid_sorted, cellid)
idx = np.clip(idx, 0, len(cid_sorted) - 1)
valid = cid_sorted[idx] == cellid
if not valid.all():
print(f"warning: dropping {(~valid).sum()} hits with unknown cellID")
pts = xyz[idx[valid], :2].astype(np.float32)
img = rasterize(pts)
kin = np.array([mom / MOM_SCALE, (eta - ETA_MID) / ETA_HALF,
np.cos(phi), np.sin(phi)], dtype=np.float32)
# -- load the model --
ckpt_path = hf_hub_download(MODEL_REPO, "best.pt")
model = RingViT(num_classes=4)
model.load_state_dict(torch.load(ckpt_path, map_location="cpu"))
model.eval()
timg = torch.from_numpy(img).unsqueeze(0).unsqueeze(0) # (1, 1, 384, 384)
tkin = torch.from_numpy(kin).unsqueeze(0) # (1, 4)
with torch.no_grad():
logits = model(timg, tkin)
probs = torch.softmax(logits, dim=1)[0]
pred = int(probs.argmax())
print(f"\npredicted: {NAMES[pred]} (p={probs[pred]:.3f})")
for i, name in enumerate(NAMES):
print(f" {name:10s} {probs[i]:.4f}")
if pid in PID2LABEL:
true_c = PID2LABEL[abs(pid)]
verdict = "CORRECT" if pred == true_c else "WRONG"
print(f"\ntrue (from filename): {NAMES[true_c]} -- {verdict}")
if __name__ == "__main__":
main()
|