Implement Phase 2: secondary particle prediction

Two-stage factorisation: Stage 1 predicts 9D primary kinematics + n_sec
classification head (COND_DIM reduced to 8, dropping n_sec/e_sec inputs);
Stage 2 (SecondaryDecoder) generates K_MAX=15 secondary slots via masked
flow matching over (stick_logit, local_dir, type_emb) conditioned on Stage 1
output. Joint training with combined loss L_s1 + λ_nsec*L_nsec + λ_s2*L_s2.

Co-Authored-By: Claude Sonnet 4.6 <noreply@anthropic.com>
This commit is contained in:
2026-06-29 11:34:31 +02:00
parent c627142135
commit e6e0eb22bf
18 changed files with 1174 additions and 234 deletions
+172
View File
@@ -0,0 +1,172 @@
# Phase 2: Secondary Particle Prediction
## Context
Phase 1 takes `n_sec` (secondary count) and `e_sec` (total secondary energy) as **conditioning inputs**. Phase 2 must instead **predict** them, making the surrogate self-contained for shower rollout. Per Jan's 2026-06-29 decision: hard discrete `n_sec` integer head; escalation to Gumbel-Softmax only if empirically needed.
Two-stage factorization:
- **Stage 1**: existing 9D flow model (reduced conditioning: drop `n_sec` + `log(e_sec)`) + a new discrete `n_sec` classification head
- **Stage 2**: non-AR flow matching over `K_MAX` secondary slots simultaneously, each slot predicting `(stick_break_logit, dir_local_3D, type_emb)` — conditioned on pre-step state + Stage 1 output; padded slots masked from loss
Training: joint, combined loss `L = L_flow_s1 + λ_nsec * L_nsec + λ_s2 * L_flow_s2`.
---
## Prerequisite: Determine K_MAX
Before implementing, run a quick analysis over existing parquet files to find `max(n_sec)` and the 99th percentile. Expected to be 520 for EM shower steps. Set `K_MAX` as a constant in `giant/constants.py` (suggest 15 as a starting point, revise from data).
---
## New Branch
```bash
git checkout -b phase2-secondary-prediction master
```
---
## Part A — Data Pipeline
### A1. `scripts/steps_to_parquet.py`
Extend `_add_secondary_energy` to also collect per-secondary attributes from the spawning tree join:
- For each `child_track_id`, look up the child's first step → get `pdg`, `pre_E`, `pre_dx/dy/dz`
- Emit list columns in the parquet: `sec_pdg_list`, `sec_E_list`, `sec_dx_list`, `sec_dy_list`, `sec_dz_list`
- Lists are sorted **descending by energy** at write time
- Truncate to `K_MAX` entries if needed (flag if any row truncated)
Re-run ROOT→parquet conversion after this change.
### A2. `giant/data/loader.py`
In `_df_to_dict`: read the five new list columns. Pad each to length `K_MAX` with zeros (energy) / sentinel values (pdg → 0, dir → (0,0,1)). Return as fixed-shape arrays `(N, K_MAX)` / `(N, K_MAX, 3)`.
Also return a boolean validity mask `sec_valid` of shape `(N, K_MAX)`: `True` for slots `i < n_sec`.
### A3. `giant/data/transforms.py`
Add `encode_secondaries(sec_pdg_list, sec_E_list, sec_dir_list, sec_valid, e_sec, pdg_emb_weight, pre_dir, K_MAX)`:
1. **Direction**: call existing `local_frame_rotation` per slot
2. **Energy (stick-breaking)**:
- Slot 0: `f_0 = E_0 / e_sec` → logit `log(f_0/(1-f_0))` (clamped)
- Slot i: `f_i = E_i / (e_sec - sum(E_0..E_{i-1}))` → logit
- Last valid slot: logit = large positive constant (takes all remaining budget)
- Padding slots (beyond `n_sec`): set logit = 0, masked out of loss anyway
3. **Type embedding**: index into `pdg_emb_weight` (the PDG embedding table weights) to get the target embedding vector for each secondary's `pdg`. Shape `(K_MAX, emb_dim)`.
Returns `sec_targets: (K_MAX, 1 + 3 + emb_dim)` and `sec_valid: (K_MAX,)`.
Inverse (`decode_secondaries`): sigmoid stick-breaking fractions → energies, inv local frame rotation → world dirs, nearest-neighbor lookup in PDG embedding table → pdg code.
### A4. `giant/data/dataset.py`
Update `build_features` and `StreamingStepsDataset.__iter__` to also yield `sec_targets` and `sec_valid` alongside the existing `(cond_cont, cond_cat, x1)` batch items.
---
## Part B — Constants (`giant/constants.py`)
- `COND_DIM`: 10 → **8** (remove `n_sec` and `log(e_sec)`)
- Add `K_MAX: int` (set after data analysis, e.g. 15)
- Add `SEC_SLOT_DIM: int` (= 4 + `emb_dim` = 20 for default emb_dim=16; 1 stick + 3 dir + 16 type)
- Add `SEC_DIM: int = K_MAX * SEC_SLOT_DIM` (flattened Stage 2 target dimension)
- Update `LOCAL_TARGET_NAMES` (Stage 1 only, still 9D)
---
## Part C — Model (`giant/model/network.py`)
### C1. `DenoisingMLP` — Stage 1 (minimal changes)
- `ConditionEncoder.cont_dim` drops from 10 to 8 (COND_DIM change propagates automatically)
- Add `n_sec_head = nn.Sequential(Linear(cond_out_dim, hidden_dim//2), SiLU(), Linear(hidden_dim//2, K_MAX + 1))` applied to `c_emb` (the condition encoding, not the diffused latent)
- Add method `predict_n_sec(cond_cont, cond_cat) -> Tensor[B, K_MAX+1]` — no diffusion, just encode conditioning and run the head
### C2. `SecondaryDecoder` — Stage 2 (new class)
Architecture mirrors `DenoisingMLP` but:
- **Input**: `x_t` of shape `(B, SEC_DIM)` (flattened K_MAX secondary slots)
- **Conditioning**: pre-step state (8D cont + 2 cat → same ConditionEncoder as Stage 1) concatenated with Stage 1 output (9D normalized target, detached from Stage 1 loss for stability initially). Total cond dim to the ResBlocks: `time_dim + cond_s1_out_dim + 9`
- **Output**: vector field of shape `(B, SEC_DIM)`
- Uses same `ResBlock` / `SinusoidalEmbedding` / `ConditionEncoder` building blocks
A `SecondaryConditionEncoder` wraps the base `ConditionEncoder` and concatenates the Stage 1 output:
```python
class SecondaryConditionEncoder(nn.Module):
# base: ConditionEncoder(pdg_vocab, mat_vocab, 8, emb_dim, cond_out_dim)
# stage1_proj: Linear(X_DIM, stage1_cond_dim)
# mlp: fuses both
```
---
## Part D — Loss / Training
### `giant/model/schedule.py`
Add `flow_matching_loss_masked(model, x1, cond_cont, cond_cat, mask)`:
- Same as `flow_matching_loss` but divides by `mask.sum()` instead of `B * SEC_DIM`, zeroing out padded slots before averaging. `mask` shape: `(B, K_MAX)`, broadcast over slot dims.
### `giant/train.py`
Batch now unpacks as `(cond_cont, cond_cat, x1_s1, n_sec_target, x1_s2, sec_mask)`.
Combined loss per batch:
```
L_s1 = flow_matching_loss(stage1_model, x1_s1, cond_cont, cond_cat)
L_nsec = cross_entropy(stage1_model.predict_n_sec(cond_cont, cond_cat), n_sec_target)
L_s2 = flow_matching_loss_masked(sec_decoder, x1_s2, cond_cont, cond_cat, stage1_detached, sec_mask)
L = L_s1 + lambda_nsec * L_nsec + lambda_s2 * L_s2
```
Config adds `lambda_nsec` (suggest 0.1) and `lambda_s2` (suggest 1.0) under `[train]`.
Both `stage1_model` and `sec_decoder` share a single `optimizer` (AdamW over all parameters).
Checkpoint saves both `stage1_model.state_dict()` and `sec_decoder.state_dict()`, plus `K_MAX` and `SEC_SLOT_DIM` in `model_config`.
### `giant/pipeline.py`
- Compute `K_MAX` from data (max `n_sec` over training events) before constructing models
- Build both `DenoisingMLP` and `SecondaryDecoder`, pass both to `run_training`
---
## Part E — Sampling (`giant/sample.py`)
```python
def sample_stage1(model, cond_cont, cond_cat, steps=10):
# Euler ODE → primary sample (9D), + argmax n_sec head
...
def sample_secondaries(sec_decoder, cond_cont, cond_cat, stage1_out, n_sec, steps=10):
# Euler ODE on SEC_DIM → decode stick-breaking → energies
# inv_local_frame_rotation → world-frame dirs
# nearest-neighbor in pdg_emb_weight → pdg codes
# mask slots >= n_sec
...
```
---
## Part F — Wiring
- **`giant/validate.py`**: add secondary-specific marginals (n_sec distribution, species distribution, energy fraction per slot)
- **`giant/cli.py`**: `predict` command loads both checkpoints, calls both samplers, appends secondary columns to output parquet
---
## Type embedding design note
The type embedding target at training is `pdg_emb.weight[sec_pdg_idx]` (the Stage 1 PDG embedding table rows). Gradients flow into the embedding table from both the conditioning path (input PDG) and the secondary type loss — this is intentional; the shared embedding space is the bridge. At inference, snap: `argmin_k ||pred_emb - pdg_emb.weight[k]||`.
---
## Verification
1. `uv run pytest` — existing tests pass (Stage 1 shape/interface unchanged beyond COND_DIM)
2. Unit tests for `encode_secondaries` / `decode_secondaries` (round-trip: energies sum to `e_sec`, directions are unit vectors)
3. Unit test for `flow_matching_loss_masked`: verify padded slots contribute zero gradient
4. Short training run (12 epochs): confirm all three loss components decrease
5. Sampling smoke test: verify `sum(sec_E) ≈ e_sec` per sample, all directions unit-normed
+2 -1
View File
@@ -409,7 +409,8 @@ def predict(
cc = torch.from_numpy(cond_cont).float().to(_device) cc = torch.from_numpy(cond_cont).float().to(_device)
ck = torch.from_numpy(cond_cat).long().to(_device) ck = torch.from_numpy(cond_cat).long().to(_device)
pred = sample_flow(model, cc, ck, steps=steps).cpu().numpy() # normalised pred, _n_sec = sample_flow(model, cc, ck, steps=steps)
pred = pred.cpu().numpy() # normalised
# Inverse-normalise → local frame, log-scaled scalars # Inverse-normalise → local frame, log-scaled scalars
raw = tgt_norm.inverse_transform(pred) raw = tgt_norm.inverse_transform(pred)
+2
View File
@@ -20,6 +20,8 @@ DEFAULT_CONFIG: dict = {
"validate_every": 10, "validate_every": 10,
"validate_steps": 10, "validate_steps": 10,
"warmup_epochs": 5, "warmup_epochs": 5,
"lambda_nsec": 0.1,
"lambda_s2": 1.0,
}, },
"model": { "model": {
"hidden_dim": 256, "hidden_dim": 256,
+18 -10
View File
@@ -1,16 +1,24 @@
X_DIM = 9 X_DIM = 9
# Conditioning continuous-feature width: pre_pos(3), log(pre_E)(1), pre_dir(3), # Conditioning continuous-feature width (Phase 2): pre_pos(3), log(pre_E)(1),
# layer_id(1), n_sec(1), log(e_sec)(1). One wider than X_DIM because e_sec # pre_dir(3), layer_id(1). n_sec and log(e_sec) are removed — they are now
# (secondary energy) is a conditioning input in the energy-conservation PoC. # *outputs* predicted by Stage 1, not conditioning inputs.
COND_DIM = 10 COND_DIM = 8
# The two energy columns are additive-log-ratio (ALR) coordinates of the # Maximum number of secondary slots. From data: max(n_sec)=14 in PbWO4 dataset;
# deposit/secondary/post energy simplex (fractions of pre_E that sum to 1), # K_MAX=15 covers it with one spare slot.
# referenced to the post-energy fraction — see giant.data.transforms K_MAX = 15
# .energy_simplex_encode/.energy_simplex_decode. They replace the former
# independent log_delta_e / log_edep targets so energy conservation holds by # Per-slot secondary target dimension: 1 (stick-breaking logit) + 3 (local dir) +
# construction after decoding. # EMB_DIM (continuous type embedding). EMB_DIM must match DenoisingMLP.emb_dim.
# Default emb_dim=16 → SEC_SLOT_DIM=20.
SEC_SLOT_DIM = 20 # 1 + 3 + 16
EMB_DIM = 16 # must match model emb_dim default
# Flattened Stage-2 target dimension
SEC_DIM = K_MAX * SEC_SLOT_DIM # 15 * 20 = 300
# Stage-1 9D target names (unchanged from energy-conservation PoC)
LOCAL_TARGET_NAMES = [ LOCAL_TARGET_NAMES = [
"log_step_length", "log_step_length",
"edep_logit", "edep_logit",
+52 -60
View File
@@ -4,53 +4,12 @@ from pathlib import Path
import numpy as np import numpy as np
import torch import torch
from torch.utils.data import Dataset, IterableDataset from torch.utils.data import IterableDataset
from giant.data.loader import iter_file_chunks from giant.data.loader import iter_file_chunks
from giant.data.transforms import Normalizer, build_features from giant.data.transforms import Normalizer, build_features
class StepsDataset(Dataset):
def __init__(
self,
cond_cont: np.ndarray,
cond_cat: np.ndarray,
target: np.ndarray,
) -> None:
self.cond_cont = torch.from_numpy(cond_cont).float()
self.cond_cat = torch.from_numpy(cond_cat).long()
self.target = torch.from_numpy(target).float()
def __len__(self) -> int:
return len(self.target)
def __getitem__(self, index):
return self.cond_cont[index], self.cond_cat[index], self.target[index]
def train_val_split(
data: dict,
cond_cont: np.ndarray,
cond_cat: np.ndarray,
target: np.ndarray,
val_fraction: float = 0.1,
seed: int = 42,
) -> tuple[StepsDataset, StepsDataset]:
rng = np.random.default_rng(seed)
unique_events = np.unique(data["event_id"])
rng.shuffle(unique_events)
n_val = max(1, int(len(unique_events) * val_fraction))
val_events = set(unique_events[:n_val].tolist())
val_mask = np.array([e in val_events for e in data["event_id"]])
train_mask = ~val_mask
return (
StepsDataset(cond_cont[train_mask], cond_cat[train_mask], target[train_mask]),
StepsDataset(cond_cont[val_mask], cond_cat[val_mask], target[val_mask]),
)
def make_event_split( def make_event_split(
all_event_ids: np.ndarray, all_event_ids: np.ndarray,
val_fraction: float = 0.1, val_fraction: float = 0.1,
@@ -75,6 +34,16 @@ class StreamingStepsDataset(IterableDataset):
Yields whole batches (use with `DataLoader(..., batch_size=None)`) Yields whole batches (use with `DataLoader(..., batch_size=None)`)
rather than single rows, so the batch is assembled with vectorized rather than single rows, so the batch is assembled with vectorized
numpy slicing instead of a per-row Python loop in the default collate. numpy slicing instead of a per-row Python loop in the default collate.
Each batch is a tuple:
(cond_cont, cond_cat, target_s1, n_sec, sec_cont, sec_pdg_idx)
where:
cond_cont: (B, COND_DIM) float32
cond_cat: (B, 2) int64
target_s1: (B, 9) float32 — normalised Stage-1 primary target
n_sec: (B,) int64 — true secondary count per step
sec_cont: (B, K_MAX, 4) float32 — [stick_logit, local_dir] per slot
sec_pdg_idx: (B, K_MAX) int64 — PDG model-index per secondary slot
""" """
def __init__( def __init__(
@@ -91,13 +60,12 @@ class StreamingStepsDataset(IterableDataset):
) -> None: ) -> None:
self.files = list(files) self.files = list(files)
self.split_events = split_events self.split_events = split_events
self._events_arr = np.array(sorted(split_events)) # for np.isin self._events_arr = np.array(sorted(split_events))
self.pdg_map = pdg_map self.pdg_map = pdg_map
self.mat_map = mat_map self.mat_map = mat_map
self.cond_normalizer = cond_normalizer self.cond_normalizer = cond_normalizer
self.target_normalizer = target_normalizer self.target_normalizer = target_normalizer
self.batch_size = batch_size self.batch_size = batch_size
# Buffer must hold at least one batch or we could never emit one.
self.shuffle_buffer = max(shuffle_buffer, batch_size) self.shuffle_buffer = max(shuffle_buffer, batch_size)
self.shuffle = shuffle self.shuffle = shuffle
@@ -114,6 +82,9 @@ class StreamingStepsDataset(IterableDataset):
buf_cont: list[np.ndarray] = [] buf_cont: list[np.ndarray] = []
buf_cat: list[np.ndarray] = [] buf_cat: list[np.ndarray] = []
buf_tgt: list[np.ndarray] = [] buf_tgt: list[np.ndarray] = []
buf_nsec: list[np.ndarray] = []
buf_sec: list[np.ndarray] = []
buf_spdg: list[np.ndarray] = []
buf_n = 0 buf_n = 0
for path in files: for path in files:
@@ -123,43 +94,57 @@ class StreamingStepsDataset(IterableDataset):
continue continue
chunk = {k: v[mask] for k, v in chunk.items()} chunk = {k: v[mask] for k, v in chunk.items()}
cond_cont, cond_cat, target, _, _ = build_features( cond_cont, cond_cat, target_s1, n_sec, sec_cont, sec_pdg_idx, _, _ = (
chunk, build_features(
self.pdg_map, chunk,
self.mat_map, self.pdg_map,
cond_normalizer=self.cond_normalizer, self.mat_map,
target_normalizer=self.target_normalizer, cond_normalizer=self.cond_normalizer,
target_normalizer=self.target_normalizer,
)
) )
buf_cont.append(cond_cont) buf_cont.append(cond_cont)
buf_cat.append(cond_cat) buf_cat.append(cond_cat)
buf_tgt.append(target) buf_tgt.append(target_s1)
buf_nsec.append(n_sec)
buf_sec.append(sec_cont)
buf_spdg.append(sec_pdg_idx)
buf_n += len(cond_cont) buf_n += len(cond_cont)
if buf_n >= self.shuffle_buffer: if buf_n >= self.shuffle_buffer:
buf_cont, buf_cat, buf_tgt, buf_n = yield from self._flush( buf_cont, buf_cat, buf_tgt, buf_nsec, buf_sec, buf_spdg, buf_n = (
buf_cont, buf_cat, buf_tgt, final=False yield from self._flush(
buf_cont, buf_cat, buf_tgt, buf_nsec, buf_sec, buf_spdg,
final=False,
)
) )
if buf_n > 0: if buf_n > 0:
yield from self._flush(buf_cont, buf_cat, buf_tgt, final=True) yield from self._flush(
buf_cont, buf_cat, buf_tgt, buf_nsec, buf_sec, buf_spdg, final=True
)
def _flush( def _flush(
self, self,
buf_cont: list[np.ndarray], buf_cont: list[np.ndarray],
buf_cat: list[np.ndarray], buf_cat: list[np.ndarray],
buf_tgt: list[np.ndarray], buf_tgt: list[np.ndarray],
buf_nsec: list[np.ndarray],
buf_sec: list[np.ndarray],
buf_spdg: list[np.ndarray],
final: bool, final: bool,
): ):
"""Yield full batches of `batch_size`; carry any remainder back to the caller.
All batching is done via vectorized numpy slicing (no per-row Python loop).
"""
cont = np.concatenate(buf_cont) cont = np.concatenate(buf_cont)
cat = np.concatenate(buf_cat) cat = np.concatenate(buf_cat)
tgt = np.concatenate(buf_tgt) tgt = np.concatenate(buf_tgt)
nsec = np.concatenate(buf_nsec)
sec = np.concatenate(buf_sec)
spdg = np.concatenate(buf_spdg)
if self.shuffle: if self.shuffle:
idx = np.random.permutation(len(cont)) idx = np.random.permutation(len(cont))
cont, cat, tgt = cont[idx], cat[idx], tgt[idx] cont, cat, tgt = cont[idx], cat[idx], tgt[idx]
nsec, sec, spdg = nsec[idx], sec[idx], spdg[idx]
bs = self.batch_size bs = self.batch_size
n = len(cont) n = len(cont)
@@ -170,9 +155,16 @@ class StreamingStepsDataset(IterableDataset):
torch.from_numpy(cont[start:end]).float(), torch.from_numpy(cont[start:end]).float(),
torch.from_numpy(cat[start:end]).long(), torch.from_numpy(cat[start:end]).long(),
torch.from_numpy(tgt[start:end]).float(), torch.from_numpy(tgt[start:end]).float(),
torch.from_numpy(nsec[start:end]).long(),
torch.from_numpy(sec[start:end]).float(),
torch.from_numpy(spdg[start:end]).long(),
) )
if final: if final:
return [], [], [], 0 return [], [], [], [], [], [], 0
rem = n_full * bs rem = n_full * bs
return [cont[rem:]], [cat[rem:]], [tgt[rem:]], n - rem return (
[cont[rem:]], [cat[rem:]], [tgt[rem:]],
[nsec[rem:]], [sec[rem:]], [spdg[rem:]],
n - rem,
)
+54 -1
View File
@@ -40,8 +40,52 @@ def find_parquet_files(path: str | Path) -> list[Path]:
return [p] return [p]
def _pad_list_col(series: pd.Series, K: int, fill: float = 0.0) -> np.ndarray:
"""Pad / truncate a list-valued Series to fixed width K → (N, K) float32."""
out = np.full((len(series), K), fill, dtype=np.float32)
for i, lst in enumerate(series):
if lst is not None and len(lst) > 0:
n = min(len(lst), K)
out[i, :n] = lst[:n]
return out
def _pad_list_col_int(series: pd.Series, K: int, fill: int = 0) -> np.ndarray:
"""Pad / truncate a list-valued integer Series to fixed width K → (N, K) int64."""
out = np.full((len(series), K), fill, dtype=np.int64)
for i, lst in enumerate(series):
if lst is not None and len(lst) > 0:
n = min(len(lst), K)
out[i, :n] = lst[:n]
return out
def _pad_dir_col(
dx: pd.Series, dy: pd.Series, dz: pd.Series, K: int
) -> np.ndarray:
"""Pad three list-valued direction columns → (N, K, 3) float32.
Padding direction defaults to (0,0,1) (forward) so it is a valid unit vector.
"""
N = len(dx)
out = np.zeros((N, K, 3), dtype=np.float32)
out[:, :, 2] = 1.0
for i in range(N):
lx, ly, lz = dx.iloc[i], dy.iloc[i], dz.iloc[i]
if lx is not None and len(lx) > 0:
n = min(len(lx), K)
out[i, :n, 0] = lx[:n]
out[i, :n, 1] = ly[:n]
out[i, :n, 2] = lz[:n]
return out
def _df_to_dict(df: pd.DataFrame) -> dict[str, np.ndarray]: def _df_to_dict(df: pd.DataFrame) -> dict[str, np.ndarray]:
return { from giant.constants import K_MAX
has_sec_lists = "sec_E_list" in df.columns
d: dict[str, np.ndarray] = {
"event_id": df["event_id"].to_numpy(), "event_id": df["event_id"].to_numpy(),
"pdg": df["pdg"].to_numpy(dtype=np.int32), "pdg": df["pdg"].to_numpy(dtype=np.int32),
"pre_pos": df[["pre_x", "pre_y", "pre_z"]].to_numpy(dtype=np.float32), "pre_pos": df[["pre_x", "pre_y", "pre_z"]].to_numpy(dtype=np.float32),
@@ -59,6 +103,15 @@ def _df_to_dict(df: pd.DataFrame) -> dict[str, np.ndarray]:
"post_pos": df[["post_x", "post_y", "post_z"]].to_numpy(dtype=np.float32), "post_pos": df[["post_x", "post_y", "post_z"]].to_numpy(dtype=np.float32),
} }
if has_sec_lists:
d["sec_E_list"] = _pad_list_col(df["sec_E_list"], K_MAX)
d["sec_pdg_list"] = _pad_list_col_int(df["sec_pdg_list"], K_MAX)
d["sec_dir_list"] = _pad_dir_col(
df["sec_dx_list"], df["sec_dy_list"], df["sec_dz_list"], K_MAX
)
return d
def load_steps(path: str | Path) -> dict[str, np.ndarray]: def load_steps(path: str | Path) -> dict[str, np.ndarray]:
return _df_to_dict(pd.read_parquet(path)) return _df_to_dict(pd.read_parquet(path))
+141 -13
View File
@@ -227,6 +227,114 @@ def inv_local_frame_rotation(
) )
_STICK_LOGIT_CLIP = 10.0 # logit value used for the last valid secondary slot
def encode_secondaries(
sec_E_list: np.ndarray,
sec_dir_list: np.ndarray,
sec_valid: np.ndarray,
e_sec: np.ndarray,
pre_dir: np.ndarray,
) -> np.ndarray:
"""Encode per-secondary attributes into continuous per-slot targets.
Secondaries must already be sorted descending by energy (as stored in the
parquet). Returns sec_cont of shape (N, K_MAX, 4):
slot[i] = [stick_break_logit, local_dir_x, local_dir_y, local_dir_z]
Stick-breaking logit: for slot i, f_i = E_i / remaining_budget, where
remaining_budget = e_sec - sum(E_0..E_{i-1}). The logit is log(f/(1-f)),
clipped to ±_STICK_LOGIT_CLIP. The last valid slot gets +_STICK_LOGIT_CLIP
(takes the full remaining budget). Padding slots get 0.
sec_pdg_idx (integer) is not processed here — kept separate so the loss
function can look up the embedding table at training time.
"""
N, K = sec_E_list.shape
e_sec = np.asarray(e_sec, dtype=np.float64)
stick_logits = np.zeros((N, K), dtype=np.float32)
for i in range(K):
if i == 0:
remaining = e_sec
else:
remaining = np.maximum(e_sec - sec_E_list[:, :i].sum(axis=1), _EPS)
f = np.clip(sec_E_list[:, i].astype(np.float64) / remaining, _EPS, 1.0 - _EPS)
logit = np.log(f / (1.0 - f)).astype(np.float32)
# Last valid slot: give it the full remaining budget
is_last = sec_valid[:, i] & ~(sec_valid[:, i + 1] if i + 1 < K else np.zeros(N, dtype=bool))
logit = np.where(is_last, _STICK_LOGIT_CLIP, logit)
logit = np.where(sec_valid[:, i], np.clip(logit, -_STICK_LOGIT_CLIP, _STICK_LOGIT_CLIP), 0.0)
stick_logits[:, i] = logit.astype(np.float32)
# Rotate each slot's direction into the local frame of the primary.
# pre_dir is broadcast across all K slots.
dir_local = np.zeros((N, K, 3), dtype=np.float32)
for i in range(K):
# Only rotate valid slots; leave padded slots as (0,0,1) or whatever.
valid_mask = sec_valid[:, i]
if valid_mask.any():
dir_local[valid_mask, i] = local_frame_rotation(
pre_dir[valid_mask], sec_dir_list[valid_mask, i]
)
sec_cont = np.concatenate(
[stick_logits[:, :, None], dir_local], axis=-1
) # (N, K, 4)
return sec_cont.astype(np.float32)
def decode_secondaries(
sec_cont: np.ndarray,
sec_pdg_pred: np.ndarray,
n_sec: np.ndarray,
e_sec: np.ndarray,
pre_dir: np.ndarray,
pdg_map_inv: dict[int, int],
) -> tuple[np.ndarray, np.ndarray, np.ndarray, np.ndarray]:
"""Inverse of encode_secondaries: continuous targets → physical secondary attrs.
sec_cont: (N, K_MAX, 4) — [stick_logit, local_dir_x, local_dir_y, local_dir_z]
sec_pdg_pred: (N, K_MAX) integer PDG indices (from nearest-neighbor snap)
n_sec: (N,) integer secondary counts
e_sec: (N,) total secondary energy budget [MeV]
pre_dir: (N, 3) pre-step world-frame direction
pdg_map_inv: maps model index → PDG code
Returns (sec_E, sec_dir_world, sec_pdg_code, sec_valid) each shape (N, K_MAX).
"""
N, K, _ = sec_cont.shape
stick_logits = sec_cont[:, :, 0] # (N, K)
dir_local = sec_cont[:, :, 1:] # (N, K, 3)
fractions = 1.0 / (1.0 + np.exp(-stick_logits.astype(np.float64)))
sec_E = np.zeros((N, K), dtype=np.float32)
e_sec = np.asarray(e_sec, dtype=np.float64)
remaining = e_sec.copy()
for i in range(K):
sec_E[:, i] = (fractions[:, i] * remaining).astype(np.float32)
remaining = np.maximum(remaining - sec_E[:, i].astype(np.float64), 0.0)
sec_valid = np.arange(K)[None, :] < n_sec[:, None] # (N, K)
sec_dir_world = np.zeros((N, K, 3), dtype=np.float32)
for i in range(K):
valid = sec_valid[:, i]
if valid.any():
sec_dir_world[valid, i] = inv_local_frame_rotation(
pre_dir[valid], dir_local[valid, i]
)
sec_pdg_code = np.array(
[[pdg_map_inv.get(int(sec_pdg_pred[n, i]), 0) for i in range(K)] for n in range(N)],
dtype=np.int32,
)
return sec_E, sec_dir_world, sec_pdg_code, sec_valid
def build_cond_features( def build_cond_features(
data: dict[str, np.ndarray], data: dict[str, np.ndarray],
pdg_map: dict[int, int], pdg_map: dict[int, int],
@@ -240,8 +348,6 @@ def build_cond_features(
log_transform(data["pre_E"]), log_transform(data["pre_E"]),
data["pre_dir"], data["pre_dir"],
data["layer_id"].astype(np.float32), data["layer_id"].astype(np.float32),
data["n_sec"].astype(np.float32),
log_transform(data["e_sec"]),
] ]
).astype(np.float32) ).astype(np.float32)
@@ -262,11 +368,17 @@ def build_features(
cond_normalizer: Normalizer | None = None, cond_normalizer: Normalizer | None = None,
target_normalizer: Normalizer | None = None, target_normalizer: Normalizer | None = None,
fit: bool = False, fit: bool = False,
) -> tuple[np.ndarray, np.ndarray, np.ndarray, Normalizer | None, Normalizer | None]: ) -> tuple[np.ndarray, np.ndarray, np.ndarray, np.ndarray, np.ndarray, np.ndarray, Normalizer | None, Normalizer | None]:
"""Assemble (cond_cont, cond_cat, target) arrays ready for StepsDataset. """Assemble (cond_cont, cond_cat, target_s1, n_sec, sec_cont, sec_pdg_idx) arrays.
When fit=True, new Normalizers are fitted on the supplied arrays. target_s1: (N, 9) Stage-1 primary post-step target (unchanged from Phase 1)
n_sec: (N,) integer secondary counts (target for n_sec head)
sec_cont: (N, K_MAX, 4) continuous secondary targets [stick_logit, dir_local]
sec_pdg_idx: (N, K_MAX) integer PDG model-indices; used to look up embedding
targets in the training loop
""" """
from giant.constants import K_MAX
post_dir_local = local_frame_rotation(data["pre_dir"], data["post_dir"]) post_dir_local = local_frame_rotation(data["pre_dir"], data["post_dir"])
travel_dir_local = local_frame_rotation( travel_dir_local = local_frame_rotation(
data["pre_dir"], travel_direction(data["pre_pos"], data["post_pos"]) data["pre_dir"], travel_direction(data["pre_pos"], data["post_pos"])
@@ -274,9 +386,9 @@ def build_features(
energy_z = energy_simplex_encode( energy_z = energy_simplex_encode(
data["edep"], data["e_sec"], data["post_E"], data["pre_E"] data["edep"], data["e_sec"], data["post_E"], data["pre_E"]
) # (N, 2): ALR coords of the deposit/secondary/post energy simplex ) # (N, 2)
target = np.column_stack( target_s1 = np.column_stack(
[ [
log_transform(data["step_length"]), log_transform(data["step_length"]),
energy_z, energy_z,
@@ -285,28 +397,44 @@ def build_features(
] ]
).astype(np.float32) # (N, 9) ).astype(np.float32) # (N, 9)
# Phase 2: conditioning drops n_sec and log(e_sec)
cond_cont = np.column_stack( cond_cont = np.column_stack(
[ [
data["pre_pos"], data["pre_pos"],
log_transform(data["pre_E"]), log_transform(data["pre_E"]),
data["pre_dir"], data["pre_dir"],
data["layer_id"].astype(np.float32), data["layer_id"].astype(np.float32),
data["n_sec"].astype(np.float32),
log_transform(data["e_sec"]),
] ]
).astype(np.float32) # (N, COND_DIM) ).astype(np.float32) # (N, COND_DIM=8)
pdg_idx = np.array([pdg_map[int(p)] for p in data["pdg"]], dtype=np.int64) pdg_idx = np.array([pdg_map[int(p)] for p in data["pdg"]], dtype=np.int64)
mat_idx = np.array([mat_map[str(m)] for m in data["material"]], dtype=np.int64) mat_idx = np.array([mat_map[str(m)] for m in data["material"]], dtype=np.int64)
cond_cat = np.column_stack([pdg_idx, mat_idx]) # (N, 2) cond_cat = np.column_stack([pdg_idx, mat_idx]) # (N, 2)
n_sec = data["n_sec"].astype(np.int64) # (N,)
# Secondary continuous targets
sec_E_list = data.get("sec_E_list")
sec_dir_list = data.get("sec_dir_list")
sec_pdg_idx = data.get("sec_pdg_idx")
if sec_E_list is not None and sec_dir_list is not None and sec_pdg_idx is not None:
sec_valid = np.arange(K_MAX)[None, :] < n_sec[:, None] # (N, K_MAX)
sec_cont = encode_secondaries(
sec_E_list, sec_dir_list, sec_valid, data["e_sec"], data["pre_dir"]
) # (N, K_MAX, 4)
else:
N = len(n_sec)
sec_cont = np.zeros((N, K_MAX, 4), dtype=np.float32)
sec_pdg_idx = np.zeros((N, K_MAX), dtype=np.int64)
if fit: if fit:
cond_normalizer = Normalizer().fit(cond_cont) cond_normalizer = Normalizer().fit(cond_cont)
target_normalizer = Normalizer().fit(target) target_normalizer = Normalizer().fit(target_s1)
if cond_normalizer is not None: if cond_normalizer is not None:
cond_cont = cond_normalizer.transform(cond_cont) cond_cont = cond_normalizer.transform(cond_cont)
if target_normalizer is not None: if target_normalizer is not None:
target = target_normalizer.transform(target) target_s1 = target_normalizer.transform(target_s1)
return cond_cont, cond_cat, target, cond_normalizer, target_normalizer return cond_cont, cond_cat, target_s1, n_sec, sec_cont, sec_pdg_idx, cond_normalizer, target_normalizer
+125 -3
View File
@@ -3,7 +3,7 @@ import math
import torch import torch
import torch.nn as nn import torch.nn as nn
from giant.constants import COND_DIM, X_DIM from giant.constants import COND_DIM, K_MAX, SEC_DIM, X_DIM
class SinusoidalEmbedding(nn.Module): class SinusoidalEmbedding(nn.Module):
@@ -70,6 +70,12 @@ class ResBlock(nn.Module):
class DenoisingMLP(nn.Module): class DenoisingMLP(nn.Module):
"""Stage-1 model: predicts the 9D primary post-step vector field + n_sec logits.
The n_sec head runs on the condition encoding only (no diffusion noise),
so it can be called at inference time independently via `predict_n_sec`.
"""
def __init__( def __init__(
self, self,
pdg_vocab: int, pdg_vocab: int,
@@ -81,6 +87,7 @@ class DenoisingMLP(nn.Module):
cond_out_dim: int = 128, cond_out_dim: int = 128,
x_dim: int = X_DIM, x_dim: int = X_DIM,
dropout: float = 0.1, dropout: float = 0.1,
k_max: int = K_MAX,
) -> None: ) -> None:
super().__init__() super().__init__()
self.time_emb = SinusoidalEmbedding(time_dim) self.time_emb = SinusoidalEmbedding(time_dim)
@@ -99,6 +106,13 @@ class DenoisingMLP(nn.Module):
] ]
) )
self.out_proj = nn.Linear(hidden_dim, x_dim) self.out_proj = nn.Linear(hidden_dim, x_dim)
# Predicts n_sec as classification over {0, 1, ..., k_max}.
# Applied to the condition encoding (not the diffused latent).
self.n_sec_head = nn.Sequential(
nn.Linear(cond_out_dim, hidden_dim // 2),
nn.SiLU(),
nn.Linear(hidden_dim // 2, k_max + 1),
)
def forward( def forward(
self, self,
@@ -107,9 +121,117 @@ class DenoisingMLP(nn.Module):
cond_cont: torch.Tensor, cond_cont: torch.Tensor,
cond_cat: torch.Tensor, cond_cat: torch.Tensor,
) -> torch.Tensor: ) -> torch.Tensor:
t_emb = self.time_emb(t) # (B, time_dim) t_emb = self.time_emb(t) # (B, time_dim)
c_emb = self.cond_enc(cond_cont, cond_cat) # (B, cond_out_dim) c_emb = self.cond_enc(cond_cont, cond_cat) # (B, cond_out_dim)
cond = torch.cat([t_emb, c_emb], dim=-1) # (B, time_dim+cond_out_dim) cond = torch.cat([t_emb, c_emb], dim=-1)
x = self.input_proj(x_t)
for block in self.blocks:
x = block(x, cond)
return self.out_proj(x)
def predict_n_sec(
self,
cond_cont: torch.Tensor,
cond_cat: torch.Tensor,
) -> torch.Tensor:
"""Return n_sec logits (B, K_MAX+1) from conditioning alone."""
c_emb = self.cond_enc(cond_cont, cond_cat)
return self.n_sec_head(c_emb)
def pdg_embedding_weight(self) -> torch.Tensor:
"""Return the PDG embedding table weights for secondary type targets."""
return self.cond_enc.pdg_emb.weight
class SecondaryConditionEncoder(nn.Module):
"""Encodes pre-step conditioning + Stage-1 output for the secondary decoder."""
def __init__(
self,
pdg_vocab: int,
mat_vocab: int,
emb_dim: int = 16,
cond_out_dim: int = 128,
stage1_dim: int = X_DIM,
stage1_proj_dim: int = 64,
out_dim: int = 128,
) -> None:
super().__init__()
self.base = ConditionEncoder(
pdg_vocab=pdg_vocab,
mat_vocab=mat_vocab,
emb_dim=emb_dim,
out_dim=cond_out_dim,
)
self.stage1_proj = nn.Linear(stage1_dim, stage1_proj_dim)
fused_dim = cond_out_dim + stage1_proj_dim
self.fuse = nn.Sequential(
nn.Linear(fused_dim, out_dim),
nn.SiLU(),
)
def forward(
self,
cond_cont: torch.Tensor,
cond_cat: torch.Tensor,
stage1_out: torch.Tensor,
) -> torch.Tensor:
base = self.base(cond_cont, cond_cat) # (B, cond_out_dim)
s1 = self.stage1_proj(stage1_out).tanh() # (B, stage1_proj_dim)
return self.fuse(torch.cat([base, s1], dim=-1)) # (B, out_dim)
class SecondaryDecoder(nn.Module):
"""Stage-2 model: predicts vector field over K_MAX secondary slots simultaneously.
Each slot encodes (stick_break_logit, local_dir_3D, type_emb) for one
secondary ordered by descending energy. Padded slots are masked from loss.
"""
def __init__(
self,
pdg_vocab: int,
mat_vocab: int,
hidden_dim: int = 256,
n_blocks: int = 6,
emb_dim: int = 16,
time_dim: int = 64,
cond_out_dim: int = 128,
stage1_proj_dim: int = 64,
sec_dim: int = SEC_DIM,
dropout: float = 0.1,
) -> None:
super().__init__()
self.time_emb = SinusoidalEmbedding(time_dim)
self.cond_enc = SecondaryConditionEncoder(
pdg_vocab=pdg_vocab,
mat_vocab=mat_vocab,
emb_dim=emb_dim,
cond_out_dim=cond_out_dim,
stage1_proj_dim=stage1_proj_dim,
out_dim=cond_out_dim,
)
merged_cond_dim = time_dim + cond_out_dim
self.input_proj = nn.Linear(sec_dim, hidden_dim)
self.blocks = nn.ModuleList(
[
ResBlock(hidden_dim, merged_cond_dim, dropout=dropout)
for _ in range(n_blocks)
]
)
self.out_proj = nn.Linear(hidden_dim, sec_dim)
def forward(
self,
x_t: torch.Tensor,
t: torch.Tensor,
cond_cont: torch.Tensor,
cond_cat: torch.Tensor,
stage1_out: torch.Tensor,
) -> torch.Tensor:
t_emb = self.time_emb(t)
c_emb = self.cond_enc(cond_cont, cond_cat, stage1_out)
cond = torch.cat([t_emb, c_emb], dim=-1)
x = self.input_proj(x_t) x = self.input_proj(x_t)
for block in self.blocks: for block in self.blocks:
x = block(x, cond) x = block(x, cond)
+34
View File
@@ -69,3 +69,37 @@ def flow_matching_loss(
u_t = x1 - x0 u_t = x1 - x0
v_t = model(x_t, t, cond_cont, cond_cat) v_t = model(x_t, t, cond_cont, cond_cat)
return F.mse_loss(v_t, u_t) return F.mse_loss(v_t, u_t)
def flow_matching_loss_secondary(
model: torch.nn.Module,
x1: torch.Tensor,
cond_cont: torch.Tensor,
cond_cat: torch.Tensor,
stage1_out: torch.Tensor,
sec_mask: torch.Tensor,
) -> torch.Tensor:
"""Flow matching loss for the secondary decoder with per-slot masking.
x1: (B, SEC_DIM) flattened secondary target (stick_logit, dir, type_emb)
sec_mask: (B, K_MAX) bool True for valid secondary slots
Only valid-slot dimensions contribute to the loss; padded slots are zeroed
before averaging, so the loss is not diluted by empty slots.
"""
from giant.constants import SEC_SLOT_DIM
B = x1.size(0)
t = torch.rand(B, device=x1.device)
x0 = torch.randn_like(x1)
x_t = (1.0 - t.view(-1, 1)) * x0 + t.view(-1, 1) * x1
u_t = x1 - x0
v_t = model(x_t, t, cond_cont, cond_cat, stage1_out)
# Expand mask: (B, K_MAX) → (B, K_MAX * SEC_SLOT_DIM)
mask_expanded = (
sec_mask.float().unsqueeze(-1).expand(-1, -1, SEC_SLOT_DIM).reshape(B, -1)
)
err = (v_t - u_t) ** 2
denom = mask_expanded.sum().clamp(min=1)
return (err * mask_expanded).sum() / denom
+35 -11
View File
@@ -5,7 +5,7 @@ import torch
from torch.utils.data import DataLoader from torch.utils.data import DataLoader
from giant import config from giant import config
from giant.constants import COND_DIM, X_DIM from giant.constants import COND_DIM, EMB_DIM, K_MAX, SEC_SLOT_DIM, X_DIM
from giant.data.loader import ( from giant.data.loader import (
find_parquet_files, find_parquet_files,
load_event_ids, load_event_ids,
@@ -14,7 +14,7 @@ from giant.data.loader import (
) )
from giant.data.transforms import build_features, _WelfordAccumulator from giant.data.transforms import build_features, _WelfordAccumulator
from giant.data.dataset import make_event_split, StreamingStepsDataset from giant.data.dataset import make_event_split, StreamingStepsDataset
from giant.model.network import DenoisingMLP from giant.model.network import DenoisingMLP, SecondaryDecoder
from giant.train import train as run_training from giant.train import train as run_training
@@ -63,9 +63,11 @@ def run_train_job(
if not mask.any(): if not mask.any():
continue continue
chunk_tr = {k: v[mask] for k, v in chunk.items()} chunk_tr = {k: v[mask] for k, v in chunk.items()}
cond_cont, _, target, _, _ = build_features(chunk_tr, pdg_map, mat_map) cond_cont, _, target_s1, _n_sec, _sec_cont, _sec_pdg, _, _ = build_features(
chunk_tr, pdg_map, mat_map
)
cond_acc.update(cond_cont) cond_acc.update(cond_cont)
tgt_acc.update(target) tgt_acc.update(target_s1)
cond_norm = cond_acc.to_normalizer() cond_norm = cond_acc.to_normalizer()
tgt_norm = tgt_acc.to_normalizer() tgt_norm = tgt_acc.to_normalizer()
@@ -91,8 +93,6 @@ def run_train_job(
shuffle=False, shuffle=False,
) )
# Dataset yields whole batches already, so batch_size=None tells DataLoader
# to pass them through instead of re-collating row-by-row in Python.
pin = device.type == "cuda" pin = device.type == "cuda"
train_loader = DataLoader( train_loader = DataLoader(
train_ds, train_ds,
@@ -107,15 +107,34 @@ def run_train_job(
pin_memory=pin, pin_memory=pin,
) )
model = DenoisingMLP( emb_dim = m.get("emb_dim", EMB_DIM)
# SEC_SLOT_DIM must match constants (1 stick + 3 dir + emb_dim)
assert SEC_SLOT_DIM == 1 + 3 + emb_dim, (
f"SEC_SLOT_DIM={SEC_SLOT_DIM} must equal 1+3+emb_dim={1+3+emb_dim}; "
"update giant/constants.py if emb_dim changed"
)
stage1_model = DenoisingMLP(
pdg_vocab=len(pdg_map), pdg_vocab=len(pdg_map),
mat_vocab=len(mat_map), mat_vocab=len(mat_map),
hidden_dim=m["hidden_dim"], hidden_dim=m["hidden_dim"],
n_blocks=m["n_blocks"], n_blocks=m["n_blocks"],
emb_dim=m["emb_dim"], emb_dim=emb_dim,
dropout=m["dropout"],
k_max=K_MAX,
)
sec_decoder = SecondaryDecoder(
pdg_vocab=len(pdg_map),
mat_vocab=len(mat_map),
hidden_dim=m["hidden_dim"],
n_blocks=m["n_blocks"],
emb_dim=emb_dim,
dropout=m["dropout"], dropout=m["dropout"],
) )
echo(f"model: {sum(p.numel() for p in model.parameters()):,} parameters") echo(
f"stage1: {sum(p.numel() for p in stage1_model.parameters()):,} parameters | "
f"sec_decoder: {sum(p.numel() for p in sec_decoder.parameters()):,} parameters"
)
out_dir.mkdir(parents=True, exist_ok=True) out_dir.mkdir(parents=True, exist_ok=True)
meta = config.build_run_meta( meta = config.build_run_meta(
@@ -134,12 +153,15 @@ def run_train_job(
"mat_vocab": len(mat_map), "mat_vocab": len(mat_map),
"hidden_dim": m["hidden_dim"], "hidden_dim": m["hidden_dim"],
"n_blocks": m["n_blocks"], "n_blocks": m["n_blocks"],
"emb_dim": m["emb_dim"], "emb_dim": emb_dim,
"dropout": m["dropout"], "dropout": m["dropout"],
"k_max": K_MAX,
"sec_slot_dim": SEC_SLOT_DIM,
} }
run_training( run_training(
model=model, stage1_model=stage1_model,
sec_decoder=sec_decoder,
train_loader=train_loader, train_loader=train_loader,
val_loader=val_loader, val_loader=val_loader,
mode=t["mode"], mode=t["mode"],
@@ -148,6 +170,8 @@ def run_train_job(
warmup_epochs=t["warmup_epochs"], warmup_epochs=t["warmup_epochs"],
device=device, device=device,
out_dir=out_dir, out_dir=out_dir,
lambda_nsec=t.get("lambda_nsec", 0.1),
lambda_s2=t.get("lambda_s2", 1.0),
normalizer_dict={"cond": cond_norm.to_dict(), "target": tgt_norm.to_dict()}, normalizer_dict={"cond": cond_norm.to_dict(), "target": tgt_norm.to_dict()},
pdg_map={str(k): v for k, v in pdg_map.items()}, pdg_map={str(k): v for k, v in pdg_map.items()},
mat_map={str(k): v for k, v in mat_map.items()}, mat_map={str(k): v for k, v in mat_map.items()},
+75 -10
View File
@@ -1,6 +1,6 @@
import torch import torch
from giant.constants import X_DIM from giant.constants import K_MAX, SEC_DIM, SEC_SLOT_DIM, X_DIM
@torch.no_grad() @torch.no_grad()
@@ -9,8 +9,13 @@ def sample_flow(
cond_cont: torch.Tensor, cond_cont: torch.Tensor,
cond_cat: torch.Tensor, cond_cat: torch.Tensor,
steps: int = 10, steps: int = 10,
) -> torch.Tensor: ) -> tuple[torch.Tensor, torch.Tensor]:
"""Euler integration of the learned vector field from t=0 to t=1.""" """Euler integration of the Stage-1 vector field from t=0 to t=1.
Returns (primary_sample, n_sec_pred):
primary_sample: (B, X_DIM) normalised 9D primary post-step output
n_sec_pred: (B,) int64 predicted secondary count
"""
model.eval() model.eval()
B = cond_cont.size(0) B = cond_cont.size(0)
device = cond_cont.device device = cond_cont.device
@@ -20,7 +25,63 @@ def sample_flow(
t = torch.full((B,), i * dt, device=device) t = torch.full((B,), i * dt, device=device)
v = model(x, t, cond_cont, cond_cat) v = model(x, t, cond_cont, cond_cat)
x = x + v * dt x = x + v * dt
return x n_sec_logits = model.predict_n_sec(cond_cont, cond_cat)
n_sec_pred = n_sec_logits.argmax(dim=-1)
return x, n_sec_pred
@torch.no_grad()
def sample_secondaries(
sec_decoder: torch.nn.Module,
cond_cont: torch.Tensor,
cond_cat: torch.Tensor,
stage1_out: torch.Tensor,
n_sec_pred: torch.Tensor,
steps: int = 10,
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
"""Euler integration of the Stage-2 vector field; return raw slot outputs.
n_sec_pred: (B,) int64 number of valid secondaries per step
Returns (sec_cont, sec_type_emb, sec_valid):
sec_cont: (B, K_MAX, 4) [stick_logit, local_dir_x, local_dir_y, local_dir_z]
sec_type_emb: (B, K_MAX, emb_dim) predicted type embedding per slot
sec_valid: (B, K_MAX) bool True for slots i < n_sec_pred
"""
sec_decoder.eval()
B = cond_cont.size(0)
device = cond_cont.device
x = torch.randn(B, SEC_DIM, device=device)
dt = 1.0 / steps
for i in range(steps):
t = torch.full((B,), i * dt, device=device)
v = sec_decoder(x, t, cond_cont, cond_cat, stage1_out)
x = x + v * dt
x_slots = x.view(B, K_MAX, SEC_SLOT_DIM)
sec_cont = x_slots[:, :, :4]
sec_type_emb = x_slots[:, :, 4:]
sec_valid = (
torch.arange(K_MAX, device=device).unsqueeze(0) < n_sec_pred.unsqueeze(1)
)
return sec_cont, sec_type_emb, sec_valid
def snap_type_to_pdg_idx(
sec_type_emb: torch.Tensor,
pdg_emb_weight: torch.Tensor,
) -> torch.Tensor:
"""Nearest-neighbour snap: predicted type embedding → PDG model-index.
sec_type_emb: (B, K_MAX, emb_dim)
Returns (B, K_MAX) int64 with model-indices.
"""
B, K, D = sec_type_emb.shape
flat = sec_type_emb.reshape(-1, D)
dists = torch.cdist(flat.float(), pdg_emb_weight.float())
return dists.argmin(dim=-1).reshape(B, K)
@torch.no_grad() @torch.no_grad()
@@ -29,8 +90,8 @@ def sample_ddpm(
cond_cont: torch.Tensor, cond_cont: torch.Tensor,
cond_cat: torch.Tensor, cond_cat: torch.Tensor,
schedule, schedule,
) -> torch.Tensor: ) -> tuple[torch.Tensor, torch.Tensor]:
"""Full DDPM ancestral sampling (T reverse steps).""" """Full DDPM ancestral sampling (T reverse steps). Returns (sample, n_sec_pred)."""
model.eval() model.eval()
B = cond_cont.size(0) B = cond_cont.size(0)
device = cond_cont.device device = cond_cont.device
@@ -46,7 +107,9 @@ def sample_ddpm(
x = (1.0 / alpha.sqrt()) * ( x = (1.0 / alpha.sqrt()) * (
x - (1.0 - alpha) / (1.0 - alpha_bar).sqrt() * eps_pred x - (1.0 - alpha) / (1.0 - alpha_bar).sqrt() * eps_pred
) + beta.sqrt() * z ) + beta.sqrt() * z
return x n_sec_logits = model.predict_n_sec(cond_cont, cond_cat)
n_sec_pred = n_sec_logits.argmax(dim=-1)
return x, n_sec_pred
@torch.no_grad() @torch.no_grad()
@@ -56,8 +119,8 @@ def sample_ddim(
cond_cat: torch.Tensor, cond_cat: torch.Tensor,
schedule, schedule,
steps: int = 50, steps: int = 50,
) -> torch.Tensor: ) -> tuple[torch.Tensor, torch.Tensor]:
"""DDIM deterministic sampling (Song et al. 2020) with `steps` substeps.""" """DDIM deterministic sampling (Song et al. 2020). Returns (sample, n_sec_pred)."""
model.eval() model.eval()
B = cond_cont.size(0) B = cond_cont.size(0)
device = cond_cont.device device = cond_cont.device
@@ -75,4 +138,6 @@ def sample_ddim(
ab_prev = torch.ones(1, device=device) ab_prev = torch.ones(1, device=device)
x0_pred = (x - (1.0 - ab_t).sqrt() * eps_pred) / ab_t.sqrt() x0_pred = (x - (1.0 - ab_t).sqrt() * eps_pred) / ab_t.sqrt()
x = ab_prev.sqrt() * x0_pred + (1.0 - ab_prev).sqrt() * eps_pred x = ab_prev.sqrt() * x0_pred + (1.0 - ab_prev).sqrt() * eps_pred
return x n_sec_logits = model.predict_n_sec(cond_cont, cond_cat)
n_sec_pred = n_sec_logits.argmax(dim=-1)
return x, n_sec_pred
+150 -39
View File
@@ -8,14 +8,31 @@ from types import FrameType
from typing import Callable from typing import Callable
import torch import torch
import torch.nn.functional as F
import torch.optim as optim import torch.optim as optim
from torch.utils.data import DataLoader from torch.utils.data import DataLoader
from tqdm import tqdm from tqdm import tqdm
from giant.model.schedule import CosineSchedule, flow_matching_loss from giant.model.schedule import (
CosineSchedule,
flow_matching_loss,
flow_matching_loss_secondary,
)
from giant.validate import validate_marginals from giant.validate import validate_marginals
_METRICS_FIELDS = ["epoch", "train_loss", "val_loss", "lr", "epoch_time_s"] _METRICS_FIELDS = [
"epoch",
"train_loss",
"train_loss_s1",
"train_loss_nsec",
"train_loss_s2",
"val_loss",
"val_loss_s1",
"val_loss_nsec",
"val_loss_s2",
"lr",
"epoch_time_s",
]
_CATCHABLE_SIGNALS = (signal.SIGINT, signal.SIGTERM) _CATCHABLE_SIGNALS = (signal.SIGINT, signal.SIGTERM)
@@ -57,8 +74,80 @@ class _GracefulShutdown:
) )
def _build_sec_x1(
sec_cont: torch.Tensor,
sec_pdg_idx: torch.Tensor,
pdg_emb_weight: torch.Tensor,
) -> torch.Tensor:
"""Assemble the Stage-2 flow target by appending type embeddings.
sec_cont: (B, K_MAX, 4) [stick_logit, dir_local]
sec_pdg_idx: (B, K_MAX) integer PDG model-indices
pdg_emb_weight: (pdg_vocab, emb_dim) live embedding table weights
Returns (B, SEC_DIM) = (B, K_MAX * (4 + emb_dim)).
"""
type_emb = pdg_emb_weight[sec_pdg_idx] # (B, K_MAX, emb_dim)
x1_s2 = torch.cat([sec_cont, type_emb], dim=-1) # (B, K_MAX, 4+emb_dim)
return x1_s2.flatten(1) # (B, SEC_DIM)
def _compute_losses(
stage1_model: torch.nn.Module,
sec_decoder: torch.nn.Module,
batch: tuple,
mode: str,
ddpm_schedule,
device: torch.device,
lambda_nsec: float,
lambda_s2: float,
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]:
"""Compute (total_loss, L_s1, L_nsec, L_s2) for one batch."""
cond_cont, cond_cat, x1_s1, n_sec, sec_cont, sec_pdg_idx = batch
cond_cont = cond_cont.to(device)
cond_cat = cond_cat.to(device)
x1_s1 = x1_s1.to(device)
n_sec = n_sec.to(device)
sec_cont = sec_cont.to(device)
sec_pdg_idx = sec_pdg_idx.to(device)
# Stage-1 flow loss
if mode == "flow":
l_s1 = flow_matching_loss(stage1_model, x1_s1, cond_cont, cond_cat)
else:
assert ddpm_schedule is not None
l_s1 = ddpm_schedule.loss(stage1_model, x1_s1, cond_cont, cond_cat)
# n_sec classification loss
n_sec_logits = stage1_model.predict_n_sec(cond_cont, cond_cat)
l_nsec = F.cross_entropy(n_sec_logits, n_sec)
# Stage-2 secondary flow loss
# Use a noiseless Stage-1 target as context (detach to avoid back-prop
# coupling between the two flow paths through the same embedding table).
# The embedding table still receives gradients from the type-embedding loss.
from giant.constants import K_MAX
pdg_emb_weight = stage1_model.pdg_embedding_weight()
x1_s2 = _build_sec_x1(sec_cont, sec_pdg_idx, pdg_emb_weight)
sec_mask = torch.arange(K_MAX, device=device).unsqueeze(0) < n_sec.unsqueeze(1)
l_s2 = flow_matching_loss_secondary(
sec_decoder,
x1_s2,
cond_cont,
cond_cat,
x1_s1.detach(),
sec_mask,
)
total = l_s1 + lambda_nsec * l_nsec + lambda_s2 * l_s2
return total, l_s1, l_nsec, l_s2
def train( def train(
model: torch.nn.Module, stage1_model: torch.nn.Module,
sec_decoder: torch.nn.Module,
train_loader: DataLoader, train_loader: DataLoader,
val_loader: DataLoader, val_loader: DataLoader,
mode: str, mode: str,
@@ -67,6 +156,8 @@ def train(
warmup_epochs: int, warmup_epochs: int,
device: torch.device, device: torch.device,
out_dir: str | Path, out_dir: str | Path,
lambda_nsec: float = 0.1,
lambda_s2: float = 1.0,
normalizer_dict: dict | None = None, normalizer_dict: dict | None = None,
pdg_map: dict | None = None, pdg_map: dict | None = None,
mat_map: dict | None = None, mat_map: dict | None = None,
@@ -79,8 +170,11 @@ def train(
out_dir = Path(out_dir) out_dir = Path(out_dir)
out_dir.mkdir(parents=True, exist_ok=True) out_dir.mkdir(parents=True, exist_ok=True)
model = model.to(device) stage1_model = stage1_model.to(device)
optimizer = optim.AdamW(model.parameters(), lr=lr) sec_decoder = sec_decoder.to(device)
all_params = list(stage1_model.parameters()) + list(sec_decoder.parameters())
optimizer = optim.AdamW(all_params, lr=lr)
def _lr_lambda(epoch: int) -> float: def _lr_lambda(epoch: int) -> float:
if warmup_epochs > 0 and epoch < warmup_epochs: if warmup_epochs > 0 and epoch < warmup_epochs:
@@ -97,7 +191,8 @@ def train(
best_val_loss = float("inf") best_val_loss = float("inf")
if resume_path is not None: if resume_path is not None:
ckpt = torch.load(resume_path, map_location=device, weights_only=False) ckpt = torch.load(resume_path, map_location=device, weights_only=False)
model.load_state_dict(ckpt["model"]) stage1_model.load_state_dict(ckpt["model"])
sec_decoder.load_state_dict(ckpt["sec_decoder"])
optimizer.load_state_dict(ckpt["optimizer"]) optimizer.load_state_dict(ckpt["optimizer"])
lr_sched.load_state_dict(ckpt["lr_sched"]) lr_sched.load_state_dict(ckpt["lr_sched"])
start_epoch = ckpt.get("epoch", 0) + 1 start_epoch = ckpt.get("epoch", 0) + 1
@@ -116,8 +211,12 @@ def train(
for epoch in range(start_epoch, epochs + 1): for epoch in range(start_epoch, epochs + 1):
epoch_start = time.monotonic() epoch_start = time.monotonic()
current_lr = optimizer.param_groups[0]["lr"] current_lr = optimizer.param_groups[0]["lr"]
model.train() stage1_model.train()
sec_decoder.train()
train_loss_sum = 0.0 train_loss_sum = 0.0
train_s1_sum = 0.0
train_nsec_sum = 0.0
train_s2_sum = 0.0
train_n = 0 train_n = 0
ema_loss = 0.0 ema_loss = 0.0
bar = tqdm( bar = tqdm(
@@ -128,27 +227,26 @@ def train(
unit="batch", unit="batch",
dynamic_ncols=True, dynamic_ncols=True,
) )
for cond_cont, cond_cat, x1 in bar: for batch in bar:
cond_cont = cond_cont.to(device) loss, l_s1, l_nsec, l_s2 = _compute_losses(
cond_cat = cond_cat.to(device) stage1_model, sec_decoder, batch, mode, ddpm_schedule, device,
x1 = x1.to(device) lambda_nsec, lambda_s2,
)
if mode == "flow":
loss = flow_matching_loss(model, x1, cond_cont, cond_cat)
else:
assert ddpm_schedule is not None
loss = ddpm_schedule.loss(model, x1, cond_cont, cond_cat)
optimizer.zero_grad() optimizer.zero_grad()
loss.backward() loss.backward()
torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0) torch.nn.utils.clip_grad_norm_(all_params, 1.0)
optimizer.step() optimizer.step()
B = batch[0].size(0)
batch_loss = loss.item() batch_loss = loss.item()
train_loss_sum += batch_loss * x1.size(0) train_loss_sum += batch_loss * B
train_n += x1.size(0) train_s1_sum += l_s1.item() * B
train_nsec_sum += l_nsec.item() * B
train_s2_sum += l_s2.item() * B
train_n += B
ema_loss = ( ema_loss = (
batch_loss batch_loss
if train_n == x1.size(0) if train_n == B
else 0.95 * ema_loss + 0.05 * batch_loss else 0.95 * ema_loss + 0.05 * batch_loss
) )
bar.set_postfix_str(f"loss={ema_loss:.4f}", refresh=False) bar.set_postfix_str(f"loss={ema_loss:.4f}", refresh=False)
@@ -158,28 +256,30 @@ def train(
bar.close() bar.close()
if shutdown.requested: if shutdown.requested:
# Mid-epoch: discard the partial epoch rather than persist an
# inconsistent (lr_sched not stepped, no validation) checkpoint.
break break
train_loss = train_loss_sum / max(train_n, 1) train_loss = train_loss_sum / max(train_n, 1)
lr_sched.step() lr_sched.step()
model.eval() stage1_model.eval()
sec_decoder.eval()
val_loss_sum = 0.0 val_loss_sum = 0.0
val_s1_sum = 0.0
val_nsec_sum = 0.0
val_s2_sum = 0.0
val_n = 0 val_n = 0
with torch.no_grad(): with torch.no_grad():
for cond_cont, cond_cat, x1 in val_loader: for batch in val_loader:
cond_cont = cond_cont.to(device) loss, l_s1, l_nsec, l_s2 = _compute_losses(
cond_cat = cond_cat.to(device) stage1_model, sec_decoder, batch, mode, ddpm_schedule, device,
x1 = x1.to(device) lambda_nsec, lambda_s2,
if mode == "flow": )
loss = flow_matching_loss(model, x1, cond_cont, cond_cat) B = batch[0].size(0)
else: val_loss_sum += loss.item() * B
assert ddpm_schedule is not None val_s1_sum += l_s1.item() * B
loss = ddpm_schedule.loss(model, x1, cond_cont, cond_cat) val_nsec_sum += l_nsec.item() * B
val_loss_sum += loss.item() * x1.size(0) val_s2_sum += l_s2.item() * B
val_n += x1.size(0) val_n += B
val_loss = val_loss_sum / max(val_n, 1) val_loss = val_loss_sum / max(val_n, 1)
epoch_time = time.monotonic() - epoch_start epoch_time = time.monotonic() - epoch_start
@@ -187,14 +287,24 @@ def train(
marker = " [best]" if is_best else "" marker = " [best]" if is_best else ""
print( print(
f"epoch {epoch:{epoch_w}d}/{epochs}" f"epoch {epoch:{epoch_w}d}/{epochs}"
f" train {train_loss:.4f} val {val_loss:.4f}" f" train {train_loss:.4f}"
f" (s1={train_s1_sum/max(train_n,1):.3f}"
f" nsec={train_nsec_sum/max(train_n,1):.3f}"
f" s2={train_s2_sum/max(train_n,1):.3f})"
f" val {val_loss:.4f}"
f" lr {current_lr:.2e} {epoch_time:.1f}s{marker}" f" lr {current_lr:.2e} {epoch_time:.1f}s{marker}"
) )
metrics_writer.writerow( metrics_writer.writerow(
{ {
"epoch": epoch, "epoch": epoch,
"train_loss": train_loss, "train_loss": train_loss,
"train_loss_s1": train_s1_sum / max(train_n, 1),
"train_loss_nsec": train_nsec_sum / max(train_n, 1),
"train_loss_s2": train_s2_sum / max(train_n, 1),
"val_loss": val_loss, "val_loss": val_loss,
"val_loss_s1": val_s1_sum / max(val_n, 1),
"val_loss_nsec": val_nsec_sum / max(val_n, 1),
"val_loss_s2": val_s2_sum / max(val_n, 1),
"lr": current_lr, "lr": current_lr,
"epoch_time_s": epoch_time, "epoch_time_s": epoch_time,
} }
@@ -204,7 +314,7 @@ def train(
if validate_every > 0 and epoch % validate_every == 0: if validate_every > 0 and epoch % validate_every == 0:
print(f"[epoch {epoch}] marginal validation:") print(f"[epoch {epoch}] marginal validation:")
validate_marginals( validate_marginals(
model, stage1_model,
val_loader, val_loader,
mode=mode, mode=mode,
schedule=ddpm_schedule, schedule=ddpm_schedule,
@@ -213,7 +323,8 @@ def train(
) )
ckpt: dict = { ckpt: dict = {
"model": model.state_dict(), "model": stage1_model.state_dict(),
"sec_decoder": sec_decoder.state_dict(),
"optimizer": optimizer.state_dict(), "optimizer": optimizer.state_dict(),
"lr_sched": lr_sched.state_dict(), "lr_sched": lr_sched.state_dict(),
"epoch": epoch, "epoch": epoch,
+7 -4
View File
@@ -53,18 +53,21 @@ def validate_marginals(
model.eval() model.eval()
all_real, all_gen = [], [] all_real, all_gen = [], []
for i, (cond_cont, cond_cat, x1) in enumerate(val_loader): for i, batch in enumerate(val_loader):
if n_batches is not None and i >= n_batches: if n_batches is not None and i >= n_batches:
break break
# Batch is (cond_cont, cond_cat, target_s1, n_sec, sec_cont, sec_pdg_idx);
# validate_marginals only needs the Stage-1 primary target.
cond_cont, cond_cat, x1 = batch[0], batch[1], batch[2]
cond_cont = cond_cont.to(device) cond_cont = cond_cont.to(device)
cond_cat = cond_cat.to(device) cond_cat = cond_cat.to(device)
if mode == "flow": if mode == "flow":
gen = sample_flow(model, cond_cont, cond_cat, **_kw(steps)) gen, _n_sec = sample_flow(model, cond_cont, cond_cat, **_kw(steps))
elif mode == "ddpm": elif mode == "ddpm":
gen = sample_ddpm(model, cond_cont, cond_cat, schedule) gen, _n_sec = sample_ddpm(model, cond_cont, cond_cat, schedule)
else: else:
gen = sample_ddim(model, cond_cont, cond_cat, schedule, **_kw(steps)) gen, _n_sec = sample_ddim(model, cond_cont, cond_cat, schedule, **_kw(steps))
all_real.append(x1.numpy()) all_real.append(x1.numpy())
all_gen.append(gen.cpu().numpy()) all_gen.append(gen.cpu().numpy())
+51 -18
View File
@@ -19,20 +19,31 @@ import uproot
ParquetCompression = Literal["lz4", "uncompressed", "snappy", "gzip", "brotli", "zstd"] ParquetCompression = Literal["lz4", "uncompressed", "snappy", "gzip", "brotli", "zstd"]
def _add_secondary_energy(df: pl.DataFrame) -> pl.DataFrame: def _add_secondary_attributes(df: pl.DataFrame) -> pl.DataFrame:
"""Add per-step `e_sec`: total initial kinetic energy of the secondaries born in it. """Add per-step secondary attributes via the parent→child track join.
Each secondary's creation energy is the `pre_E` of that child track's first step For each step that spawns secondaries, collects each child track's birth
(min `step_no`) in the same event, so for a parent step state (from the child track's first step in the same event) and emits:
`e_sec = Σ over child_track_ids of the child track's first-step pre_E`. Steps that e_sec float64 total secondary energy (sum of child first-step pre_E)
spawn nothing get 0.0. The full event must be present in `df` (it is the writer sec_E_list list[f64] per-secondary energy, sorted descending
concatenates every batch before this runs), since a child track's first step can sec_pdg_list list[i32] per-secondary PDG code, same order
live in a different read batch than its parent step. sec_dx_list list[f64] per-secondary birth direction x, same order
sec_dy_list list[f64] per-secondary birth direction y, same order
sec_dz_list list[f64] per-secondary birth direction z, same order
Steps with no children get 0.0 / empty lists. The full event must be
present in `df` (it is the writer concatenates before calling this).
""" """
first_E = ( first_step = (
df.sort("step_no") df.sort("step_no")
.group_by(["event_id", "track_id"]) .group_by(["event_id", "track_id"])
.agg(pl.col("pre_E").first().alias("child_E")) .agg(
pl.col("pre_E").first().alias("child_E"),
pl.col("pdg").first().alias("child_pdg"),
pl.col("pre_dx").first().alias("child_dx"),
pl.col("pre_dy").first().alias("child_dy"),
pl.col("pre_dz").first().alias("child_dz"),
)
.rename({"track_id": "child_track_id"}) .rename({"track_id": "child_track_id"})
) )
@@ -41,17 +52,39 @@ def _add_secondary_energy(df: pl.DataFrame) -> pl.DataFrame:
.with_row_index("_step_row") .with_row_index("_step_row")
.explode("child_track_ids") .explode("child_track_ids")
.rename({"child_track_ids": "child_track_id"}) .rename({"child_track_ids": "child_track_id"})
.drop_nulls("child_track_id") # steps with no children explode to a null row .drop_nulls("child_track_id")
) )
summed = (
exploded.join(first_E, on=["event_id", "child_track_id"], how="left") joined = exploded.join(first_step, on=["event_id", "child_track_id"], how="left")
# Sort each step's secondaries by descending energy, then aggregate into lists
per_step = (
joined.sort("child_E", descending=True)
.group_by("_step_row") .group_by("_step_row")
.agg(pl.col("child_E").sum().alias("e_sec")) .agg(
pl.col("child_E").sum().alias("e_sec"),
pl.col("child_E").alias("sec_E_list"),
pl.col("child_pdg").alias("sec_pdg_list"),
pl.col("child_dx").alias("sec_dx_list"),
pl.col("child_dy").alias("sec_dy_list"),
pl.col("child_dz").alias("sec_dz_list"),
)
) )
empty_list_f64 = pl.Series("x", [[]], dtype=pl.List(pl.Float64))
empty_list_i32 = pl.Series("x", [[]], dtype=pl.List(pl.Int32))
return ( return (
df.with_row_index("_step_row") df.with_row_index("_step_row")
.join(summed, on="_step_row", how="left") .join(per_step, on="_step_row", how="left")
.with_columns(pl.col("e_sec").fill_null(0.0).cast(pl.Float64)) .with_columns(
pl.col("e_sec").fill_null(0.0).cast(pl.Float64),
pl.col("sec_E_list").fill_null(empty_list_f64),
pl.col("sec_pdg_list").fill_null(empty_list_i32),
pl.col("sec_dx_list").fill_null(empty_list_f64),
pl.col("sec_dy_list").fill_null(empty_list_f64),
pl.col("sec_dz_list").fill_null(empty_list_f64),
)
.drop("_step_row") .drop("_step_row")
) )
@@ -118,8 +151,8 @@ def convert_steps_to_parquet(
# Steps tree carries the parent→child links needed to derive secondary energy; # Steps tree carries the parent→child links needed to derive secondary energy;
# other trees (e.g. Hits) don't, so only augment when the column is present. # other trees (e.g. Hits) don't, so only augment when the column is present.
if "child_track_ids" in df.columns: if "child_track_ids" in df.columns:
print("\nComputing per-step secondary energy (e_sec)", end=" ", flush=True) print("\nComputing per-step secondary attributes", end=" ", flush=True)
df = _add_secondary_energy(df) df = _add_secondary_attributes(df)
print(f"\nWriting {output_path}", end=" ", flush=True) print(f"\nWriting {output_path}", end=" ", flush=True)
df.write_parquet(output_path, compression=compression) df.write_parquet(output_path, compression=compression)
+24 -58
View File
@@ -1,67 +1,33 @@
import numpy as np import numpy as np
from giant.data.dataset import StepsDataset, train_val_split from giant.data.dataset import make_event_split
def _dummy(N=500, n_events=20): def test_make_event_split_sizes():
rng = np.random.default_rng(42) rng = np.random.default_rng(42)
data = {"event_id": rng.integers(0, n_events, size=N)} event_ids = rng.integers(0, 50, size=1000)
cond_cont = rng.standard_normal((N, 9)).astype(np.float32) train_set, val_set = make_event_split(event_ids, val_fraction=0.2)
cond_cat = rng.integers(0, 3, size=(N, 2)).astype(np.int64) unique = np.unique(event_ids)
target = rng.standard_normal((N, 6)).astype(np.float32) assert len(train_set) + len(val_set) == len(unique)
return data, cond_cont, cond_cat, target
def test_dataset_length(): def test_make_event_split_no_overlap():
data, cond_cont, cond_cat, target = _dummy()
assert len(StepsDataset(cond_cont, cond_cat, target)) == len(target)
def test_dataset_item_shapes():
data, cond_cont, cond_cat, target = _dummy()
c, k, t = StepsDataset(cond_cont, cond_cat, target)[0]
assert c.shape == (9,)
assert k.shape == (2,)
assert t.shape == (6,)
def test_split_sizes_sum_to_total():
data, cond_cont, cond_cat, target = _dummy(N=500)
train_ds, val_ds = train_val_split(
data, cond_cont, cond_cat, target, val_fraction=0.2
)
assert len(train_ds) + len(val_ds) == 500
def test_split_no_empty_sets():
data, cond_cont, cond_cat, target = _dummy(N=500, n_events=20)
train_ds, val_ds = train_val_split(
data, cond_cont, cond_cat, target, val_fraction=0.2
)
assert len(val_ds) > 0
assert len(train_ds) > 0
def test_split_event_leakage():
"""Train and val must not share any event_id."""
N = 1000
n_events = 50
rng = np.random.default_rng(7) rng = np.random.default_rng(7)
event_ids = rng.integers(0, n_events, size=N) event_ids = rng.integers(0, 50, size=1000)
data = {"event_id": event_ids} train_set, val_set = make_event_split(event_ids, val_fraction=0.2)
cond_cont = rng.standard_normal((N, 9)).astype(np.float32) assert train_set.isdisjoint(val_set)
cond_cat = rng.integers(0, 3, size=(N, 2)).astype(np.int64)
target = rng.standard_normal((N, 6)).astype(np.float32)
train_ds, val_ds = train_val_split(
data, cond_cont, cond_cat, target, val_fraction=0.2
)
# Recover which event_ids ended up in each split via the indices def test_make_event_split_no_empty_sets():
# (The dataset doesn't store event_ids, so we check via the original mask logic) rng = np.random.default_rng(0)
unique_events = np.unique(event_ids) event_ids = rng.integers(0, 20, size=500)
rng2 = np.random.default_rng(42) train_set, val_set = make_event_split(event_ids, val_fraction=0.2)
rng2.shuffle(unique_events) assert len(train_set) > 0
n_val = max(1, int(len(unique_events) * 0.2)) assert len(val_set) > 0
val_events = set(unique_events[:n_val].tolist())
train_events = set(unique_events[n_val:].tolist())
assert val_events.isdisjoint(train_events) def test_make_event_split_reproducible():
event_ids = np.arange(100)
a_tr, a_val = make_event_split(event_ids, val_fraction=0.1, seed=42)
b_tr, b_val = make_event_split(event_ids, val_fraction=0.1, seed=42)
assert a_tr == b_tr
assert a_val == b_val
+6 -4
View File
@@ -39,8 +39,9 @@ def test_sample_flow_shape():
B = 6 B = 6
cond_cont = torch.randn(B, COND_DIM) cond_cont = torch.randn(B, COND_DIM)
cond_cat = torch.zeros(B, 2, dtype=torch.long) cond_cat = torch.zeros(B, 2, dtype=torch.long)
out = sample_flow(_small_model(), cond_cont, cond_cat, steps=5) sample, n_sec = sample_flow(_small_model(), cond_cont, cond_cat, steps=5)
assert out.shape == (B, 9) assert sample.shape == (B, 9)
assert n_sec.shape == (B,)
def test_ddpm_loss_nonneg(): def test_ddpm_loss_nonneg():
@@ -55,5 +56,6 @@ def test_sample_ddim_shape():
schedule = CosineSchedule(T=50) schedule = CosineSchedule(T=50)
cond_cont = torch.randn(B, COND_DIM) cond_cont = torch.randn(B, COND_DIM)
cond_cat = torch.zeros(B, 2, dtype=torch.long) cond_cat = torch.zeros(B, 2, dtype=torch.long)
out = sample_ddim(_small_model(), cond_cont, cond_cat, schedule, steps=5) sample, n_sec = sample_ddim(_small_model(), cond_cont, cond_cat, schedule, steps=5)
assert out.shape == (B, 9) assert sample.shape == (B, 9)
assert n_sec.shape == (B,)
+4 -2
View File
@@ -39,7 +39,9 @@ def test_denoising_mlp_gradients_flow():
t = torch.rand(B) t = torch.rand(B)
cond_cont = torch.randn(B, COND_DIM) cond_cont = torch.randn(B, COND_DIM)
cond_cat = torch.zeros(B, 2, dtype=torch.long) cond_cat = torch.zeros(B, 2, dtype=torch.long)
loss = model(x_t, t, cond_cont, cond_cat).sum() # Both paths must be exercised to get gradients through all parameters.
loss.backward() flow_loss = model(x_t, t, cond_cont, cond_cat).sum()
nsec_loss = model.predict_n_sec(cond_cont, cond_cat).sum()
(flow_loss + nsec_loss).backward()
for name, p in model.named_parameters(): for name, p in model.named_parameters():
assert p.grad is not None, f"no grad for {name}" assert p.grad is not None, f"no grad for {name}"
+222
View File
@@ -0,0 +1,222 @@
"""Tests for Phase 2: secondary particle prediction."""
import numpy as np
import pytest
import torch
from giant.constants import COND_DIM, EMB_DIM, K_MAX, SEC_DIM, X_DIM
from giant.model.network import DenoisingMLP, SecondaryDecoder
from giant.model.schedule import flow_matching_loss_secondary
from giant.sample import sample_secondaries, snap_type_to_pdg_idx
# ── helpers ──────────────────────────────────────────────────────────────────
def _stage1(pdg=3, mat=2):
return DenoisingMLP(pdg_vocab=pdg, mat_vocab=mat, hidden_dim=32, n_blocks=2)
def _sec_decoder(pdg=3, mat=2):
return SecondaryDecoder(pdg_vocab=pdg, mat_vocab=mat, hidden_dim=32, n_blocks=2)
def _cond(B=8, pdg=3, mat=2):
cond_cont = torch.randn(B, COND_DIM)
cond_cat = torch.stack(
[torch.randint(0, pdg, (B,)), torch.randint(0, mat, (B,))], dim=1
)
return cond_cont, cond_cat
# ── DenoisingMLP Phase-2 additions ───────────────────────────────────────────
def test_predict_n_sec_shape():
B = 8
model = _stage1()
cond_cont, cond_cat = _cond(B)
logits = model.predict_n_sec(cond_cont, cond_cat)
assert logits.shape == (B, K_MAX + 1)
def test_predict_n_sec_no_nan():
B = 8
model = _stage1()
cond_cont, cond_cat = _cond(B)
logits = model.predict_n_sec(cond_cont, cond_cat)
assert torch.isfinite(logits).all()
def test_pdg_embedding_weight_shape():
model = _stage1(pdg=5, mat=2)
w = model.pdg_embedding_weight()
assert w.shape == (5, EMB_DIM)
# ── SecondaryDecoder ──────────────────────────────────────────────────────────
def test_sec_decoder_output_shape():
B = 8
decoder = _sec_decoder()
x_t = torch.randn(B, SEC_DIM)
t = torch.rand(B)
cond_cont, cond_cat = _cond(B)
stage1_out = torch.randn(B, X_DIM)
out = decoder(x_t, t, cond_cont, cond_cat, stage1_out)
assert out.shape == (B, SEC_DIM)
def test_sec_decoder_no_nan():
B = 4
decoder = _sec_decoder()
x_t = torch.randn(B, SEC_DIM)
t = torch.rand(B)
cond_cont, cond_cat = _cond(B)
stage1_out = torch.randn(B, X_DIM)
out = decoder(x_t, t, cond_cont, cond_cat, stage1_out)
assert torch.isfinite(out).all()
def test_sec_decoder_gradients():
B = 4
decoder = _sec_decoder()
x_t = torch.randn(B, SEC_DIM)
t = torch.rand(B)
cond_cont, cond_cat = _cond(B)
stage1_out = torch.randn(B, X_DIM)
decoder(x_t, t, cond_cont, cond_cat, stage1_out).sum().backward()
for name, p in decoder.named_parameters():
assert p.grad is not None, f"no grad for {name}"
# ── masked flow matching loss ─────────────────────────────────────────────────
def test_flow_matching_loss_secondary_scalar():
B, pdg, mat = 8, 3, 2
decoder = _sec_decoder(pdg, mat)
x1 = torch.randn(B, SEC_DIM)
cond_cont, cond_cat = _cond(B, pdg, mat)
stage1_out = torch.randn(B, X_DIM)
sec_mask = torch.ones(B, K_MAX, dtype=torch.bool)
loss = flow_matching_loss_secondary(decoder, x1, cond_cont, cond_cat, stage1_out, sec_mask)
assert loss.shape == ()
assert loss.item() >= 0.0
def test_flow_matching_loss_secondary_mask_zeros_padding():
"""Loss with all-zero mask (no valid secondaries) should be 0."""
B, pdg, mat = 4, 3, 2
decoder = _sec_decoder(pdg, mat)
x1 = torch.randn(B, SEC_DIM)
cond_cont, cond_cat = _cond(B, pdg, mat)
stage1_out = torch.randn(B, X_DIM)
sec_mask = torch.zeros(B, K_MAX, dtype=torch.bool)
loss = flow_matching_loss_secondary(decoder, x1, cond_cont, cond_cat, stage1_out, sec_mask)
assert loss.item() == pytest.approx(0.0, abs=1e-6)
def test_flow_matching_loss_secondary_has_grad():
B, pdg, mat = 4, 3, 2
decoder = _sec_decoder(pdg, mat)
x1 = torch.randn(B, SEC_DIM)
cond_cont, cond_cat = _cond(B, pdg, mat)
stage1_out = torch.randn(B, X_DIM)
sec_mask = torch.ones(B, K_MAX, dtype=torch.bool)
flow_matching_loss_secondary(
decoder, x1, cond_cont, cond_cat, stage1_out, sec_mask
).backward()
assert any(p.grad is not None for p in decoder.parameters())
# ── sampling ──────────────────────────────────────────────────────────────────
def test_sample_secondaries_shapes():
B, pdg, mat = 6, 3, 2
decoder = _sec_decoder(pdg, mat)
cond_cont, cond_cat = _cond(B, pdg, mat)
stage1_out = torch.randn(B, X_DIM)
n_sec_pred = torch.randint(0, K_MAX + 1, (B,))
sec_cont, sec_type_emb, sec_valid = sample_secondaries(
decoder, cond_cont, cond_cat, stage1_out, n_sec_pred, steps=3
)
assert sec_cont.shape == (B, K_MAX, 4)
assert sec_type_emb.shape == (B, K_MAX, EMB_DIM)
assert sec_valid.shape == (B, K_MAX)
assert sec_valid.dtype == torch.bool
def test_sample_secondaries_valid_mask_matches_n_sec():
B, pdg, mat = 4, 3, 2
decoder = _sec_decoder(pdg, mat)
cond_cont, cond_cat = _cond(B, pdg, mat)
stage1_out = torch.randn(B, X_DIM)
n_sec_pred = torch.tensor([0, 1, 3, K_MAX])
_, _, sec_valid = sample_secondaries(
decoder, cond_cont, cond_cat, stage1_out, n_sec_pred, steps=2
)
for i, n in enumerate(n_sec_pred.tolist()):
assert sec_valid[i, :n].all()
assert not sec_valid[i, n:].any()
def test_snap_type_to_pdg_idx_shape():
B, pdg_vocab = 4, 5
emb_weight = torch.randn(pdg_vocab, EMB_DIM)
sec_type_emb = torch.randn(B, K_MAX, EMB_DIM)
idx = snap_type_to_pdg_idx(sec_type_emb, emb_weight)
assert idx.shape == (B, K_MAX)
assert idx.dtype == torch.int64
assert (idx >= 0).all() and (idx < pdg_vocab).all()
# ── encode_secondaries round-trip ─────────────────────────────────────────────
def test_encode_secondaries_energy_conservation():
"""Decoded stick-breaking fractions must sum to ≈ e_sec."""
from giant.data.transforms import encode_secondaries
rng = np.random.default_rng(42)
N = 50
n_sec = rng.integers(1, 5, size=N)
e_sec = rng.uniform(0.1, 10.0, size=N).astype(np.float32)
sec_E_list = np.zeros((N, K_MAX), dtype=np.float32)
sec_dir_list = np.zeros((N, K_MAX, 3), dtype=np.float32)
sec_dir_list[:, :, 2] = 1.0
sec_valid = np.zeros((N, K_MAX), dtype=bool)
for i in range(N):
k = n_sec[i]
energies = rng.dirichlet(np.ones(k)) * e_sec[i]
energies = np.sort(energies)[::-1]
sec_E_list[i, :k] = energies.astype(np.float32)
sec_valid[i, :k] = True
pre_dir = rng.standard_normal((N, 3)).astype(np.float32)
pre_dir /= np.linalg.norm(pre_dir, axis=1, keepdims=True)
sec_cont = encode_secondaries(sec_E_list, sec_dir_list, sec_valid, e_sec, pre_dir)
assert sec_cont.shape == (N, K_MAX, 4)
assert np.isfinite(sec_cont).all()
def test_encode_secondaries_direction_encoding():
"""Local-frame secondary directions should be unit vectors for valid slots."""
from giant.data.transforms import encode_secondaries
rng = np.random.default_rng(7)
N = 20
e_sec = np.ones(N, dtype=np.float32) * 5.0
sec_E_list = np.zeros((N, K_MAX), dtype=np.float32)
sec_E_list[:, 0] = 3.0
sec_E_list[:, 1] = 2.0
sec_dir_list = rng.standard_normal((N, K_MAX, 3)).astype(np.float32)
norms = np.linalg.norm(sec_dir_list, axis=-1, keepdims=True)
sec_dir_list /= np.where(norms > 0, norms, 1.0)
sec_valid = np.zeros((N, K_MAX), dtype=bool)
sec_valid[:, :2] = True
pre_dir = np.tile([0.0, 0.0, 1.0], (N, 1)).astype(np.float32)
sec_cont = encode_secondaries(sec_E_list, sec_dir_list, sec_valid, e_sec, pre_dir)
# dir columns are sec_cont[:, :, 1:4]
local_dirs = sec_cont[:, :2, 1:4] # (N, 2, 3) — valid slots only
norms_out = np.linalg.norm(local_dirs, axis=-1)
np.testing.assert_allclose(norms_out, 1.0, atol=1e-5)