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.
244 lines
8.4 KiB
Python
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()
|