在类中包装 ONNX 模型(基础篇)
📘

在类中包装 ONNX 模型(基础篇)

第 1/3 篇

把 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 里的成员如何承接品种、周期与会话生命周期。

MQL5 / C++
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 触发一次,回测和实盘节奏一致。

MQL5 / C++
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% 左右。

MQL5 / C++
  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);
    }
  };

常见问题

类封装能把模型加载、心跳节拍、输入输出收口到一处,多模型复用同一父类,后期换模型或加回归预测只改子类不改调用层。
推理时类别映射错乱,投票结果偏向某一类;需核对训练时的类别顺序与代码里写死的底数是否一致再上线。
小布可读取你上传的模型结构说明,自动生成类骨架与心跳控制建议,并在品种页标注模型推理延迟和异常。
在类里用定时器或外部节拍函数周期性喂入最新行情,模型推理放到独立调用点,主循环只做调度不阻塞。
在公共父类加统一推理接口,子类把回归输出归一化成伪三类概率返回,上层按分类逻辑取用即可。