用于时间序列挖掘的数据标签(第 6 部分):使用 ONNX 在 EA 中应用和测试(基础篇)
「在EA里直接跑ONNX模型」
把训练好的 ONNX 模型塞进 MT5 的 EA 里做推理,是这条时间序列挖掘链路落到实盘验证的关键一步。MQL5 自 2875 版本起内置了 ONNX 运行时,不需要外部桥接,直接用代码加载模型就能在 OnTick 里拿预测值。 下面这段是最小可跑的加载与推理骨架,先加载模型句柄,再按输入维度喂数据、取输出。注意输入张量形状必须和导出模型时一致,否则会返回无效句柄。 高风险提示:外汇与贵金属杠杆品种波动剧烈,模型在历史样本上表现不等于未来概率,任何信号都只能作为辅助,实盘前务必在策略测试器用真实点差回测。
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——里面记着最佳模型路径和对应得分,加载模型时直接读它。 下面这段就是训练时定义检查点回调的典型写法,只给了文件名模板、没给路径,所以文件归到了默认根目录。
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。外汇与贵金属相关的预测模型接入后请务必小资金验证,这类时序推断在极端行情下失准概率显著抬高。
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 环境配置——外汇与贵金属行情高波动,模型仅提供概率倾向,实盘前务必在策略测试器跑历史回测。
# 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): &class="macro">#x27;&class="macro">#x27;&class="macro">#x27; rewrite dataset <span class="keyword">class</span> &class="macro">#x27;&class="macro">#x27;&class="macro">#x27; def to_dataloader(self, train: <span class="keyword">class="type">bool</span> = True, batch_size: <span class="keyword">class="type">int</span> = <span class="number">class="num">64</span>, batch_sampler: Sampler | str = None, shuffle:<span class="keyword">class="type">bool</span>=False, drop_last:<span class="keyword">class="type">bool</span>=False, **kwargs) -> DataLoader: default_kwargs = dict( shuffle=shuffle, # drop_last=train and len(self) > batch_size, drop_last=drop_last, # collate_fn=self._collate_fn, batch_size=batch_size, batch_sampler=batch_sampler, ) default_kwargs.update(kwargs) kwargs = default_kwargs # print(kwargs[&class="macro">#x27;drop_last&class="macro">#x27;]) <span class="keyword">if</span> kwargs["batch_sampler"] is not None: sampler = kwargs["batch_sampler"] <span class="keyword">if</span> isinstance(sampler, str): <span class="keyword">if</span> sampler == "synchronized": kwargs["batch_sampler"] = TimeSynchronizedBatchSampler( SequentialSampler(self), batch_size=kwargs["batch_size"], shuffle=kwargs["shuffle"], drop_last=kwargs["drop_last"], )
◍ 把 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,是验证整条管道的第一步。
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,