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:
+13
-14
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user