Files
sanguo_vnpy_v2/scripts/shadow_desk/spike_p0_checkpoint_replay.py
T

183 lines
7.0 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.
# -*- coding: utf-8 -*-
"""影子柜台 P0 对账 spikedocs/design/paper-shadow-desk-design.md §6 P0)。
验证目标:checkpoint 续跑(分段跑,段间用 initial_positions + 段末现金恢复)
与全量重放(一次跑完整个区间)账目一致。
判定标准:
- 期末持仓逐只证券数量完全一致;均价差 < 1e-6(相对)
- 期末现金 / 总值差 < 1e-6(相对)
- 两边净值曲线在共同交易日上的偏差 < 1e-6(相对)
跨除权:区间拉长到年,段边界自然落在除权日前后的概率极高;可通过
--split 显式把边界压在已知除权日上(前复权因子是否随段变化由对账暴露)。
用法(VPS,真实数据):
cd C:\\sanguo_vnpy_v2
C:\\Python310\\python.exe -X utf8 scripts\\shadow_desk\\spike_p0_checkpoint_replay.py ^
--strategy all_weather --start 2024-01-01 --end 2024-12-31 --cash 1000000
输出末行 SPIKE_P0_PASS / SPIKE_P0_FAIL。
"""
from __future__ import annotations
import argparse
import json
import sys
TOL = 1e-6 # 相对容差
def _run_segment(params: dict) -> dict:
"""跑一段回测,返回含 final_portfolio 的结果 dict。"""
from sanguo_portfolio.runner_backtest import run_backtest_json
return run_backtest_json(params)
def _positions_map(final: dict | None) -> dict[str, dict]:
fp = (final or {}).get("final_portfolio") or {}
return {p["security"]: p for p in fp.get("positions", [])}
def _rel_diff(a: float, b: float) -> float:
denom = max(abs(a), abs(b), 1e-12)
return abs(a - b) / denom
def compare(full: dict, segmented: dict) -> list[str]:
"""返回问题清单(空 = 通过)。"""
problems: list[str] = []
fpf = (full.get("final_portfolio") or {})
fps = (segmented.get("final_portfolio") or {})
if not fpf or not fps:
return [f"缺 final_portfolio: full={bool(fpf)} segmented={bool(fps)}"]
# 1. 期末持仓
pf, ps = _positions_map(full), _positions_map(segmented)
for code in sorted(set(pf) | set(ps)):
a, b = pf.get(code), ps.get(code)
if a is None or b is None:
problems.append(f"持仓 {code}: 单边缺失 full={a is not None} seg={b is not None}")
continue
if int(a["amount"]) != int(b["amount"]):
problems.append(f"持仓 {code}: 数量 full={a['amount']} seg={b['amount']}")
if _rel_diff(a["avg_cost"], b["avg_cost"]) > TOL:
problems.append(
f"持仓 {code}: 均价 full={a['avg_cost']:.6f} seg={b['avg_cost']:.6f}")
# 2. 现金 / 总值
for key in ("cash", "total_value"):
d = _rel_diff(fpf[key], fps[key])
if d > TOL:
problems.append(f"{key}: full={fpf[key]:.4f} seg={fps[key]:.4f} rel_diff={d:.2e}")
# 3. 净值曲线共同交易日
ef = {p["date"]: float(p["equity"]) for p in full.get("equity_curve") or []}
es = {p["date"]: float(p["equity"]) for p in segmented.get("equity_curve") or []}
common = sorted(set(ef) & set(es))
if not common:
problems.append("净值曲线无共同交易日")
worst, worst_d = "", 0.0
for d in common:
rd = _rel_diff(ef[d], es[d])
if rd > worst_d:
worst, worst_d = d, rd
if worst_d > TOL:
problems.append(f"净值曲线最大偏差 {worst_d:.2e} @ {worst}")
return problems
def main() -> int:
ap = argparse.ArgumentParser(description="影子柜台 P0 对账 spike")
ap.add_argument("--strategy", default="all_weather")
ap.add_argument("--start", default="2024-01-01")
ap.add_argument("--end", default="2024-12-31")
ap.add_argument("--cash", type=float, default=1_000_000.0)
ap.add_argument("--benchmark", default="000300.XSHG")
ap.add_argument("--max-pool", type=int, default=30)
ap.add_argument(
"--split", nargs="*", default=[],
help="段边界日期(升序, 2 个 = 3 段)。缺省自动取 1/3, 2/3 处的月初。")
ap.add_argument("--provider", default="unified")
ap.add_argument("--provider-config", default="{}")
ap.add_argument("--slippage", type=float, default=0.0)
args = ap.parse_args()
base = {
"strategy": args.strategy,
"start_date": args.start,
"end_date": args.end,
"initial_cash": args.cash,
"benchmark": args.benchmark,
"max_pool": args.max_pool,
"commission_rate": 0.0003,
"stamp_duty_rate": 0.001,
"min_commission": 5.0,
"slippage": args.slippage,
"provider": args.provider,
"provider_config": args.provider_config,
}
splits = args.split or _auto_splits(args.start, args.end)
print(f"[spike] 全量重放 {args.start} ~ {args.end} ...", flush=True)
full = _run_segment(base)
# 分段续跑: 每段用上一段期末现金 + 期末持仓恢复
bounds = [args.start] + list(splits) + [args.end]
carry_cash, carry_positions = args.cash, None
seg_results = []
for i in range(len(bounds) - 1):
seg = dict(base)
seg["start_date"], seg["end_date"] = bounds[i], bounds[i + 1]
seg["initial_cash"] = carry_cash
if carry_positions:
seg["initial_positions"] = carry_positions
print(f"[spike] 段{i + 1}/{len(bounds) - 1}: {seg['start_date']} ~ {seg['end_date']} "
f"cash={carry_cash:.2f} positions={len(carry_positions or [])}", flush=True)
r = _run_segment(seg)
seg_results.append(r)
fp = r.get("final_portfolio") or {}
carry_cash = fp.get("cash", 0.0)
carry_positions = [
{"security": p["security"], "amount": p["amount"], "avg_cost": p["avg_cost"]}
for p in fp.get("positions", [])
]
last_seg = seg_results[-1]
# 分段侧期末 end_date 是最后一段区间,与全量一致
problems = compare(full, last_seg)
fpf = (full.get("final_portfolio") or {})
print("\n[spike] ===== 对账结果 =====", flush=True)
print(f" 全量期末: 现金={fpf.get('cash', 0):.2f} 总值={fpf.get('total_value', 0):.2f} "
f"持仓数={len(fpf.get('positions', []))}", flush=True)
if problems:
for p in problems:
print(f"{p}", flush=True)
print("SPIKE_P0_FAIL", flush=True)
return 1
print(" ✅ 期末持仓逐只一致; 现金/总值/净值曲线相对偏差 < 1e-6", flush=True)
print("SPIKE_P0_PASS", flush=True)
return 0
def _auto_splits(start: str, end: str) -> list[str]:
"""缺省段边界: 区间 1/3 与 2/3 处最近月初(月度调仓策略段边界取月初更干净)。"""
from datetime import date
def to_date(s: str) -> date:
y, m, d = map(int, s.split("-"))
return date(y, m, d)
s, e = to_date(start), to_date(end)
total = (e - s).days
marks = []
for frac in (1 / 3, 2 / 3):
target = s + __import__("datetime").timedelta(days=int(total * frac))
# 回退到月初
first = target.replace(day=1)
marks.append(first.isoformat())
return marks
if __name__ == "__main__":
sys.exit(main())