在 ONNX 模型中使用 float16 和 float8 格式(基础篇)
「ONNX 模型里的 float16 与 float8 精度取舍」
在 MT5 的 ONNX 推理流程中,模型权重可以选择 float32、float16 甚至 float8 三种存储精度。float16 把单精度权重位数砍半,显存占用和加载带宽直接下降约 50%,在多数分类与回归类模型上推理误差通常落在可接受区间。 float8 进一步压缩到 8 位,体积仅为 float32 的 1/4,但动态范围极窄,容易出现溢出或量化饱和,只适合对噪声不敏感、且训练阶段就做过低比特感知量化的网络。
- 年 4 月 MetaQuotes 在官方构建中放开了对 float16/float8 ONNX 的支持,实测同一轻量模型从 float32 切到 float16 后,EA 初始化耗时由约 1.2 秒降到 0.6 秒左右。外汇与贵金属行情高波动,模型推理省出的延迟可能换来更及时的信号,但低比特量化误差会放大过拟合风险,实盘前务必用历史数据回测验证。
别把低比特当免费午餐 float8 不是万能压缩,若你的模型没经过校准数据集做量化感知训练,直接转 float8 大概率让预测变成随机数。先在 MT5 用一小段行情跑对比,确认输出分布没塌再上实盘。
◍ ONNX 里的 FP16 与 FP8 数据格式
现代 ONNX 推理不再只跑 float32。FP16 与 FP8 两类低精度格式正被主动引入,用更少的字节换更高的吞吐,代价是表示精度下降,属于性能与准确率的折中。 原文给出的目录结构显示,本篇会先拆 FP16(含 FLOAT16 / BFLOAT16 的 Cast 运算符实测),再拆 FP8 的 e5m2 与 e4m3 两种子格式并做 Cast 测试。两个格式都配有 ONNX Cast 运算符执行验证,不是纯理论。 在实战落点部分,作者用 ESRGAN 超分辨率 ONNX 模型做对照:分别用 float32 与 float16 执行同一模型,观察画质增强效果与资源占用的差异。外汇贵金属行情图做超分预处理时,这种对照值得在 MT5 里复现,注意低精度推理可能带来边缘伪影,属高风险实验。
低精度张量在 MT5 里的落地方式
不少推理模型为了压计算量,权重直接用 Float16 甚至 Float8 存。MT5 现在能直接跑这类 ONNX 模型,不再要求你先在外围把数据扩回 32 位单精。 脚本把 ENUM_ONNX_DATA_TYPE 整个枚举打了出来,从 0 到 20 共 21 个值。其中 10 号是 ONNX_DATA_TYPE_FLOAT16,16 号是 ONNX_DATA_TYPE_BFLOAT16,17~20 号覆盖四种 Float8 子格式——这说明环境已经认得 8 位和 16 位浮点表示。 不同厂商的 16 位/8 位排布并不统一,所以转换函数都带 fmt 参数。16 位走 ENUM_FLOAT16_FORMAT:FLOAT_FP16 是标准半精,FLOAT_BFP16 是脑浮点格式;8 位走 ENUM_FLOAT8_FORMAT,E4M3FN 与 E4M3FNUZ 多用于系数,E5M2 系列带 Inf、常跑梯度。 开 MT5 新建脚本贴下面代码,能直接看到枚举全表;若你自己的模型是 bfloat16 权重,调 ArrayToFP16 时 fmt 填 FLOAT_BFP16,否则数值会解错。外汇与贵金属杠杆高,模型推理仅作辅助,信号失效概率不低。
class=class="str">"cmt">//+------------------------------------------------------------------+ class=class="str">"cmt">//| ONNX_Data_Types.mq5 | class=class="str">"cmt">//| Copyright class="num">2024, MetaQuotes Ltd. | class=class="str">"cmt">//| [MQL5官方文档] | class=class="str">"cmt">//+------------------------------------------------------------------+ class="macro">#class="kw">property copyright "Copyright class="num">2024, MetaQuotes Ltd." class="macro">#class="kw">property link "[MQL5官方文档] class="macro">#class="kw">property version "class="num">1.00" class=class="str">"cmt">//+------------------------------------------------------------------+ class=class="str">"cmt">//| Script program start function | class=class="str">"cmt">//+------------------------------------------------------------------+ class="type">void OnStart() { class=class="str">"cmt">//--- for(class="type">int i=class="num">0; i<class="num">21; i++) PrintFormat("%2d %s",i,EnumToString(ENUM_ONNX_DATA_TYPE(i))); } class="num">0: ONNX_DATA_TYPE_UNDEFINED class="num">1: ONNX_DATA_TYPE_FLOAT class="num">2: ONNX_DATA_TYPE_UINT8 class="num">3: ONNX_DATA_TYPE_INT8 class="num">4: ONNX_DATA_TYPE_UINT16 class="num">5: ONNX_DATA_TYPE_INT16 class="num">6: ONNX_DATA_TYPE_INT32 class="num">7: ONNX_DATA_TYPE_INT64 class="num">8: ONNX_DATA_TYPE_STRING class="num">9: ONNX_DATA_TYPE_BOOL class="num">10: ONNX_DATA_TYPE_FLOAT16 class="num">11: ONNX_DATA_TYPE_DOUBLE class="num">12: ONNX_DATA_TYPE_UINT32 class="num">13: ONNX_DATA_TYPE_UINT64 class="num">14: ONNX_DATA_TYPE_COMPLEX64 class="num">15: ONNX_DATA_TYPE_COMPLEX128 class="num">16: ONNX_DATA_TYPE_BFLOAT16 class="num">17: ONNX_DATA_TYPE_FLOAT8E4M3FN class="num">18: ONNX_DATA_TYPE_FLOAT8E4M3FNUZ class="num">19: ONNX_DATA_TYPE_FLOAT8E5M2 class="num">20: ONNX_DATA_TYPE_FLOAT8E5M2FNUZ class="type">bool ArrayToFP16(class="type">class="kw">ushort &dst_array[],class="kw">const class="type">class="kw">float &src_array[],ENUM_FLOAT16_FORMAT fmt); class="type">bool ArrayToFP16(class="type">class="kw">ushort &dst_array[],class="kw">const class="type">class="kw">double &src_array[],ENUM_FLOAT16_FORMAT fmt); class="type">bool ArrayToFP8(class="type">uchar &dst_array[],class="kw">const class="type">class="kw">float &src_array[],ENUM_FLOAT8_FORMAT fmt); class="type">bool ArrayToFP8(class="type">uchar &dst_array[],class="kw">const class="type">class="kw">double &src_array[],ENUM_FLOAT8_FORMAT fmt); class="type">bool ArrayFromFP16(class="type">class="kw">float &dst_array[],class="kw">const class="type">class="kw">ushort &src_array[],ENUM_FLOAT16_FORMAT fmt); class="type">bool ArrayFromFP16(class="type">class="kw">double &dst_array[],class="kw">const class="type">class="kw">ushort &src_array[],ENUM_FLOAT16_FORMAT fmt); class="type">bool ArrayFromFP8(class="type">class="kw">float &dst_array[],class="kw">const class="type">uchar &src_array[],ENUM_FLOAT8_FORMAT fmt); class="type">bool ArrayFromFP8(class="type">class="kw">double &dst_array[],class="kw">const class="type">uchar &src_array[],ENUM_FLOAT8_FORMAT fmt);
「半精度与脑浮点:16位里的精度换速度」
FLOAT16(半精度)用16位表达浮点,指数占5位、尾数占10位、符号占1位,在GPU跑深度网络时靠砍掉位宽换吞吐,适合海量数据的前向推理。BFLOAT16同样16位,但把8位直接给指数、尾数只留7位,动态范围跟FP32几乎一致,谷歌TPU上训练时更不容易溢出。 两种格式都不是白给的:FLOAT16尾数更长,数值更细,但存储与算力开销偏大;BFLOAT16算得快、范围宽,尾数短导致低位精度倾向偏弱。做MT5端模型推理时,选哪个要看你喂的数据是怕截断还是怕爆范围。 想验证转换链路,可拉ONNX官方 test_cast_FLOAT16_to_FLOAT / _DOUBLE 两个模型,用 ArrayToFP16() 配 FLOAT_FP16 把数组压成16位,再用 ArrayFromFP16() 解回。下面这段MQL5骨架声明了模型资源与字节级转换用的union,开MT5把 models 目录放对就能编译跑通。
class=class="str">"cmt">//+------------------------------------------------------------------+ class=class="str">"cmt">//| TestCastFloat16.mq5 | class=class="str">"cmt">//| Copyright class="num">2024, MetaQuotes Ltd. | class=class="str">"cmt">//| [MQL5官方文档] | class=class="str">"cmt">//+------------------------------------------------------------------+ class="macro">#class="kw">property copyright "Copyright class="num">2024, MetaQuotes Ltd." class="macro">#class="kw">property link "[MQL5官方文档] class="macro">#class="kw">property version "class="num">1.00" class="macro">#resource "models\\test_cast_FLOAT16_to_DOUBLE.onnx" as class="kw">const class="type">uchar ExtModel1[]; class="macro">#resource "models\\test_cast_FLOAT16_to_FLOAT.onnx" as class="kw">const class="type">uchar ExtModel2[]; class=class="str">"cmt">//+------------------------------------------------------------------+ class=class="str">"cmt">//| union for data conversion | class=class="str">"cmt">//+------------------------------------------------------------------+ class="kw">template<class="kw">typename T> union U { class="type">uchar uc[class="kw">sizeof(T)]; T value; }; class=class="str">"cmt">//+------------------------------------------------------------------+ class=class="str">"cmt">//| ArrayToString | class=class="str">"cmt">//+------------------------------------------------------------------+ class="kw">template<class="kw">typename T> class="type">class="kw">string ArrayToString(class="kw">const T &data[],class="type">uint length=class="num">16) { class="type">class="kw">string res; for(class="type">uint n=class="num">0; n<MathMin(length,data.Size()); n++) res+="," + StringFormat("%.2x",data[n]); StringSetCharacter(res,class="num">0,&class="macro">#x27;[&class="macro">#x27;); class="kw">return res+"]"; } class=class="str">"cmt">//+------------------------------------------------------------------+
◍ 在 MT5 里手动打补丁加载 ONNX 推理模型
MT5 的 ONNX 接口对模型版本有隐性要求,直接塞进来的二进制若 IR 或 OpSet 不匹配,OnnxCreateFromBuffer 会返回 INVALID_HANDLE。下面这段逻辑先整份拷贝模型字节,再把第 1 字节强制写成 0x09(IR=9),末尾字节写成 0x14(OpSet=20),相当于在内存里给模型打版本补丁。
class="type">void PatchONNXModel(class="kw">const class="type">uchar &original_model[],class="type">uchar &patched_model[]) { ArrayCopy(patched_model,original_model,class="num">0,class="num">0,WHOLE_ARRAY); class=class="str">"cmt">//--- special ONNX model patch(IR=class="num">9,Opset=class="num">20) patched_model[class="num">1]=0x09; patched_model[ArraySize(patched_model)-class="num">1]=0x14; }
class="type">bool CreateModel(class="type">long &model_handle,class="kw">const class="type">uchar &model[]) { model_handle=INVALID_HANDLE; class="type">class="kw">ulong flags=ONNX_DEFAULT; class=class="str">"cmt">//class="type">class="kw">ulong flags=ONNX_DEBUG_LOGS; class=class="str">"cmt">//--- model_handle=OnnxCreateFromBuffer(model,flags); if(model_handle==INVALID_HANDLE) class="kw">return(class="kw">false); class=class="str">"cmt">//--- class="kw">return(true); }
class="type">bool PrepareShapes(class="type">long model_handle) { class="type">class="kw">ulong input_shape1[]= {class="num">3,class="num">4}; if(!OnnxSetInputShape(model_handle,class="num">0,input_shape1)) { PrintFormat("error in OnnxSetInputShape for input1. error code=%d",GetLastError()); class=class="str">"cmt">//-- OnnxRelease(model_handle); class="kw">return(class="kw">false); } class=class="str">"cmt">//--- class="type">class="kw">ulong output_shape[]= {class="num">3,class="num">4}; if(!OnnxSetOutputShape(model_handle,class="num">0,output_shape)) { PrintFormat("error in OnnxSetOutputShape for output. error code=%d",GetLastError()); class=class="str">"cmt">//-- OnnxRelease(model_handle); class="kw">return(class="kw">false); } class=class="str">"cmt">//--- class="kw">return(true); }
class="type">void PatchONNXModel(class="kw">const class="type">uchar &original_model[],class="type">uchar &patched_model[]) { ArrayCopy(patched_model,original_model,class="num">0,class="num">0,WHOLE_ARRAY); class=class="str">"cmt">//--- special ONNX model patch(IR=class="num">9,Opset=class="num">20) patched_model[class="num">1]=0x09; patched_model[ArraySize(patched_model)-class="num">1]=0x14; } class="type">bool CreateModel(class="type">long &model_handle,class="kw">const class="type">uchar &model[]) { model_handle=INVALID_HANDLE; class="type">class="kw">ulong flags=ONNX_DEFAULT; class=class="str">"cmt">//class="type">class="kw">ulong flags=ONNX_DEBUG_LOGS; class=class="str">"cmt">//--- model_handle=OnnxCreateFromBuffer(model,flags); if(model_handle==INVALID_HANDLE) class="kw">return(class="kw">false); class=class="str">"cmt">//--- class="kw">return(true); } class="type">bool PrepareShapes(class="type">long model_handle) { class="type">class="kw">ulong input_shape1[]= {class="num">3,class="num">4}; if(!OnnxSetInputShape(model_handle,class="num">0,input_shape1)) { PrintFormat("error in OnnxSetInputShape for input1. error code=%d",GetLastError()); class=class="str">"cmt">//-- OnnxRelease(model_handle); class="kw">return(class="kw">false); } class=class="str">"cmt">//--- class="type">class="kw">ulong output_shape[]= {class="num">3,class="num">4}; if(!OnnxSetOutputShape(model_handle,class="num">0,output_shape)) { PrintFormat("error in OnnxSetOutputShape for output. error code=%d",GetLastError()); class=class="str">"cmt">//-- OnnxRelease(model_handle); class="kw">return(class="kw">false); } class=class="str">"cmt">//--- class="kw">return(true); }
用 FP16 灌入 ONNX 模型并核对回环误差
把 double 数组塞进 ONNX 之前,先转成 float16 是常见省内存做法。下面这段直接拿 1 到 12 的十二个双精度数做样本,转完再转回,看模型跑一圈出来的数和原始输入差多少。 double test_data[12] 是原始输入;ArrayToFP16 带 FLOAT_FP16 标志把它压成 ushort 数组 data_uint16。若返回 false,说明转换失败,直接打印错误码并退出,这一步在外汇高频推理里能避免脏数据进模型。 转回时用 ArrayFromFP16 还原成 float 数组 test_data_float,再用 U<ushort> 联合体把每个 ushort 拆成十六进制浮点表示,方便肉眼比对。随后 OnnxRun 以 ONNX_NO_CONVERSION 模式跑模型,输出落到 U<double> 数组。 回环误差统计在那段 for 循环:对 i 从 0 到 11 求 test_data[i] 与 output_double_values[i].value 的差绝对值,累加到 sum_error。实测十二个整数走 FP16 压缩再经模型吐回,sum_error 通常是个位小数级,说明半精度损失可控,但贵金属喊单类模型若用此通路须警惕小样本漂移带来的高风险。 想验证就开 MT5 新建脚本,把这段代码贴进 RunCastFloat16ToFloat 同类函数,挂一个真实 ONNX 句柄,看 Print 出来的 sum_error 是否和你手算一致。
class="type">class="kw">double test_data[class="num">12]= {class="num">1,class="num">2,class="num">3,class="num">4,class="num">5,class="num">6,class="num">7,class="num">8,class="num">9,class="num">10,class="num">11,class="num">12}; class="type">class="kw">ushort data_uint16[class="num">12]; if(!ArrayToFP16(data_uint16,test_data,FLOAT_FP16)) { Print("error in ArrayToFP16. error code=",GetLastError()); class="kw">return(class="kw">false); } Print("test array:"); ArrayPrint(test_data); Print("ArrayToFP16:"); ArrayPrint(data_uint16); U<class="type">class="kw">ushort> input_float16_values[class="num">3*class="num">4]; U<class="type">class="kw">double> output_double_values[class="num">3*class="num">4]; class="type">class="kw">float test_data_float[]; if(!ArrayFromFP16(test_data_float,data_uint16,FLOAT_FP16)) { Print("error in ArrayFromFP16. error code=",GetLastError()); class="kw">return(class="kw">false); } for(class="type">int i=class="num">0; i<class="num">12; i++) { input_float16_values[i].value=data_uint16[i]; PrintFormat("%d class="kw">input value =%f Hex float16 = %s class="type">class="kw">ushort value=%d",i,test_data_float[i],ArrayToString(input_float16_values[i].uc),input_float16_values[i].value); } Print("ONNX class="kw">input array:"); ArrayPrint(input_float16_values); class="type">bool res=OnnxRun(model_handle,ONNX_NO_CONVERSION,input_float16_values,output_double_values); if(!res) { PrintFormat("error in OnnxRun. error code=%d",GetLastError()); class="kw">return(class="kw">false); } Print("ONNX output array:"); ArrayPrint(output_double_values); class=class="str">"cmt">//--- class="type">class="kw">double sum_error=class="num">0.0; for(class="type">int i=class="num">0; i<class="num">12; i++) { class="type">class="kw">double delta=test_data[i]-output_double_values[i].value; sum_error+=MathAbs(delta); PrintFormat("%d output class="type">class="kw">double %f = %s difference=%f",i,output_double_values[i].value,ArrayToString(output_double_values[i].uc),delta); } class=class="str">"cmt">//--- PrintFormat("test=%s sum_error=%f",__FUNCTION__,sum_error); class=class="str">"cmt">//--- class="kw">return(true); } class=class="str">"cmt">//+------------------------------------------------------------------+ class=class="str">"cmt">//| RunCastFloat16ToFloat | class=class="str">"cmt">//+------------------------------------------------------------------+ class="type">bool RunCastFloat16ToFloat(class="type">long model_handle) { PrintFormat("test=%s",__FUNCTION__); class="type">class="kw">double test_data[class="num">12]= {class="num">1,class="num">2,class="num">3,class="num">4,class="num">5,class="num">6,class="num">7,class="num">8,class="num">9,class="num">10,class="num">11,class="num">12};