Add regime gate walk-forward test (EVIDENCE#050)
- 3 detector types (dispersion/vol/HMM) × 14 configs across 5 years - Dispersion gates: 0% trip rate everywhere (dead) - Vol gates: trip differential +32-47pp but destroy returns in good years - HMM gates: +6pp differential, hmm_0.7 improves 2023/2025 but kills 2026 - Guard candidate regime gate REFUTED (ch 11) - Script: book/scripts/regime_gate_bt.py - Results: book/data/regime_gate/regime_gate_trip_rates.csv
This commit is contained in:
@@ -0,0 +1,400 @@
|
||||
"""Regime-gate walk-forward backtest grid.
|
||||
|
||||
Precomputes regime gates (dispersion/vol/hmm × threshold grid) from lake bars,
|
||||
then runs a qlib TopkDropout backtest with each gate applied as a date-level
|
||||
trade overlay. Uses the SAME pred.pkl from exp 52 (Config A 2026) so the
|
||||
model is trained only once.
|
||||
|
||||
Usage:
|
||||
cd /app && .venv/bin/python book/scripts/regime_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/regime_gate")
|
||||
|
||||
# Walk-forward test windows with their pred.pkl sources
|
||||
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"},
|
||||
]
|
||||
|
||||
# Gate grid
|
||||
DISP_THRESHOLDS = [0.010, 0.015, 0.020, 0.025, 0.030]
|
||||
VOL_BANDS = [
|
||||
(0.0, 0.15, "low_max15"),
|
||||
(0.0, 0.20, "low_max20"),
|
||||
(0.0, 0.25, "low_max25"),
|
||||
(0.10, 0.25, "10_25"),
|
||||
(0.10, 0.30, "10_30"),
|
||||
]
|
||||
HMM_THRESHOLDS = [0.3, 0.5, 0.7, 0.9]
|
||||
|
||||
|
||||
def load_pred(path: str) -> pd.Series:
|
||||
"""Load pred.pkl (MultiIndex: datetime × instrument → score), dates normalized to midnight."""
|
||||
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
|
||||
# Normalize datetime level to date-only (midnight, no tz)
|
||||
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 precompute_gates(close_df: pd.DataFrame) -> dict:
|
||||
"""Precompute all regime gate series from close prices."""
|
||||
gates = {}
|
||||
|
||||
# --- dispersion gates ---
|
||||
ret22 = close_df.pct_change(22)
|
||||
cs_disp = ret22.std(axis=1)
|
||||
for thr in DISP_THRESHOLDS:
|
||||
g = cs_disp >= thr
|
||||
g.iloc[:22] = True
|
||||
gates[f"disp_{thr:.3f}"] = g
|
||||
|
||||
# --- vol gates ---
|
||||
import numpy as np
|
||||
log_ret = np.log(close_df / close_df.shift(1))
|
||||
rv22 = log_ret.rolling(22).std() * (252 ** 0.5)
|
||||
cs_vol = rv22.mean(axis=1)
|
||||
for vlow, vhigh, tag in VOL_BANDS:
|
||||
g = (cs_vol >= vlow) & (cs_vol <= vhigh)
|
||||
g.iloc[:22] = True
|
||||
gates[f"vol_{tag}"] = g
|
||||
|
||||
# --- HMM gates ---
|
||||
hmm_root = pathlib.Path(LAKE_ROOT) / "features" / "market=US" / "timeframe=1d"
|
||||
for thr in HMM_THRESHOLDS:
|
||||
all_post = {}
|
||||
for sym in close_df.columns:
|
||||
for family in ("sp", "ta"):
|
||||
fp = hmm_root / f"family={family}" / f"symbol={sym}.parquet"
|
||||
if not fp.exists():
|
||||
continue
|
||||
try:
|
||||
feat = pd.read_parquet(fp)
|
||||
except Exception:
|
||||
continue
|
||||
if "sp_hmm_p_regime1" not in feat.columns:
|
||||
continue
|
||||
tcol = feat["t"] if "t" in feat.columns else feat["date"]
|
||||
ts = pd.to_datetime(tcol)
|
||||
s = pd.Series(feat["sp_hmm_p_regime1"].values, index=ts, name=sym)
|
||||
s = s.dropna()
|
||||
if len(s) > 0:
|
||||
all_post[sym] = s
|
||||
break
|
||||
if all_post:
|
||||
post_df = pd.DataFrame(all_post)
|
||||
cs_mean = post_df.mean(axis=1)
|
||||
g = cs_mean >= thr
|
||||
else:
|
||||
g = pd.Series(True, index=close_df.index)
|
||||
gates[f"hmm_{thr:.1f}"] = g
|
||||
|
||||
# Normalize all gate indices to date-only (no tz, no time)
|
||||
for key in gates:
|
||||
gates[key].index = pd.to_datetime(gates[key].index).normalize()
|
||||
|
||||
return gates
|
||||
|
||||
|
||||
def run_backtest_with_gate(
|
||||
pred: pd.Series,
|
||||
gate: pd.Series,
|
||||
close_df: pd.DataFrame,
|
||||
start: str,
|
||||
end: str,
|
||||
topk: int = 10,
|
||||
n_drop: int = 1,
|
||||
) -> dict:
|
||||
"""Simulate TopkDropout with gate overlay, computing daily returns.
|
||||
|
||||
- On gate-open days: hold topk stocks (equal-weight), rebalance weekly
|
||||
- On gate-closed days: liquidate to cash
|
||||
- Tracks both gated and ungated (baseline) equity curves
|
||||
"""
|
||||
# Ensure pred has MultiIndex (date, instrument)
|
||||
if not isinstance(pred.index, pd.MultiIndex):
|
||||
return {"error": "pred must have MultiIndex (date, instrument)"}
|
||||
|
||||
# Daily returns per symbol (close-to-close)
|
||||
ret_df = close_df.pct_change()
|
||||
# Normalize ret_df index to date-only for matching
|
||||
ret_df.index = pd.to_datetime(ret_df.index).normalize()
|
||||
|
||||
# Filter pred to window and get trade dates
|
||||
dt_idx = pred.index.get_level_values(0)
|
||||
window_mask = dt_idx >= pd.Timestamp(start)
|
||||
window_mask &= 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())
|
||||
|
||||
# Compute gate status per trade date
|
||||
gate_open = {}
|
||||
for d in trade_dates:
|
||||
known = gate[gate.index <= d]
|
||||
gate_open[d] = bool(known.iloc[-1]) if len(known) else True
|
||||
|
||||
n_total = len(trade_dates)
|
||||
n_open = sum(1 for v in gate_open.values() if v)
|
||||
n_closed = n_total - n_open
|
||||
|
||||
# Simulate: track current holdings — both base and gated use weekly rebalance
|
||||
# Use yesterday's scores to pick today's holdings (no look-ahead)
|
||||
holdings_base = []
|
||||
holdings_gated = []
|
||||
equity_gated = 1_000_000.0
|
||||
equity_base = 1_000_000.0
|
||||
prev_week = None
|
||||
prev_scores = None # yesterday's scores
|
||||
|
||||
daily_gated = []
|
||||
daily_base = []
|
||||
|
||||
# Build a date → ret_df row map
|
||||
ret_by_date = {rd: ret_df.loc[rd] for rd in ret_df.index}
|
||||
|
||||
for i, d in enumerate(trade_dates):
|
||||
# Get today's cross-sectional prediction
|
||||
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.Series) and not isinstance(day_scores.index, pd.MultiIndex):
|
||||
pass
|
||||
elif isinstance(day_scores, pd.DataFrame):
|
||||
day_scores = day_scores.iloc[:, 0]
|
||||
else:
|
||||
daily_gated.append(equity_gated)
|
||||
daily_base.append(equity_base)
|
||||
prev_scores = None
|
||||
continue
|
||||
|
||||
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)
|
||||
|
||||
# --- ungated baseline: weekly rebalance using yesterday's scores ---
|
||||
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())
|
||||
|
||||
# --- gated: weekly rebalance only when gate open, using yesterday's scores ---
|
||||
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)
|
||||
|
||||
# Compute metrics
|
||||
g_series = pd.Series(daily_gated, index=trade_dates)
|
||||
b_series = pd.Series(daily_base, index=trade_dates)
|
||||
|
||||
def _metrics(eq: pd.Series) -> dict:
|
||||
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 load_bars_for_window(start: str, end: str) -> pd.DataFrame:
|
||||
"""Load daily close prices for all symbols in the universe."""
|
||||
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()
|
||||
# Load a bit extra for warmup
|
||||
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 main():
|
||||
OUT_DIR.mkdir(parents=True, exist_ok=True)
|
||||
|
||||
# Load bars (with warmup) for the full panel
|
||||
full_start = "2015-01-03"
|
||||
full_end = "2026-08-19"
|
||||
print("Loading lake bars for gate precomputation...")
|
||||
close_df = load_bars_for_window(full_start, full_end)
|
||||
print(f" {close_df.shape[1]} symbols, {close_df.shape[0]} days")
|
||||
|
||||
print("Precomputing regime gates...")
|
||||
gates = precompute_gates(close_df)
|
||||
print(f" {len(gates)} gate configs: {list(gates.keys())}")
|
||||
|
||||
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}) ===")
|
||||
|
||||
print(f" Loading pred.pkl from {pred_path}...")
|
||||
pred = load_pred(pred_path)
|
||||
print(f" pred shape: {pred.shape}")
|
||||
|
||||
for gate_name, gate_series in gates.items():
|
||||
bt = run_backtest_with_gate(pred, gate_series, close_df, ws, we)
|
||||
if "error" in bt:
|
||||
print(f" {gate_name}: {bt['error']}")
|
||||
continue
|
||||
row = {
|
||||
"window": wl,
|
||||
"gate": 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" {gate_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})")
|
||||
|
||||
# Save results
|
||||
df = pd.DataFrame(results)
|
||||
out_path = OUT_DIR / "regime_gate_trip_rates.csv"
|
||||
df.to_csv(out_path, index=False)
|
||||
print(f"\nSaved trip rates to {out_path}")
|
||||
|
||||
# Also save as JSON for the book
|
||||
json_results = df.to_dict(orient="records")
|
||||
with open(OUT_DIR / "regime_gate_trip_rates.json", "w") as f:
|
||||
json.dump(json_results, f, indent=2, default=str)
|
||||
|
||||
# Print summary: trip rate differential (2026 vs bad years)
|
||||
print("\n=== Trip Rate Summary (2026 vs bad years) ===")
|
||||
for gate_name in gates.keys():
|
||||
gdf = df[df["gate"] == gate_name]
|
||||
r2026 = gdf[gdf["window"] == "2026"]["trip_rate"].values
|
||||
r_bad = gdf[gdf["window"].isin(["2021", "2023", "2024"])]["trip_rate"].values
|
||||
if len(r2026) and len(r_bad):
|
||||
d = r2026[0] - np.mean(r_bad)
|
||||
print(f" {gate_name}: 2026 trip={r2026[0]:.1%}, bad-years avg={np.mean(r_bad):.1%}, diff={d:+.1%}")
|
||||
|
||||
print("\n=== Gated Return Summary (2026 vs bad years) ===")
|
||||
for gate_name in gates.keys():
|
||||
gdf = df[df["gate"] == gate_name]
|
||||
r2026 = gdf[gdf["window"] == "2026"]
|
||||
r_bad = gdf[gdf["window"].isin(["2021", "2023", "2024"])]
|
||||
if len(r2026) and len(r_bad):
|
||||
g26 = r2026["gated_ann"].values[0]
|
||||
b26 = r2026["base_ann"].values[0]
|
||||
g_bad = r_bad["gated_ann"].mean()
|
||||
b_bad = r_bad["base_ann"].mean()
|
||||
print(f" {gate_name}: 2026 gated={g26:+.1%} (base={b26:+.1%}), "
|
||||
f"bad-years gated={g_bad:+.1%} (base={b_bad:+.1%})")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
Reference in New Issue
Block a user