You cannot select more than 25 topics Topics must start with a letter or number, can include dashes ('-') and can be up to 35 characters long.
fithia2/apps/tools/parking_strategy_probe.py

219 lines
7.7 KiB
Python

"""Probe cash parking variants on top of a strategy manifest.
Examples:
python -m apps.tools.parking_strategy_probe \
--config configs/experiments/return_max_long_v7.70.json \
--start 2022 --end 2026 \
--period-start y2026=2026-01-01 \
--preset ve_10_10 --preset vt_24_t13
python -m apps.tools.parking_strategy_probe \
--config configs/experiments/return_max_long_v7.70.json \
--start 2022 --end 2026 \
--override dynamic_vol_24='{"cash_parking_enabled": true, "cash_parking_symbol": "dynamic", "cash_parking_gate_mode": "volatility", "cash_parking_gate_vol_lookback": 20, "cash_parking_gate_vol_threshold": 0.24}'
"""
from __future__ import annotations
import argparse
import datetime as dt
import json
from dataclasses import dataclass
from typing import Any
from apps.backtester.run import (
BacktestRunner,
_build_merged_snapshot_store,
load_manifest,
resolve_config,
)
from libs.common.logging import configure_logging
@dataclass(frozen=True)
class VariantSpec:
name: str
preset: str | None
risk_overrides: dict[str, Any]
def _parse_date(value: str, *, is_end: bool = False) -> dt.date:
parts = value.split("-")
if len(parts) == 1 and len(value) == 4 and value.isdigit():
year = int(value)
return dt.date(year, 12, 31) if is_end else dt.date(year, 1, 1)
if len(parts) == 2 and all(part.isdigit() for part in parts):
year = int(parts[0])
month = int(parts[1])
if is_end:
next_month = dt.date(year + (month // 12), (month % 12) + 1, 1)
return next_month - dt.timedelta(days=1)
return dt.date(year, month, 1)
return dt.date.fromisoformat(value)
def _parse_period_start(value: str) -> tuple[str, dt.date]:
if "=" in value:
label, raw_date = value.split("=", 1)
else:
label, raw_date = value, value
return label, _parse_date(raw_date)
def _parse_override(value: str) -> VariantSpec:
if "=" not in value:
raise argparse.ArgumentTypeError("--override format: name=<json>")
name, payload = value.split("=", 1)
try:
risk_overrides = json.loads(payload)
except json.JSONDecodeError as exc:
raise argparse.ArgumentTypeError(f"invalid JSON for override '{name}': {exc}") from exc
if not isinstance(risk_overrides, dict):
raise argparse.ArgumentTypeError("override JSON must decode to an object")
return VariantSpec(name=name, preset=None, risk_overrides=risk_overrides)
def _period_stats(curve: list[Any], start_date: dt.date) -> dict[str, Any] | None:
points = [(state.date, float(state.equity)) for state in curve if state.date >= start_date]
if not points:
return None
start_equity = points[0][1]
peak = start_equity
max_dd_pct = 0.0
for _, equity in points:
peak = max(peak, equity)
if peak > 0:
max_dd_pct = max(max_dd_pct, (peak - equity) / peak * 100.0)
end_equity = points[-1][1]
return {
"start_date": points[0][0].isoformat(),
"end_date": points[-1][0].isoformat(),
"return_pct": (end_equity / start_equity - 1.0) * 100.0 if start_equity else 0.0,
"max_dd_pct": max_dd_pct,
}
def _build_variants(args: argparse.Namespace) -> list[VariantSpec]:
variants = [VariantSpec(name="no_parking", preset=None, risk_overrides={})]
for preset in args.preset or []:
variants.append(VariantSpec(name=preset, preset=preset, risk_overrides={}))
variants.extend(args.override or [])
return variants
def main() -> None:
parser = argparse.ArgumentParser(description="Probe strategy cash parking variants")
parser.add_argument("--config", required=True, help="Experiment manifest JSON path")
parser.add_argument("--start", required=True, help="Start date (YYYY, YYYY-MM, YYYY-MM-DD)")
parser.add_argument("--end", required=True, help="End date (YYYY, YYYY-MM, YYYY-MM-DD)")
parser.add_argument(
"--period-start",
action="append",
default=[],
help="Sub-period start. Format: label=YYYY-MM-DD (repeatable).",
)
parser.add_argument(
"--preset",
action="append",
default=[],
help="Parking preset name to evaluate (repeatable).",
)
parser.add_argument(
"--override",
action="append",
type=_parse_override,
default=[],
help="Named risk override. Format: name=<json> (repeatable).",
)
parser.add_argument(
"--json",
action="store_true",
help="Emit machine-readable JSON only.",
)
parser.add_argument(
"--capital",
type=float,
default=10_000.0,
help="Initial equity for the probe (default: 10000).",
)
args = parser.parse_args()
configure_logging("WARNING")
start_date = _parse_date(args.start)
end_date = _parse_date(args.end, is_end=True)
period_starts = [_parse_period_start(value) for value in args.period_start]
variants = _build_variants(args)
manifest = load_manifest(args.config)
base_config = resolve_config(manifest)
store = _build_merged_snapshot_store(
manifest,
base_config,
snapshot_dir_override=None,
).slice_by_date_range(start_date, end_date)
rows: list[dict[str, Any]] = []
for variant in variants:
config = resolve_config(manifest)
if variant.preset:
config.risk.cash_parking_preset = variant.preset
config.risk.apply_parking_preset()
for key, value in variant.risk_overrides.items():
setattr(config.risk, key, value)
runner = BacktestRunner(
manifest=manifest,
config=config,
store=store,
initial_equity=args.capital,
split_name="parking_probe",
)
result = runner.run(output_root=None)
row = {
"variant": variant.name,
"full_return_pct": round(result.metrics.total_return_pct or 0.0, 2),
"full_dd_pct": round(result.metrics.max_drawdown_pct or 0.0, 2),
"full_sharpe": round(result.metrics.sharpe_ratio or 0.0, 3),
"trade_count": result.metrics.trade_count,
"last_date": runner._equity_curve[-1].date.isoformat() if runner._equity_curve else None,
}
for label, period_start in period_starts:
stats = _period_stats(runner._equity_curve, period_start)
row[label] = {
"return_pct": round(stats["return_pct"], 2) if stats else None,
"max_dd_pct": round(stats["max_dd_pct"], 2) if stats else None,
"start_date": stats["start_date"] if stats else None,
"end_date": stats["end_date"] if stats else None,
}
rows.append(row)
if args.json:
print(json.dumps(rows, ensure_ascii=False, indent=2))
return
period_labels = [label for label, _ in period_starts]
header = ["variant", "full_return", "full_dd", "full_sharpe", "trades", "last_date"]
for label in period_labels:
header.extend([f"{label}_return", f"{label}_dd"])
print("\t".join(header))
for row in rows:
fields = [
row["variant"],
f"{row['full_return_pct']:.2f}",
f"{row['full_dd_pct']:.2f}",
f"{row['full_sharpe']:.3f}",
str(row["trade_count"]),
str(row["last_date"] or "-"),
]
for label in period_labels:
stats = row.get(label) or {}
fields.extend([
"-" if stats.get("return_pct") is None else f"{stats['return_pct']:.2f}",
"-" if stats.get("max_dd_pct") is None else f"{stats['max_dd_pct']:.2f}",
])
print("\t".join(fields))
if __name__ == "__main__":
main()