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

244 lines
8.4 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 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)
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())
manifest["snapshot_id"] = output_dir.name
(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()