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 <noreply@anthropic.com>
This commit is contained in:
2026-06-29 14:17:40 +02:00
parent e6e0eb22bf
commit b25c6967ab
+13 -14
View File
@@ -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)