294 lines
11 KiB
Python
294 lines
11 KiB
Python
"""Signal-quality gate walk-forward backtest.
|
|
|
|
Gates trades based on whether the model's recent topk predictions were correct
|
|
(hit rate). This is a retrospective gate — it measures prediction accuracy,
|
|
not market state.
|
|
|
|
Usage:
|
|
cd /app && .venv/bin/python book/scripts/signal_quality_gate_bt.py
|
|
"""
|
|
|
|
from __future__ import annotations
|
|
|
|
import json
|
|
import pathlib
|
|
import sys
|
|
import time
|
|
|
|
import numpy as np
|
|
import pandas as pd
|
|
|
|
LAKE_ROOT = "/home/data/lake"
|
|
MARKET = "US"
|
|
OUT_DIR = pathlib.Path("/app/experiments/book/data/signal_quality_gate")
|
|
|
|
WINDOWS = [
|
|
{"label": "2026", "start": "2026-01-04", "end": "2026-08-19",
|
|
"pred": f"{LAKE_ROOT}/mlruns/52/9f98ea5c550a409f87b56a6cd8fee343/artifacts/pred.pkl"},
|
|
{"label": "2025", "start": "2025-01-02", "end": "2025-12-31",
|
|
"pred": f"{LAKE_ROOT}/mlruns/52/fe96741654df4780957a3a949999ae6a/artifacts/pred.pkl"},
|
|
{"label": "2024", "start": "2024-01-02", "end": "2024-12-31",
|
|
"pred": f"{LAKE_ROOT}/mlruns/52/71ed5bfa9984490f8bba8b222f7acc39/artifacts/pred.pkl"},
|
|
{"label": "2023", "start": "2023-01-03", "end": "2023-12-29",
|
|
"pred": f"{LAKE_ROOT}/mlruns/56/8ca46e554311444c9a42637a788226e8/artifacts/pred.pkl"},
|
|
{"label": "2021", "start": "2021-01-04", "end": "2021-12-31",
|
|
"pred": f"{LAKE_ROOT}/mlruns/56/4e0700ddab2a4e108b46efece7346ee3/artifacts/pred.pkl"},
|
|
]
|
|
|
|
# Signal-quality gate configs: (lookback_days, threshold, name)
|
|
SIGNAL_GATE_CONFIGS = [
|
|
(5, 0.50, "hitrate_5d_0.50"),
|
|
(5, 0.60, "hitrate_5d_0.60"),
|
|
(5, 0.70, "hitrate_5d_0.70"),
|
|
(10, 0.50, "hitrate_10d_0.50"),
|
|
(10, 0.60, "hitrate_10d_0.60"),
|
|
(10, 0.70, "hitrate_10d_0.70"),
|
|
(20, 0.40, "hitrate_20d_0.40"),
|
|
(20, 0.50, "hitrate_20d_0.50"),
|
|
(20, 0.60, "hitrate_20d_0.60"),
|
|
]
|
|
|
|
|
|
def load_pred(path: str) -> pd.Series:
|
|
df = pd.read_pickle(path)
|
|
if isinstance(df, pd.DataFrame):
|
|
if "score" in df.columns:
|
|
s = df["score"]
|
|
else:
|
|
s = df.iloc[:, 0]
|
|
else:
|
|
s = df
|
|
idx = s.index
|
|
new_dt = pd.to_datetime(idx.get_level_values(0)).normalize()
|
|
s.index = pd.MultiIndex.from_arrays([new_dt, idx.get_level_values(1)], names=idx.names)
|
|
return s
|
|
|
|
|
|
def load_bars_for_window(start: str, end: str) -> pd.DataFrame:
|
|
from tac_qlib.data.config import LakeConfig, resolve_lake_root
|
|
cfg = LakeConfig(resolve_lake_root(LAKE_ROOT), MARKET)
|
|
sp = cfg.lake_root / "symbols.parquet"
|
|
if sp.exists():
|
|
syms = pd.read_parquet(sp)
|
|
col = "symbol" if "symbol" in syms.columns else syms.columns[0]
|
|
symbols = sorted(syms[col].astype(str).str.upper().tolist())
|
|
else:
|
|
return pd.DataFrame()
|
|
closes = {}
|
|
for sym in symbols:
|
|
p = cfg.bar_path("1d", sym)
|
|
if not p.exists():
|
|
continue
|
|
try:
|
|
df = pd.read_parquet(p)
|
|
except Exception:
|
|
continue
|
|
if not len(df):
|
|
continue
|
|
tcol = df["t"] if "t" in df.columns else df["date"]
|
|
ts = pd.to_datetime(tcol)
|
|
df = df.assign(_t=ts).set_index("_t").sort_index()
|
|
warmup_start = pd.Timestamp(start) - pd.Timedelta(days=60)
|
|
df = df.loc[warmup_start:end]
|
|
if len(df) >= 22:
|
|
closes[sym] = df["c"]
|
|
return pd.DataFrame(closes)
|
|
|
|
|
|
def compute_hit_rate_series(
|
|
pred: pd.Series, ret_df: pd.DataFrame, topk: int = 10, lookback: int = 10,
|
|
) -> pd.Series:
|
|
dt_idx = pred.index.get_level_values(0)
|
|
trade_dates = sorted(dt_idx.unique())
|
|
hit_rates = {}
|
|
for i in range(1, len(trade_dates)):
|
|
prev_date = trade_dates[i - 1]
|
|
curr_date = trade_dates[i]
|
|
try:
|
|
prev_scores = pred.loc[prev_date]
|
|
except KeyError:
|
|
continue
|
|
if isinstance(prev_scores, pd.DataFrame):
|
|
prev_scores = prev_scores.iloc[:, 0]
|
|
prev_scores = prev_scores.dropna().sort_values(ascending=False)
|
|
topk_syms = list(prev_scores.index[:topk])
|
|
if curr_date not in ret_df.index:
|
|
continue
|
|
today_ret = ret_df.loc[curr_date]
|
|
topk_rets = today_ret.reindex(topk_syms).dropna()
|
|
if len(topk_rets) > 0:
|
|
hit_rate = (topk_rets > 0).mean()
|
|
hit_rates[curr_date] = hit_rate
|
|
hit_series = pd.Series(hit_rates)
|
|
if len(hit_series) == 0:
|
|
return hit_series
|
|
rolling_hr = hit_series.rolling(lookback, min_periods=max(1, lookback // 2)).mean()
|
|
return rolling_hr
|
|
|
|
|
|
def run_backtest(pred, hit_rate, close_df, start, end, topk=10, threshold=0.5):
|
|
if not isinstance(pred.index, pd.MultiIndex):
|
|
return {"error": "pred must have MultiIndex"}
|
|
ret_df = close_df.pct_change()
|
|
ret_df.index = pd.to_datetime(ret_df.index).normalize()
|
|
dt_idx = pred.index.get_level_values(0)
|
|
window_mask = (dt_idx >= pd.Timestamp(start)) & (dt_idx <= pd.Timestamp(end))
|
|
window_pred = pred.loc[window_mask]
|
|
if len(window_pred) == 0:
|
|
return {"error": "no pred data in window"}
|
|
trade_dates = sorted(dt_idx[window_mask].unique())
|
|
gate_open = {}
|
|
for d in trade_dates:
|
|
known = hit_rate[hit_rate.index <= d]
|
|
if len(known) > 0 and not pd.isna(known.iloc[-1]):
|
|
gate_open[d] = bool(known.iloc[-1] >= threshold)
|
|
else:
|
|
gate_open[d] = True
|
|
n_total = len(trade_dates)
|
|
n_open = sum(1 for v in gate_open.values() if v)
|
|
n_closed = n_total - n_open
|
|
holdings_base = []
|
|
holdings_gated = []
|
|
equity_gated = 1_000_000.0
|
|
equity_base = 1_000_000.0
|
|
prev_week = None
|
|
prev_scores = None
|
|
daily_gated = []
|
|
daily_base = []
|
|
ret_by_date = {rd: ret_df.loc[rd] for rd in ret_df.index}
|
|
for d in trade_dates:
|
|
try:
|
|
day_scores = window_pred.loc[d]
|
|
except KeyError:
|
|
daily_gated.append(equity_gated)
|
|
daily_base.append(equity_base)
|
|
prev_scores = None
|
|
continue
|
|
if isinstance(day_scores, pd.DataFrame):
|
|
day_scores = day_scores.iloc[:, 0]
|
|
day_scores = day_scores.dropna().sort_values(ascending=False)
|
|
if len(day_scores) == 0:
|
|
daily_gated.append(equity_gated)
|
|
daily_base.append(equity_base)
|
|
prev_scores = None
|
|
continue
|
|
ret_row = ret_by_date.get(d)
|
|
if ret_row is None:
|
|
daily_gated.append(equity_gated)
|
|
daily_base.append(equity_base)
|
|
prev_scores = day_scores
|
|
continue
|
|
cur_week = (d.isocalendar()[0], d.isocalendar()[1]) if hasattr(d, 'isocalendar') else None
|
|
gate_val = gate_open.get(d, True)
|
|
if cur_week != prev_week or not holdings_base:
|
|
if prev_scores is not None:
|
|
holdings_base = list(prev_scores.index[:topk])
|
|
if holdings_base:
|
|
base_rets = ret_row.reindex(holdings_base).dropna()
|
|
if len(base_rets) > 0:
|
|
equity_base *= (1 + base_rets.mean())
|
|
if gate_val:
|
|
if cur_week != prev_week or not holdings_gated:
|
|
if prev_scores is not None:
|
|
holdings_gated = list(prev_scores.index[:topk])
|
|
if holdings_gated:
|
|
hold_rets = ret_row.reindex(holdings_gated).dropna()
|
|
if len(hold_rets) > 0:
|
|
equity_gated *= (1 + hold_rets.mean())
|
|
else:
|
|
holdings_gated = []
|
|
prev_week = cur_week
|
|
prev_scores = day_scores
|
|
daily_gated.append(equity_gated)
|
|
daily_base.append(equity_base)
|
|
g_series = pd.Series(daily_gated, index=trade_dates)
|
|
b_series = pd.Series(daily_base, index=trade_dates)
|
|
def _metrics(eq):
|
|
if len(eq) < 2:
|
|
return {"ann_return": 0, "sharpe": 0, "maxDD": 0}
|
|
rets = eq.pct_change().dropna()
|
|
ann_ret = float((eq.iloc[-1] / eq.iloc[0]) ** (252 / max(len(eq), 1)) - 1)
|
|
vol = float(rets.std() * (252 ** 0.5)) if len(rets) > 1 else 0
|
|
sharpe = ann_ret / vol if vol > 0 else 0
|
|
peak = eq.cummax()
|
|
dd = (eq - peak) / peak
|
|
maxDD = float(dd.min())
|
|
return {"ann_return": round(ann_ret, 6), "sharpe": round(sharpe, 4), "maxDD": round(maxDD, 6)}
|
|
base_m = _metrics(b_series)
|
|
gated_m = _metrics(g_series)
|
|
return {
|
|
"trade_dates": n_total,
|
|
"gate_open_days": n_open,
|
|
"gate_closed_days": n_closed,
|
|
"trip_rate": round(n_closed / n_total, 4) if n_total else 0,
|
|
"base": base_m,
|
|
"gated": gated_m,
|
|
}
|
|
|
|
|
|
def main():
|
|
OUT_DIR.mkdir(parents=True, exist_ok=True)
|
|
full_start = "2015-01-03"
|
|
full_end = "2026-08-19"
|
|
print("Loading lake bars...")
|
|
close_df = load_bars_for_window(full_start, full_end)
|
|
print(f" {close_df.shape[1]} symbols, {close_df.shape[0]} days")
|
|
ret_df = close_df.pct_change()
|
|
ret_df.index = pd.to_datetime(ret_df.index).normalize()
|
|
results = []
|
|
for window in WINDOWS:
|
|
wl, ws, we = window["label"], window["start"], window["end"]
|
|
pred_path = window["pred"]
|
|
print(f"\n=== Window {wl} ({ws} to {we}) ===")
|
|
pred = load_pred(pred_path)
|
|
print(f" pred shape: {pred.shape}")
|
|
hit_rates = {}
|
|
for lookback, _, name in SIGNAL_GATE_CONFIGS:
|
|
if lookback not in hit_rates:
|
|
hr = compute_hit_rate_series(pred, ret_df, topk=10, lookback=lookback)
|
|
hit_rates[lookback] = hr
|
|
print(f" lookback={lookback}: {len(hr)} days with hit rates")
|
|
for lookback, threshold, name in SIGNAL_GATE_CONFIGS:
|
|
hr = hit_rates[lookback]
|
|
bt = run_backtest(pred, hr, close_df, ws, we, topk=10, threshold=threshold)
|
|
if "error" in bt:
|
|
print(f" {name}: {bt['error']}")
|
|
continue
|
|
row = {
|
|
"window": wl,
|
|
"gate": name,
|
|
"start": ws,
|
|
"end": we,
|
|
"trade_dates": bt["trade_dates"],
|
|
"gate_open": bt["gate_open_days"],
|
|
"gate_closed": bt["gate_closed_days"],
|
|
"trip_rate": bt["trip_rate"],
|
|
"base_ann": bt["base"]["ann_return"],
|
|
"base_sharpe": bt["base"]["sharpe"],
|
|
"base_maxDD": bt["base"]["maxDD"],
|
|
"gated_ann": bt["gated"]["ann_return"],
|
|
"gated_sharpe": bt["gated"]["sharpe"],
|
|
"gated_maxDD": bt["gated"]["maxDD"],
|
|
}
|
|
results.append(row)
|
|
print(f" {name}: trip={bt['trip_rate']:.1%}, "
|
|
f"base={bt['base']['ann_return']:+.1%} (Sharpe {bt['base']['sharpe']:.2f}), "
|
|
f"gated={bt['gated']['ann_return']:+.1%} (Sharpe {bt['gated']['sharpe']:.2f})")
|
|
df = pd.DataFrame(results)
|
|
out_path = OUT_DIR / "signal_quality_gate_results.csv"
|
|
df.to_csv(out_path, index=False)
|
|
with open(OUT_DIR / "signal_quality_gate_results.json", "w") as f:
|
|
json.dump(df.to_dict(orient="records"), f, indent=2, default=str)
|
|
print(f"\nSaved to {out_path}")
|
|
print("\n=== Summary: Gated Return by Window ===")
|
|
for gate_name in df["gate"].unique():
|
|
gdf = df[df["gate"] == gate_name]
|
|
print(f"\n{gate_name}:")
|
|
for _, r in gdf.iterrows():
|
|
print(f" {r['window']}: base={r['base_ann']:+.1%}, gated={r['gated_ann']:+.1%}, "
|
|
f"trip={r['trip_rate']:.0%}, diff={r['gated_ann']-r['base_ann']:+.1%}pp")
|
|
|
|
|
|
if __name__ == "__main__":
|
|
main()
|