Files
foxhunt/scripts/validation/check_tier3.py
jgrusewski 0d373da490 plan5(task4): tiered-exit validation script suite (tier1/2/3 checks)
Creates scripts/validation/ with per-tier exit checks consuming the
aggregate JSON from scripts/aggregate-multi-seed-metrics.py (P5T1B):

  check_tier1.py — convergence (std/mean ≤ 0.15 on val_sharpe /
    avg_q_value / train_loss; avg_q_value max ≤ 500 fold-1 explosion
    guard; placeholders for Q-saturation + hot-path-DtoH per spec).
  check_tier2.py — behavioural (val_trades_per_bar ≥ 0.005,
    val_active_frac > 0.2, dir argmax entropy > 0.8·log4 with
    val_dir_entropy primary + val_dir_dist_* fallback).
  check_tier3.py — profitability (val_sharpe_annualised > 1.0 with
    val_sharpe per-bar fallback, val_win_rate ≥ 0.52 gated on
    >500 trades, val_profit_factor mean ≥ 1.1 AND cross-seed std < 0.3).
  check_all_tiers.py — subprocess wrapper, exits 0 only if all pass.

Stdlib-only (statistics / argparse / json / subprocess) — no new deps.
Defensive missing-metric handling: each check FAILs with an explanatory
message when its required aggregate key is absent rather than silently
passing, so missing HEALTH_DIAG metrics are surfaced loudly.

Test harness scripts/validation/tests/test_tier_checks.sh exercises
good + bad fixtures across all four scripts and against the wrapper.

Audit row added to docs/dqn-wire-up-audit.md documenting the suite +
the deferred metrics list (val_trades_per_bar, val_active_frac,
val_dir_entropy/_dist_*, val_sharpe_annualised, val_win_rate,
val_profit_factor, val_trade_count) that HEALTH_DIAG must emit before
tiers 2/3 can ever PASS on real data — tracked for Plan 5 Task 5
pre-flight wire-up.

Co-Authored-By: Claude Opus 4.7 (1M context) <noreply@anthropic.com>
2026-04-26 12:33:28 +02:00

186 lines
6.4 KiB
Python
Executable File
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.
#!/usr/bin/env python3
"""Tier 3 profitability exit check (Plan 5 Task 4).
Consumes the aggregated metrics JSON produced by
`scripts/aggregate-multi-seed-metrics.py` (Plan 5 Task 1B.2).
Pass criteria (per spec §2 Tier 3 — profitability):
1. `val_sharpe_annualised` mean > 1.0 on stable epochs.
If only `val_sharpe` (per-bar) is in the aggregate, fall back to it
as a proxy and document the difference in the PASS/FAIL message.
2. `val_win_rate` mean ≥ 0.52 — only enforced when the corresponding
trade-count series indicates > 500 trades/bar over stable epochs.
Below 500-trade gate, win-rate is too noisy to assess.
3. `val_profit_factor` mean ≥ 1.1 AND across-seed std < 0.3 on stable
epochs (the std is the across-seed dispersion baked into each
epoch entry by the aggregator).
Exits 0 on PASS, 1 on FAIL. Each check prints a `PASS:` / `FAIL:` line.
"""
from __future__ import annotations
import argparse
import json
import sys
from statistics import mean
def _stable(series, warmup_end):
return [x for x in series if x["epoch"] >= warmup_end]
def check_sharpe(agg, warmup_end):
annualised = agg.get("val_sharpe_annualised")
if annualised:
stable = _stable(annualised, warmup_end)
if not stable:
return False, (
f"No stable epochs for val_sharpe_annualised past "
f"warmup_end={warmup_end}"
)
m = mean(x["mean"] for x in stable)
if m <= 1.0:
return False, f"val_sharpe_annualised mean = {m:.4f} <= 1.0"
return True, f"val_sharpe_annualised mean = {m:.4f} > 1.0"
# Fallback: per-bar val_sharpe — note proxy in message.
proxy = agg.get("val_sharpe")
if proxy:
stable = _stable(proxy, warmup_end)
if not stable:
return False, (
f"No stable epochs for val_sharpe past warmup_end={warmup_end}"
)
m = mean(x["mean"] for x in stable)
# Per-bar sharpe is much smaller than annualised; use the same > 1.0
# threshold but flag the proxy substitution loud.
if m <= 1.0:
return False, (
f"val_sharpe (per-bar PROXY for annualised) mean = {m:.4f} "
f"<= 1.0 — annualised metric missing from aggregate"
)
return True, (
f"val_sharpe (per-bar PROXY for annualised) mean = {m:.4f} > 1.0"
)
return False, (
"neither val_sharpe_annualised nor val_sharpe present in aggregate "
"— required for Tier 3 profitability check"
)
def check_win_rate(agg, warmup_end):
series = agg.get("val_win_rate")
if not series:
return False, (
"val_win_rate missing from aggregate — required for Tier 3"
)
stable = _stable(series, warmup_end)
if not stable:
return False, (
f"No stable epochs for val_win_rate past warmup_end={warmup_end}"
)
# Trade-count gate. Use val_trades_per_bar mean × an assumed bar count
# of 500; if explicit val_trade_count is present, prefer that.
trade_count_series = agg.get("val_trade_count")
if trade_count_series:
tc_stable = _stable(trade_count_series, warmup_end)
if tc_stable:
tc = mean(x["mean"] for x in tc_stable)
if tc <= 500:
return True, (
f"val_win_rate skipped — trade count {tc:.0f} <= 500 "
f"(too noisy)"
)
else:
# Approximate from per-bar trade rate × stable epoch span; if the
# rate is below ~0.001 over even a long horizon the count is below
# 500. Conservative — only skip if we can prove count <= 500.
tpb_series = agg.get("val_trades_per_bar")
if tpb_series:
tpb_stable = _stable(tpb_series, warmup_end)
if tpb_stable:
tpb = mean(x["mean"] for x in tpb_stable)
# Without knowing the eval window length, assume a
# conservative 5e5 bars (a single day at 1Hz). If even
# that yields <= 500 trades, skip.
if tpb * 5e5 <= 500:
return True, (
f"val_win_rate skipped — estimated trade count "
f"{tpb * 5e5:.0f} <= 500 (too noisy)"
)
m = mean(x["mean"] for x in stable)
if m < 0.52:
return False, f"val_win_rate mean = {m:.4f} < 0.52"
return True, f"val_win_rate mean = {m:.4f} >= 0.52"
def check_profit_factor(agg, warmup_end):
series = agg.get("val_profit_factor")
if not series:
return False, (
"val_profit_factor missing from aggregate — required for Tier 3"
)
stable = _stable(series, warmup_end)
if not stable:
return False, (
f"No stable epochs for val_profit_factor past warmup_end={warmup_end}"
)
m = mean(x["mean"] for x in stable)
s = mean(x["std"] for x in stable)
if m < 1.1:
return False, f"val_profit_factor mean = {m:.4f} < 1.1"
if s >= 0.3:
return False, (
f"val_profit_factor cross-seed std = {s:.4f} >= 0.3 (unstable)"
)
return True, (
f"val_profit_factor mean = {m:.4f} >= 1.1, std = {s:.4f} < 0.3"
)
def main():
ap = argparse.ArgumentParser(
description="Tier 3 profitability exit check (Plan 5 Task 4).",
)
ap.add_argument(
"metrics_json",
help="Aggregated metrics JSON (output of aggregate-multi-seed-metrics.py).",
)
ap.add_argument(
"--warmup-end",
type=int,
default=15,
help="Epoch at which warmup ends (stable epochs are > this); default 15.",
)
args = ap.parse_args()
try:
with open(args.metrics_json) as fh:
data = json.load(fh)
except (OSError, json.JSONDecodeError) as exc:
print(f"FAIL: cannot read {args.metrics_json}: {exc}")
sys.exit(1)
agg = data.get("aggregates")
if not isinstance(agg, dict):
print("FAIL: missing or malformed `aggregates` block in metrics JSON")
sys.exit(1)
checks = [
check_sharpe(agg, args.warmup_end),
check_win_rate(agg, args.warmup_end),
check_profit_factor(agg, args.warmup_end),
]
all_pass = all(ok for ok, _ in checks)
for ok, msg in checks:
print(("PASS" if ok else "FAIL") + ": " + msg)
sys.exit(0 if all_pass else 1)
if __name__ == "__main__":
main()