You cannot select more than 25 topics Topics must start with a letter or number, can include dashes ('-') and can be up to 35 characters long.
fithia2/scripts/enrich_peer_surprise_featur...

282 lines
9.9 KiB
Python

"""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()