价格行为分析工具包开发(第 35 部分):预测模型训练与部署·进阶篇
「信号落地:箭头、声光与下单的衔接」
收到服务端信号后,EA 先在图表上清理带指定前缀的旧对象,再按信号类型画箭头:买信号用 Wingdings 码 233、卖信号 234、平仓 158,颜色分别走 ColorBuy / ColorSell / ColorClose。箭头锚在 TimeCurrent() 的 BID 价上,创建成功会播放 alert.wav,视觉与听觉同时给交易者一个可验证的触发点。 SL/TP 水平线仅在 DrawSLTPLines 开启且 m.sl / m.tp 大于 0 时绘制,分别挂 OBJ_HLINE 到对应价格。若 EnableTrading 为真,逻辑会先 PositionSelect(_Symbol) 判断有无持仓:无仓时 SIG_BUY / SIG_SELL 各用 FixedLots 开仓并带入 sl、tp;有仓且收到 SIG_CLOSE 则按 SlippagePoints 平仓。外汇与贵金属杠杆高,实盘开启自动交易前建议先在 MT5 策略测试器跑一遍信号去重与对象清理逻辑。 后端 ML 引擎用 Python 写,训练时 train() 会直接丢弃异常行,回测默认窗口 30 天,SL/TP 优先取 ATR 失败则回退固定值;这些设定决定了信号里 sl/tp 字段可能为空,EA 端 m.sl>0 的判断就是防画线报错的关键。
out = StringToDouble(StringSubstr(txt, p + StringLen(key))); } class="type">void ActOnSignal(class="kw">const SServerMsg &m) { class="kw">static ESignal last = SIG_WAIT; if(m.code == SIG_WAIT || m.code == last) class="kw">return; last = m.code; class=class="str">"cmt">// remove old objects for(class="type">int i=ObjectsTotal(class="num">0)-class="num">1;i>=class="num">0;--i) if(StringFind(ObjectName(class="num">0,i),objPrefix)==class="num">0) ObjectDelete(class="num">0,ObjectName(class="num">0,i)); class=class="str">"cmt">// draw arrow class="type">int arrow = (m.code==SIG_BUY ? class="num">233 : m.code==SIG_SELL ? class="num">234 : class="num">158); class="type">class="kw">color clr = (m.code==SIG_BUY ? ColorBuy : m.code==SIG_SELL ? ColorSell : ColorClose); class="type">class="kw">string id = objPrefix + "Arr_" + TimeToString(TimeCurrent(),TIME_SECONDS); class="type">class="kw">double y = SymbolInfoDouble(_Symbol, SYMBOL_BID); if(ObjectCreate(class="num">0,id,OBJ_ARROW,class="num">0,TimeCurrent(),y)) { ObjectSetInteger(class="num">0,id,OBJPROP_ARROWCODE,arrow); ObjectSetInteger(class="num">0,id,OBJPROP_COLOR,clr); ObjectSetInteger(class="num">0,id,OBJPROP_WIDTH,ArrowSize); PlaySound("alert.wav"); } class=class="str">"cmt">// draw SL/TP lines if(DrawSLTPLines && m.sl>class="num">0) ObjectCreate(class="num">0,objPrefix+"SL_"+id,OBJ_HLINE,class="num">0,class="num">0,m.sl); if(DrawSLTPLines && m.tp>class="num">0) ObjectCreate(class="num">0,objPrefix+"TP_"+id,OBJ_HLINE,class="num">0,class="num">0,m.tp); class=class="str">"cmt">// execute trade if(EnableTrading) { class="type">bool hasPos = PositionSelect(_Symbol); if(m.code==SIG_BUY && !hasPos) trade.Buy(FixedLots,_Symbol,class="num">0,m.sl,m.tp); if(m.code==SIG_SELL && !hasPos) trade.Sell(FixedLots,_Symbol,class="num">0,m.sl,m.tp); if(m.code==SIG_CLOSE&& hasPos) trade.PositionClose(_Symbol,SlippagePoints); } }
◍ 合成指数行情采集与 Prophet 增量建模
这套采集脚本面向 Boom 900、Crash 1000 以及 Volatility 75 (1s) 三类合成指数,LOOKAHEAD 固定为 10 分钟,阈值 THRESH_LABEL=0.0015(即 0.15%)用来给训练集打标签,STEP_SECONDS=60 代表实盘每 60 秒落一条样本。 ATR 周期取 14,止损倍率 SL_MULT=1.0、止盈倍率 TP_MULT=2.0,当 ATR 计算失效时回退到 ATR_FALLBACK_P=0.002 作为波动率代理,这些都是后续小布盯盘模块直接读取的常量。 CSV_HEADER 定义了 12 个字段:timestamp、symbol、price、spike_mag、macd、rsi、atr、slope、env_low、env_up、delta、label,意味着每条记录同时保留动量、超买超卖与包络通道偏移,方便用监督学习区分噪音 spike 与真实突破。 Prophet 建模走的是懒加载:prophet_delta() 发现某品种缓存为空且序列长度 ≥20 时,才丢到守护线程里 _compile_prophet 拟合,避免主线程阻塞;拟合后写入 _PROP 字典并打时间戳,下次同品种直接复用。 MT5 连接由 init_mt5() 负责,先用默认参数 mt5.initialize() 试连,失败再带 TERM_PATH、LOGIN、PASSWORD、SERVER 显式初始化,任何失败直接 sys.exit 并打印 last_error(),实盘跑之前务必先确认这五个变量已填好。
class="kw">import os,sys,time,logging,warnings,argparse,threading,io class="kw">import class="type">class="kw">datetime as dt from pathlib class="kw">import Path class="kw">import numpy as np, pandas as pd, ta, joblib, pytz from flask class="kw">import Flask, request, jsonify, abort from prophet class="kw">import Prophet from pykalman class="kw">import KalmanFilter class="kw">import MetaTrader5 as mt5 warnings.filterwarnings("ignore") logging.basicConfig(level=logging.INFO, format="%(asctime)s %(levelname)-7s %(message)s", datefmt="%H:%M:%S") Path(MODEL_DIR).mkdir(parents=True, exist_ok=True) os.chdir(BASE_DIR) _mt5_lock = threading.Lock() def init_mt5(): if mt5.initialize(): class="kw">return if not mt5.initialize(path=TERM_PATH, login=LOGIN, password=PASSWORD, server=SERVER): sys.exit(f"MT5 init failed {mt5.last_error()}") def ensure_symbol(sym): class="kw">return mt5.symbol_select(sym, True) _PROP_LOCK = threading.Lock() _PROP = {} # sym -> (model, timestamp) or None def _compile_prophet(df, sym): mdl = Prophet(daily_seasonality=False, weekly_seasonality=False) mdl.fit(df) with _PROP_LOCK: _PROP[sym] = (mdl, time.time()) def prophet_delta(prices, times, sym): if len(prices) < class="num">20: class="kw">return class="num">0.0 with _PROP_LOCK: entry = _PROP.get(sym) if entry is None: _PROP[sym] = None df = pd.DataFrame({"ds": pd.to_datetime(times, unit=&class="macro">#x27;s&class="macro">#x27;), "y": prices}) threading.Thread(target=_compile_prophet, args=(df, sym), daemon=True).start() class="kw">return class="num">0.0 mdl, ts = entry
把异常波动和标签写进特征行
这段逻辑干的事很直接:用 z-score、MACD 柱差、RSI 和 Prophet 残差拼出一个复合异动分数,再按未来 N 根 K 线的相对涨跌贴上 BUY / SELL / WAIT 标签。外汇与贵金属波动受杠杆放大,按此逻辑回测仅代表历史样本倾向,实盘高风险。 z_spike 取最近 20 根收益率序列,用 (末值-均值)/标准差 算 z,绝对值超 2.5 才认作尖峰;combo_spike 把 z、MACD 差、近 4 根价差归一后相加,阈值 3.0 以上才算复合异动。gen_row 里若 high/low 缺失则 ATR 填 0,否则用 average_true_range 取末值。 标签侧以 LOOKAHEAD 根后的收盘价算相对变化 ch,超 THRESH_LABEL 为 BUY、低于负阈值为 SELL,否则 WAIT。把这段代码直接丢进你的 Python 特征工程脚本,调一下 combo_spike 的 3.0 和 z_spike 的 2.5,看样本命中率怎么漂。
if time.time() - ts > class="num">3600: with _PROP_LOCK: _PROP[sym] = None class="kw">return class="num">0.0 fut = mdl.make_future_dataframe(periods=class="num">1, freq=&class="macro">#x27;s&class="macro">#x27;) class="kw">return class="type">class="kw">float(mdl.predict(fut).iloc[-class="num">1]["yhat"] - prices[-class="num">1]) def z_spike(prices, win=class="num">20): if len(prices) < win: class="kw">return False, class="num">0.0 r = np.diff(prices[-win:]) z = (r[-class="num">1] - r.mean())/(r.std()+class="num">1e-6) class="kw">return abs(z) > class="num">2.5, class="type">class="kw">float(z) def macd_div(prices): if len(prices) < class="num">35: class="kw">return class="num">0.0 class="kw">return class="type">class="kw">float(ta.trend.macd_diff(pd.Series(prices)).iloc[-class="num">1]) def rsi_val(prices, l=class="num">14): if len(prices) < l+class="num">1: class="kw">return class="num">50.0 class="kw">return class="type">class="kw">float(ta.momentum.rsi(pd.Series(prices), l).iloc[-class="num">1]) def combo_spike(prices): _, z = z_spike(prices) m = macd_div(prices) v = prices[-class="num">1] - prices[-class="num">4] if len(prices) >= class="num">4 else class="num">0.0 s = abs(z) + abs(m) + abs(v)/(np.std(prices[-class="num">20:])+class="num">1e-6) class="kw">return s > class="num">3.0, s def append_rows(rows): if not rows: class="kw">return pd.DataFrame(rows, columns=CSV_HEADER)\ .to_csv(CSV_FILE, mode="a", index=False, header=not Path(CSV_FILE).exists()) def gen_row(i, closes, times, sym, highs=None, lows=None): if i < LOOKAHEAD or i+LOOKAHEAD >= len(closes): class="kw">return None seq = closes[:i] _, mag = combo_spike(seq) atr = ta.volatility.average_true_range(pd.Series(highs[:i+class="num">1]), pd.Series(lows[:i+class="num">1]), pd.Series(seq)).iloc[-class="num">1] if highs else class="num">0.0 row = [ times[i], sym, closes[i], mag, macd_div(seq), rsi_val(seq), atr, class="num">0.0, class="num">0.0, class="num">0.0, prophet_delta(seq, times[:i], sym) ] ch = (closes[i+LOOKAHEAD] - closes[i]) / closes[i] row.append("BUY" if ch > THRESH_LABEL else "SELL" if ch < -THRESH_LABEL else "WAIT") class="kw">return row def collect_loop(): if not Path(CSV_FILE).exists(): append_rows([]) last = {}
「实时抓取与模型训练的衔接写法」
这段逻辑把 MT5 的实时分钟线拉取、历史回补和梯度提升模型训练串在了一条流水线上。实时循环里用 copy_rates_from_pos 取 M1 数据,LOOKAHEAD+1 根bar 的长度保证能算出未来偏移标签;若某品种最后一根 bar 的时间没变就跳过,避免重复写行。 历史补数走 copy_rates_range,传入带 UTC 的起止时间,取回后同时抽出 close/high/low 三个序列,再用列表推导批量生成特征行。实盘里若某品种返回 0 行会直接 return,不会污染 CSV。 训练侧用 StandardScaler 加 GradientBoostingClassifier(n_estimators=400,learning_rate=0.05,max_depth=3,random_state=42)封成 Pipeline。train_models 读 CSV 后按品种切分,样本少于 400 行就跳过——外汇与贵金属波动随机性高,小样本拟合极易过拟合,这个阈值算是一道底线。 开 MT5 把 SYMBOLS 换成你盯的 XAUUSD/EURUSD,STEP_SECONDS 设成 60,就能边跑边落盘验证特征行是否随 tick 更新。
print("Collecting… CTRL-C to stop") init_mt5() class="kw">while True: for sym in SYMBOLS: if not ensure_symbol(sym): class="kw">continue bars = mt5.copy_rates_from_pos(sym, mt5.TIMEFRAME_M1, class="num">0, LOOKAHEAD+class="num">1) if bars is None or len(bars) < LOOKAHEAD+class="num">1: class="kw">continue if last.get(sym) == bars[-class="num">1][&class="macro">#x27;time&class="macro">#x27;]: class="kw">continue last[sym] = bars[-class="num">1][&class="macro">#x27;time&class="macro">#x27;] closes = bars[&class="macro">#x27;close&class="macro">#x27;].tolist() times = bars[&class="macro">#x27;time&class="macro">#x27;].tolist() row = gen_row(len(closes)-LOOKAHEAD-class="num">1, closes, times, sym) if row: append_rows([row]) time.sleep(STEP_SECONDS) def history_from_mt5(sym, start, end): init_mt5() r = mt5.copy_rates_range(sym, mt5.TIMEFRAME_M1, start.replace(tzinfo=UTC), end.replace(tzinfo=UTC)) if r is None or len(r)==class="num">0: class="kw">return closes, times = r[&class="macro">#x27;close&class="macro">#x27;].tolist(), r[&class="macro">#x27;time&class="macro">#x27;].tolist() highs, lows = r[&class="macro">#x27;high&class="macro">#x27;].tolist(), r[&class="macro">#x27;low&class="macro">#x27;].tolist() rows = [gen_row(i, closes, times, sym, highs, lows) for i in range(len(closes)-LOOKAHEAD) if gen_row(i, closes, times, sym, highs, lows)] append_rows([rw for rw in rows if rw]) print(sym, "imported", len(rows), "rows") def build_pipe(X, y): pipe = Pipeline([ ("sc", StandardScaler()), ("gb", GradientBoostingClassifier(n_estimators=class="num">400, learning_rate=class="num">0.05, max_depth=class="num">3, random_state=class="num">42)) ]) class="kw">return pipe.fit(X, y) def train_models(): df = pd.read_csv(CSV_FILE) df = df.dropna(subset=FEATURES) for sym in SYMBOLS: d = df[df.symbol == sym] if len(d) < class="num">400: class="kw">continue
◍ 把训练好的模型挂成在线信号接口
模型训练完不是终点。上面这段把每个品种的管道模型用 joblib 落盘成 品种名.pkl,全局模型存为单一文件,之后 Flask 服务直接按 symbol 加载,无需重训。
服务开了三个 POST 入口:/upload_history 接收 MT5 导出的 close/time(缺省用 close 填 high/low),/upload_spike_csv 吃 EA 吐的逗号分隔 spikes,/analyze 实时算特征后调 predict_proba 出 signal、sl、tp 和 strength(取三类概率最大值)。
回测函数 backtest_one 离线复用同一套开平仓逻辑,保证样本外检验和实盘推断一致;info() 打印总行数与标签分布,并遍历模型目录报出每个 pkl 的特征数(来自 named_steps['sc'].n_features_in_)。
命令行入口用 subparsers 切成 collect / history / train / backtest / serve / info 六种模式。开 MT5 导出一段 EURUSD 历史,跑 python main.py serve 再 POST 到 /analyze,就能看到该品种当前的概率强度,外汇贵金属波动剧烈,信号仅作概率参考。
model = build_pipe(d[FEATURES], d.label.map({"WAIT":class="num">0,"BUY":class="num">1,"SELL":class="num">2})) joblib.dump(model, Path(MODEL_DIR)/f"{sym.replace(&class="macro">#x27; &class="macro">#x27;,&class="macro">#x27;_&class="macro">#x27;)}.pkl") global_model = build_pipe(df[FEATURES], df.label.map({"WAIT":class="num">0,"BUY":class="num">1,"SELL":class="num">2})) joblib.dump(global_model, GLOBAL_PKL) app = Flask(__name__) app.config["MAX_CONTENT_LENGTH"] = class="num">32*class="num">1024*class="num">1024 @app.route("/upload_history", methods=["POST"]) def upload_history(): j = request.get_json(force=True) close, ts = np.array(j["close"]), np.array(j["time"],dtype=class="type">int) high = np.array(j.get("high", close)) low = np.array(j.get("low", close)) df = pd.DataFrame({"timestamp": ts, "price": close}) # compute features as in gen_row… append_rows(df.assign(symbol=j["symbol"]).values.tolist()) class="kw">return jsonify(status="ok", rows_written=len(df)) @app.route("/upload_spike_csv", methods=["POST"]) def upload_spike_csv(): j = request.get_json(force=True) df_ea = pd.read_csv(io.StringIO(j.get("csv","")), sep=",") # map EA columns → CSV_HEADER append_rows(mapped_rows) class="kw">return jsonify(status="ok", rows_written=len(mapped_rows)) @app.route("/analyze", methods=["POST"]) def api_analyze(): j = request.get_json(force=True) mdl = load_model(j["symbol"]) feats = [...] # compute from j["prices"], j["timestamps"] proba = mdl.predict_proba([feats])[class="num">0] signal = decide_open(proba[class="num">1], proba[class="num">2], j["symbol"]) # build sl, tp, manage _trades… class="kw">return jsonify(signal=signal, sl=sl, tp=tp, strength=max(proba)) def backtest_one(sym, df): mdl = load_model(sym) for i in range(len(df)): feats = [...] # offline feature calcs pr = mdl.predict_proba([feats])[class="num">0] # open/close logic identical to /analyze class="kw">return trades def info(): df = pd.read_csv(CSV_FILE) print("Rows:", len(df), "Labels:", df.label.value_counts()) for pkl in Path(MODEL_DIR).glob("*.pkl"): mdl = joblib.load(pkl) print(pkl.name, "features", mdl.named_steps["sc"].n_features_in_) if __name__ == "__main__": parser = argparse.ArgumentParser() subs = parser.add_subparsers(dest="mode", required=True) subs.add_parser("collect") subs.add_parser("history") subs.add_parser("train") subs.add_parser("backtest") subs.add_parser("serve") subs.add_parser("info") args = parser.parse_args()
把调度入口跑通才算闭环
这段 Python 调度分支是整条 MT5 数据链路的开关,没把它跑起来,前面采集、训练、回测都只是散件。 if args.mode == "collect": 收到 collect 指令就先 init_mt5() 登录终端,再进 collect_loop() 持续拉 tick;elif 收到 history 就走 history_cli(args) 按参数补历史;train 调 train_models() 跑模型;backtest 交 backtest_cli(args) 做推演。 serve 模式会 init_mt5() 后把 Flask 绑到 0.0.0.0:5000 且开 threaded=True,意味着本机以外也能连推理接口;info 则直接打环境信息。 开命令行敲 python main.py --mode info 先确认终端版本和品种权限,再切 collect 看 MT5 终端日志里 tick 数是否每秒增长,外汇与贵金属杠杆高,实盘前务必在策略测试器用历史数据验证一遍。
if args.mode == "collect": init_mt5(); collect_loop() elif args.mode == "history": history_cli(args) elif args.mode == "train": train_models() elif args.mode == "backtest": backtest_cli(args) elif args.mode == "serve": init_mt5(); app.run("class="num">0.0.class="num">0.0", class="num">5000, threaded=True) elif args.mode == "info": info()