diff --git a/examples/04_training/01_train_dynedge.py b/examples/04_training/01_train_dynedge.py index cf45f11a3..a2cc886e4 100644 --- a/examples/04_training/01_train_dynedge.py +++ b/examples/04_training/01_train_dynedge.py @@ -24,7 +24,7 @@ # Constants features = FEATURES.PROMETHEUS -truth = TRUTH.PROMETHEUS +truth = TRUTH.PROMETHEUS_LEGACY def main( diff --git a/examples/04_training/02_train_tito_model.py b/examples/04_training/02_train_tito_model.py index 4dfac2f49..8a96605b3 100644 --- a/examples/04_training/02_train_tito_model.py +++ b/examples/04_training/02_train_tito_model.py @@ -26,7 +26,7 @@ # Constants features = FEATURES.PROMETHEUS -truth = TRUTH.PROMETHEUS +truth = TRUTH.PROMETHEUS_LEGACY def main( diff --git a/examples/04_training/05_train_RNN_TITO.py b/examples/04_training/05_train_RNN_TITO.py index 948d4a3c8..0eee3063b 100644 --- a/examples/04_training/05_train_RNN_TITO.py +++ b/examples/04_training/05_train_RNN_TITO.py @@ -30,7 +30,7 @@ # Constants features = FEATURES.PROMETHEUS -truth = TRUTH.PROMETHEUS +truth = TRUTH.PROMETHEUS_LEGACY def main( diff --git a/examples/04_training/06_train_icemix_model.py b/examples/04_training/06_train_icemix_model.py index b76c91aa4..7395f7c8f 100644 --- a/examples/04_training/06_train_icemix_model.py +++ b/examples/04_training/06_train_icemix_model.py @@ -32,7 +32,7 @@ # Constants features = FEATURES.PROMETHEUS -truth = TRUTH.PROMETHEUS +truth = TRUTH.PROMETHEUS_LEGACY def main( diff --git a/examples/04_training/07_train_normalizing_flow.py b/examples/04_training/07_train_normalizing_flow.py index b78b34607..56bf0ff35 100644 --- a/examples/04_training/07_train_normalizing_flow.py +++ b/examples/04_training/07_train_normalizing_flow.py @@ -30,7 +30,7 @@ # Constants features = FEATURES.PROMETHEUS -truth = TRUTH.PROMETHEUS +truth = TRUTH.PROMETHEUS_LEGACY def main( diff --git a/examples/04_training/08_train_grit_model.py b/examples/04_training/08_train_grit_model.py index 24a5394a2..da18d99d8 100644 --- a/examples/04_training/08_train_grit_model.py +++ b/examples/04_training/08_train_grit_model.py @@ -24,7 +24,7 @@ # Constants features = FEATURES.PROMETHEUS -truth = TRUTH.PROMETHEUS +truth = TRUTH.PROMETHEUS_LEGACY def main( diff --git a/src/graphnet/data/constants.py b/src/graphnet/data/constants.py index 702974da9..3b8736e0b 100644 --- a/src/graphnet/data/constants.py +++ b/src/graphnet/data/constants.py @@ -87,7 +87,21 @@ class TRUTH: ] DEEPCORE = ICECUBE86 UPGRADE = DEEPCORE + # Truth schema written by current Prometheus versions; the default + # column list of `PrometheusTruthExtractor`. PROMETHEUS = [ + "interaction", + "initial_state_energy", + "initial_state_type", + "initial_state_zenith", + "initial_state_azimuth", + "initial_state_x", + "initial_state_y", + "initial_state_z", + ] + # Truth schema written by older Prometheus versions, e.g. the bundled + # example database `prometheus-events.db`. + PROMETHEUS_LEGACY = [ "injection_energy", "injection_type", "injection_interaction_type", diff --git a/src/graphnet/data/extractors/prometheus/prometheus_extractor.py b/src/graphnet/data/extractors/prometheus/prometheus_extractor.py index b9481d483..f3427a92d 100644 --- a/src/graphnet/data/extractors/prometheus/prometheus_extractor.py +++ b/src/graphnet/data/extractors/prometheus/prometheus_extractor.py @@ -1,9 +1,10 @@ """Parquet Extractor for conversion of simulation files from PROMETHEUS.""" -from typing import List +from typing import List, Optional import pandas as pd import numpy as np +from graphnet.data.constants import TRUTH from graphnet.data.extractors import Extractor @@ -52,23 +53,24 @@ class PrometheusTruthExtractor(PrometheusExtractor): This Extractor will "initial_state" i.e. neutrino truth. """ - def __init__(self, table_name: str = "mc_truth") -> None: + def __init__( + self, + table_name: str = "mc_truth", + columns: Optional[List[str]] = None, + ) -> None: """Construct PrometheusTruthExtractor. Args: table_name: Name of the table in the parquet files that contain event-level truth. Defaults to "mc_truth". + columns: Columns to extract from the table. Defaults to + `TRUTH.PROMETHEUS`, the truth fields written by every + Prometheus injection type. Override to include truth-level + data not present per default, e.g. LeptonInjector's + "bjorken_x", "bjorken_y" and "column_depth". """ - columns = [ - "interaction", - "initial_state_energy", - "initial_state_type", - "initial_state_zenith", - "initial_state_azimuth", - "initial_state_x", - "initial_state_y", - "initial_state_z", - ] + if columns is None: + columns = list(TRUTH.PROMETHEUS) super().__init__(extractor_name=table_name, columns=columns) diff --git a/src/graphnet/data/readers/prometheus_reader.py b/src/graphnet/data/readers/prometheus_reader.py index 3f97c4468..3bf2cc0e0 100644 --- a/src/graphnet/data/readers/prometheus_reader.py +++ b/src/graphnet/data/readers/prometheus_reader.py @@ -22,17 +22,31 @@ def __call__(self, file_path: str) -> List[OrderedDict]: Returns: Extracted data. + + Raises: + ValueError: If a table required by one of the configured + extractors does not exist in the file. """ # Open file outputs = [] file = pd.read_parquet(file_path) + extractors: List[PrometheusExtractor] = [] + missing_tables = [] + for extractor in self._extractors: + assert isinstance(extractor, PrometheusExtractor) + extractors.append(extractor) + if extractor._table not in file.columns: + missing_tables.append(extractor._table) + if missing_tables: + raise ValueError( + f"Table(s) {missing_tables} not found in {file_path}. " + f"Available tables: {list(file.columns)}." + ) for k in range(len(file)): # Loop over events in file extracted_event = OrderedDict() - for extractor in self._extractors: - assert isinstance(extractor, PrometheusExtractor) - if extractor._table in file.columns: - output = extractor(file[extractor._table][k]) - extracted_event[extractor._extractor_name] = output + for extractor in extractors: + output = extractor(file[extractor._table][k]) + extracted_event[extractor._extractor_name] = output outputs.append(extracted_event) return outputs diff --git a/src/graphnet/datasets/prometheus_datasets.py b/src/graphnet/datasets/prometheus_datasets.py index 8dc028278..3dcc555b1 100644 --- a/src/graphnet/datasets/prometheus_datasets.py +++ b/src/graphnet/datasets/prometheus_datasets.py @@ -8,7 +8,7 @@ from graphnet.training.labels import Direction, Track from graphnet.data import ERDAHostedDataset -from graphnet.data.constants import FEATURES +from graphnet.data.constants import FEATURES, TRUTH from graphnet.data.utilities import query_database @@ -18,16 +18,7 @@ class PublicPrometheusDataset(ERDAHostedDataset): # Static Member Variables: _pulsemaps = ["photons"] _truth_table = "mc_truth" - _event_truth = [ - "interaction", - "initial_state_energy", - "initial_state_type", - "initial_state_zenith", - "initial_state_azimuth", - "initial_state_x", - "initial_state_y", - "initial_state_z", - ] + _event_truth = TRUTH.PROMETHEUS _pulse_truth = None _features = FEATURES.PROMETHEUS diff --git a/tests/data/test_datamodule.py b/tests/data/test_datamodule.py index f35784644..1d159af97 100644 --- a/tests/data/test_datamodule.py +++ b/tests/data/test_datamodule.py @@ -85,7 +85,7 @@ def dataset_setup(dataset_ref: pytest.FixtureRequest) -> tuple: dataset_kwargs = { "truth_table": "mc_truth", "pulsemaps": "total", - "truth": TRUTH.PROMETHEUS, + "truth": TRUTH.PROMETHEUS_LEGACY, "features": FEATURES.PROMETHEUS, "path": data_path, "graph_definition": graph_definition, diff --git a/tests/data/test_prometheus_reader.py b/tests/data/test_prometheus_reader.py new file mode 100644 index 000000000..12d2f96da --- /dev/null +++ b/tests/data/test_prometheus_reader.py @@ -0,0 +1,55 @@ +"""Tests for the PrometheusReader.""" + +import os +from typing import List + +import pytest + +from graphnet.constants import TEST_DATA_DIR +from graphnet.data.constants import TRUTH +from graphnet.data.extractors.prometheus import ( + PrometheusExtractor, + PrometheusFeatureExtractor, + PrometheusTruthExtractor, +) +from graphnet.data.readers import PrometheusReader + +FILE_PATH = os.path.join( + TEST_DATA_DIR, "prometheus", "22980001_photons.parquet" +) + + +def test_prometheus_reader_extracts_configured_tables() -> None: + """Reader extracts every configured table when all are present.""" + reader = PrometheusReader() + reader.set_extractors( + [PrometheusTruthExtractor(), PrometheusFeatureExtractor()] + ) + events = reader(FILE_PATH) + assert len(events) > 0 + assert set(events[0].keys()) == {"mc_truth", "photons"} + assert set(TRUTH.PROMETHEUS) <= set(events[0]["mc_truth"].keys()) + + +def test_prometheus_truth_extractor_columns_override() -> None: + """Truth extractor extracts only the overridden columns.""" + reader = PrometheusReader() + extractors: List[PrometheusExtractor] = [ + PrometheusTruthExtractor(columns=["initial_state_energy"]) + ] + reader.set_extractors(extractors) + events = reader(FILE_PATH) + assert list(events[0]["mc_truth"].keys()) == ["initial_state_energy"] + + +def test_prometheus_reader_raises_on_missing_table() -> None: + """Reader raises if an extractor's table is not in the file.""" + reader = PrometheusReader() + reader.set_extractors( + [ + PrometheusTruthExtractor(), + PrometheusFeatureExtractor(table_name="not_a_table"), + ] + ) + with pytest.raises(ValueError, match="not_a_table"): + reader(FILE_PATH)