利用回归衡量度评估 ONNX 模型·进阶篇
(2/3)· 从 MAE 到 RMSLE 七种衡量度的语义边界与 MQL5 落地,多数人在损失函数和衡量度之间混淆了优化目标与品质准则
◍ 模型类的初始化与释放骨架
在 MT5 里封装 ONNX 推理模型,第一步是把句柄生命周期管起来。下面这段基类代码给出了最小可用骨架:取模型名、虚 Init 占位、CheckInit 做符号周期校验并加载模型、Shutdown 释放会话。 CheckInit 会先比对 symbol 和 period 是否和类内成员 m_symbol、m_period 一致,不一致直接 PrintFormat 报错并返回 false。一致时才调用 OnnxCreateFromBuffer 从静态字节数组 model[] 建会话,失败则打印 GetLastError() 后返回 false,成功返回 true。 Shutdown 里判断 m_handle 不是 INVALID_HANDLE 就调 OnnxRelease 并置回无效句柄,避免 EA 重载或退出时显存/内存泄漏。外汇与贵金属行情波动剧烈、杠杆高风险大,这类本地推理模块加载失败时必须阻断后续 OnTick,否则可能用空模型跑出无意义信号。 别把虚 Init 当摆设 基类 Init 默认返回 false,意味着派生类不重写就永远初始化不通过。写自己的黄金 15 分钟模型时,记得在派生类里把权重路径和输入输出张量名填实,否则 CheckInit 过了也跑不出推理。
class="type">class="kw">string GetModelName(class="type">void) { class="kw">return(m_name); } class="kw">virtual class="type">bool Init(class="kw">const class="type">class="kw">string symbol, class="kw">const ENUM_TIMEFRAMES period) { class="kw">return(false); } class="type">bool CheckInit(class="kw">const class="type">class="kw">string symbol, class="kw">const ENUM_TIMEFRAMES period,class="kw">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="type">void Shutdown(class="type">void) { if(m_handle!=INVALID_HANDLE) { OnnxRelease(m_handle); m_handle=INVALID_HANDLE; } } class="kw">virtual class="type">bool CheckOnTick(class="type">void) {
新K线与预测分类的底层判定
这段 CTestEngine 派生类的代码负责两件事:判定是否进入新Bar,以及把回归预测价转成涨跌分类。外汇与贵金属波动剧烈,任何预测都只是概率倾向,实盘前务必在 MT5 策略测试器核查逻辑。 新Bar检查先用 TimeCurrent() 与 m_next_bar 比较,未到时间直接 return(false) 跳过。到时间后把 m_next_bar 对齐到当前周期整点:减去对周期秒数的取余,再加上一个周期秒数,这样下一根Bar的开盘时刻就锁定了。 PredictPrice 是虚函数桩,默认返回 DBL_MAX 表示无模型可用。真正落地分类在 PredictClass:先把传入时间按周期取整,调 PredictPrice 拿到预测价,若返回 DBL_MAX 则直接 return(-1) 放弃分类。 接着用 CopyClose 取该时间点起两根收盘线,取前一根 prev_price。用 delta = prev_price - predicted_price 衡量偏离,当 fabs(delta) <= m_class_delta 归为 PRICE_SAME,delta<0 为 PRICE_UP,否则 PRICE_DOWN。最后 probabilities 向量清零,仅把命中分类的概率置 1.0——意味着此桩实现是确定性分类,不带模糊概率。
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) | class=class="str">"cmt">//+------------------------------------------------------------------+ class="kw">virtual class="type">class="kw">double PredictPrice(class="type">class="kw">datetime date) { 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">class="kw">datetime date,vector& probabilities) { date-=date%PeriodSeconds(m_period); class="type">class="kw">double predicted_price=PredictPrice(date); if(predicted_price==DBL_MAX) class="kw">return(-class="num">1); class="type">class="kw">double last_close[class="num">2]; if(CopyClose(m_symbol,m_period,date,class="num">2,last_close)!=class="num">2) class="kw">return(-class="num">1); class="type">class="kw">double prev_price=last_close[class="num">0]; class=class="str">"cmt">//--- classify predicted price movement class="type">int predicted_class=-class="num">1; class="type">class="kw">double delta=prev_price-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">//--- set predicted probability as class="num">1.0 probabilities.Fill(class="num">0); if(predicted_class<(class="type">int)probabilities.Size()) probabilities[predicted_class]=class="num">1; class=class="str">"cmt">//--- and class="kw">return predicted class class="kw">return(predicted_class); } };
「EURUSD日线10根K线的回归包装类」
第一个落地模型叫 model.eurusd.D1.10.onnx,是用 EURUSD 日线连续 10 根 OHLC 价格序列训出来的回归模型,思路和之前公开过的共享项目里的初版一致。 喂给模型前必须把 10 根 OHLC 做标准化:每根序列相对均价的偏离,除以该序列自身标准差。这样序列被压到均值 0、离散度 1 的区间里,训练时的收敛性会明显提升。 下面这段 MQH 是直接在 MT5 里调用该 ONNX 的包装类。构造函数写死品种 EURUSD、周期 PERIOD_D1、样本数 10;Init 里除了校验符号周期,还显式把输入张量形状设成 {1,10,4}——batch 为 1,序列长 10,OHLC 四通道。 外汇与贵金属属高风险品种,模型输出只是概率倾向,实盘前务必在策略测试器里跑一遍验证。
class=class="str">"cmt">//+------------------------------------------------------------------+ class=class="str">"cmt">//| ModelEurusdD1_10.mqh | class=class="str">"cmt">//| Copyright class="num">2023, MetaQuotes Ltd. | class=class="str">"cmt">//| [MQL5官方文档] | class=class="str">"cmt">//+------------------------------------------------------------------+ class="macro">#include "ModelSymbolPeriod.mqh" class="macro">#resource "Python/model.eurusd.D1.class="num">10.onnx" as class="type">uchar model_eurusd_D1_10[] class=class="str">"cmt">//+------------------------------------------------------------------+ class=class="str">"cmt">//| ONNX-model wrapper class | class=class="str">"cmt">//+------------------------------------------------------------------+ class CModelEurusdD1_10 : class="kw">public CModelSymbolPeriod { class="kw">private: class="type">int m_sample_size; class="kw">public: class=class="str">"cmt">//+------------------------------------------------------------------+ class=class="str">"cmt">//| Constructor | class=class="str">"cmt">//+------------------------------------------------------------------+ CModelEurusdD1_10(class="type">void) : CModelSymbolPeriod("EURUSD",PERIOD_D1) { m_name="D1_10"; m_sample_size=class="num">10; } class=class="str">"cmt">//+------------------------------------------------------------------+ class=class="str">"cmt">//| ONNX-model initialization | class=class="str">"cmt">//+------------------------------------------------------------------+ class="kw">virtual class="type">bool Init(class="kw">const class="type">class="kw">string symbol, class="kw">const ENUM_TIMEFRAMES period) { class=class="str">"cmt">//--- check symbol, period, create model if(!CModelSymbolPeriod::CheckInit(symbol,period,model_eurusd_D1_10)) { Print("model_eurusd_D1_10 : initialization error"); class="kw">return(false); } class=class="str">"cmt">//--- since not all sizes defined in the class="kw">input tensor we must set them explicitly class=class="str">"cmt">//--- first index - batch size, second index - series size, third index - number of series(OHLC) class="kw">const class="type">long input_shape[] = {class="num">1,m_sample_size,class="num">4}; if(!OnnxSetInputShape(m_handle,class="num">0,input_shape)) { Print("model_eurusd_D1_10 : OnnxSetInputShape error ",GetLastError()); class="kw">return(false); }
◍ 给 ONNX 模型喂归一化 OHLC 的实操细节
加载 EURUSD D1 的 ONNX 模型时,输出张量不会自动带全维度,必须手动指定形状。代码里用 output_shape[] = {1,1} 表示批次为 1、预测价格数为 1,且批次维度必须和输入张量一致,否则 OnnxSetOutputShape 会返回 false 并打印错误码。
预测函数里先静态开好 input_data(m_sample_size,4) 和 output_data(1),其中 m_sample_size 是回看窗口长度。从 MT5 拷最近 m_sample_size 根日线 OHLC 到 rates 矩阵,注意 date-=date%PeriodSeconds(m_period) 把时间对齐到周期起点,避免跨周期取错 bar。
归一化按列做:先 rates.Mean(1) 和 rates.Std(1) 拿到每列(O/H/L/C)的均值与标准差,铺成 mm、ms 两个同形矩阵,再把 rates 转置成纵向 OHLC 向量后减均值除标准差。模型输入要求纵向向量,这步转置不能省。
最后 input_data.Assign(x_norm) 后调 OnnxRun 跑推理,外汇与贵金属杠杆高、模型预测仅代表统计倾向,实盘前务必在 MT5 策略测试器用历史数据验证归一化与输出维度是否匹配。
class=class="str">"cmt">//--- since not all sizes defined in the output tensor we must set them explicitly class=class="str">"cmt">//--- first index - batch size, must match the batch size of the class="kw">input tensor class=class="str">"cmt">//--- second index - number of predicted prices class="kw">const class="type">long output_shape[] = {class="num">1,class="num">1}; if(!OnnxSetOutputShape(m_handle,class="num">0,output_shape)) { Print("model_eurusd_D1_10 : OnnxSetOutputShape error ",GetLastError()); class="kw">return(false); } class=class="str">"cmt">//--- ok class="kw">return(true); } class=class="str">"cmt">//+------------------------------------------------------------------+ class=class="str">"cmt">//| Predict price | class=class="str">"cmt">//+------------------------------------------------------------------+ class="kw">virtual class="type">class="kw">double PredictPrice(class="type">class="kw">datetime date) { class="kw">static matrixf input_data(m_sample_size,class="num">4); class=class="str">"cmt">// matrix for prepared class="kw">input data class="kw">static vectorf output_data(class="num">1); class=class="str">"cmt">// vector to get result class="kw">static matrix mm(m_sample_size,class="num">4); class=class="str">"cmt">// matrix of horizontal vectors Mean class="kw">static matrix ms(m_sample_size,class="num">4); class=class="str">"cmt">// matrix of horizontal vectors Std class="kw">static matrix x_norm(m_sample_size,class="num">4); class=class="str">"cmt">// matrix for prices normalize class=class="str">"cmt">//--- prepare class="kw">input data matrix rates; class=class="str">"cmt">//--- request last bars date-=date%PeriodSeconds(m_period); if(!rates.CopyRates(m_symbol,m_period,COPY_RATES_OHLC,date-class="num">1,m_sample_size)) class="kw">return(DBL_MAX); class=class="str">"cmt">//--- get series Mean vector m=rates.Mean(class="num">1); class=class="str">"cmt">//--- get series Std vector s=rates.Std(class="num">1); class=class="str">"cmt">//--- prepare matrices for prices normalization for(class="type">int i=class="num">0; i<m_sample_size; i++) { mm.Row(m,i); ms.Row(s,i); } class=class="str">"cmt">//--- the class="kw">input of the model must be a set of vertical OHLC vectors x_norm=rates.Transpose(); class=class="str">"cmt">//--- normalize prices x_norm-=mm; x_norm/=ms; class=class="str">"cmt">//--- run the inference input_data.Assign(x_norm); if(!OnnxRun(m_handle,ONNX_NO_CONVERSION,input_data,output_data))
反归一化取回预测价格
模型推理结束后,输出值仍处于归一化空间,必须还原成真实报价才能用于下单或画线。上面这段代码在输出数组首位乘上缩放系数 s[3] 并加上偏移 m[3],就是把网络输出反归一化。 若还原出的 predicted 无效,函数直接 return DBL_MAX,调用层可用这个值判断预测失败,避免在 MT5 里用脏数据发单。外汇与贵金属波动剧烈,这类边界判断能降低误触发风险。 把这段逻辑接进你已有的 EA 预测函数,在 return 前打印一次 predicted 与 DBL_MAX 的比较结果,就能在策略测试器里确认反归一化是否生效。
class="kw">return(DBL_MAX); class=class="str">"cmt">//--- denormalize the price from the output value class="type">class="kw">double predicted=output_data[class="num">0]*s[class="num">3]+m[class="num">3]; class=class="str">"cmt">//--- class="kw">return prediction class="kw">return(predicted); } }; class=class="str">"cmt">//+------------------------------------------------------------------+
「EURUSD D1 三十根收盘价的回归包装」
第二个要落地的模型叫 model.eurusd.D1.30.onnx,它吃的是 EURUSD 日线最近 30 根收盘价,外加两条周期分别为 21 和 34 的简单移动平均线。和前几个类一样,Init 里直接调基类的 CheckInit,由基类去开 ONNX 会话,并把输入输出的张量尺寸显式钉死,避免推理时维度对不上。 PredictPrice 负责把 30 根历史收盘价和算好的双均线喂进去,常规化方式和训练时保持一致,否则归一化偏移会让输出偏离训练分布。这个模型最早是为“在类里包 ONNX”写的,本文把它从分类任务改成了回归,用来直接估价格而不是判方向。 外汇和贵金属杠杆高、跳空频繁,这类回归输出只是概率倾向,不能直接当进场依据,开 MT5 加载资源跑一遍才知道在你broker点差下漂多少。
class=class="str">"cmt">//+------------------------------------------------------------------+ class=class="str">"cmt">//| ModelEurusdD1_30.mqh | class=class="str">"cmt">//| Copyright class="num">2023, MetaQuotes Ltd. | class=class="str">"cmt">//| [MQL5官方文档] | class=class="str">"cmt">//+------------------------------------------------------------------+ class="macro">#include "ModelSymbolPeriod.mqh" class="macro">#resource "Python/model.eurusd.D1.class="num">30.onnx" as class="type">uchar model_eurusd_D1_30[] class=class="str">"cmt">//+------------------------------------------------------------------+ class=class="str">"cmt">//| ONNX-model wrapper class | class=class="str">"cmt">//+------------------------------------------------------------------+ class CModelEurusdD1_30 : class="kw">public CModelSymbolPeriod { class="kw">private: class="type">int m_sample_size; class="type">int m_fast_period; class="type">int m_slow_period; class="type">int m_sma_fast; class="type">int m_sma_slow; class="kw">public: class=class="str">"cmt">//+------------------------------------------------------------------+ class=class="str">"cmt">//| Constructor | class=class="str">"cmt">//+------------------------------------------------------------------+ CModelEurusdD1_30(class="type">void) : CModelSymbolPeriod("EURUSD",PERIOD_D1) { m_name="D1_30"; m_sample_size=class="num">30; m_fast_period=class="num">21; m_slow_period=class="num">34; m_sma_fast=INVALID_HANDLE; m_sma_slow=INVALID_HANDLE; } class=class="str">"cmt">//+------------------------------------------------------------------+ class=class="str">"cmt">//| ONNX-model initialization | class=class="str">"cmt">//+------------------------------------------------------------------+ class="kw">virtual class="type">bool Init(class="kw">const class="type">class="kw">string symbol, class="kw">const ENUM_TIMEFRAMES period) { class=class="str">"cmt">//--- check symbol, period, create model if(!CModelSymbolPeriod::CheckInit(symbol,period,model_eurusd_D1_30)) { Print("model_eurusd_D1_30 : initialization error"); class="kw">return(false); } class=class="str">"cmt">//--- since not all sizes defined in the class="kw">input tensor we must set them explicitly