Files
Chan/research/step46_engine_parity.py
T
jackyu66gitandCursor 0b4d7693b8 缠论引擎提速 2.6x,瓶颈是逐行 Series 查找而非指标计算
原以为浪费在 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>
2026-08-28 04:04:44 +08:00

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()