"""Convert the Steps tree from a ROOT file to Parquet. See `uv run dwarf convert --help` for the CLI. """ from pathlib import Path from typing import Literal, cast import awkward as ak import polars as pl import uproot ParquetCompression = Literal["lz4", "uncompressed", "snappy", "gzip", "brotli", "zstd"] def _add_secondary_attributes(df: pl.DataFrame) -> tuple[pl.DataFrame, int]: """Add per-step secondary attributes via the parent→child track join. For each step that spawns secondaries, collects each child track's birth state (from the child track's first step in the same event) and emits: e_sec float64 — total secondary energy (sum of child first-step pre_E) sec_E_list list[f64] — per-secondary energy, sorted descending sec_pdg_list list[i32] — per-secondary PDG code, same order sec_dx_list list[f64] — per-secondary birth direction x, same order sec_dy_list list[f64] — per-secondary birth direction y, same order sec_dz_list list[f64] — per-secondary birth direction z, same order Steps with no children get 0.0 / empty lists. The full event must be present in `df` (it is — the writer concatenates before calling this). A listed child_track_id can fail to match any row in `first_step` — the child track never took a recorded step (e.g. absorbed below the tracking threshold at birth). Such orphans carry no physical secondary data, so they're dropped from child_track_ids/sec_*_list rather than left as nulls: a null in a float32 list silently becomes NaN once the parquet round-trips through the loader (`giant/data/loader.py:_pad_list_col`), and that NaN poisons every later secondary slot in the same step via the cumulative-sum "remaining budget" in `encode_secondaries`. Returns (df, n_orphaned) — the caller uses the count to report/aggregate across files rather than relying solely on the printed message here. """ first_step = ( df.sort("step_no") .group_by(["event_id", "track_id"]) .agg( pl.col("pre_E").first().alias("child_E"), pl.col("pdg").first().alias("child_pdg"), pl.col("pre_dx").first().alias("child_dx"), pl.col("pre_dy").first().alias("child_dy"), pl.col("pre_dz").first().alias("child_dz"), ) .rename({"track_id": "child_track_id"}) ) child_track_id_dtype = cast(pl.List, df.schema["child_track_ids"]).inner exploded = ( df.select(["event_id", "child_track_ids"]) .with_row_index("_step_row") .explode("child_track_ids") .rename({"child_track_ids": "child_track_id"}) .drop_nulls("child_track_id") ) joined = exploded.join(first_step, on=["event_id", "child_track_id"], how="left") n_orphaned = joined["child_E"].null_count() if n_orphaned: print( f" dropping {n_orphaned} orphaned child_track_id(s) with no " "recorded first step (absorbed below tracking threshold?)" ) joined = joined.drop_nulls("child_E") # Sort each step's secondaries by descending energy, then aggregate into lists per_step = ( joined.sort("child_E", descending=True) .group_by("_step_row") .agg( pl.col("child_track_id").alias("child_track_ids"), pl.col("child_E").sum().alias("e_sec"), pl.col("child_E").alias("sec_E_list"), pl.col("child_pdg").alias("sec_pdg_list"), pl.col("child_dx").alias("sec_dx_list"), pl.col("child_dy").alias("sec_dy_list"), pl.col("child_dz").alias("sec_dz_list"), ) ) empty_list_f64 = pl.Series("x", [[]], dtype=pl.List(pl.Float64)) empty_list_i32 = pl.Series("x", [[]], dtype=pl.List(pl.Int32)) empty_list_child_id = pl.Series("x", [[]], dtype=pl.List(child_track_id_dtype)) out = ( df.drop("child_track_ids") .with_row_index("_step_row") .join(per_step, on="_step_row", how="left") .with_columns( pl.col("child_track_ids").fill_null(empty_list_child_id), pl.col("e_sec").fill_null(0.0).cast(pl.Float64), pl.col("sec_E_list").fill_null(empty_list_f64), pl.col("sec_pdg_list").fill_null(empty_list_i32), pl.col("sec_dx_list").fill_null(empty_list_f64), pl.col("sec_dy_list").fill_null(empty_list_f64), pl.col("sec_dz_list").fill_null(empty_list_f64), ) .drop("_step_row") ) return out, n_orphaned def _batch_to_polars(batch: ak.Array) -> pl.DataFrame: """Convert one awkward-array batch to a Polars DataFrame. Flat numeric/string fields are converted via numpy; variable-length fields (like child_track_ids) fall back to Python lists so polars stores them as List columns — a type parquet understands natively. """ col_dict: dict = {} for field in ak.fields(batch): arr = batch[field] if arr.ndim == 1 and not isinstance(arr.layout, ak.contents.ListOffsetArray): col_dict[field] = ak.to_numpy(arr) else: col_dict[field] = ak.to_list(arr) return pl.DataFrame(col_dict) def convert_steps_to_parquet( root_path: str | Path, output_path: str | Path | None = None, batch_size: str = "100 MB", tree_name: str = "Steps", compression: ParquetCompression = "snappy", ) -> tuple[Path, int]: """Read *tree_name* from *root_path* and write it to a Parquet file. Reads in batches of *batch_size* so that peak ROOT-deserialization memory stays bounded. All batches are collected as Polars DataFrames and written in a single pass at the end (Polars' parquet writer does not support row-group appending without pyarrow). Parameters ---------- root_path: Input ROOT file. output_path: Output Parquet file. Defaults to *root_path* with .parquet suffix. batch_size: Uproot read batch size — an uproot size string ("100 MB") or integer row count (500_000). tree_name: Name of the TTree inside the ROOT file. compression: Parquet compression codec (snappy | lz4 | zstd | gzip | none). Returns (output_path, n_orphaned) — n_orphaned is the count of dropped orphaned child_track_ids (see `_add_secondary_attributes`), 0 if the tree has no child_track_ids column at all. Callers converting many files use it to aggregate a total instead of grepping the printed per-file message. """ root_path = Path(root_path) if output_path is None: output_path = root_path.with_suffix(".parquet") else: output_path = Path(output_path) with uproot.open(root_path) as f: tree = f[tree_name] n_entries = tree.num_entries print(f"Reading '{tree_name}' from {root_path.name} ({n_entries} entries)") batches: list[pl.DataFrame] = [] rows_done = 0 for batch in tree.iterate(library="ak", step_size=batch_size): batches.append(_batch_to_polars(batch)) rows_done += len(batch) print(f" {rows_done:,} / {n_entries:,} rows read", end="\r", flush=True) df = pl.concat(batches) # Steps tree carries the parent→child links needed to derive secondary energy; # other trees (e.g. Hits) don't, so only augment when the column is present. n_orphaned = 0 if "child_track_ids" in df.columns: print("\nComputing per-step secondary attributes …", end=" ", flush=True) df, n_orphaned = _add_secondary_attributes(df) print(f"\nWriting {output_path} …", end=" ", flush=True) df.write_parquet(output_path, compression=compression) print(f"done ({output_path.stat().st_size / 1e6:.1f} MB)") return output_path, n_orphaned