From b25c6967abd7e24c88e57e4fae8eda303287768f Mon Sep 17 00:00:00 2001 From: Lars Bogner Date: Mon, 29 Jun 2026 14:17:40 +0200 Subject: [PATCH] Fix giant.analysis import after Phase 2 dataset API changes train_val_split was removed from giant.data.dataset in favor of make_event_split + StreamingStepsDataset (event-based split, streaming batches), and build_features grew secondary-prediction outputs. make_val_loader and collect_samples still referenced the old API. Co-Authored-By: Claude Sonnet 4.6 --- giant/analysis.py | 27 +++++++++++++-------------- 1 file changed, 13 insertions(+), 14 deletions(-) diff --git a/giant/analysis.py b/giant/analysis.py index 0a5bd01..b32b992 100644 --- a/giant/analysis.py +++ b/giant/analysis.py @@ -105,11 +105,10 @@ from giant.constants import ( PREDICT_SCHEMA_VERSION, PREDICT_SCHEMA_VERSION_KEY, ) -from giant.data.dataset import train_val_split -from giant.data.loader import find_parquet_files, load_steps +from giant.data.dataset import make_event_split, StreamingStepsDataset +from giant.data.loader import find_parquet_files, load_event_ids from giant.data.transforms import ( Normalizer, - build_features, energy_simplex_decode, inv_log_transform, reconstruct_post_pos, @@ -275,20 +274,20 @@ def make_val_loader( very large datasets, build a StreamingStepsDataset directly (see giant.pipeline). """ files = find_parquet_files(data) - chunks = [load_steps(f) for f in files] - full = {k: np.concatenate([c[k] for c in chunks]) for k in chunks[0]} + all_event_ids = np.concatenate([load_event_ids(f) for f in files]) + _, val_events = make_event_split(all_event_ids, val_fraction=val_fraction, seed=seed) - cond_cont, cond_cat, target, _, _ = build_features( - full, - bundle.pdg_map, - bundle.mat_map, + val_ds = StreamingStepsDataset( + files=files, + split_events=val_events, + pdg_map=bundle.pdg_map, + mat_map=bundle.mat_map, cond_normalizer=bundle.cond_normalizer, target_normalizer=bundle.target_normalizer, + batch_size=batch_size, + shuffle=False, ) - _, val_ds = train_val_split( - full, cond_cont, cond_cat, target, val_fraction=val_fraction, seed=seed - ) - return DataLoader(val_ds, batch_size=batch_size, shuffle=False) + return DataLoader(val_ds, batch_size=None) def _to_raw_targets( @@ -316,7 +315,7 @@ def collect_samples( steps_kw = {} if steps is None else {"steps": steps} cond_list, real_list, gen_list = [], [], [] - for i, (cond_cont, cond_cat, x1) in enumerate(val_loader): + for i, (cond_cont, cond_cat, x1, _n_sec, _sec_cont, _sec_pdg) in enumerate(val_loader): if n_batches is not None and i >= n_batches: break cond_cont = cond_cont.to(device)