Files
tac-exp-dev/book/scripts/regime_gate_bt.py
T
zhaoli 718b048df9 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
2026-08-20 22:14:31 +00:00

401 lines
14 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
"""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()