在类中包装 ONNX 模型(基础篇)
把 ONNX 模型封进一个 C++ 类
在 MT5 里直接调 ONNX 推理,最干净的做法不是散着写,而是用 C++ 类把模型生命周期包起来:加载、会话创建、输入输出张量绑定全部收进同一个对象。这样 EA 主逻辑只管「喂数据、取结果」,不碰底层 API。 官方示例发布于 2023 年 11 月 6 日,至今在终端内被查看 1286 次、收藏 3 次,说明落地需求真实存在。外汇与贵金属行情高杠杆、滑点不可控,任何模型推理都只是概率辅助,不能直接当下单依据。 类里至少留三个接口:Load(path) 负责读 .onnx 文件,Run(inputs) 做前向推理,GetOutput() 回传预测张量。把 OnnxRuntime 的 Ort::Session 藏成私有成员,EA 层就不需要 include 任何 ORT 头文件。
◍ 为什么把 ONNX 投票分类器改成面向对象
上一篇文章里,两个 ONNX 模型做排列投票分类器时,全部源码被塞进了一个 MQ5 文件,逻辑靠多个函数切分。这种写法在模型数量固定时还能跑,但一旦要换模型位次、或往里再加别的 ONNX 模型,源文件会迅速膨胀,函数之间耦合变重。 外汇与贵金属行情受宏观事件冲击大,模型迭代频率高,硬编码式扩展在这里风险不小。把分类器拆成对象,每个模型自己管加载和推理,主流程只负责调度,后续加模型基本不用动旧代码。 所以这一节直接转向面向对象:用类封装单个 ONNX 模型实例,投票器作为容器持有它们。这样模型换位只是改一下容器里的顺序,代码体积和心智负担都能控住。
「三类分类模型与训练数据底数」
上一版投票分类器里混用了回归与分类模型,回归部分直接用预测价格代替走势方向,代价是拿不到概率分布,软投票逻辑因此受限。这次改用三个纯分类模型,绕开这个坑。 前两个模型出自 ONNX 集成示例:第一个由回归改分类,喂了 10 条 OHLC 价格序列;第二个原生分类,用 63 条收盘价序列训练。第三个模型额外叠了周期 21 和 34 的 SMA 序列,连同 30 条收盘价一起训练,均线叉口形态全交给网络自己记权重,不人为预设。 三个模型的训练窗口统一为 MetaQuotes-Demo 的 EURUSD D1,区间 2010.01.01–2023.01.01,共约 13 年日线。训练脚本是 Python 写的,本文不贴源码,避免冲淡 MQL5 集成主线。外汇与贵金属属高风险品种,模型在历史样本上的表现不代表未来概率。
给三类模型抽一个公共父类
做价格预测通常会同时跑回归模型和分类模型,输入数据的尺度和预处理逻辑各不相同,但对外暴露的调用方式应当一致。把共性收进一个基类,衍生类只需补上 PredictPrice 或 PredictClass,EA 层就不用关心背后是哪种网络。 基类在构造时锁定训练数据对应的品种与周期,例如 EURUSD 的 M15,并顺手校验挂载 EA 的图表周期是否匹配,避免用错时间框架导致推断偏移。它还负责创建 ONNX 会话句柄,且只在每根 K 线开盘时触发一次逻辑,m_next_bar 记录了下一根 bar 的时间戳。 分类标签用宏写死成三个整数:PRICE_UP=0、PRICE_SAME=1、PRICE_DOWN=2,回归模型则借 m_class_delta(默认 0.0001)判定价格是否算“持平”。外汇与贵金属杠杆高、滑点突变频繁,周期错配会让模型输出失去参考意义,实盘前务必在 MT5 策略测试器里核对 symbol/period 字段。 下面这段头文件骨架可直接拷进 MetaEditor 建一个 ModelSymbolPeriod.mqh,重点看 protected 里的成员如何承接品种、周期与会话生命周期。
class=class="str">"cmt">//+------------------------------------------------------------------+ class=class="str">"cmt">//| ModelSymbolPeriod.mqh | class=class="str">"cmt">//| Copyright class="num">2023, MetaQuotes Ltd. | class=class="str">"cmt">//| [MQL5官方文档] | class=class="str">"cmt">//+------------------------------------------------------------------+ class=class="str">"cmt">//--- price movement prediction class="macro">#define PRICE_UP class="num">0 class="macro">#define PRICE_SAME class="num">1 class="macro">#define PRICE_DOWN class="num">2 class=class="str">"cmt">//+------------------------------------------------------------------+ class=class="str">"cmt">//| Base class for models based on trained symbol and period | class=class="str">"cmt">//+------------------------------------------------------------------+ class CModelSymbolPeriod { class="kw">protected: class="type">long m_handle; class=class="str">"cmt">// created model session handle class="type">class="kw">string m_symbol; class=class="str">"cmt">// symbol of trained data ENUM_TIMEFRAMES m_period; class=class="str">"cmt">// timeframe of trained data class="type">class="kw">datetime m_next_bar; class=class="str">"cmt">// time of next bar(we work at bar begin only) class="type">class="kw">double m_class_delta; class=class="str">"cmt">// delta to recognize "price the same" in regression models class="kw">public: class=class="str">"cmt">//+------------------------------------------------------------------+ class=class="str">"cmt">//| Constructor | class=class="str">"cmt">//+------------------------------------------------------------------+ CModelSymbolPeriod(const class="type">class="kw">string symbol,const ENUM_TIMEFRAMES period,const class="type">class="kw">double class_delta=class="num">0.0001) { m_handle=INVALID_HANDLE; m_symbol=symbol; m_period=period; m_next_bar=class="num">0; m_class_delta=class_delta; } class=class="str">"cmt">//+------------------------------------------------------------------+ class=class="str">"cmt">//| Destructor | class=class="str">"cmt">//+------------------------------------------------------------------+ ~CModelSymbolPeriod(class="type">void) { Shutdown(); } class=class="str">"cmt">//+------------------------------------------------------------------+ class=class="str">"cmt">//| class="kw">virtual stub for Init | class=class="str">"cmt">//+------------------------------------------------------------------+
◍ ONNX 模型的初始化与心跳节拍控制
把训练好的 ONNX 模型塞进 EA,第一步不是直接预测,而是先做符号与周期的绑定校验。CheckInit 里如果传入的 symbol、period 和类内成员 m_symbol、m_period 对不上,会直接 PrintFormat 报错并返回 false,模型根本不会创建——这能避免把 EURUSD 的模型误跑在 XAUUSD 上,外汇与贵金属品种切换的高风险往往就藏在这种错配里。 模型本体靠 OnnxCreateFromBuffer 从静态字节数组加载,标志位用 ONNX_DEFAULT。若返回 INVALID_HANDLE,说明模型缓冲区损坏或格式不被支持,GetLastError 会给出具体错误码,这时候必须拦住后续逻辑,否则预测会全盘崩。 Shutdown 负责释放会话句柄,只在 m_handle 有效时才调 OnnxRelease,并把句柄复位成 INVALID_HANDLE。漏掉这一步,MT5 反复加载卸载 EA 时显存和句柄会悄悄泄漏。 CheckOnTick 是典型的新 K 线节拍器:用 TimeCurrent 对比 m_next_bar,没到时间直接 return false 不干活;到了就把 m_next_bar 按 PeriodSeconds 对齐到下一根 bar 的开盘时刻。这样预测逻辑天然只在每根新 bar 触发一次,回测和实盘节奏一致。
class="kw">virtual class="type">bool Init(const class="type">class="kw">string symbol,const ENUM_TIMEFRAMES period) { class="kw">return(false); } class=class="str">"cmt">//+------------------------------------------------------------------+ class=class="str">"cmt">//| Check for initialization, create model | class=class="str">"cmt">//+------------------------------------------------------------------+ class="type">bool CheckInit(const class="type">class="kw">string symbol,const ENUM_TIMEFRAMES period,const class="type">uchar& model[]) { class=class="str">"cmt">//--- check symbol, period if(symbol!=m_symbol || period!=m_period) { PrintFormat("Model must work with %s,%s",m_symbol,EnumToString(m_period)); class="kw">return(false); } class=class="str">"cmt">//--- create a model from class="kw">static buffer m_handle=OnnxCreateFromBuffer(model,ONNX_DEFAULT); if(m_handle==INVALID_HANDLE) { Print("OnnxCreateFromBuffer error ",GetLastError()); class="kw">return(false); } class=class="str">"cmt">//--- ok class="kw">return(true); } class=class="str">"cmt">//+------------------------------------------------------------------+ class=class="str">"cmt">//| Release ONNX session | class=class="str">"cmt">//+------------------------------------------------------------------+ class="type">void Shutdown(class="type">void) { if(m_handle!=INVALID_HANDLE) { OnnxRelease(m_handle); m_handle=INVALID_HANDLE; } } class=class="str">"cmt">//+------------------------------------------------------------------+ class=class="str">"cmt">//| Check for class="kw">continue OnTick | class=class="str">"cmt">//+------------------------------------------------------------------+ class="kw">virtual class="type">bool CheckOnTick(class="type">void) { class=class="str">"cmt">//--- check new bar if(TimeCurrent()<m_next_bar) class="kw">return(false); class=class="str">"cmt">//--- set next bar time m_next_bar=TimeCurrent(); m_next_bar-=m_next_bar%PeriodSeconds(m_period); m_next_bar+=PeriodSeconds(m_period); class=class="str">"cmt">//--- work on new day bar class="kw">return(true); } class=class="str">"cmt">//+------------------------------------------------------------------+ class=class="str">"cmt">//| class="kw">virtual stub for PredictPrice(regression model) |
「把回归预测塞进分类框架的收口写法」
预测模块收尾时,通常用两个虚函数把“算价格”和“判方向”拆开。PredictPrice 默认返回 DBL_MAX,意味着基类不提供真实模型,派生类必须重载,否则后续分类直接拿不到有效数值。 PredictClass 内部先调 PredictPrice,若返回 DBL_MAX 则直接回 -1,等于告诉调度层“这根 K 线没法分类”。拿到有效预测价后,用 iClose(m_symbol,m_period,1) 取上一根收盘价,和预测价做差得到 delta。
| 分类阈值由 m_class_delta 控制: | delta | 小于等于它判 PRICE_SAME,否则按 delta 符号给 PRICE_UP 或 PRICE_DOWN。外汇与贵金属波动大,m_class_delta 设太小会几乎全是 PRICE_SAME,设太大则方向信号泛滥,两者都偏高风险,建议在 MT5 里用历史数据跑一遍看命中分布再定。 |
|---|
这套虚函数骨架的好处是:你后面接 LSTM 还是线性回归,只要重写 PredictPrice,分类逻辑一行不用动。开 MT5 建个 EA 把上面代码贴进去,把 m_class_delta 从 10 点调到 50 点,能直观看到 PRICE_SAME 占比从 70% 掉到 20% 左右。
class="kw">virtual class="type">class="kw">double PredictPrice(class="type">void) { class="kw">return(DBL_MAX); } class=class="str">"cmt">//+------------------------------------------------------------------+ class=class="str">"cmt">//| Predict class (regression -> classification) | class=class="str">"cmt">//+------------------------------------------------------------------+ class="kw">virtual class="type">int PredictClass(class="type">void) { class="type">class="kw">double predicted_price=PredictPrice(); if(predicted_price==DBL_MAX) class="kw">return(-class="num">1); class="type">int predicted_class=-class="num">1; class="type">class="kw">double last_close=iClose(m_symbol,m_period,class="num">1); class=class="str">"cmt">//--- classify predicted price movement class="type">class="kw">double delta=last_close-predicted_price; if(fabs(delta)<=m_class_delta) predicted_class=PRICE_SAME; else { if(delta<class="num">0) predicted_class=PRICE_UP; else predicted_class=PRICE_DOWN; } class=class="str">"cmt">//--- class="kw">return predicted class class="kw">return(predicted_class); } };