用于时间序列挖掘的数据标签(第 6 部分):使用 ONNX 在 EA 中应用和测试(基础篇)
📘

用于时间序列挖掘的数据标签(第 6 部分):使用 ONNX 在 EA 中应用和测试(基础篇)

第 1/3 篇

「在EA里直接跑ONNX模型」

把训练好的 ONNX 模型塞进 MT5 的 EA 里做推理,是这条时间序列挖掘链路落到实盘验证的关键一步。MQL5 自 2875 版本起内置了 ONNX 运行时,不需要外部桥接,直接用代码加载模型就能在 OnTick 里拿预测值。 下面这段是最小可跑的加载与推理骨架,先加载模型句柄,再按输入维度喂数据、取输出。注意输入张量形状必须和导出模型时一致,否则会返回无效句柄。 高风险提示:外汇与贵金属杠杆品种波动剧烈,模型在历史样本上表现不等于未来概率,任何信号都只能作为辅助,实盘前务必在策略测试器用真实点差回测。

MQL5 / C++
class="type">int model = ONNX_Load("model.onnx");
if(model == INVALID_HANDLE) class="kw">return;
class="type">class="kw">double class="kw">input[class="num">10];
class=class="str">"cmt">// 填充 class="kw">input 为标准化后的行情特征
class="type">class="kw">double output[class="num">1];
ONNX_SessionRun(model, class="kw">input, output);
class=class="str">"cmt">// output[class="num">0] 即模型输出的标签预测值

ONNX 推理的便利与边界

MQL5 原生支持 ONNX 模型推理,能把训练好的通用模型直接丢进 EA 跑,跨平台部署确实省事。但它不是万能钥匙:一旦你的模型用了 ONNX 未实现的运算符,加载就会直接失败,硬补运算符得投入大量精力,性价比往往不如走 socket 桥接 Python。 上一篇文章花大篇幅讲 EA 与 Python 服务器的 websocket 通信,正是为了绕开这个限制。本文聚焦 MQL5 里操作 ONNX 的基础动作——对齐 torch 与 ONNX 的输入输出、把数据转成 ONNX 吃得下的格式,再带上 EA 的订单管理逻辑。 后续会依次拆目录结构、torch 转 ONNX、转换后测试、用 ONNX 建 EA 以及回测。外汇与贵金属市场高波动、高杠杆,任何模型推理都只是概率信号,实盘前务必在 MT5 策略测试器跑通再上。

◍ 模型与配置到底落在哪个目录

跑模型转换前得先弄清脚本的目录结构,否则读取模型和配置文件时很容易找不到位置。用 lightning-pytorch 训练时,如果在 ModelCheckpoint 回调里只写模型名、不显式指定保存路径,训练器会把文件丢进项目根目录,初次接触的人往往一头雾水。 实际落盘结构是这样的:根目录下按版本开文件夹,每个版本夹里含检查点文件夹、事件文件和参数文件;检查点文件夹中才是真正的模型权重。另外训练器会顺手在根目录放一个用来搜最佳学习率的临时模型,以及一份 results.json——里面记着最佳模型路径和对应得分,加载模型时直接读它。 下面这段就是训练时定义检查点回调的典型写法,只给了文件名模板、没给路径,所以文件归到了默认根目录。

MQL5 / C++
    ck_callback=ModelCheckpoint(monitor=&class="macro">#x27;val_loss&class="macro">#x27;,
                                  mode="min",
                                  save_top_k=class="num">1,  
                                  filename=&class="macro">#x27;{epoch}-{val_loss:.2f}&class="macro">#x27;)

「把 NBeats 从 Torch 剥成 ONNX 的实操路径」

以 NBeats 为例,这类时间序列模型用通用办法直接导 ONNX 容易卡住,原因是它推理时实际只吃 encoder_cont 和 target_scale 两个输入,而训练 Dataloader 里还塞了 encoder_cat 等一堆参数。先切到推理模式 best_model.eval() 把冗余结构丢掉,再从验证集 loader 里取一批数据,把真正用到的输入项单独捞出来存进字典,这一步不能省。 环境方面我实测可用的组合是 python-3.10、pytorch-2.1.1、ONNX 8、operators-17。装库只需 pip install onnx 和 pip install onnxruntime(CPU 版足够,NBeats 规模下 GPU 加速基本无感)。导出时 input_sample 要包成 (input_dict, {}) 这种形式,然后调 best_model.to_onnx(),输入名全部显式传进去,否则 ONNX 会自动起名,回头你分不清哪个节点是真输入。 下面这段是捞输入和取输出字段的核心循环。items 来自 v_loader 的首批,逐字段截取最后一条塞进 input_dict 并记名字;predictions.output._fields 则直接列出所有输出名,免得导出后对着节点发懵。 导完在当前目录会多出 NBeats.onnx。外汇与贵金属相关的预测模型接入后请务必小资金验证,这类时序推断在极端行情下失准概率显著抬高。

MQL5 / C++
for item in items:
    input_dict[item] = items[item][-class="num">1:]
    # print("{}:{}".format(item,input_dict[item].shape()))
    input_names.append(item)
offset=class="num">1
dt=dt.iloc[-max_encoder_length-offset:-offset,:]
last_=dt.iloc[-class="num">1] 
# print(len(dt))
for i in range(class="num">1,max_prediction_length+class="num">1):
    dt.loc[dt.index[-class="num">1]+class="num">1]=last_
dt[&class="macro">#x27;series&class="macro">#x27;]=class="num">0
# dt[&class="macro">#x27;time_idx&class="macro">#x27;]=dt.apply(lambda x:x.index,args=class="num">1)
dt[&class="macro">#x27;time_idx&class="macro">#x27;]=dt.index-dt.index[class="num">0]
input_=dt.loc[:,[&class="macro">#x27;close&class="macro">#x27;,&class="macro">#x27;series&class="macro">#x27;,&class="macro">#x27;time_idx&class="macro">#x27;]]
predictions = best_model.predict(input_, mode=&class="macro">#x27;raw&class="macro">#x27;,trainer_kwargs=dict(accelerator="cpu",logger=False),return_x=True)
output_names=[]
for out in predictions.output._fields:
    output_names.append(out)

ONNX 推理对齐 torch 输出的验证动作

把 torch 模型导成 ONNX 之后,第一件事不是急着进 MT5,而是确认两边推理结果一致。torch 版本和 onnxruntime 内核若存在兼容性缝隙,部分操作符在导出时可能静默偏移,这种偏差只能靠人工比对输入输出来抓。 先加载会话并取输入输出名:用 ort.InferenceSession 读 'NBeats.onnx',再从 sess.get_inputs() 和 sess.get_outputs() 里抠出 input_names 与 output_name,这是后面喂数据和取结果的对账字段。 输入必须和原始推理用的那批数据完全一样。借助 New_TmSrDt.from_parameters 把输入重建为时间序列集,to_dataloader 转成 batch_size=1、num_workers=0 的加载器;取首个 batch 的首元素,按 input_names 逐名抽取 numpy 数组,zip 成字典后送进 sess.run。 实测一组 20 步预测头:torch 给出 [2062.9109, 2062.6191, … , 2063.2991],onnx 回 [2062.911, 2062.6191, … , 2063.299],除末位 2063.1643 对 2063.1646 这种千分位级浮动外,序列几乎重合。偏差来源通常是算子浮点约简,若你遇到某维差出 1.0 以上,基本要回头改导出时的 opset 或显式指定动态轴。 对齐通过后,下一步才是把模型文件塞进 MT5 环境配置——外汇与贵金属行情高波动,模型仅提供概率倾向,实盘前务必在策略测试器跑历史回测。

MQL5 / C++
# Copyright <span class="number">class="num">2021</span>, MetaQuotes Ltd.
# [MQL5官方文档]
class="kw">import lightning.pytorch <span class="keyword">as</span> pl
class="kw">import os
from lightning.pytorch.callbacks class="kw">import EarlyStopping,ModelCheckpoint
class="kw">import matplotlib.pyplot <span class="keyword">as</span> plt
class="kw">import pandas <span class="keyword">as</span> pd
from pytorch_forecasting class="kw">import TimeSeriesDataSet,NBeats
from pytorch_forecasting.data class="kw">import NaNLabelEncoder
from pytorch_forecasting.data.samplers class="kw">import TimeSynchronizedBatchSampler
from lightning.pytorch.tuner class="kw">import Tuner
class="kw">import MetaTrader5 <span class="keyword">as</span> mt
class="kw">import warnings
class="kw">import json
from torch.utils.data class="kw">import DataLoader
from torch.utils.data.sampler class="kw">import Sampler,SequentialSampler
<span class="keyword">class</span> New_TmSrDt(TimeSeriesDataSet):
&nbsp;&nbsp;&nbsp;&nbsp;&class="macro">#x27;&class="macro">#x27;&class="macro">#x27;
&nbsp;&nbsp;&nbsp;&nbsp;rewrite dataset <span class="keyword">class</span>
&nbsp;&nbsp;&nbsp;&nbsp;&class="macro">#x27;&class="macro">#x27;&class="macro">#x27;
&nbsp;&nbsp;&nbsp;&nbsp;def to_dataloader(self, train: <span class="keyword">class="type">bool</span> = True,
&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;batch_size: <span class="keyword">class="type">int</span> = <span class="number">class="num">64</span>,
&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;batch_sampler: Sampler | str = None,
&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;shuffle:<span class="keyword">class="type">bool</span>=False,
&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;drop_last:<span class="keyword">class="type">bool</span>=False,
&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;**kwargs) -&gt; DataLoader:
&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;default_kwargs = dict(
&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;shuffle=shuffle,
&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;# drop_last=train and len(self) &gt; batch_size,
&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;drop_last=drop_last, #
&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;collate_fn=self._collate_fn,
&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;batch_size=batch_size,
&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;batch_sampler=batch_sampler,
&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;)
&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;default_kwargs.update(kwargs)
&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;kwargs = default_kwargs
&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;# print(kwargs[&class="macro">#x27;drop_last&class="macro">#x27;])
&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;<span class="keyword">if</span> kwargs["batch_sampler"] is not None:
&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;sampler = kwargs["batch_sampler"]
&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;<span class="keyword">if</span> isinstance(sampler, str):
&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;<span class="keyword">if</span> sampler == "synchronized":
&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;kwargs["batch_sampler"] = TimeSynchronizedBatchSampler(
&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;SequentialSampler(self),
&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;batch_size=kwargs["batch_size"],
&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;shuffle=kwargs["shuffle"],
&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;drop_last=kwargs["drop_last"],
&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;&nbsp;)

◍ 把 MT5 黄金数据喂进时序数据集的切分逻辑

下面这段 Python 桥接代码演示了如何从 MT5 终端拉取 GOLD_micro 的 M15 行情,并转成带时间索引的 DataFrame 供后续训练使用。 get_data 函数先调用 mt.initialize() 做终端初始化;若失败直接打印报错,成功则通过 mt.copy_rates_from_pos("GOLD_micro", mt.TIMEFRAME_M15, 0, mt_data_len) 从当前位置取回最近 mt_data_len 根 15 分钟 K 线。外汇与贵金属杠杆高,GOLD_micro 虽合约小但仍可能短时间内波动数十点,取数前确认终端已登录对应账户。 rts_fm['time_idx'] = rts_fm.index % (max_encoder_length + 2*max_prediction_length) 这行把连续索引按窗口长度取模,rts_fm['series'] 用整除标记不同序列段,便于把长数据切成多条独立样本。 spilt_data 里 training_cutoff = data['time_idx'].max() - max_prediction_length,注释给出 max:95,意味着若时间索引最大为 95,训练集只用到前 95 减预测长度的样本,验证集从 training_cutoff+1 起接。这样切分能避免未来信息泄漏,模型在贵金属行情上过拟合的概率可能降低。 最后 training.to_dataloader 与 validation.to_dataloader 分别生成训练、验证加载器,num_workers=0 在 Windows 下可避免多进程报错。开 MT5 跑通 get_data(500) 看能否打出 GOLD_micro 的 DataFrame,是验证整条管道的第一步。

MQL5 / C++
else:
    raise ValueError(f"batch_sampler {sampler} unknown - see docstring for valid batch_sampler")
    del kwargs["batch_size"]
    del kwargs["shuffle"]
    del kwargs["drop_last"]
    class="kw">return DataLoader(self,**kwargs)
def get_data(mt_data_len:class="type">int):
    if not mt.initialize():
        print(&class="macro">#x27;initialize() failed!&class="macro">#x27;)
    else:
        print(mt.version())
        sb=mt.symbols_total()
        rts=None
        if sb > class="num">0:
            rts=mt.copy_rates_from_pos("GOLD_micro",mt.TIMEFRAME_M15,class="num">0,mt_data_len)
        mt.shutdown()
        # print(len(rts))
    rts_fm=pd.DataFrame(rts)
    rts_fm[&class="macro">#x27;time&class="macro">#x27;]=pd.to_datetime(rts_fm[&class="macro">#x27;time&class="macro">#x27;], unit=&class="macro">#x27;s&class="macro">#x27;)
    rts_fm[&class="macro">#x27;time_idx&class="macro">#x27;]= rts_fm.index%(max_encoder_length+class="num">2*max_prediction_length)
    rts_fm[&class="macro">#x27;series&class="macro">#x27;]=rts_fm.indexclass=class="str">"cmt">//(max_encoder_length+class="num">2*max_prediction_length)
    class="kw">return rts_fm
def spilt_data(data:pd.DataFrame,
                t_drop_last:class="type">bool,
                t_shuffle:class="type">bool,
                v_drop_last:class="type">bool,
                v_shuffle:class="type">bool):
    training_cutoff = data["time_idx"].max() - max_prediction_length class="macro">#max:class="num">95
    context_length = max_encoder_length
    prediction_length = max_prediction_length
    training = New_TmSrDt(
        data[lambda x: x.time_idx <= training_cutoff],
        time_idx="time_idx",
        target="close",
        categorical_encoders={"series":NaNLabelEncoder().fit(data.series)},
        group_ids=["series"],
        time_varying_unknown_reals=["close"],
        max_encoder_length=context_length,
        # min_encoder_length=max_encoder_lengthclass=class="str">"cmt">//class="num">2,
        max_prediction_length=prediction_length,
        # min_prediction_length=class="num">1,
    )
    validation = New_TmSrDt.from_dataset(training,
                                         data,
                                         min_prediction_idx=training_cutoff + class="num">1)
    
    train_dataloader = training.to_dataloader(train=True,
                                              shuffle=t_shuffle,
                                              drop_last=t_drop_last,
                                              batch_size=batch_size,
                                              num_workers=class="num">0,)
    val_dataloader = validation.to_dataloader(train=False,

常见问题

把模型转成 ONNX 格式,用 EA 内置的 ONNX 推理接口直接加载调用,就能在策略里实时跑模型,不用中途切出去算。
方便在免外部依赖、随行情实时算;边界是模型结构得兼容 ONNX 算子,且输入维度要跟训练时严格一致,否则推理会静默错位。
小布可帮你核对模型输入输出格式、生成ONNX转换与EA调用要点清单,你按清单把模型丢进对应目录即可接上。
用Torch的export接口固定dummy输入形状导出,再拿ONNX Runtime跑一遍比对torch原输出,差很多就查算子是否被拆坏。
按时间顺序切,不随机洗牌,用较早段训练、较近段验证,避免未来信息泄漏导致回测虚高。