原以为浪费在 add_indicators 算了太多用不到的指标,实测它只占全量构建的 1.3%——talib 是向量化 C 代码,便宜。真正的两处: cal_kl_data 占 96%:每根 K 线 df.iloc[i] 新建一个 40 列 Series,再在其上做 几十次逐键查找。改为预取 ndarray 后 2 万根 1946ms → 824ms。 ChanKLC.cal_all_ema_status 占 25%:每次合并 KLU 都立即重算,而它产出的 ema_status / ema52_pos / ema52_status 全仓无任何读取方(含前端)。改为惰性 求值,保留属性形式以防将来有人读。顺带删掉 get_klc_list 里累加一整轮后直接 丢弃的 ema_up_list / ema_down_list。 另加 TF_DF(lean=True):只构建到中枢,跳过线段/走势中枢/MACD 状态机——这些 只服务 bsp_list 与 web 展示,笔和中枢不依赖。研究与实盘走这条快 3.6x。 结果 2 万根 5m:full 1946 → 754ms,lean → 543ms。 step46_engine_parity.py 是配套的安全网,改引擎前先跑一次 --save。它对 KLC 端点与分型、笔起止价与 is_sure、中枢 zg/zd/available_ts/阶梯、信号全部输出列, 以及 26 个被下游消费的 dataframe 列取哈希。本次三处改动逐步验证,另用 git stash 切回改动前代码在 20 万根 × 5 用例上做了跨版本逐位对拍,全部一致; 增量路径与 web API 也各验一遍。 Co-authored-by: Cursor <cursoragent@cursor.com>
201 lines
7.4 KiB
Python
201 lines
7.4 KiB
Python
"""Step 46:缠论引擎的逐位对拍基线。
|
|
|
|
重构引擎前先固化一份指纹,改完再对一次。没有它,任何"精简"都无法证明
|
|
没有改变行为——而 HANDOFF 里所有回测数字都绑定当前实现,**行为变化是静默的**:
|
|
不报错、不崩溃,只是信号悄悄变了一批。
|
|
|
|
指纹覆盖三条链路各自依赖的东西:
|
|
|
|
结构 klu / klc / bi / bi_zs / seg 的数量与关键端点
|
|
信号 中枢表(zg/zd/available_ts) 与 fast_bsp3 的全部输出列
|
|
数值 dataframe 上被下游真正消费的列(逐位比较)
|
|
|
|
用法:
|
|
python step46_engine_parity.py --save # 改动前,存基线
|
|
python step46_engine_parity.py --check # 改动后,对比
|
|
"""
|
|
from __future__ import annotations
|
|
|
|
import argparse
|
|
import hashlib
|
|
import json
|
|
import sys
|
|
import warnings
|
|
from pathlib import Path
|
|
|
|
import numpy as np
|
|
import pandas as pd
|
|
|
|
warnings.filterwarnings("ignore")
|
|
HERE = Path(__file__).resolve().parent
|
|
sys.path.insert(0, str(HERE))
|
|
sys.path.insert(0, str(HERE.parent))
|
|
pd.set_option("display.width", 300)
|
|
|
|
BASELINE = HERE / "out" / "step46_baseline.json"
|
|
|
|
# 覆盖面:多币多级别,1m 用较短窗口以免跑太久
|
|
CASES = [
|
|
("BTC/USDT:USDT", "1m", 20_000),
|
|
("BTC/USDT:USDT", "5m", 20_000),
|
|
("ETH/USDT:USDT", "5m", 20_000),
|
|
("SOL/USDT:USDT", "15m", 20_000),
|
|
("XRP/USDT:USDT", "30m", 20_000),
|
|
]
|
|
|
|
# 下游真正消费的列(见 §5.5 的列使用扫描)。精简若动到这些,必须体现在指纹里。
|
|
CONSUMED = ["open", "high", "low", "close", "volume", "atr",
|
|
"macd", "macdsignal", "macdhist",
|
|
"ema5", "ema13", "ema24", "ema26", "ema52", "ema104", "ema156", "ema208", "ema7",
|
|
"rsi", "volume_ratio",
|
|
"bb2633upper", "bb2633lower", "bb2633middle",
|
|
"bbp30", "bbp120", "bbp365"]
|
|
|
|
|
|
def _h(arr) -> str:
|
|
a = np.asarray(arr, dtype=np.float64)
|
|
a = np.nan_to_num(a, nan=-9.87654321e30, posinf=1e300, neginf=-1e300)
|
|
return hashlib.sha256(a.tobytes()).hexdigest()[:16]
|
|
|
|
|
|
def fingerprint(pair: str, tf: str, rows: int) -> dict:
|
|
from chanlun import TF_DF
|
|
from chanlun.analysis.fast_bsp import (
|
|
add_zone_ladder, build_htf_zones, find_fast_bsp3,
|
|
)
|
|
from lib.data import fetch_ohlcv
|
|
|
|
df = fetch_ohlcv(pair, tf, rows)
|
|
chan = TF_DF(df, 1, tf)
|
|
cdf = chan.dataframe
|
|
|
|
fp: dict = {"n_rows": int(len(cdf))}
|
|
|
|
# --- 结构 ---
|
|
fp["n_klu"] = len(getattr(chan, "klu_list", []) or [])
|
|
fp["n_klc"] = len(getattr(chan, "klc_list", []) or [])
|
|
fp["n_bi"] = len(getattr(chan, "bi_list", []) or [])
|
|
fp["n_seg"] = len(getattr(chan, "seg_list", []) or [])
|
|
fp["n_bi_zs"] = len(getattr(chan, "bi_zs_list", []) or [])
|
|
fp["n_zs"] = len(getattr(chan, "zs_list", []) or [])
|
|
fp["n_bsp"] = len(getattr(chan, "bsp_list", []) or [])
|
|
|
|
# KLC 端点(包含关系的结果,最容易被指标改动影响)
|
|
klc = getattr(chan, "klc_list", []) or []
|
|
fp["klc_high"] = _h([k.high for k in klc])
|
|
fp["klc_low"] = _h([k.low for k in klc])
|
|
fp["klc_fx"] = _h([float(getattr(k.fx, "value", 0) or 0) for k in klc])
|
|
|
|
# 笔端点
|
|
bi = getattr(chan, "bi_list", []) or []
|
|
fp["bi_start"] = _h([float(getattr(b, "start_price", 0) or 0) for b in bi])
|
|
fp["bi_end"] = _h([float(getattr(b, "end_price", 0) or 0) for b in bi])
|
|
fp["bi_sure"] = _h([1.0 if getattr(b, "is_sure", False) else 0.0 for b in bi])
|
|
|
|
# --- 信号 ---
|
|
zones = build_htf_zones(cdf, tf, chan=chan)
|
|
fp["n_zones"] = int(len(zones))
|
|
if len(zones):
|
|
zl = add_zone_ladder(zones.reset_index(drop=True))
|
|
fp["zone_zg"] = _h(zl["zg"])
|
|
fp["zone_zd"] = _h(zl["zd"])
|
|
fp["zone_avail"] = _h(zl["available_ts"])
|
|
fp["zone_ladder"] = _h(zl["z_above"].astype(float) * 2 + zl["z_below"].astype(float))
|
|
sig = find_fast_bsp3(cdf, zl)
|
|
fp["n_sig"] = int(len(sig))
|
|
for c in ("entry_idx", "direction", "bo_idx", "pb_idx", "lag", "depth", "zone_i"):
|
|
if c in sig.columns:
|
|
fp[f"sig_{c}"] = _h(sig[c])
|
|
else:
|
|
fp["n_sig"] = 0
|
|
|
|
# --- 数值列(只对下游消费的列逐位比较)---
|
|
for c in CONSUMED:
|
|
fp[f"col_{c}"] = _h(cdf[c]) if c in cdf.columns else "MISSING"
|
|
fp["_all_columns"] = sorted(map(str, cdf.columns))
|
|
return fp
|
|
|
|
|
|
def collect() -> dict:
|
|
out = {}
|
|
for pair, tf, rows in CASES:
|
|
key = f"{pair.split('/')[0]}_{tf}"
|
|
print(f" 计算 {key} ...", flush=True)
|
|
try:
|
|
out[key] = fingerprint(pair, tf, rows)
|
|
except Exception as e:
|
|
out[key] = {"error": repr(e)[:200]}
|
|
print(f" 失败: {e!r}", flush=True)
|
|
return out
|
|
|
|
|
|
def main() -> None:
|
|
ap = argparse.ArgumentParser()
|
|
ap.add_argument("--save", action="store_true", help="存为基线")
|
|
ap.add_argument("--check", action="store_true", help="与基线对比")
|
|
ap.add_argument("--rows", type=int, default=0, help="覆盖各用例根数(验证大样本时用)")
|
|
ap.add_argument("--out", default="", help="基线文件名,便于大小样本各存一份")
|
|
args = ap.parse_args()
|
|
|
|
global BASELINE, CASES
|
|
if args.out:
|
|
BASELINE = HERE / "out" / args.out
|
|
if args.rows:
|
|
CASES = [(p, tf, args.rows) for p, tf, _ in CASES]
|
|
if not (args.save or args.check):
|
|
ap.error("需要 --save 或 --check")
|
|
|
|
BASELINE.parent.mkdir(exist_ok=True)
|
|
print(f"[引擎对拍] {len(CASES)} 个用例\n")
|
|
cur = collect()
|
|
|
|
if args.save:
|
|
BASELINE.write_text(json.dumps(cur, ensure_ascii=False, indent=1))
|
|
print(f"\n基线已存:{BASELINE}")
|
|
for k, v in cur.items():
|
|
if "error" in v:
|
|
continue
|
|
print(f" {k}: klu {v['n_klu']} klc {v['n_klc']} bi {v['n_bi']} "
|
|
f"中枢 {v['n_zones']} 信号 {v['n_sig']} 列数 {len(v['_all_columns'])}")
|
|
return
|
|
|
|
if not BASELINE.exists():
|
|
print(f"基线不存在:{BASELINE},先跑 --save")
|
|
return
|
|
old = json.loads(BASELINE.read_text())
|
|
|
|
print("\n" + "=" * 90)
|
|
bad = 0
|
|
for key in sorted(set(old) | set(cur)):
|
|
o, n = old.get(key), cur.get(key)
|
|
if o is None or n is None:
|
|
print(f"❌ {key}: 用例缺失")
|
|
bad += 1
|
|
continue
|
|
# 列集合单独看:删列是预期内的,不算行为变化
|
|
o_cols, n_cols = set(o.get("_all_columns", [])), set(n.get("_all_columns", []))
|
|
diffs = [k for k in o if k != "_all_columns" and o.get(k) != n.get(k)]
|
|
dropped, added = sorted(o_cols - n_cols), sorted(n_cols - o_cols)
|
|
# 被删列在指纹里会变成 MISSING,若该列本就不被消费则无害
|
|
harmful = [d for d in diffs if not (d.startswith("col_") and n.get(d) == "MISSING"
|
|
and d[4:] not in CONSUMED)]
|
|
if not harmful:
|
|
print(f"✅ {key}: 行为一致"
|
|
+ (f"(删列 {len(dropped)} 个)" if dropped else ""))
|
|
else:
|
|
bad += 1
|
|
print(f"❌ {key}: {len(harmful)} 项不一致")
|
|
for d in harmful[:12]:
|
|
print(f" {d}: {o.get(d)} → {n.get(d)}")
|
|
if dropped:
|
|
print(f" 删掉的列: {dropped}")
|
|
if added:
|
|
print(f" 新增的列: {added}")
|
|
|
|
print("\n" + ("✅ 全部用例行为一致,可以放心继续" if not bad
|
|
else f"❌ {bad} 个用例有行为变化——**回测数字已失效,不要继续**"))
|
|
|
|
|
|
if __name__ == "__main__":
|
|
main()
|