Handle material column as string type

material is a string literal (e.g. "G4_PbWO4"), not an integer. Store as
object array and key mat_map on str throughout loader and transforms.

Co-Authored-By: Claude Sonnet 4.6 <noreply@anthropic.com>
This commit is contained in:
2026-06-17 14:29:11 +02:00
parent 78c2a61ecd
commit 877f83a5cd
2 changed files with 11 additions and 11 deletions
+7 -7
View File
@@ -23,7 +23,7 @@ def _df_to_dict(df: pd.DataFrame) -> dict[str, np.ndarray]:
"pre_pos": df[["pre_x", "pre_y", "pre_z"]].to_numpy(dtype=np.float32), "pre_pos": df[["pre_x", "pre_y", "pre_z"]].to_numpy(dtype=np.float32),
"pre_energy": df["pre_energy"].to_numpy(dtype=np.float32), "pre_energy": df["pre_energy"].to_numpy(dtype=np.float32),
"pre_dir": df[["pre_dir_x", "pre_dir_y", "pre_dir_z"]].to_numpy(dtype=np.float32), "pre_dir": df[["pre_dir_x", "pre_dir_y", "pre_dir_z"]].to_numpy(dtype=np.float32),
"material": df["material"].to_numpy(dtype=np.int32), "material": df["material"].to_numpy(dtype=object),
"layer_id": df["layer_id"].to_numpy(dtype=np.int32), "layer_id": df["layer_id"].to_numpy(dtype=np.int32),
"n_sec": df["child_track_ids"].apply(len).to_numpy(dtype=np.int32), "n_sec": df["child_track_ids"].apply(len).to_numpy(dtype=np.int32),
"step_length": df["step_length"].to_numpy(dtype=np.float32), "step_length": df["step_length"].to_numpy(dtype=np.float32),
@@ -64,7 +64,7 @@ def _cond_df_to_dict(df: pd.DataFrame) -> dict[str, np.ndarray]:
"pre_pos": df[["pre_x", "pre_y", "pre_z"]].to_numpy(dtype=np.float32), "pre_pos": df[["pre_x", "pre_y", "pre_z"]].to_numpy(dtype=np.float32),
"pre_energy": df["pre_energy"].to_numpy(dtype=np.float32), "pre_energy": df["pre_energy"].to_numpy(dtype=np.float32),
"pre_dir": df[["pre_dir_x", "pre_dir_y", "pre_dir_z"]].to_numpy(dtype=np.float32), "pre_dir": df[["pre_dir_x", "pre_dir_y", "pre_dir_z"]].to_numpy(dtype=np.float32),
"material": df["material"].to_numpy(dtype=np.int32), "material": df["material"].to_numpy(dtype=object),
"layer_id": df["layer_id"].to_numpy(dtype=np.int32), "layer_id": df["layer_id"].to_numpy(dtype=np.int32),
"n_sec": df["child_track_ids"].apply(len).to_numpy(dtype=np.int32), "n_sec": df["child_track_ids"].apply(len).to_numpy(dtype=np.int32),
} }
@@ -79,9 +79,9 @@ def iter_cond_chunks(path: str | Path) -> Iterator[dict[str, np.ndarray]]:
def build_index_maps( def build_index_maps(
data: dict[str, np.ndarray], data: dict[str, np.ndarray],
) -> tuple[dict[int, int], dict[int, int]]: ) -> tuple[dict[int, int], dict[str, int]]:
pdg_vals = sorted(int(v) for v in np.unique(data["pdg"])) pdg_vals = sorted(int(v) for v in np.unique(data["pdg"]))
mat_vals = sorted(int(v) for v in np.unique(data["material"])) mat_vals = sorted(str(v) for v in np.unique(data["material"]))
return ( return (
{v: i for i, v in enumerate(pdg_vals)}, {v: i for i, v in enumerate(pdg_vals)},
{v: i for i, v in enumerate(mat_vals)}, {v: i for i, v in enumerate(mat_vals)},
@@ -90,14 +90,14 @@ def build_index_maps(
def build_index_maps_from_files( def build_index_maps_from_files(
files: list[Path], files: list[Path],
) -> tuple[dict[int, int], dict[int, int]]: ) -> tuple[dict[int, int], dict[str, int]]:
"""Scan only pdg and material columns across all files (2-column read).""" """Scan only pdg and material columns across all files (2-column read)."""
pdg_vals: set[int] = set() pdg_vals: set[int] = set()
mat_vals: set[int] = set() mat_vals: set[str] = set()
for path in files: for path in files:
df = pd.read_parquet(path, columns=["pdg", "material"]) df = pd.read_parquet(path, columns=["pdg", "material"])
pdg_vals.update(int(v) for v in df["pdg"].unique()) pdg_vals.update(int(v) for v in df["pdg"].unique())
mat_vals.update(int(v) for v in df["material"].unique()) mat_vals.update(str(v) for v in df["material"].unique())
return ( return (
{v: i for i, v in enumerate(sorted(pdg_vals))}, {v: i for i, v in enumerate(sorted(pdg_vals))},
{v: i for i, v in enumerate(sorted(mat_vals))}, {v: i for i, v in enumerate(sorted(mat_vals))},
+4 -4
View File
@@ -123,7 +123,7 @@ def inv_local_frame_rotation(pre_dir: np.ndarray, post_dir_local: np.ndarray) ->
def build_cond_features( def build_cond_features(
data: dict[str, np.ndarray], data: dict[str, np.ndarray],
pdg_map: dict[int, int], pdg_map: dict[int, int],
mat_map: dict[int, int], mat_map: dict[str, int],
cond_normalizer: "Normalizer | None" = None, cond_normalizer: "Normalizer | None" = None,
) -> tuple[np.ndarray, np.ndarray]: ) -> tuple[np.ndarray, np.ndarray]:
"""Build conditioning arrays only — no target, no post-step variables.""" """Build conditioning arrays only — no target, no post-step variables."""
@@ -136,7 +136,7 @@ def build_cond_features(
]).astype(np.float32) ]).astype(np.float32)
pdg_idx = np.array([pdg_map[int(p)] for p in data["pdg"]], dtype=np.int64) pdg_idx = np.array([pdg_map[int(p)] for p in data["pdg"]], dtype=np.int64)
mat_idx = np.array([mat_map[int(m)] for m in data["material"]], dtype=np.int64) mat_idx = np.array([mat_map[str(m)] for m in data["material"]], dtype=np.int64)
cond_cat = np.column_stack([pdg_idx, mat_idx]) cond_cat = np.column_stack([pdg_idx, mat_idx])
if cond_normalizer is not None: if cond_normalizer is not None:
@@ -148,7 +148,7 @@ def build_cond_features(
def build_features( def build_features(
data: dict[str, np.ndarray], data: dict[str, np.ndarray],
pdg_map: dict[int, int], pdg_map: dict[int, int],
mat_map: dict[int, int], mat_map: dict[str, int],
cond_normalizer: Normalizer | None = None, cond_normalizer: Normalizer | None = None,
target_normalizer: Normalizer | None = None, target_normalizer: Normalizer | None = None,
fit: bool = False, fit: bool = False,
@@ -175,7 +175,7 @@ def build_features(
]).astype(np.float32) # (N, 9) ]).astype(np.float32) # (N, 9)
pdg_idx = np.array([pdg_map[int(p)] for p in data["pdg"]], dtype=np.int64) pdg_idx = np.array([pdg_map[int(p)] for p in data["pdg"]], dtype=np.int64)
mat_idx = np.array([mat_map[int(m)] for m in data["material"]], dtype=np.int64) mat_idx = np.array([mat_map[str(m)] for m in data["material"]], dtype=np.int64)
cond_cat = np.column_stack([pdg_idx, mat_idx]) # (N, 2) cond_cat = np.column_stack([pdg_idx, mat_idx]) # (N, 2)
if fit: if fit: