From 903227d2df46ef92cf3092c1a6bfd1a3fd0a6041 Mon Sep 17 00:00:00 2001 From: Lars Bogner Date: Mon, 22 Jun 2026 09:06:05 +0200 Subject: [PATCH] Skip rows with unknown PDG codes during predict Checkpoints have a fixed PDG embedding vocab sized at train time, so a code unseen during training has no embedding index. Drop those rows and report a per-code skip count instead of raising a KeyError. --- giant/cli.py | 24 +++++++++++++++++++++++- 1 file changed, 23 insertions(+), 1 deletion(-) diff --git a/giant/cli.py b/giant/cli.py index 9cc42a5..3e7bf5f 100644 --- a/giant/cli.py +++ b/giant/cli.py @@ -1,3 +1,4 @@ +from collections import Counter from enum import Enum from pathlib import Path from typing import Optional @@ -344,13 +345,26 @@ def predict( writer: pq.ParquetWriter | None = None total = 0 + skipped = 0 + unknown_pdg_counts: Counter[int] = Counter() total_rows = sum(pq.ParquetFile(path).metadata.num_rows for path in files) chunk_iter = iter_file_chunks if coord == Coord.local else iter_cond_chunks bar = tqdm(total=total_rows, desc="predict", unit="row", dynamic_ncols=True) for path in files: for chunk in chunk_iter(path): + N_in = len(chunk["event_id"]) + + pdg_mask = np.array([int(p) in pdg_map for p in chunk["pdg"]]) + if not pdg_mask.all(): + unknown_pdg_counts.update(int(p) for p in chunk["pdg"][~pdg_mask]) + chunk = {k: v[pdg_mask] for k, v in chunk.items()} + N = len(chunk["event_id"]) + skipped += N_in - N + if N == 0: + bar.update(N_in) + continue if coord == Coord.local: cond_cont, cond_cat, target_raw, _, _ = build_features( @@ -458,12 +472,20 @@ def predict( writer = pq.ParquetWriter(out, table.schema) writer.write_table(table) total += N - bar.update(N) + bar.update(N_in) bar.close() if writer is not None: writer.close() + if skipped: + codes = ", ".join( + f"{pdg} ({count})" for pdg, count in sorted(unknown_pdg_counts.items()) + ) + typer.echo( + f"warning: skipped {skipped:,} row(s) with unknown PDG code(s): {codes}", + err=True, + ) typer.echo(f"wrote {total:,} rows → {out}")