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_tier2_features.py

254 lines
8.9 KiB
Python

"""Enrich snapshot with Tier 2 features: FINRA short ratio, Hurst, Entropy, Sector momentum.
Usage:
PYTHONUNBUFFERED=1 uv run python3 scripts/enrich_tier2_features.py \
--input data/parquet/midlarge-liquid-long-v1_bucketfix_full_audit \
--output data/parquet/midlarge-liquid-long-v1_bucketfix_full_audit_tier2
"""
from __future__ import annotations
import argparse
import datetime as dt
import json
import sys
import time
from pathlib import Path
import pyarrow as pa
import pyarrow.parquet as pq
import requests
sys.path.insert(0, str(Path(__file__).resolve().parent.parent))
from libs.features.market_features import pre_event_entropy, pre_event_hurst
from libs.oracle_client.models import PriceBar
ORACLE_URL = "http://localhost:18001"
FEATURES = ["pre_event_hurst_60d", "pre_event_entropy_60d", "pre_event_short_ratio", "pre_event_sector_momentum_20d"]
SECTOR_ETF = {
"Technology": "XLK", "Health Care": "XLV", "Financials": "XLF",
"Consumer Discretionary": "XLY", "Industrials": "XLI",
"Energy": "XLE", "Utilities": "XLU", "Real Estate": "XLRE",
"Materials": "XLB", "Communication Services": "XLC",
"Consumer Staples": "XLP", "Healthcare": "XLV",
}
# Cache sector ETF bars and SPY bars
_sector_etf_cache: dict[str, list[dict]] = {}
_spy_cache: list[dict] = []
def _fetch_bars_raw(ticker: str, start: str, end: str) -> list[dict]:
resp = requests.get(f"{ORACLE_URL}/api/v1/price/data/{ticker}",
params={"start_date": start, "end_date": end}, timeout=30)
if resp.status_code != 200:
return []
data = resp.json()
return data.get("bars") or data.get("data") or []
def fetch_bars(ticker: str, event_date: str) -> list[PriceBar]:
from datetime import datetime, timedelta
end_dt = datetime.strptime(event_date, "%Y-%m-%d")
start_dt = end_dt - timedelta(days=120)
bars_raw = _fetch_bars_raw(ticker, start_dt.strftime("%Y-%m-%d"), event_date)
bars = []
for b in bars_raw:
try:
bars.append(PriceBar(date=b["date"], open=float(b.get("open", 0)),
high=float(b.get("high", 0)), low=float(b.get("low", 0)),
close=float(b.get("close", 0)), volume=int(b.get("volume", 0))))
except (KeyError, ValueError, TypeError):
continue
return bars
def fetch_short_ratio(ticker: str, event_date: str) -> float | None:
"""Fetch average short ratio over 5 days before event."""
resp = requests.get(f"{ORACLE_URL}/api/v1/finra/short-ratio/{ticker}",
params={"days": 30}, timeout=30)
if resp.status_code != 200:
return None
data = resp.json()
points = data.get("data") or []
if not points:
return None
# Filter to dates before event
before = [p for p in points if p.get("date", "") < event_date and p.get("short_percent") is not None]
before.sort(key=lambda p: p["date"], reverse=True)
recent = before[:5]
if not recent:
return None
return sum(p["short_percent"] for p in recent) / len(recent)
def fetch_sector_momentum(sector: str, event_date: str) -> float | None:
"""Compute sector ETF 20d return minus SPY 20d return."""
from datetime import datetime, timedelta
etf = SECTOR_ETF.get(sector)
if not etf:
return None
end_dt = datetime.strptime(event_date, "%Y-%m-%d")
start_dt = end_dt - timedelta(days=45)
start_str = start_dt.strftime("%Y-%m-%d")
# Fetch sector ETF (cached)
cache_key = f"{etf}_{start_str}_{event_date}"
if cache_key not in _sector_etf_cache:
_sector_etf_cache[cache_key] = _fetch_bars_raw(etf, start_str, event_date)
etf_bars = _sector_etf_cache[cache_key]
etf_dated = {b["date"]: float(b.get("close", 0)) for b in etf_bars if b.get("close")}
etf_dates = sorted(etf_dated.keys())
if event_date not in etf_dated or len(etf_dates) < 21:
return None
idx = etf_dates.index(event_date)
if idx < 20:
return None
etf_ret = (etf_dated[etf_dates[idx]] - etf_dated[etf_dates[idx - 20]]) / etf_dated[etf_dates[idx - 20]]
# Fetch SPY (cached)
spy_key = f"SPY_{start_str}_{event_date}"
if spy_key not in _sector_etf_cache:
_sector_etf_cache[spy_key] = _fetch_bars_raw("SPY", start_str, event_date)
spy_bars = _sector_etf_cache[spy_key]
spy_dated = {b["date"]: float(b.get("close", 0)) for b in spy_bars if b.get("close")}
spy_dates = sorted(spy_dated.keys())
if event_date not in spy_dated:
return None
spy_idx = spy_dates.index(event_date)
if spy_idx < 20:
return None
spy_ret = (spy_dated[spy_dates[spy_idx]] - spy_dated[spy_dates[spy_idx - 20]]) / spy_dated[spy_dates[spy_idx - 20]]
return etf_ret - spy_ret
def compute_features_for_event(ticker: str, event_date: str, sector: str) -> dict[str, float | None]:
bars = fetch_bars(ticker, event_date)
result: dict[str, float | None] = {}
# Hurst & Entropy from price bars
result["pre_event_hurst_60d"] = pre_event_hurst(bars, event_date, 60) if bars else None
result["pre_event_entropy_60d"] = pre_event_entropy(bars, event_date, 60) if bars else None
# FINRA short ratio
try:
result["pre_event_short_ratio"] = fetch_short_ratio(ticker, event_date)
except Exception:
result["pre_event_short_ratio"] = None
# Sector momentum
try:
result["pre_event_sector_momentum_20d"] = fetch_sector_momentum(sector, event_date)
except Exception:
result["pre_event_sector_momentum_20d"] = None
return result
def enrich_split(input_path: Path, output_path: Path):
table = pq.read_table(input_path)
n = len(table)
if n == 0 or "ticker" not in table.column_names or "event_date" not in table.column_names:
print(" Empty or schema-less split, copying")
pq.write_table(table, output_path)
return
existing_cols = set(table.column_names)
if all(f in existing_cols for f in FEATURES):
null_counts = {f: table.column(f).null_count for f in FEATURES}
if all(v < n * 0.3 for v in null_counts.values()):
print(f" Already enriched ({null_counts}), copying")
pq.write_table(table, output_path)
return
tickers = table.column("ticker").to_pylist()
event_dates = table.column("event_date").to_pylist()
# Get sector info if available
sectors = table.column("asset_type_proxy").to_pylist() if "asset_type_proxy" in table.column_names else ["UNKNOWN"] * n
# Actually sector might be elsewhere — check for common sector columns
for col_name in ["sector", "asset_type_proxy"]:
if col_name in table.column_names:
sectors = table.column(col_name).to_pylist()
break
results = {f: [None] * n for f in FEATURES}
success = 0
for i in range(n):
ticker = tickers[i]
event_date = str(event_dates[i])
sector = str(sectors[i]) if sectors[i] else "UNKNOWN"
if i % 100 == 0:
print(f" {i}/{n} ({success} ok)...")
try:
feats = compute_features_for_event(ticker, event_date, sector)
for f in FEATURES:
results[f][i] = feats.get(f)
if feats.get("pre_event_hurst_60d") is not None:
success += 1
except Exception as e:
if i < 5:
print(f" Error {ticker} {event_date}: {e}")
print(f" Done: {success}/{n}")
for f in FEATURES:
arr = pa.array(results[f], type=pa.float64())
if f in existing_cols:
idx = table.column_names.index(f)
table = table.set_column(idx, f, arr)
else:
table = table.append_column(f, arr)
pq.write_table(table, output_path)
print(f" Written to {output_path}")
def main():
parser = argparse.ArgumentParser()
parser.add_argument("--input", required=True)
parser.add_argument("--output", required=True)
args = parser.parse_args()
input_dir = Path(args.input)
output_dir = Path(args.output)
output_dir.mkdir(parents=True, exist_ok=True)
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 []) + FEATURES))
(output_dir / "manifest.json").write_text(json.dumps(manifest, indent=2))
for split in ["train", "valid", "test"]:
input_path = input_dir / f"{split}.parquet"
if not input_path.exists():
continue
output_path = output_dir / f"{split}.parquet"
print(f"Enriching {split}...")
enrich_split(input_path, output_path)
print("All done!")
if __name__ == "__main__":
main()