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.
282 lines
9.9 KiB
Python
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()
|