"""Enrich snapshot with PIT-safe trailing catalyst persistence features. This derives simple trailing event-density features from the snapshot rows themselves, using only events that happened strictly before the current row. All splits are loaded together so valid/test rows can safely use earlier train-period events. Usage: uv run python3 scripts/enrich_catalyst_persistence_features.py \ --input data/parquet/midlarge-liquid-long-v1_bucketfix_full_audit_canonical \ --output data/parquet/midlarge-liquid-long-v1_bucketfix_full_audit_canonical_cp """ from __future__ import annotations import argparse import datetime as dt import json from collections import defaultdict, deque from pathlib import Path from typing import Any import pandas as pd import pyarrow as pa import pyarrow.parquet as pq LOOKBACK_WINDOWS = (20, 60) PERSISTENCE_FEATURES = [ "prior_catalyst_count_20d", "prior_catalyst_count_60d", "prior_catalyst_type_diversity_20d", "prior_catalyst_type_diversity_60d", ] def _parse_date(value: Any) -> dt.date | None: if value is None: return None if isinstance(value, dt.datetime): return value.date() if isinstance(value, dt.date): return value text = str(value).strip() if not text: return None try: return dt.date.fromisoformat(text[:10]) except ValueError: return None def _load_all_rows(input_dir: Path) -> pd.DataFrame: frames: list[pd.DataFrame] = [] split_order = {"train": 0, "valid": 1, "test": 2} for split in ("train", "valid", "test"): parquet_path = input_dir / f"{split}.parquet" if not parquet_path.exists(): continue df = pq.read_table(parquet_path).to_pandas() df["__split_name"] = split df["__split_order"] = split_order[split] df["__split_row_idx"] = range(len(df)) frames.append(df) if not frames: return pd.DataFrame() combined = pd.concat(frames, ignore_index=True) combined["__event_date_obj"] = combined["event_date"].map(_parse_date) combined = combined.sort_values( by=["ticker", "__event_date_obj", "__split_order", "__split_row_idx"], kind="stable", na_position="last", ).reset_index(drop=True) return combined def compute_catalyst_persistence_features(combined: pd.DataFrame) -> pd.DataFrame: if combined.empty: out = combined.copy() for feature in PERSISTENCE_FEATURES: out[feature] = pd.Series(dtype="float64") return out queues: dict[int, dict[str, deque[tuple[dt.date, str]]]] = { window: defaultdict(deque) for window in LOOKBACK_WINDOWS } feature_values = {feature: [None] * len(combined) for feature in PERSISTENCE_FEATURES} for idx, row in combined.iterrows(): ticker = str(row.get("ticker") or "").upper() event_date = row.get("__event_date_obj") event_type = str(row.get("event_type") or "UNKNOWN") or "UNKNOWN" if not ticker or event_date is None: continue for window in LOOKBACK_WINDOWS: dq = queues[window][ticker] cutoff = event_date - dt.timedelta(days=window) while dq and dq[0][0] < cutoff: dq.popleft() feature_values[f"prior_catalyst_count_{window}d"][idx] = float(len(dq)) feature_values[f"prior_catalyst_type_diversity_{window}d"][idx] = float( len({prior_type for _, prior_type in dq}) ) for window in LOOKBACK_WINDOWS: queues[window][ticker].append((event_date, event_type)) out = combined.copy() for feature, values in feature_values.items(): out[feature] = values return out def enrich_snapshot_dir(input_dir: Path, output_dir: Path) -> None: output_dir.mkdir(parents=True, exist_ok=True) combined = _load_all_rows(input_dir) if combined.empty: manifest_src = input_dir / "manifest.json" if manifest_src.exists(): output_dir.joinpath("manifest.json").write_text(manifest_src.read_text()) return enriched = compute_catalyst_persistence_features(combined) manifest_src = input_dir / "manifest.json" if manifest_src.exists(): manifest = json.loads(manifest_src.read_text()) source_snapshot_id = manifest.get("snapshot_id") manifest["snapshot_id"] = output_dir.name manifest["output_dir"] = str(output_dir.resolve()) manifest["created_at_utc"] = dt.datetime.now(dt.UTC).isoformat() manifest["enrichment_source_snapshot_id"] = source_snapshot_id manifest["export_enrichments"] = sorted( set((manifest.get("export_enrichments") or []) + PERSISTENCE_FEATURES) ) output_dir.joinpath("manifest.json").write_text(json.dumps(manifest, indent=2)) for split in ("train", "valid", "test"): split_df = enriched.loc[enriched["__split_name"] == split].copy() if split_df.empty: continue split_df = split_df.drop( columns=["__split_name", "__split_order", "__split_row_idx", "__event_date_obj"], errors="ignore", ) split_table = pa.Table.from_pandas(split_df, preserve_index=False) pq.write_table(split_table, output_dir / f"{split}.parquet") def main() -> None: parser = argparse.ArgumentParser() parser.add_argument("--input", required=True) parser.add_argument("--output", required=True) args = parser.parse_args() enrich_snapshot_dir(Path(args.input), Path(args.output)) if __name__ == "__main__": main()