183 lines
7.0 KiB
Python
183 lines
7.0 KiB
Python
# -*- coding: utf-8 -*-
|
||
"""影子柜台 P0 对账 spike(docs/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())
|