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)