Back to home page

EIC code displayed by LXR

 
 

    


File indexing completed on 2026-09-01 09:34:17

0001 #!/usr/bin/env python3
0002 """Trace CKF trajectories to MC particles for a reconstructed DIS file set.
0003 
0004 The expensive PODIO reader is opened once per input file.  Parallelism is
0005 therefore across independent files, rather than across event ranges within one
0006 file.  This is substantially more efficient for the 100-file DIS samples.
0007 
0008 Run this script inside a matching EIC environment. This example processes a
0009 numbered reconstructed-file campaign:
0010 
0011     python track2particle.py \
0012         --config epic_noise \
0013         --input-root /path/to/rootfiles \
0014         --file-start 1 --file-stop 100 \
0015         --events 99 --workers 8 \
0016         --output-root /path/to/track2particle_output \
0017         --eta-range -1 1 0.1 \
0018         --require-all-events
0019 
0020 Omit the momentum options to use the shared defaults from
0021 ``epic_analysis_base.py``. Use a new output root when changing an event-level
0022 cut, because that cut affects the saved per-file particle chunks.
0023 """
0024 
0025 from __future__ import annotations
0026 
0027 import argparse
0028 import contextlib
0029 from concurrent.futures import ProcessPoolExecutor, as_completed
0030 import io
0031 import json
0032 from multiprocessing import get_context
0033 from pathlib import Path
0034 import resource
0035 import signal
0036 import sys
0037 import time
0038 
0039 import awkward as ak
0040 import numpy as np
0041 import pandas as pd
0042 
0043 import epic_analysis_base as ana
0044 
0045 
0046 # Relative defaults make a cloned copy usable without editing personal paths.
0047 DEFAULT_INPUT_ROOT = Path("rootfiles")
0048 DEFAULT_OUTPUT_ROOT = Path("track2particle_output")
0049 DEFAULT_MODULE_DIRECTORY = Path(__file__).resolve().parent
0050 PER_RUN_DIRECTORY = "per_run_analysis"
0051 
0052 DEFAULT_CUTS = {
0053     "momentum_min": ana.TRACK_MOM_MIN,  # GeV
0054     "vertex_r_max": ana.VERTEX_CUT_R_MAX,  # mm
0055     "vertex_z_max": ana.VERTEX_CUT_Z_MAX,  # mm
0056     "track_hit_count_min": ana.TRACK_HIT_COUNT_MIN,
0057     "track_hit_fraction_min": ana.TRACK_HIT_FRACTION_MIN,
0058 }
0059 
0060 DEFAULT_MOMENTUM_BINS = np.array([0.0, 0.5, 1.0, 5.0, 1000.0])
0061 DEFAULT_ETA_BINS = np.round(np.arange(-4.0, 4.01, 0.1), 1)
0062 TRACK_CLASS_ORDER = ("good_signal", "background_track", "fake_or_ghost")
0063 TRACK_CLASS_LABELS = {
0064     "good_signal": "good signal",
0065     "background_track": "background track",
0066     "fake_or_ghost": "fake/ghost",
0067 }
0068 
0069 
0070 def normalize_eta_bins(eta_bins):
0071     """Return increasing eta-bin edges.
0072 
0073     A three-value, non-increasing input such as ``[-1, 1, 0.1]`` is
0074     interpreted as ``[eta_min, eta_max, step]``. An increasing sequence is
0075     treated as explicit bin edges.
0076     """
0077     eta_bins = np.asarray(eta_bins, dtype=float)
0078     if (
0079         len(eta_bins) == 3
0080         and eta_bins[0] < eta_bins[1]
0081         and eta_bins[2] > 0
0082         and not np.all(np.diff(eta_bins) > 0)
0083     ):
0084         eta_min, eta_max, eta_step = eta_bins
0085         eta_bins = np.arange(eta_min, eta_max, eta_step)
0086         eta_bins = np.append(eta_bins, eta_max)
0087 
0088     if (
0089         eta_bins.ndim != 1
0090         or len(eta_bins) < 2
0091         or not np.all(np.isfinite(eta_bins))
0092         or not np.all(np.diff(eta_bins) > 0)
0093     ):
0094         raise ValueError(
0095             "eta_bins must be increasing bin edges or "
0096             "[eta_min, eta_max, step]"
0097         )
0098     return eta_bins
0099 
0100 
0101 class EventTimeout(Exception):
0102     """Raised when one event exceeds the configured analysis time limit."""
0103 
0104 
0105 def stop_slow_event(signum, frame):
0106     """Convert a worker alarm into an exception handled per event."""
0107     del signum, frame
0108     raise EventTimeout("event exceeded the configured time limit")
0109 
0110 
0111 def get_output_directory(input_file, output_root):
0112     """Return the per-file output directory for one reconstructed input.
0113 
0114     New analyses live under ``per_run_analysis``. The legacy flat location is
0115     returned only when it already exists, which keeps older output roots
0116     readable without affecting the new layout.
0117     """
0118     file_tag = Path(input_file).name
0119     for suffix in (".root", ".edm4eic", ".edm4hep"):
0120         if file_tag.endswith(suffix):
0121             file_tag = file_tag[: -len(suffix)]
0122     output_root = Path(output_root)
0123     output_directory = output_root / PER_RUN_DIRECTORY / file_tag
0124     legacy_directory = output_root / file_tag
0125     if not output_directory.exists() and legacy_directory.exists():
0126         return legacy_directory
0127     return output_directory
0128 
0129 
0130 def manifest_chunk_paths(manifest, manifest_file, column):
0131     """Resolve unique chunk paths from one manifest column.
0132 
0133     Current manifests store paths relative to their per-file directory, so an
0134     output tree can be moved intact. Absolute paths from older manifests remain
0135     supported.
0136     """
0137     manifest_directory = Path(manifest_file).parent
0138     paths = []
0139     for value in manifest[column].dropna().unique():
0140         if not str(value):
0141             continue
0142         path = Path(value)
0143         paths.append(path if path.is_absolute() else manifest_directory / path)
0144     return paths
0145 
0146 
0147 def build_input_files(config, input_root, file_start, file_stop):
0148     """Build and validate the expected numbered reconstructed-file paths."""
0149     input_directory = Path(input_root) / config
0150     input_files = [
0151         input_directory
0152         / (
0153             "rec_pythia8NCDIS_10x275_minQ2=1_beamEffects_"
0154             f"xAngle=-0.025_hiDiv_1.{file_index:04d}_{config}_n99_skip0.root"
0155         )
0156         for file_index in range(file_start, file_stop + 1)
0157     ]
0158     missing_files = [path for path in input_files if not path.is_file()]
0159     if missing_files:
0160         preview = "\n".join(f"  {path}" for path in missing_files[:10])
0161         if len(missing_files) > 10:
0162             preview += f"\n  ... and {len(missing_files) - 10} more"
0163         raise FileNotFoundError(
0164             f"Missing {len(missing_files)} expected input files:\n{preview}"
0165         )
0166     return input_files
0167 
0168 
0169 def manifest_is_reusable(manifest_file, first_event, number_of_events):
0170     """Check whether a prior manifest covers exactly the requested event range."""
0171     manifest_file = Path(manifest_file)
0172     if not manifest_file.is_file():
0173         return False
0174     try:
0175         manifest = pd.read_csv(manifest_file)
0176     except Exception:
0177         return False
0178     required_columns = {
0179         "event",
0180         "status",
0181         "trajectory_file",
0182         "particle_file",
0183     }
0184     if not required_columns.issubset(manifest.columns):
0185         return False
0186     expected_events = np.arange(
0187         first_event,
0188         first_event + number_of_events,
0189         dtype=int,
0190     )
0191     actual_events = np.sort(manifest["event"].to_numpy(dtype=int))
0192     if not np.array_equal(actual_events, expected_events):
0193         return False
0194 
0195     # A manifest is authoritative only if every chunk path it names still
0196     # exists. Empty paths are allowed for events that produced no output rows.
0197     for column in ("trajectory_file", "particle_file"):
0198         paths = manifest_chunk_paths(manifest, manifest_file, column)
0199         if any(not path.is_file() for path in paths):
0200             return False
0201     return True
0202 
0203 
0204 def settings_are_reusable(settings_file, requested_settings):
0205     """Return whether saved processing settings match the requested values."""
0206     settings_file = Path(settings_file)
0207     if not settings_file.is_file():
0208         return False
0209     try:
0210         with settings_file.open() as stream:
0211             saved_settings = json.load(stream)
0212     except (OSError, ValueError):
0213         return False
0214     return saved_settings == requested_settings
0215 
0216 
0217 def analyze_input_file(task):
0218     """Analyze one complete input file inside one clean worker process."""
0219     input_file = Path(task["input_file"])
0220     output_root = Path(task["output_root"])
0221     module_directory = Path(task["module_directory"])
0222     first_event = int(task["first_event"])
0223     number_of_events = int(task["number_of_events"])
0224     timeout_seconds = int(task["timeout_seconds"])
0225     cuts = dict(task["cuts"])
0226     resume = bool(task["resume"])
0227 
0228     output_directory = get_output_directory(input_file, output_root)
0229     chunk_directory = output_directory / "chunks"
0230     manifest_file = output_directory / "manifest.csv"
0231     settings_file = output_directory / "analysis_settings.json"
0232     worker_log = output_directory / "analysis.log"
0233     output_directory.mkdir(parents=True, exist_ok=True)
0234     chunk_directory.mkdir(parents=True, exist_ok=True)
0235 
0236     requested_settings = {
0237         "schema_version": 1,
0238         "input_file": str(input_file.resolve()),
0239         "first_event": first_event,
0240         "number_of_events": number_of_events,
0241         "cuts": cuts,
0242     }
0243 
0244     if resume and manifest_file.exists() and not settings_are_reusable(
0245         settings_file,
0246         requested_settings,
0247     ):
0248         raise ValueError(
0249             "Existing manifest was produced with different or unrecorded "
0250             f"analysis settings: {manifest_file}. Use a different "
0251             "--output-root for the new cuts."
0252         )
0253 
0254     if (
0255         resume
0256         and settings_are_reusable(settings_file, requested_settings)
0257         and manifest_is_reusable(
0258             manifest_file,
0259             first_event,
0260             number_of_events,
0261         )
0262     ):
0263         manifest = pd.read_csv(manifest_file)
0264         return {
0265             "input_file": str(input_file),
0266             "output_directory": str(output_directory),
0267             "manifest_file": str(manifest_file),
0268             "status": "reused",
0269             "events_ok": int((manifest["status"] == "ok").sum()),
0270             "events_problem": int((manifest["status"] != "ok").sum()),
0271             "elapsed_seconds": 0.0,
0272             "max_rss_mb": 0.0,
0273         }
0274 
0275     if manifest_file.exists() and not resume:
0276         raise FileExistsError(
0277             f"Output already exists; use --resume or a different output root: "
0278             f"{manifest_file}"
0279         )
0280 
0281     start_time = time.monotonic()
0282     trajectory_tables = []
0283     particle_tables = []
0284     status_rows = []
0285 
0286     with settings_file.open("w") as stream:
0287         json.dump(requested_settings, stream, indent=2, sort_keys=True)
0288         stream.write("\n")
0289 
0290     # Keep verbose PODIO/helper output in a file beside this input's manifest.
0291     # The parent process prints only one concise progress line per input file.
0292     with worker_log.open("w") as log_stream:
0293         with contextlib.redirect_stdout(log_stream), contextlib.redirect_stderr(
0294             log_stream
0295         ):
0296             print(f"Input: {input_file}")
0297             print(f"UTC start: {time.strftime('%Y-%m-%dT%H:%M:%SZ', time.gmtime())}")
0298 
0299             if str(module_directory) not in sys.path:
0300                 sys.path.insert(0, str(module_directory))
0301             import epic_analysis_base as ana_worker
0302             import epic_analysis_podio as pod_worker
0303 
0304             uproot_tree = ana_worker.read_ur(str(input_file), "events")
0305             available_events = int(uproot_tree.num_entries)
0306             stop_event = min(
0307                 first_event + number_of_events,
0308                 available_events,
0309             )
0310             event_numbers = np.arange(first_event, stop_event, dtype=int)
0311             if len(event_numbers) != number_of_events:
0312                 raise ValueError(
0313                     f"{input_file}: requested {number_of_events} events from "
0314                     f"{first_event}, but only {len(event_numbers)} are available"
0315                 )
0316 
0317             podio_events = pod_worker.read_podio(str(input_file))
0318             signal.signal(signal.SIGALRM, stop_slow_event)
0319 
0320             for event_number in event_numbers:
0321                 try:
0322                     if timeout_seconds > 0:
0323                         signal.alarm(timeout_seconds)
0324 
0325                     event_number = int(event_number)
0326                     event = podio_events[event_number]
0327 
0328                     # Suppress repetitive per-event helper messages while
0329                     # retaining exceptions and the event status in manifest.csv.
0330                     quiet_output = io.StringIO()
0331                     with contextlib.redirect_stdout(quiet_output):
0332                         mc_particles = ak.to_dataframe(
0333                             ana_worker.get_part(
0334                                 uproot_tree,
0335                                 entry_start=event_number,
0336                                 entry_stop=event_number + 1,
0337                                 kprimary=1,
0338                             )
0339                         ).reset_index()
0340                         track_parameters = ak.to_dataframe(
0341                             ana_worker.get_params(
0342                                 uproot_tree,
0343                                 "CentralCKFTrackParameters",
0344                                 entry_start=event_number,
0345                                 entry_stop=event_number + 1,
0346                             )
0347                         ).reset_index()
0348                         (
0349                             selected_particles,
0350                             trajectories,
0351                         ) = pod_worker.get_part_traj_counts(
0352                             event,
0353                             mc_particles,
0354                             ksignal=1,
0355                         )
0356 
0357                     track_parameters = track_parameters.rename(
0358                         columns={
0359                             "subentry": "traj_id",
0360                             "mom": "reco_mom",
0361                             "eta": "reco_eta",
0362                             "pt": "reco_pt",
0363                         }
0364                     )[["traj_id", "reco_mom", "reco_eta", "reco_pt"]]
0365 
0366                     particle_cut = (
0367                         (
0368                             selected_particles["vertex_r"].abs()
0369                             < cuts["vertex_r_max"]
0370                         )
0371                         & (
0372                             selected_particles["vertex.z"].abs()
0373                             < cuts["vertex_z_max"]
0374                         )
0375                         & (
0376                             selected_particles["mom"]
0377                             > cuts["momentum_min"]
0378                         )
0379                     )
0380                     selected_particles = selected_particles[
0381                         particle_cut
0382                     ].copy()
0383                     particle_id_column = (
0384                         "orig_subentry"
0385                         if "orig_subentry" in selected_particles.columns
0386                         else "subentry"
0387                     )
0388                     selected_particles["particle_id"] = selected_particles[
0389                         particle_id_column
0390                     ].astype(int)
0391 
0392                     trajectories = trajectories.merge(
0393                         track_parameters,
0394                         on="traj_id",
0395                         how="left",
0396                         validate="one_to_one",
0397                     )
0398                     trajectories["event"] = event_number
0399                     selected_particles["event"] = event_number
0400 
0401                     if not trajectories.empty:
0402                         trajectory_tables.append(trajectories)
0403                     if not selected_particles.empty:
0404                         particle_tables.append(selected_particles)
0405 
0406                     status_rows.append(
0407                         {
0408                             "event": event_number,
0409                             "status": "ok",
0410                             "n_trajectories": int(len(trajectories)),
0411                             "n_selected_particles": int(
0412                                 len(selected_particles)
0413                             ),
0414                             "error": "",
0415                         }
0416                     )
0417                 except EventTimeout:
0418                     status_rows.append(
0419                         {
0420                             "event": int(event_number),
0421                             "status": "timeout",
0422                             "n_trajectories": 0,
0423                             "n_selected_particles": 0,
0424                             "error": (
0425                                 f"event exceeded the {timeout_seconds}-second "
0426                                 "time limit"
0427                             ),
0428                         }
0429                     )
0430                 except Exception as error:
0431                     status_rows.append(
0432                         {
0433                             "event": int(event_number),
0434                             "status": "error",
0435                             "n_trajectories": 0,
0436                             "n_selected_particles": 0,
0437                             "error": f"{type(error).__name__}: {error}",
0438                         }
0439                     )
0440                 finally:
0441                     signal.alarm(0)
0442 
0443             first_processed_event = int(event_numbers[0])
0444             last_processed_event = int(event_numbers[-1])
0445             chunk_label = (
0446                 f"{first_processed_event:06d}_{last_processed_event:06d}"
0447             )
0448 
0449             trajectory_file = ""
0450             if trajectory_tables:
0451                 trajectory_file = (
0452                     chunk_directory / f"trajectories_{chunk_label}.csv.gz"
0453                 )
0454                 pd.concat(
0455                     trajectory_tables,
0456                     ignore_index=True,
0457                 ).to_csv(trajectory_file, index=False)
0458 
0459             particle_file = ""
0460             if particle_tables:
0461                 particle_file = (
0462                     chunk_directory / f"particles_{chunk_label}.csv.gz"
0463                 )
0464                 pd.concat(
0465                     particle_tables,
0466                     ignore_index=True,
0467                 ).to_csv(particle_file, index=False)
0468 
0469             # Save portable paths relative to this per-file directory. The
0470             # manifest may then be moved together with its chunks.
0471             trajectory_manifest_path = (
0472                 str(trajectory_file.relative_to(output_directory))
0473                 if trajectory_file
0474                 else ""
0475             )
0476             particle_manifest_path = (
0477                 str(particle_file.relative_to(output_directory))
0478                 if particle_file
0479                 else ""
0480             )
0481             for row in status_rows:
0482                 row["trajectory_file"] = (
0483                     trajectory_manifest_path
0484                 )
0485                 row["particle_file"] = particle_manifest_path
0486 
0487             manifest = pd.DataFrame(status_rows).sort_values("event")
0488             manifest.to_csv(manifest_file, index=False)
0489             print(f"UTC end: {time.strftime('%Y-%m-%dT%H:%M:%SZ', time.gmtime())}")
0490 
0491     elapsed_seconds = time.monotonic() - start_time
0492     max_rss_mb = resource.getrusage(resource.RUSAGE_SELF).ru_maxrss / 1024.0
0493     return {
0494         "input_file": str(input_file),
0495         "output_directory": str(output_directory),
0496         "manifest_file": str(manifest_file),
0497         "status": "processed",
0498         "events_ok": int((manifest["status"] == "ok").sum()),
0499         "events_problem": int((manifest["status"] != "ok").sum()),
0500         "elapsed_seconds": elapsed_seconds,
0501         "max_rss_mb": max_rss_mb,
0502     }
0503 
0504 
0505 def process_input_files(
0506     input_files,
0507     output_root,
0508     module_directory,
0509     first_event,
0510     number_of_events,
0511     timeout_seconds,
0512     cuts,
0513     number_of_workers,
0514     resume,
0515 ):
0516     """Process independent input files concurrently."""
0517     tasks = [
0518         {
0519             "input_file": str(input_file),
0520             "output_root": str(output_root),
0521             "module_directory": str(module_directory),
0522             "first_event": first_event,
0523             "number_of_events": number_of_events,
0524             "timeout_seconds": timeout_seconds,
0525             "cuts": cuts,
0526             "resume": resume,
0527         }
0528         for input_file in input_files
0529     ]
0530     worker_count = min(max(1, int(number_of_workers)), len(tasks))
0531     print(
0532         f"Processing {len(tasks)} files with {worker_count} file workers; "
0533         f"{number_of_events} events per file"
0534     )
0535 
0536     results = []
0537     if worker_count == 1:
0538         for task_index, task in enumerate(tasks, start=1):
0539             result = analyze_input_file(task)
0540             results.append(result)
0541             print_file_progress(task_index, len(tasks), result)
0542         return results
0543 
0544     # The standalone parent never imports PODIO or ROOT. Forked workers
0545     # therefore start without inherited readers and import them independently.
0546     with ProcessPoolExecutor(
0547         max_workers=worker_count,
0548         mp_context=get_context("fork"),
0549     ) as executor:
0550         futures = {
0551             executor.submit(analyze_input_file, task): task["input_file"]
0552             for task in tasks
0553         }
0554         for completed_count, future in enumerate(
0555             as_completed(futures),
0556             start=1,
0557         ):
0558             input_file = futures[future]
0559             try:
0560                 result = future.result()
0561             except Exception as error:
0562                 result = {
0563                     "input_file": input_file,
0564                     "output_directory": str(
0565                         get_output_directory(input_file, output_root)
0566                     ),
0567                     "manifest_file": "",
0568                     "status": "worker_error",
0569                     "events_ok": 0,
0570                     "events_problem": number_of_events,
0571                     "elapsed_seconds": np.nan,
0572                     "max_rss_mb": np.nan,
0573                     "error": f"{type(error).__name__}: {error}",
0574                 }
0575             results.append(result)
0576             print_file_progress(completed_count, len(tasks), result)
0577     return results
0578 
0579 
0580 def print_file_progress(completed_count, total_count, result):
0581     """Print one concise progress line for a completed input file."""
0582     print(
0583         f"[{completed_count:03d}/{total_count:03d}] "
0584         f"{Path(result['input_file']).name}: {result['status']}, "
0585         f"ok={result['events_ok']}, problem={result['events_problem']}, "
0586         f"time={result['elapsed_seconds']:.1f}s, "
0587         f"maxRSS={result['max_rss_mb']:.0f} MB",
0588         flush=True,
0589     )
0590     if result.get("error"):
0591         print(f"    {result['error']}", flush=True)
0592 
0593 
0594 def load_active_chunks(manifest, manifest_file=None):
0595     """Load chunk files named by successful rows in one manifest.
0596 
0597     ``manifest_file`` is required for portable relative chunk paths. It is
0598     optional only for compatibility with callers holding old absolute paths.
0599     """
0600     good_events = manifest[manifest["status"] == "ok"]
0601     manifest_file = Path(manifest_file or "manifest.csv")
0602     trajectory_files = manifest_chunk_paths(
0603         good_events, manifest_file, "trajectory_file"
0604     )
0605     particle_files = manifest_chunk_paths(
0606         good_events, manifest_file, "particle_file"
0607     )
0608     trajectories = (
0609         pd.concat(
0610             [pd.read_csv(path) for path in trajectory_files],
0611             ignore_index=True,
0612         )
0613         if trajectory_files
0614         else pd.DataFrame()
0615     )
0616     particles = (
0617         pd.concat(
0618             [pd.read_csv(path) for path in particle_files],
0619             ignore_index=True,
0620         )
0621         if particle_files
0622         else pd.DataFrame()
0623     )
0624     return trajectories, particles
0625 
0626 
0627 def select_valid_tracks(trajectories, cuts):
0628     """Classify reconstructed candidate tracks by their dominant hit source.
0629 
0630     Every trajectory hit contributes to ``total_count``. Noise or invalid
0631     hits have particle ID -1, while detector-background hits retain their
0632     real particle ID and generator status. A track is matched only when a
0633     real particle supplies strictly more than half of all trajectory hits.
0634     """
0635     if trajectories.empty:
0636         return pd.DataFrame(), pd.DataFrame()
0637 
0638     valid_tracks = trajectories[
0639         trajectories["total_count"].abs() >= cuts["track_hit_count_min"]
0640     ].copy()
0641     has_real_source = valid_tracks["most_common_source"] >= 0
0642     has_majority = (
0643         valid_tracks["max_fraction"]
0644         > cuts["track_hit_fraction_min"]
0645     )
0646     from_signal = valid_tracks["part_status"].isin([1, 2])
0647     valid_tracks["is_majority_matched"] = has_real_source & has_majority
0648     valid_tracks["is_good_signal"] = (
0649         valid_tracks["is_majority_matched"] & from_signal
0650     )
0651     valid_tracks["track_class"] = np.select(
0652         [
0653             valid_tracks["is_good_signal"],
0654             valid_tracks["is_majority_matched"],
0655         ],
0656         TRACK_CLASS_ORDER[:2],
0657         default=TRACK_CLASS_ORDER[2],
0658     )
0659     good_signal_tracks = valid_tracks[
0660         valid_tracks["is_good_signal"]
0661     ].copy()
0662     return valid_tracks, good_signal_tracks
0663 
0664 
0665 def mark_particles_with_tracks(particles, valid_tracks):
0666     """Mark particles supplying a strict majority of a candidate track's hits."""
0667     particles = particles.copy()
0668     if particles.empty:
0669         particles["has_valid_track"] = pd.Series(dtype=bool)
0670         return particles
0671 
0672     majority_tracks = valid_tracks[
0673         valid_tracks["is_majority_matched"]
0674     ]
0675     valid_particle_keys = (
0676         majority_tracks[["event", "most_common_source"]]
0677         .rename(columns={"most_common_source": "particle_id"})
0678         .drop_duplicates()
0679         if not majority_tracks.empty
0680         else pd.DataFrame(columns=["event", "particle_id"])
0681     )
0682     particles = particles.merge(
0683         valid_particle_keys.assign(has_valid_track=True),
0684         on=["event", "particle_id"],
0685         how="left",
0686     )
0687     particles["has_valid_track"] = (
0688         particles["has_valid_track"].fillna(False).astype(bool)
0689     )
0690     return particles
0691 
0692 
0693 def build_binned_performance(
0694     data,
0695     passed_column,
0696     momentum_column,
0697     eta_column,
0698     metric_name,
0699     momentum_bins,
0700     eta_bins,
0701 ):
0702     """Return one row per momentum/eta bin with a 68% Wilson interval."""
0703     eta_bins = normalize_eta_bins(eta_bins)
0704     required_columns = {passed_column, momentum_column, eta_column}
0705     missing_columns = required_columns - set(data.columns)
0706     if missing_columns:
0707         raise KeyError(
0708             f"{metric_name}: missing columns {sorted(missing_columns)}"
0709         )
0710 
0711     finite_data = data.dropna(
0712         subset=[momentum_column, eta_column]
0713     ).copy()
0714     finite_data[passed_column] = finite_data[passed_column].astype(bool)
0715     denominator, _, _ = np.histogram2d(
0716         finite_data[momentum_column],
0717         finite_data[eta_column],
0718         bins=[momentum_bins, eta_bins],
0719     )
0720     numerator, _, _ = np.histogram2d(
0721         finite_data.loc[finite_data[passed_column], momentum_column],
0722         finite_data.loc[finite_data[passed_column], eta_column],
0723         bins=[momentum_bins, eta_bins],
0724     )
0725 
0726     rows = []
0727     for momentum_index in range(len(momentum_bins) - 1):
0728         for eta_index in range(len(eta_bins) - 1):
0729             total = int(denominator[momentum_index, eta_index])
0730             passed = int(numerator[momentum_index, eta_index])
0731             value = np.nan
0732             interval_low = np.nan
0733             interval_high = np.nan
0734             if total > 0:
0735                 value = passed / total
0736                 scale = 1.0 + 1.0 / total
0737                 center = (value + 0.5 / total) / scale
0738                 half_width = (
0739                     np.sqrt(
0740                         value * (1.0 - value) / total
0741                         + 0.25 / total**2
0742                     )
0743                     / scale
0744                 )
0745                 interval_low = max(0.0, center - half_width)
0746                 interval_high = min(1.0, center + half_width)
0747 
0748             rows.append(
0749                 {
0750                     "metric": metric_name,
0751                     "momentum_bin": momentum_index,
0752                     "momentum_low_GeV": momentum_bins[momentum_index],
0753                     "momentum_high_GeV": momentum_bins[
0754                         momentum_index + 1
0755                     ],
0756                     "eta_bin": eta_index,
0757                     "eta_low": eta_bins[eta_index],
0758                     "eta_high": eta_bins[eta_index + 1],
0759                     "eta_center": 0.5
0760                     * (eta_bins[eta_index] + eta_bins[eta_index + 1]),
0761                     "numerator": passed,
0762                     "denominator": total,
0763                     "value": value,
0764                     "interval_low": interval_low,
0765                     "interval_high": interval_high,
0766                 }
0767             )
0768     return pd.DataFrame(rows)
0769 
0770 
0771 def _plot_fraction_distribution(
0772     denominator_data,
0773     numerator_data,
0774     metric_label,
0775     beam,
0776     setting,
0777     nhit_cut,
0778     bins,
0779     momentum_bins,
0780     output_directory,
0781 ):
0782     """Plot numerator and denominator eta spectra in each momentum bin."""
0783     import matplotlib
0784 
0785     matplotlib.use("Agg")
0786     import matplotlib.pyplot as plt
0787     import seaborn as sns
0788 
0789     bins = normalize_eta_bins(bins)
0790     eta_min, eta_max = bins[0], bins[-1]
0791     n_panels = max(1, len(momentum_bins) - 1)
0792     ncols = 2 if n_panels > 1 else 1
0793     nrows = int(np.ceil(n_panels / ncols))
0794     fig, axes = plt.subplots(
0795         nrows,
0796         ncols,
0797         sharex=True,
0798         sharey=True,
0799         figsize=(3 * ncols, 3 * nrows),
0800     )
0801     axes = np.atleast_1d(axes).ravel()
0802     total_denominator = 0
0803     total_numerator = 0
0804 
0805     for panel, (mom_min, mom_max) in enumerate(
0806         zip(momentum_bins[:-1], momentum_bins[1:])
0807     ):
0808         ax = axes[panel]
0809         denominator = denominator_data[
0810             (denominator_data.mom >= mom_min)
0811             & (denominator_data.mom < mom_max)
0812             & (denominator_data.eta >= eta_min)
0813             & (denominator_data.eta <= eta_max)
0814         ]
0815         numerator = numerator_data[
0816             (numerator_data.mom >= mom_min)
0817             & (numerator_data.mom < mom_max)
0818             & (numerator_data.eta >= eta_min)
0819             & (numerator_data.eta <= eta_max)
0820         ]
0821         sns.histplot(denominator, x="eta", bins=bins, ax=ax)
0822         sns.histplot(numerator, x="eta", bins=bins, ax=ax)
0823 
0824         denominator_count = len(denominator)
0825         numerator_count = len(numerator)
0826         total_denominator += denominator_count
0827         total_numerator += numerator_count
0828         fraction = (
0829             numerator_count / denominator_count
0830             if denominator_count
0831             else np.nan
0832         )
0833         panel_label = "Eff" if metric_label == "Efficiency" else metric_label
0834         ax.text(
0835             0.15,
0836             0.8,
0837             f"{mom_min}<p<{mom_max} GeV\n{panel_label}: "
0838             f"{fraction:.3f}={numerator_count}/{denominator_count}",
0839             fontsize=10,
0840             transform=ax.transAxes,
0841             ha="left",
0842         )
0843         ax.set_xlim(eta_min, eta_max)
0844         ax.set_xlabel(r"$\eta$")
0845 
0846     for ax in axes[len(momentum_bins) - 1 :]:
0847         ax.set_visible(False)
0848 
0849     total_fraction = (
0850         total_numerator / total_denominator
0851         if total_denominator
0852         else np.nan
0853     )
0854     fig.suptitle(
0855         f"{metric_label} ({nhit_cut} hits) | {beam}, {setting} | "
0856         f"total={total_fraction:.3f} "
0857         f"({total_numerator}/{total_denominator})"
0858     )
0859     plt.tight_layout(rect=[0, 0, 1, 0.96])
0860     output_directory = Path(output_directory)
0861     output_directory.mkdir(parents=True, exist_ok=True)
0862     metric_tag = "eff" if metric_label == "Efficiency" else "purity"
0863     fig.savefig(
0864         output_directory / f"bg_{beam}_{setting}_{nhit_cut}_{metric_tag}.png"
0865     )
0866     plt.close(fig)
0867     return total_fraction, total_numerator, total_denominator
0868 
0869 
0870 def plot_purity_distribution(
0871     valid_track_params,
0872     good_signal_track_params,
0873     beam,
0874     setting,
0875     nhit_cut,
0876     bins=None,
0877     mom_bin=None,
0878     outdir=None,
0879 ):
0880     """Plot good-signal tracks over all candidate tracks versus eta."""
0881     bins = np.arange(-4, 4.01, 0.1) if bins is None else bins
0882     momentum_bins = (
0883         DEFAULT_MOMENTUM_BINS if mom_bin is None else np.asarray(mom_bin)
0884     )
0885     return _plot_fraction_distribution(
0886         valid_track_params,
0887         good_signal_track_params,
0888         "Purity",
0889         beam,
0890         setting,
0891         nhit_cut,
0892         bins,
0893         momentum_bins,
0894         outdir,
0895     )
0896 
0897 
0898 def plot_efficiency_distribution(
0899     particles,
0900     beam,
0901     setting,
0902     nhit_cut,
0903     bins=None,
0904     mom_bin=None,
0905     outdir=None,
0906 ):
0907     """Plot reconstructed signal particles over selected particles versus eta."""
0908     bins = np.arange(-4, 4.01, 0.1) if bins is None else bins
0909     momentum_bins = (
0910         DEFAULT_MOMENTUM_BINS if mom_bin is None else np.asarray(mom_bin)
0911     )
0912     particles_with_track = particles[particles["has_valid_track"]]
0913     return _plot_fraction_distribution(
0914         particles,
0915         particles_with_track,
0916         "Efficiency",
0917         beam,
0918         setting,
0919         nhit_cut,
0920         bins,
0921         momentum_bins,
0922         outdir,
0923     )
0924 
0925 
0926 def plot_track_quality_metrics(
0927     track_class_bins,
0928     tracks,
0929     momentum_bins,
0930     analysis_label,
0931     output_directory,
0932 ):
0933     """Plot track-class fractions and the physical hit-count distribution."""
0934     import matplotlib
0935 
0936     matplotlib.use("Agg")
0937     import matplotlib.pyplot as plt
0938     from matplotlib.lines import Line2D
0939     import seaborn as sns
0940 
0941     class_colors = sns.color_palette(n_colors=len(TRACK_CLASS_ORDER))
0942     panel_count = len(momentum_bins) - 1
0943     ncols = 2 if panel_count > 1 else 1
0944     nrows = int(np.ceil(panel_count / ncols))
0945     fig, axes = plt.subplots(
0946         nrows,
0947         ncols,
0948         sharex=True,
0949         sharey=True,
0950         figsize=(3.5 * ncols, 3 * nrows),
0951     )
0952     axes = np.atleast_1d(axes).ravel()
0953     for momentum_index, ax in enumerate(axes[:panel_count]):
0954         panel_data = track_class_bins[
0955             track_class_bins["momentum_bin"] == momentum_index
0956         ]
0957         for class_name in TRACK_CLASS_ORDER:
0958             class_data = panel_data[panel_data["metric"] == class_name]
0959             ax.plot(
0960                 class_data["eta_center"],
0961                 class_data["value"],
0962                 marker=".",
0963                 linewidth=1,
0964                 label=TRACK_CLASS_LABELS[class_name],
0965             )
0966         ax.set_ylim(0, 1.05)
0967         ax.set_xlabel(r"$\eta$")
0968         ax.set_ylabel("track fraction")
0969         ax.set_title(
0970             f"{momentum_bins[momentum_index]} < p < "
0971             f"{momentum_bins[momentum_index + 1]} GeV"
0972         )
0973     for ax in axes[panel_count:]:
0974         ax.set_visible(False)
0975     axes[0].legend(fontsize=8)
0976     fig.suptitle(f"Track classification | {analysis_label}")
0977     fig.tight_layout(rect=[0, 0, 1, 0.96])
0978     fig.savefig(Path(output_directory) / "track_class_fractions.png")
0979     plt.close(fig)
0980 
0981     fig, axes = plt.subplots(
0982         nrows,
0983         ncols,
0984         sharex=True,
0985         sharey=True,
0986         figsize=(3.5 * ncols, 3 * nrows),
0987     )
0988     axes = np.atleast_1d(axes).ravel()
0989     for momentum_index, ax in enumerate(axes[:panel_count]):
0990         momentum_low = momentum_bins[momentum_index]
0991         momentum_high = momentum_bins[momentum_index + 1]
0992         panel_tracks = tracks[
0993             (tracks["reco_mom"] >= momentum_low)
0994             & (tracks["reco_mom"] < momentum_high)
0995         ].copy()
0996         if panel_tracks.empty:
0997             ax.text(
0998                 0.5,
0999                 0.5,
1000                 "no tracks",
1001                 ha="center",
1002                 va="center",
1003                 transform=ax.transAxes,
1004             )
1005             ax.set_xlim(3, 15)
1006             ax.set_xticks(np.arange(3, 16, 2))
1007             ax.set_xlabel("hits per trajectory")
1008             ax.set_ylabel("tracks")
1009             ax.set_title(f"{momentum_low} < p < {momentum_high} GeV")
1010             continue
1011         panel_tracks["track_class_label"] = panel_tracks[
1012             "track_class"
1013         ].map(TRACK_CLASS_LABELS)
1014         sns.histplot(
1015             panel_tracks,
1016             x="total_count",
1017             hue="track_class_label",
1018             hue_order=[TRACK_CLASS_LABELS[name] for name in TRACK_CLASS_ORDER],
1019             palette=class_colors,
1020             discrete=True,
1021             element="step",
1022             fill=False,
1023             common_norm=False,
1024             legend=False,
1025             ax=ax,
1026         )
1027         ax.set_xlim(3, 15)
1028         ax.set_xticks(np.arange(3, 16, 2))
1029         ax.set_xlabel("hits per trajectory")
1030         ax.set_ylabel("tracks")
1031         ax.set_title(f"{momentum_low} < p < {momentum_high} GeV")
1032 
1033     for ax in axes[panel_count:]:
1034         ax.set_visible(False)
1035     legend_handles = [
1036         Line2D(
1037             [0],
1038             [0],
1039             color=class_colors[index],
1040         label=TRACK_CLASS_LABELS[class_name],
1041         )
1042         for index, class_name in enumerate(TRACK_CLASS_ORDER)
1043     ]
1044     fig.legend(
1045         handles=legend_handles,
1046         loc="upper center",
1047         bbox_to_anchor=(0.5, 0.94),
1048         ncol=len(TRACK_CLASS_ORDER),
1049         frameon=False,
1050         fontsize=8,
1051     )
1052     fig.suptitle(
1053         f"Candidate-track hit counts | {analysis_label}",
1054         y=0.995,
1055     )
1056     fig.tight_layout(rect=[0, 0, 1, 0.88])
1057     fig.savefig(Path(output_directory) / "hits_per_trajectory.png")
1058     plt.close(fig)
1059 
1060 
1061 def aggregate_results(
1062     input_files,
1063     output_root,
1064     analysis_label,
1065     beam_label,
1066     cuts,
1067     momentum_bins,
1068     eta_bins,
1069     performance_tag=None,
1070 ):
1071     """Summarize manifests/chunks and write combined performance products."""
1072     eta_bins = normalize_eta_bins(eta_bins)
1073     summary_rows = []
1074     particle_frames = []
1075     track_frames = []
1076 
1077     for input_file in input_files:
1078         output_directory = get_output_directory(input_file, output_root)
1079         manifest_file = output_directory / "manifest.csv"
1080         if not manifest_file.is_file():
1081             raise FileNotFoundError(
1082                 f"Missing manifest needed for aggregation: {manifest_file}"
1083             )
1084         manifest = pd.read_csv(manifest_file)
1085         trajectories, particles = load_active_chunks(manifest, manifest_file)
1086         valid_tracks, good_signal_tracks = select_valid_tracks(
1087             trajectories,
1088             cuts,
1089         )
1090         particles = mark_particles_with_tracks(particles, valid_tracks)
1091 
1092         n_particles = len(particles)
1093         n_particles_with_track = int(particles["has_valid_track"].sum())
1094         n_valid_tracks = len(valid_tracks)
1095         n_good_signal_tracks = len(good_signal_tracks)
1096         n_background_tracks = int(
1097             valid_tracks["track_class"].eq("background_track").sum()
1098         )
1099         n_fake_or_ghost_tracks = int(
1100             valid_tracks["track_class"].eq("fake_or_ghost").sum()
1101         )
1102         summary = {
1103             "input_file": str(input_file),
1104             "events_requested": int(len(manifest)),
1105             "events_ok": int((manifest["status"] == "ok").sum()),
1106             "events_timeout": int(
1107                 (manifest["status"] == "timeout").sum()
1108             ),
1109             "events_error": int(
1110                 manifest["status"].isin(["error", "worker_error"]).sum()
1111             ),
1112             "selected_particles": int(n_particles),
1113             "particles_with_valid_track": int(n_particles_with_track),
1114             "valid_tracks": int(n_valid_tracks),
1115             "good_signal_tracks": int(n_good_signal_tracks),
1116             "background_tracks": n_background_tracks,
1117             "fake_or_ghost_tracks": n_fake_or_ghost_tracks,
1118             "efficiency": (
1119                 n_particles_with_track / n_particles
1120                 if n_particles
1121                 else np.nan
1122             ),
1123             "purity": (
1124                 n_good_signal_tracks / n_valid_tracks
1125                 if n_valid_tracks
1126                 else np.nan
1127             ),
1128             "background_track_fraction": (
1129                 n_background_tracks / n_valid_tracks
1130                 if n_valid_tracks
1131                 else np.nan
1132             ),
1133             "fake_or_ghost_fraction": (
1134                 n_fake_or_ghost_tracks / n_valid_tracks
1135                 if n_valid_tracks
1136                 else np.nan
1137             ),
1138             "mean_hits_per_trajectory": (
1139                 valid_tracks["total_count"].mean()
1140                 if n_valid_tracks
1141                 else np.nan
1142             ),
1143             "median_hits_per_trajectory": (
1144                 valid_tracks["total_count"].median()
1145                 if n_valid_tracks
1146                 else np.nan
1147             ),
1148             "output_directory": str(output_directory),
1149         }
1150         summary_rows.append(summary)
1151         pd.DataFrame([summary]).to_csv(
1152             output_directory / "summary.csv",
1153             index=False,
1154         )
1155 
1156         if not particles.empty:
1157             particle_frames.append(
1158                 particles.assign(source_file=str(input_file))
1159             )
1160         if not valid_tracks.empty:
1161             track_frames.append(
1162                 valid_tracks.assign(source_file=str(input_file))
1163             )
1164 
1165     combined_particles = (
1166         pd.concat(particle_frames, ignore_index=True)
1167         if particle_frames
1168         else pd.DataFrame()
1169     )
1170     combined_valid_tracks = (
1171         pd.concat(track_frames, ignore_index=True)
1172         if track_frames
1173         else pd.DataFrame()
1174     )
1175     if combined_particles.empty or combined_valid_tracks.empty:
1176         raise RuntimeError(
1177             "Aggregation found no selected particles or valid tracks"
1178         )
1179 
1180     efficiency_bins = build_binned_performance(
1181         combined_particles,
1182         passed_column="has_valid_track",
1183         momentum_column="mom",
1184         eta_column="eta",
1185         metric_name="efficiency",
1186         momentum_bins=momentum_bins,
1187         eta_bins=eta_bins,
1188     )
1189     purity_bins = build_binned_performance(
1190         combined_valid_tracks,
1191         passed_column="is_good_signal",
1192         momentum_column="reco_mom",
1193         eta_column="reco_eta",
1194         metric_name="purity",
1195         momentum_bins=momentum_bins,
1196         eta_bins=eta_bins,
1197     )
1198 
1199     performance_directory = Path(output_root) / analysis_label / "performance"
1200     if performance_tag:
1201         performance_directory = performance_directory / performance_tag
1202     performance_directory.mkdir(parents=True, exist_ok=True)
1203     pd.DataFrame(summary_rows).to_csv(
1204         performance_directory / "file_summary.csv",
1205         index=False,
1206     )
1207     efficiency_bins.to_csv(
1208         performance_directory / "efficiency_bins.csv",
1209         index=False,
1210     )
1211     purity_bins.to_csv(
1212         performance_directory / "purity_bins.csv",
1213         index=False,
1214     )
1215 
1216     eta_min, eta_max = eta_bins[0], eta_bins[-1]
1217     momentum_min, momentum_max = momentum_bins[0], momentum_bins[-1]
1218     tracks_in_analysis_range = combined_valid_tracks[
1219         (combined_valid_tracks["reco_eta"] >= eta_min)
1220         & (combined_valid_tracks["reco_eta"] <= eta_max)
1221         & (combined_valid_tracks["reco_mom"] >= momentum_min)
1222         & (combined_valid_tracks["reco_mom"] < momentum_max)
1223     ].copy()
1224     track_columns = [
1225         "source_file",
1226         "event",
1227         "traj_id",
1228         "reco_mom",
1229         "reco_eta",
1230         "reco_pt",
1231         "total_count",
1232         "max_count",
1233         "max_fraction",
1234         "most_common_source",
1235         "part_status",
1236         "track_class",
1237     ]
1238     tracks_in_analysis_range[track_columns].to_csv(
1239         performance_directory / "hits_per_trajectory.csv",
1240         index=False,
1241     )
1242 
1243     track_class_bin_frames = []
1244     for class_name in TRACK_CLASS_ORDER:
1245         class_data = combined_valid_tracks.assign(
1246             in_track_class=combined_valid_tracks["track_class"].eq(class_name)
1247         )
1248         track_class_bin_frames.append(
1249             build_binned_performance(
1250                 class_data,
1251                 passed_column="in_track_class",
1252                 momentum_column="reco_mom",
1253                 eta_column="reco_eta",
1254                 metric_name=class_name,
1255                 momentum_bins=momentum_bins,
1256                 eta_bins=eta_bins,
1257             )
1258         )
1259     track_class_bins = pd.concat(track_class_bin_frames, ignore_index=True)
1260     track_class_bins.to_csv(
1261         performance_directory / "track_class_bins.csv",
1262         index=False,
1263     )
1264 
1265     candidate_tracks = len(tracks_in_analysis_range)
1266     track_class_counts = tracks_in_analysis_range[
1267         "track_class"
1268     ].value_counts()
1269     track_class_summary = {
1270         "eta_min": eta_min,
1271         "eta_max": eta_max,
1272         "candidate_tracks": candidate_tracks,
1273         "good_signal_tracks": int(track_class_counts.get("good_signal", 0)),
1274         "background_tracks": int(track_class_counts.get("background_track", 0)),
1275         "fake_or_ghost_tracks": int(track_class_counts.get("fake_or_ghost", 0)),
1276         "mean_hits_per_trajectory": tracks_in_analysis_range[
1277             "total_count"
1278         ].mean(),
1279         "median_hits_per_trajectory": tracks_in_analysis_range[
1280             "total_count"
1281         ].median(),
1282         "hits_per_trajectory_q10": tracks_in_analysis_range[
1283             "total_count"
1284         ].quantile(0.10),
1285         "hits_per_trajectory_q90": tracks_in_analysis_range[
1286             "total_count"
1287         ].quantile(0.90),
1288     }
1289     for class_name, count_name in (
1290         ("good_signal", "purity"),
1291         ("background_track", "background_track_fraction"),
1292         ("fake_or_ghost", "fake_or_ghost_fraction"),
1293     ):
1294         count = int(track_class_counts.get(class_name, 0))
1295         track_class_summary[count_name] = (
1296             count / candidate_tracks if candidate_tracks else np.nan
1297         )
1298     pd.DataFrame([track_class_summary]).to_csv(
1299         performance_directory / "track_class_summary.csv",
1300         index=False,
1301     )
1302     plot_track_quality_metrics(
1303         track_class_bins,
1304         tracks_in_analysis_range,
1305         momentum_bins,
1306         analysis_label,
1307         performance_directory,
1308     )
1309 
1310     valid_track_params = combined_valid_tracks.rename(
1311         columns={
1312             "reco_mom": "mom",
1313             "reco_eta": "eta",
1314             "reco_pt": "pt",
1315         }
1316     )
1317     good_signal_track_params = valid_track_params[
1318         valid_track_params["is_good_signal"]
1319     ].copy()
1320     nhit_cut = cuts["track_hit_count_min"]
1321     (
1322         efficiency_total,
1323         efficiency_numerator,
1324         efficiency_denominator,
1325     ) = plot_efficiency_distribution(
1326         combined_particles,
1327         beam_label,
1328         analysis_label,
1329         nhit_cut,
1330         bins=eta_bins,
1331         mom_bin=momentum_bins,
1332         outdir=performance_directory,
1333     )
1334     (
1335         purity_total,
1336         purity_numerator,
1337         purity_denominator,
1338     ) = plot_purity_distribution(
1339         valid_track_params,
1340         good_signal_track_params,
1341         beam_label,
1342         analysis_label,
1343         nhit_cut,
1344         bins=eta_bins,
1345         mom_bin=momentum_bins,
1346         outdir=performance_directory,
1347     )
1348 
1349     performance_summary = pd.DataFrame(
1350         [
1351             {
1352                 "metric": "efficiency",
1353                 "numerator": efficiency_numerator,
1354                 "denominator": efficiency_denominator,
1355                 "value": efficiency_total,
1356             },
1357             {
1358                 "metric": "purity",
1359                 "numerator": purity_numerator,
1360                 "denominator": purity_denominator,
1361                 "value": purity_total,
1362             },
1363             {
1364                 "metric": "background_track_fraction",
1365                 "numerator": track_class_summary["background_tracks"],
1366                 "denominator": candidate_tracks,
1367                 "value": track_class_summary["background_track_fraction"],
1368             },
1369             {
1370                 "metric": "fake_or_ghost_fraction",
1371                 "numerator": track_class_summary["fake_or_ghost_tracks"],
1372                 "denominator": candidate_tracks,
1373                 "value": track_class_summary["fake_or_ghost_fraction"],
1374             },
1375         ]
1376     )
1377     performance_summary["analysis_label"] = analysis_label
1378     performance_summary["input_files"] = len(input_files)
1379     performance_summary.to_csv(
1380         performance_directory / "performance_summary.csv",
1381         index=False,
1382     )
1383     return performance_summary, performance_directory
1384 
1385 
1386 def parse_arguments():
1387     """Parse and validate the command-line interface."""
1388     parser = argparse.ArgumentParser(
1389         description=(
1390             "Analyze reconstructed DIS files with one PODIO reader per file."
1391         )
1392     )
1393     parser.add_argument(
1394         "--config",
1395         default=None,
1396         help=(
1397             "Detector/campaign tag in numbered input filenames, e.g. "
1398             "no_L4_HD4_noise. Required unless --input-file is used."
1399         ),
1400     )
1401     parser.add_argument(
1402         "--input-file",
1403         action="append",
1404         type=Path,
1405         default=None,
1406         help=(
1407             "Exact reconstructed ROOT file to process or aggregate. Repeat "
1408             "for multiple files; bypasses --config and the numbered pattern."
1409         ),
1410     )
1411     parser.add_argument(
1412         "--input-root",
1413         type=Path,
1414         default=DEFAULT_INPUT_ROOT,
1415         help="Directory containing one subdirectory per configuration.",
1416     )
1417     parser.add_argument(
1418         "--output-root",
1419         type=Path,
1420         default=DEFAULT_OUTPUT_ROOT,
1421         help="Root directory for per-file chunks and combined performance.",
1422     )
1423     parser.add_argument("--file-start", type=int, default=1)
1424     parser.add_argument("--file-stop", type=int, default=100)
1425     parser.add_argument("--first-event", type=int, default=0)
1426     parser.add_argument("--events", type=int, default=99)
1427     parser.add_argument(
1428         "--workers",
1429         type=int,
1430         default=4,
1431         help="Number of input files processed concurrently.",
1432     )
1433     parser.add_argument("--timeout-seconds", type=int, default=120)
1434     parser.add_argument("--beam-label", default="10x275")
1435     parser.add_argument(
1436         "--analysis-label",
1437         default=None,
1438         help="Combined-output label; defaults to --config.",
1439     )
1440     parser.add_argument(
1441         "--eta-range",
1442         nargs=3,
1443         type=float,
1444         metavar=("MIN", "MAX", "STEP"),
1445         default=None,
1446         help="Eta range for combined tables and plots, e.g. -1 1 0.1.",
1447     )
1448     parser.add_argument(
1449         "--momentum-min",
1450         type=float,
1451         default=None,
1452         help=(
1453             "Minimum generated-particle momentum in GeV. The default comes "
1454             "from epic_analysis_base.TRACK_MOM_MIN."
1455         ),
1456     )
1457     parser.add_argument(
1458         "--momentum-bins",
1459         nargs="+",
1460         type=float,
1461         default=None,
1462         metavar="EDGE",
1463         help=(
1464             "Increasing momentum-bin edges in GeV. Defaults to the script's "
1465             "standard bins."
1466         ),
1467     )
1468     parser.add_argument(
1469         "--performance-tag",
1470         default=None,
1471         help=(
1472             "Optional subdirectory under performance/. With --eta-range, "
1473             "a non-overwriting eta tag is generated automatically."
1474         ),
1475     )
1476     parser.add_argument(
1477         "--module-directory",
1478         type=Path,
1479         default=DEFAULT_MODULE_DIRECTORY,
1480     )
1481     parser.add_argument(
1482         "--resume",
1483         action="store_true",
1484         help="Reuse complete existing manifests instead of overwriting them.",
1485     )
1486     parser.add_argument(
1487         "--require-all-events",
1488         action="store_true",
1489         help="Exit with an error if any requested event times out or fails.",
1490     )
1491     mode = parser.add_mutually_exclusive_group()
1492     mode.add_argument(
1493         "--process-only",
1494         action="store_true",
1495         help="Generate per-file manifests/chunks without combined plots.",
1496     )
1497     mode.add_argument(
1498         "--aggregate-only",
1499         action="store_true",
1500         help="Reuse existing manifests/chunks and only rebuild summaries.",
1501     )
1502     args = parser.parse_args()
1503 
1504     if args.file_start < 1 or args.file_stop < args.file_start:
1505         parser.error("--file-start/--file-stop define an invalid range")
1506     if args.first_event < 0 or args.events < 1:
1507         parser.error("--first-event must be >=0 and --events must be >=1")
1508     if args.workers < 1:
1509         parser.error("--workers must be >=1")
1510     if args.momentum_min is not None and (
1511         not np.isfinite(args.momentum_min) or args.momentum_min < 0
1512     ):
1513         parser.error("--momentum-min must be finite and >=0")
1514     if args.momentum_bins is not None:
1515         momentum_bins = np.asarray(args.momentum_bins, dtype=float)
1516         if (
1517             len(momentum_bins) < 2
1518             or not np.all(np.isfinite(momentum_bins))
1519             or not np.all(np.diff(momentum_bins) > 0)
1520         ):
1521             parser.error("--momentum-bins must be finite increasing edges")
1522     if args.input_file is None and args.config is None:
1523         parser.error("provide --config or at least one --input-file")
1524     return args
1525 
1526 
1527 def value_tag(value):
1528     """Return a filename-safe compact representation of one bin value."""
1529     return f"{value:g}".replace("-", "m").replace(".", "p")
1530 
1531 
1532 def main():
1533     """Process reconstructed files and/or aggregate their saved chunks."""
1534     args = parse_arguments()
1535     cuts = dict(DEFAULT_CUTS)
1536     if args.momentum_min is not None:
1537         cuts["momentum_min"] = args.momentum_min
1538     momentum_bins = (
1539         np.asarray(args.momentum_bins, dtype=float)
1540         if args.momentum_bins is not None
1541         else DEFAULT_MOMENTUM_BINS
1542     )
1543     if args.input_file:
1544         input_files = [path.resolve() for path in args.input_file]
1545         missing_files = [path for path in input_files if not path.is_file()]
1546         if missing_files:
1547             raise FileNotFoundError(
1548                 "Missing explicit input file(s):\n"
1549                 + "\n".join(f"  {path}" for path in missing_files)
1550             )
1551         inferred_label = input_files[0].parent.name
1552     else:
1553         input_files = build_input_files(
1554             args.config,
1555             args.input_root,
1556             args.file_start,
1557             args.file_stop,
1558         )
1559         inferred_label = args.config
1560 
1561     analysis_label = args.analysis_label or args.config or inferred_label
1562     eta_bins = (
1563         normalize_eta_bins(args.eta_range)
1564         if args.eta_range is not None
1565         else DEFAULT_ETA_BINS
1566     )
1567     performance_tag = args.performance_tag
1568     if performance_tag is None and args.eta_range is not None:
1569         performance_tag = (
1570             f"eta_{value_tag(eta_bins[0])}_{value_tag(eta_bins[-1])}_"
1571             f"step_{value_tag(args.eta_range[2])}"
1572         )
1573     args.output_root.mkdir(parents=True, exist_ok=True)
1574 
1575     print(f"Configuration: {args.config or inferred_label}")
1576     print(f"Input files:   {len(input_files)}")
1577     if not args.input_file:
1578         print(f"File range:    {args.file_start} through {args.file_stop}")
1579     print(f"Event range:   {args.first_event} through "
1580           f"{args.first_event + args.events - 1}")
1581     print(f"Output root:   {args.output_root}")
1582     print(f"Cuts:          {cuts}")
1583     print(f"Momentum bins: {momentum_bins.tolist()} GeV")
1584 
1585     if not args.aggregate_only:
1586         run_results = process_input_files(
1587             input_files=input_files,
1588             output_root=args.output_root,
1589             module_directory=args.module_directory,
1590             first_event=args.first_event,
1591             number_of_events=args.events,
1592             timeout_seconds=args.timeout_seconds,
1593             cuts=cuts,
1594             number_of_workers=args.workers,
1595             resume=args.resume,
1596         )
1597         run_summary = pd.DataFrame(run_results)
1598         if len(input_files) == 1:
1599             run_summary_file = (
1600                 get_output_directory(input_files[0], args.output_root)
1601                 / "run_summary.csv"
1602             )
1603         else:
1604             run_summary_file = (
1605                 args.output_root / f"{analysis_label}_run_summary.csv"
1606             )
1607         run_summary.to_csv(run_summary_file, index=False)
1608         print(f"Saved run summary: {run_summary_file}")
1609 
1610         worker_failures = run_summary["status"].eq("worker_error")
1611         if worker_failures.any():
1612             failed_files = run_summary.loc[
1613                 worker_failures,
1614                 "input_file",
1615             ].tolist()
1616             raise RuntimeError(
1617                 f"{len(failed_files)} file workers failed; rerun with "
1618                 "--resume after inspecting the per-file logs"
1619             )
1620         if args.require_all_events and run_summary["events_problem"].sum() > 0:
1621             problem_count = int(run_summary["events_problem"].sum())
1622             raise RuntimeError(
1623                 f"{problem_count} requested event(s) failed or timed out; "
1624                 "inspect the per-file manifests before aggregation"
1625             )
1626 
1627     if not args.process_only:
1628         performance_summary, performance_directory = aggregate_results(
1629             input_files=input_files,
1630             output_root=args.output_root,
1631             analysis_label=analysis_label,
1632             beam_label=args.beam_label,
1633             cuts=cuts,
1634             momentum_bins=momentum_bins,
1635             eta_bins=eta_bins,
1636             performance_tag=performance_tag,
1637         )
1638         print(f"Saved performance products: {performance_directory}")
1639         print(performance_summary.to_string(index=False))
1640 
1641 
1642 if __name__ == "__main__":
1643     main()