"""Enrich snapshot with PIT-safe peer-relative earnings surprise features. This derives same-sector peer baseline features from prior earnings events only. The script loads all snapshot splits together so valid/test rows can safely use earlier train-period peer events, then writes enriched split files back out. Usage: uv run python3 scripts/enrich_peer_surprise_features.py \ --input data/parquet/midlarge-liquid-long-v1_bucketfix_full_audit_canonical \ --output data/parquet/midlarge-liquid-long-v1_bucketfix_full_audit_canonical_peer """ from __future__ import annotations import argparse import datetime as dt import json import math from collections import defaultdict from dataclasses import dataclass from pathlib import Path from typing import Any import pandas as pd import pyarrow as pa import pyarrow.parquet as pq LOOKBACK_DAYS = 365 PEER_FEATURES = [ "sector", "peer_sector_event_count_365d", "peer_sector_surprise_median_365d", "peer_sector_surprise_mean_365d", "peer_sector_surprise_pos_rate_365d", "peer_relative_surprise_pct_365d", "peer_sector_sue_hist_mean_4q_median_365d", "peer_sector_sue_hist_mean_4q_mean_365d", "peer_relative_sue_hist_mean_4q_365d", "peer_sector_sue_hist_pos_rate_4q_mean_365d", ] @dataclass(slots=True) class PeerHistoryRow: event_date: dt.date ticker: str earnings_surprise_pct: float | None sue_hist_mean_4q: float | None sue_hist_pos_rate_4q: float | None def _sector_cache_path() -> Path: return Path("data/cache/sector_cache.json") def _load_sector_cache() -> dict[str, str]: path = _sector_cache_path() if not path.exists(): return {} try: raw = json.loads(path.read_text()) except Exception: return {} return {str(k).upper(): str(v) for k, v in raw.items() if v} 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 _to_float(value: Any) -> float | None: if value is None: return None try: out = float(value) except (TypeError, ValueError): return None if math.isnan(out): return None return out def _mean(values: list[float]) -> float | None: if not values: return None return sum(values) / len(values) def _median(values: list[float]) -> float | None: if not values: return None ordered = sorted(values) mid = len(ordered) // 2 if len(ordered) % 2 == 1: return ordered[mid] return (ordered[mid - 1] + ordered[mid]) / 2.0 def _positive_rate(values: list[float]) -> float | None: if not values: return None return sum(1 for value in values if value > 0.0) / len(values) def _load_all_rows(input_dir: Path, sector_cache: dict[str, str]) -> 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() if df.empty: df["sector"] = pd.Series(dtype="object") df["__split_name"] = split df["__split_order"] = split_order[split] df["__split_row_idx"] = range(len(df)) if "sector" not in df.columns: df["sector"] = pd.NA ticker_series = df.get("ticker", pd.Series([None] * len(df))) df["sector"] = [ str(existing).strip() if isinstance(existing, str) and existing.strip() else sector_cache.get(str(ticker or "").upper(), "UNKNOWN") for existing, ticker in zip(df["sector"], ticker_series, strict=False) ] 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=["__event_date_obj", "__split_order", "__split_row_idx"], kind="stable", na_position="last", ).reset_index(drop=True) return combined def compute_peer_features(combined: pd.DataFrame, lookback_days: int = LOOKBACK_DAYS) -> pd.DataFrame: if combined.empty: out = combined.copy() for feature in PEER_FEATURES: if feature not in out.columns: out[feature] = pd.NA return out histories: dict[str, list[PeerHistoryRow]] = defaultdict(list) feature_values = {feature: [None] * len(combined) for feature in PEER_FEATURES} for idx, row in combined.iterrows(): sector = str(row.get("sector") or "UNKNOWN") ticker = str(row.get("ticker") or "").upper() event_type = str(row.get("event_type") or "") event_date = row.get("__event_date_obj") current_surprise = _to_float(row.get("earnings_surprise_pct")) current_hist_mean = _to_float(row.get("sue_hist_mean_4q")) current_hist_pos_rate = _to_float(row.get("sue_hist_pos_rate_4q")) feature_values["sector"][idx] = sector if event_type == "earnings_release" and event_date is not None and sector and sector != "UNKNOWN": eligible_peers: list[PeerHistoryRow] = [] cutoff = event_date - dt.timedelta(days=lookback_days) for prior in histories.get(sector, []): if prior.event_date < cutoff: continue if prior.event_date >= event_date: continue if prior.ticker == ticker: continue eligible_peers.append(prior) surprise_values = [ value for value in (_to_float(item.earnings_surprise_pct) for item in eligible_peers) if value is not None ] hist_mean_values = [ value for value in (_to_float(item.sue_hist_mean_4q) for item in eligible_peers) if value is not None ] hist_pos_rate_values = [ value for value in (_to_float(item.sue_hist_pos_rate_4q) for item in eligible_peers) if value is not None ] surprise_median = _median(surprise_values) hist_mean_median = _median(hist_mean_values) feature_values["peer_sector_event_count_365d"][idx] = float(len(surprise_values)) feature_values["peer_sector_surprise_median_365d"][idx] = surprise_median feature_values["peer_sector_surprise_mean_365d"][idx] = _mean(surprise_values) feature_values["peer_sector_surprise_pos_rate_365d"][idx] = _positive_rate(surprise_values) feature_values["peer_relative_surprise_pct_365d"][idx] = ( current_surprise - surprise_median if current_surprise is not None and surprise_median is not None else None ) feature_values["peer_sector_sue_hist_mean_4q_median_365d"][idx] = hist_mean_median feature_values["peer_sector_sue_hist_mean_4q_mean_365d"][idx] = _mean(hist_mean_values) feature_values["peer_relative_sue_hist_mean_4q_365d"][idx] = ( current_hist_mean - hist_mean_median if current_hist_mean is not None and hist_mean_median is not None else None ) feature_values["peer_sector_sue_hist_pos_rate_4q_mean_365d"][idx] = _mean(hist_pos_rate_values) histories[sector].append( PeerHistoryRow( event_date=event_date, ticker=ticker, earnings_surprise_pct=current_surprise, sue_hist_mean_4q=current_hist_mean, sue_hist_pos_rate_4q=current_hist_pos_rate, ) ) 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) sector_cache = _load_sector_cache() combined = _load_all_rows(input_dir, sector_cache) 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_peer_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 []) + PEER_FEATURES)) output_dir.joinpath("manifest.json").write_text(json.dumps(manifest, indent=2)) for split in ("train", "valid", "test"): split_df = enriched[enriched["__split_name"] == split].copy() split_df = split_df.sort_values(by="__split_row_idx", kind="stable").drop( columns=["__split_name", "__split_order", "__split_row_idx", "__event_date_obj"], errors="ignore", ) table = pa.Table.from_pandas(split_df, preserve_index=False) pq.write_table(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()