"""ORB strategy overfitting analysis with streaming intraday simulation.""" from __future__ import annotations import argparse import asyncio import json import statistics import sys import time from pathlib import Path from typing import Any import yaml from apps.intraday_bt.oracle import make_intraday_oracle_client from rich import box from rich.console import Console from rich.panel import Panel from rich.progress import BarColumn, Progress, SpinnerColumn, TextColumn, TimeElapsedColumn from libs.backtest.domain import SplitResult, WalkForwardSummary from libs.common.config import get_settings from libs.intraday.domain import ORBStrategyParams from libs.oracle_client import OracleClient from apps.intraday_bt.orb_research import ( build_orb_research_context, force_simple_returns, generate_walk_forward_windows, resolve_orb_config, simulate_orb_period, split_trading_days, ) from apps.intraday_bt.run import _latest_backtest_date _console = Console(width=120) _DEFAULT_SPLIT_DATE = "2026-01-01" _DEFAULT_START_DATE = "2022-01-01" _WF_TRAIN_DAYS = 252 _WF_TEST_DAYS = 63 async def run_is_oos_test( context, client: OracleClient, config, split_date: str, ) -> dict[str, Any]: train_days, test_days = split_trading_days(context.trading_days, split_date=split_date) if not train_days or not test_days: return {"verdict": "SKIP", "notes": "Not enough data to split"} orb_params = config.orb_strategy or ORBStrategyParams() is_metrics = await simulate_orb_period(context, client, orb_params, train_days, run_id="is") oos_metrics = await simulate_orb_period(context, client, orb_params, test_days, run_id="oos") is_sharpe = is_metrics.sharpe_ratio or 0.0 oos_sharpe = oos_metrics.sharpe_ratio or 0.0 retention = 0.0 if is_sharpe <= 0 else oos_sharpe / is_sharpe if retention >= 0.60: verdict = "PASS" elif retention >= 0.40: verdict = "WARN" else: verdict = "FAIL" return { "verdict": verdict, "is_sharpe": round(is_sharpe, 3), "oos_sharpe": round(oos_sharpe, 3), "retention_pct": round(retention * 100, 1), "is_period": f"{train_days[0]} → {train_days[-1]} ({len(train_days)} days)", "oos_period": f"{test_days[0]} → {test_days[-1]} ({len(test_days)} days)", } def summarize_is_oos_from_results( is_result: SplitResult, oos_result: SplitResult, *, is_period: str | None = None, oos_period: str | None = None, ) -> dict[str, Any]: """Reuse already-computed train/test split results for IS/OOS retention.""" is_sharpe = is_result.sharpe_ratio or 0.0 oos_sharpe = oos_result.sharpe_ratio or 0.0 retention = 0.0 if is_sharpe <= 0 else oos_sharpe / is_sharpe if retention >= 0.60: verdict = "PASS" elif retention >= 0.40: verdict = "WARN" else: verdict = "FAIL" payload = { "verdict": verdict, "is_sharpe": round(is_sharpe, 3), "oos_sharpe": round(oos_sharpe, 3), "retention_pct": round(retention * 100, 1), "source": "split_results", } if is_period is not None: payload["is_period"] = is_period if oos_period is not None: payload["oos_period"] = oos_period return payload async def run_walk_forward_test( context, client: OracleClient, config, train_days: int = _WF_TRAIN_DAYS, test_days: int = _WF_TEST_DAYS, ) -> dict[str, Any]: windows = generate_walk_forward_windows(context.trading_days, train_days, test_days) if len(windows) < 2: return {"verdict": "SKIP", "notes": f"Need ≥ {train_days + 2 * test_days} days, got {len(context.trading_days)}"} orb_params = config.orb_strategy or ORBStrategyParams() sharpes: list[float] = [] for i, (_, wf_test) in enumerate(windows, start=1): metrics = await simulate_orb_period(context, client, orb_params, wf_test, run_id=f"wf{i:02d}") sharpe = metrics.sharpe_ratio or 0.0 sharpes.append(sharpe) sys.stdout.write(f"\r WF window {i}/{len(windows)}: test Sharpe={sharpe:.2f} ") sys.stdout.flush() print() mean_sr = statistics.mean(sharpes) std_sr = statistics.stdev(sharpes) if len(sharpes) > 1 else 0.0 cv = std_sr / abs(mean_sr) if abs(mean_sr) > 0.01 else float("inf") n_positive = sum(1 for sharpe in sharpes if sharpe > 0) if mean_sr > 0.5 and cv < 0.80: verdict = "PASS" elif mean_sr > 0 and cv < 1.5: verdict = "WARN" else: verdict = "FAIL" return { "verdict": verdict, "n_windows": len(windows), "mean_sharpe": round(mean_sr, 3), "std_sharpe": round(std_sr, 3), "cv": round(cv, 3), "n_positive": n_positive, "window_sharpes": [round(s, 3) for s in sharpes], } def summarize_walk_forward_test_from_summary(wf_summary: WalkForwardSummary) -> dict[str, Any]: """Reuse an already-built walk-forward summary for the overfit verdict.""" sharpes = [fold.test_metrics.sharpe_ratio or 0.0 for fold in wf_summary.folds] if len(sharpes) < 2: return { "verdict": "SKIP", "notes": f"Need ≥ 2 folds, got {len(sharpes)}", } mean_sr = statistics.mean(sharpes) std_sr = statistics.stdev(sharpes) if len(sharpes) > 1 else 0.0 cv = std_sr / abs(mean_sr) if abs(mean_sr) > 0.01 else float("inf") n_positive = sum(1 for sharpe in sharpes if sharpe > 0) if mean_sr > 0.5 and cv < 0.80: verdict = "PASS" elif mean_sr > 0 and cv < 1.5: verdict = "WARN" else: verdict = "FAIL" return { "verdict": verdict, "n_windows": len(sharpes), "mean_sharpe": round(mean_sr, 3), "std_sharpe": round(std_sr, 3), "cv": round(cv, 3), "n_positive": n_positive, "window_sharpes": [round(s, 3) for s in sharpes], "source": "walk_forward_summary", } async def run_param_plateau_test( context, client: OracleClient, config, quick: bool = False, param_names: list[str] | None = None, ) -> dict[str, Any]: orb_params = config.orb_strategy or ORBStrategyParams() n_values = 3 if quick else 5 params_to_test = [ ("atr_stop_multiplier", orb_params.atr_stop_multiplier, 0.4), ("breakeven_at_r", orb_params.breakeven_at_r, 0.4), ("trailing_stop_atr_multiplier", orb_params.trailing_stop_atr_multiplier, 0.4), ] if param_names is not None: selected = set(param_names) params_to_test = [item for item in params_to_test if item[0] in selected] results = [] for param_name, base_value, spread in params_to_test: lo = base_value * (1 - spread) hi = base_value * (1 + spread) test_values = [lo + (hi - lo) * i / (n_values - 1) for i in range(n_values)] sharpes: list[float] = [] for value in test_values: modified = orb_params.model_copy(update={param_name: round(value, 6)}) metrics = await simulate_orb_period( context, client, modified, context.trading_days, run_id=f"{param_name[:4]}_{value:.4f}", ) sharpe = metrics.sharpe_ratio or 0.0 sharpes.append(sharpe) sys.stdout.write(f"\r {param_name}={value:.4f} → Sharpe={sharpe:.2f} ") sys.stdout.flush() print() mean_sharpe = statistics.mean(sharpes) plateau = max(0.0, 1.0 - statistics.stdev(sharpes) / abs(mean_sharpe)) if len(sharpes) > 1 and abs(mean_sharpe) > 0.01 else 0.0 verdict = "PASS" if plateau >= 0.70 else "WARN" if plateau >= 0.40 else "FAIL" results.append({ "param": param_name, "base_value": base_value, "test_values": [round(v, 5) for v in test_values], "sharpe_values": [round(s, 3) for s in sharpes], "plateau": round(plateau, 3), "verdict": verdict, }) verdicts = [result["verdict"] for result in results] if all(v == "PASS" for v in verdicts): overall = "PASS" elif "FAIL" in verdicts: overall = "FAIL" else: overall = "WARN" return {"verdict": overall, "params": results} async def run_permutation_test( context, client: OracleClient, config, n_permutations: int = 30, ) -> dict[str, Any]: orb_params = config.orb_strategy or ORBStrategyParams() observed = await simulate_orb_period(context, client, orb_params, context.trading_days, run_id="perm_obs") observed_sharpe = observed.sharpe_ratio or 0.0 _console.print(f" [dim]Observed Sharpe (real ranking): {observed_sharpe:.3f}[/dim]") null_sharpes: list[float] = [] with Progress( SpinnerColumn(), TextColumn("[progress.description]{task.description}"), BarColumn(), "{task.completed}/{task.total}", TimeElapsedColumn(), console=_console, transient=True, ) as progress: task = progress.add_task("Permutation test", total=n_permutations) for i in range(n_permutations): metrics = await simulate_orb_period( context, client, orb_params, context.trading_days, run_id=f"perm{i:03d}", shuffle_candidates_seed=42 + i, ) null_sharpes.append(metrics.sharpe_ratio or 0.0) progress.advance(task) p_value = sum(1 for sr in null_sharpes if sr >= observed_sharpe) / max(len(null_sharpes), 1) null_sorted = sorted(null_sharpes) p95_idx = min(len(null_sorted) - 1, int(0.95 * (len(null_sorted) - 1))) if null_sorted else 0 null_p95 = null_sorted[p95_idx] if null_sorted else 0.0 null_median = statistics.median(null_sharpes) if null_sharpes else 0.0 if p_value < 0.05: verdict = "PASS" elif p_value < 0.20: verdict = "WARN" else: verdict = "FAIL" return { "verdict": verdict, "observed_sharpe": round(observed_sharpe, 3), "n_permutations": n_permutations, "null_median": round(null_median, 3), "null_p95": round(null_p95, 3), "p_value": round(p_value, 4), "null_sharpes": [round(sr, 3) for sr in null_sharpes], } def _vc(verdict: str) -> str: return {"PASS": "green", "WARN": "yellow", "FAIL": "red", "SKIP": "dim"}.get(verdict, "white") def _verdict_tag(verdict: str) -> str: color = _vc(verdict) return f"[{color}][{verdict}][/{color}]" def _print_report( config_path: Path, test1: dict[str, Any], test2: dict[str, Any], test3: dict[str, Any], test4: dict[str, Any], elapsed: float, ) -> None: weights = {"PASS": 1.0, "WARN": 0.5, "FAIL": 0.0, "SKIP": None} scores = [weights[t.get("verdict", "SKIP")] for t in [test1, test2, test3, test4]] scores = [score for score in scores if score is not None] overall_score = int(statistics.mean(scores) * 100) if scores else 0 pass_count = sum(1 for t in [test1, test2, test3, test4] if t.get("verdict") == "PASS") overall_verdict = "PASS" if pass_count >= 3 else "WARN" if pass_count >= 2 else "FAIL" color = _vc(overall_verdict) _console.print() _console.print( Panel( f"[bold]ORB OVERFITTING ANALYSIS[/bold]\n" f"Config: [cyan]{config_path}[/cyan]\n\n" f"Overall: [{color} bold]{overall_verdict}[/{color} bold] " f"Score: [bold]{overall_score}/100[/bold] ({elapsed:.0f}s)", box=box.DOUBLE, width=100, ) ) for label, payload in [ ("1. IS/OOS Retention", test1), ("2. Walk-Forward Stability", test2), ("3. Parameter Plateau", test3), ("4. Candidate Ranking Permutation", test4), ]: _console.print(f"\n [bold]{label}[/bold] {_verdict_tag(payload.get('verdict', 'SKIP'))}") _console.print(f" {json.dumps(payload, ensure_ascii=True)}") async def _async_main(args: argparse.Namespace) -> int: t0 = time.time() config_path, config = resolve_orb_config(args.config) config = force_simple_returns(config) settings = get_settings() async with make_intraday_oracle_client(settings) as client: context = await build_orb_research_context( config, args.start, args.end or _latest_backtest_date().isoformat(), client, print_progress=True, ) skip = set(args.skip.split(",")) if args.skip else set() test1: dict[str, Any] = {"verdict": "SKIP"} test2: dict[str, Any] = {"verdict": "SKIP"} test3: dict[str, Any] = {"verdict": "SKIP"} test4: dict[str, Any] = {"verdict": "SKIP"} if "is_oos" not in skip: _console.print("\n[bold][1/4] IS/OOS Retention...[/bold]") test1 = await run_is_oos_test(context, client, config, args.split_date) if "wf" not in skip: _console.print("\n[bold][2/4] Walk-Forward Stability...[/bold]") wf_train = 126 if args.quick else _WF_TRAIN_DAYS wf_test = 42 if args.quick else _WF_TEST_DAYS test2 = await run_walk_forward_test(context, client, config, train_days=wf_train, test_days=wf_test) if "plateau" not in skip: _console.print("\n[bold][3/4] Parameter Plateau...[/bold]") test3 = await run_param_plateau_test(context, client, config, quick=args.quick) if "perm" not in skip: _console.print(f"\n[bold][4/4] Candidate Permutation (N={args.permutations})...[/bold]") test4 = await run_permutation_test(context, client, config, n_permutations=args.permutations) elapsed = time.time() - t0 _print_report(config_path, test1, test2, test3, test4, elapsed) if args.output_json: report = { "config": str(config_path), "period": f"{context.trading_days[0]} → {context.trading_days[-1]}", "split_date": args.split_date, "elapsed_seconds": round(elapsed, 1), "tests": { "is_oos": test1, "walk_forward": test2, "param_plateau": test3, "permutation": test4, }, } Path(args.output_json).parent.mkdir(parents=True, exist_ok=True) Path(args.output_json).write_text(json.dumps(report, indent=2)) _console.print(f"\n[dim]Report saved → {args.output_json}[/dim]") return 0 def main() -> None: parser = argparse.ArgumentParser( prog="fithia2 intraday-overfit-check", description="ORB strategy overfitting analysis (IS/OOS, WFV, param plateau, permutation)", ) parser.add_argument("--config", "-c", required=True, help="YAML config path or strategy slug") parser.add_argument("--split-date", default=_DEFAULT_SPLIT_DATE, help=f"IS/OOS split date (default: {_DEFAULT_SPLIT_DATE})") parser.add_argument("--start", default=_DEFAULT_START_DATE, help=f"Start of backtest window (default: {_DEFAULT_START_DATE})") parser.add_argument("--end", default=None, help="End of backtest window (default: latest available)") parser.add_argument("--quick", action="store_true", help="Quick mode: fewer WF windows, 3-point plateau, 15 permutations") parser.add_argument("--permutations", type=int, default=50, help="Candidate permutation test iterations") parser.add_argument("--skip", default="", help="Comma-separated tests to skip: is_oos,wf,plateau,perm") parser.add_argument("--output-json", default=None, help="Save report as JSON to this path") args = parser.parse_args() if args.quick: args.permutations = 15 raise SystemExit(asyncio.run(_async_main(args))) if __name__ == "__main__": main()