神经网络变得简单(第 75 部分):提升轨迹预测模型的性能·综合运用
(3/3)·复杂模型拖慢实时决策?借鉴自动驾驶基线的剪枝思路,把预测成本压下来而不牺牲轨迹质量
「训练循环里的模型落盘与采样细节」
在训练函数退出前,代码会把四个神经网络权重分别存盘:状态编码器、端点编码器、主编码器和端点输出网络,文件名后缀固定为 StEnc.nnw、EndEnc.nnw、Enc.nnw、Endp.nnw,概率网络另存为 Prob.nnw。Save 末位参数传 true 表示强制覆盖旧权重,外汇与贵金属行情高波动,重训前先确认这些 .nnw 是否已备份,避免覆盖后无法回滚。 Train 函数开头先用 GetProbTrajectories(Buffer, 0.9) 拿到概率轨迹向量,阈值 0.9 决定采样偏向。主循环受 Iterations 上限、IsStopped 与 Stop 三重控制,任意一次前向传播失败就把 Stop 置真并 break,训练可能在中途静默中止。 每轮迭代先用 SampleTrajectory 按概率抽一条轨迹,batch 大小 = GPTBars + 48。起始 state 用 MathRand 平方再除以 32767 平方做偏置采样,使小序号区间被抽中的概率更低;若算出的 state ≤ 0 就 iter-- 并重抽。端点 end 取 state+batch 与 Buffer 总长减 PrecoderBars 的较小值,保证不越界。 前向阶段先 BLEncoder.feedForward 吃单帧状态,再 BLEndpoints.feedForward 接编码器输出,最后 BLProbability.feedForward 接同一编码器。任一层返回 false 就打印函数名与行号并停训,开 MT5 跑这套时建议把 Experts 日志打开,否则训练崩在哪一层的哪一行很难查。
StateEncoder.Save(FileName + "StEnc.nnw", class="num">0, class="num">0, class="num">0, TimeCurrent(), true); EndpointEncoder.Save(FileName + "EndEnc.nnw", class="num">0, class="num">0, class="num">0, TimeCurrent(), true); BLEncoder.Save(FileName + "Enc.nnw", class="num">0, class="num">0, class="num">0, TimeCurrent(), true); BLEndpoints.Save(FileName + "Endp.nnw", class="num">0, class="num">0, class="num">0, TimeCurrent(), true); BLProbability.Save(FileName + "Prob.nnw", class="num">0, class="num">0, class="num">0, TimeCurrent(), true); } class="kw">delete Result; class="kw">delete OpenCL; } class="type">void Train(class="type">void) { class=class="str">"cmt">//--- vector<class="type">float> probability = GetProbTrajectories(Buffer, class="num">0.9); vector<class="type">float> result, target; matrix<class="type">float> targets, temp_m; class="type">bool Stop = false; class=class="str">"cmt">//--- class="type">uint ticks = GetTickCount(); for(class="type">int iter = class="num">0; (iter < Iterations && !IsStopped() && !Stop); iter ++) { class="type">int tr = SampleTrajectory(probability); class="type">int batch = GPTBars + class="num">48; class="type">int state = (class="type">int)((MathRand() * MathRand() / MathPow(class="num">32767, class="num">2)) * (Buffer[tr].Total - class="num">2 - PrecoderBars - batch)); if(state <= class="num">0) { iter--; class="kw">continue; } BLEncoder.Clear(); BLEndpoints.Clear(); class="type">int end = MathMin(state + batch, Buffer[tr].Total - PrecoderBars); for(class="type">int i = state; i < end; i++) { bState.AssignArray(Buffer[tr].States[i].state); class=class="str">"cmt">//--- Trajectory if(!BLEncoder.feedForward((CBufferFloat*)GetPointer(bState), class="num">1, false, (CBufferFloat*)NULL)) { PrintFormat("%s -> %d", __FUNCTION__, __LINE__); Stop = true; break; } if(!BLEndpoints.feedForward((CNet*)GetPointer(BLEncoder), -class="num">1, (CBufferFloat*)NULL)) { PrintFormat("%s -> %d", __FUNCTION__, __LINE__); Stop = true; break; } if(!BLProbability.feedForward((CNet*)GetPointer(BLEncoder), -class="num">1,
◍ 端点预测里的矩阵整形与极值翻转
这段逻辑在做 BLEndpoints 预测结果的二次加工:先拿指针失败就打印函数行号并停循环,随后用 Zeros 建一个 PrecoderBars×3 的浮点矩阵装目标序列。 对每个 t 把状态向量Assign进 target,若长度超出 BarDescr 就做 Row→Reshape(BarDescr)→Resize(,3) 的降维截取,只留最后一行,再写回 targets 的第 t 行。 用 Col(0).CumSum() 做累计后,把三列重算成累计基准上的偏移;方向由 state[8] 与 state[7] 的大小决定,extr 取累计列的 ArgMax 或 ArgMin。若 extr==0 则反向再取一次极值,并把 targets 截断到 extr+1 行。 direct>=0 时水平取 Max 且第2列取 Col(2).Min(),否则水平取 Min 且第1列取 Col(1).Max();最后把 BLEndpoints 结果 reshape 成 NForecast×3,做平方误差矩阵并取垂直求和的最小位置 pos。外汇与贵金属行情受杠杆影响大,这类推断仅作概率参考,实盘前请在 MT5 用样本数据单步验证。
(CNet*)GetPointer(BLEndpoints))) { PrintFormat("%s -> %d", __FUNCTION__, __LINE__); Stop = true; break; } targets = matrix<class="type">float>::Zeros(PrecoderBars, class="num">3); for(class="type">int t = class="num">0; t < PrecoderBars; t++) { target.Assign(Buffer[tr].States[i + class="num">1 + t].state); if(target.Size() > BarDescr) { matrix<class="type">float> temp(class="num">1, target.Size()); temp.Row(target, class="num">0); temp.Reshape(target.Size() / BarDescr, BarDescr); temp.Resize(temp.Rows(), class="num">3); target = temp.Row(temp.Rows() - class="num">1); } targets.Row(target, t); } target = targets.Col(class="num">0).CumSum(); targets.Col(target, class="num">0); targets.Col(target + targets.Col(class="num">1), class="num">1); targets.Col(target + targets.Col(class="num">2), class="num">2); class="type">int direct = (Buffer[tr].States[i].state[class="num">8] >= Buffer[tr].States[i].state[class="num">7] ? class="num">1 : -class="num">1); class="type">ulong extr=(direct>class="num">0 ? target.ArgMax() : target.ArgMin()); if(extr==class="num">0) { direct=-direct; extr=(direct>class="num">0 ? target.ArgMax() : target.ArgMin()); } targets.Resize(extr+class="num">1, class="num">3); if(direct >= class="num">0) { target = targets.Max(AXIS_HORZ); target[class="num">2] = targets.Col(class="num">2).Min(); } else { target = targets.Min(AXIS_HORZ); target[class="num">1] = targets.Col(class="num">1).Max(); } BLEndpoints.getResults(result); targets.Reshape(class="num">1, result.Size()); targets.Row(result, class="num">0); targets.Reshape(NForecast, class="num">3); temp_m = targets; for(class="type">int i = class="num">0; i < class="num">3; i++) temp_m.Col(temp_m.Col(i) - target[i], i); temp_m = MathPow(temp_m, class="num">2.0f); class="type">ulong pos = temp_m.Sum(AXIS_VERT).ArgMin();
反向传播里塞进账户状态与周期相位
这段训练循环把三个子网络的反向传播串在一起:端点网络、编码器、概率网络。任何一处 backProp 返回失败就立刻打印函数名加行号、置 Stop 并 break,实盘回测时若日志频繁出现这类输出,说明梯度链在某层断掉了,优先查输入 buffer 是否为空。
概率网络的反向传播依赖一个 one-hot 形式的 bProbs:先清零成 NForecast 长度的零向量,再把当前 pos 位置写成 1 并 BufferWrite,相当于告诉网络“这一步就该押这个仓位标号”。
账户特征 bAccount 的构造值得照搬:用上一根 K 线的余额、净值做差分和比值,把绝对资金量归一化成相对变化率,避免不同账户规模下梯度尺度漂移。第 8 个字段 account[7] 是时间戳,代码把它除以 2023→2024 的秒数得到年相位,再取 MathSin;另两个分别除以月线、周线秒数取 Cos,把日历周期编码进状态向量。
[CODE]
targets.Row(target, pos);
Result.AssignArray(targets);
if(!BLEndpoints.backProp(Result, (CBufferFloat*)NULL))
{
PrintFormat("%s -> %d", __FUNCTION__, __LINE__);
Stop = true;
break;
}
if(!BLEncoder.backPropGradient((CBufferFloat*)NULL))
{
PrintFormat("%s -> %d", __FUNCTION__, __LINE__);
Stop = true;
break;
}
bProbs.AssignArray(vector<float>::Zeros(NForecast));
bProbs.Update((int)pos, 1);
bProbs.BufferWrite();
if(!BLProbability.backProp(GetPointer(bProbs), GetPointer(BLEndpoints)))
{
PrintFormat("%s -> %d", __FUNCTION__, __LINE__);
Stop = true;
break;
}
//--- Policy
float PrevBalance = Buffer[tr].States[MathMax(i - 1, 0)].account[0];
float PrevEquity = Buffer[tr].States[MathMax(i - 1, 0)].account[1];
bAccount.Clear();
bAccount.Add((Buffer[tr].States[i].account[0] - PrevBalance) / PrevBalance);
bAccount.Add(Buffer[tr].States[i].account[1] / PrevBalance);
bAccount.Add((Buffer[tr].States[i].account[1] - PrevEquity) / PrevEquity);
bAccount.Add(Buffer[tr].States[i].account[2]);
bAccount.Add(Buffer[tr].States[i].account[3]);
bAccount.Add(Buffer[tr].States[i].account[4] / PrevBalance);
bAccount.Add(Buffer[tr].States[i].account[5] / PrevBalance);
bAccount.Add(Buffer[tr].States[i].account[6] / PrevBalance);
double time = (double)Buffer[tr].States[i].account[7];
double x = time / (double)(D'2024.01.01' - D'2023.01.01');
bAccount.Add((float)MathSin(x != 0 ? 2.0 * M_PI * x : 0));
x = time / (double)PeriodSeconds(PERIOD_MN1);
bAccount.Add((float)MathCos(x != 0 ? 2.0 * M_PI * x : 0));
x = time / (double)PeriodSeconds(PERIOD_W1);
[/CODE]
逐行拆一下关键行:
targets.Row(target, pos); 把当前样本标签写进 targets 矩阵第 pos 行。
Result.AssignArray(targets); 把标签矩阵挂到 Result 供端点网络反传。
if(!BLEndpoints.backProp(...)) 端点网络反传,NULL 表示不传外部梯度。
bProbs.Update((int)pos, 1); 在 pos 位写 1,构造 one-hot 监督信号。
float PrevBalance = ... MathMax(i-1,0) 取前一根或第 0 根余额,防越界。
bAccount.Add((... - PrevBalance)/PrevBalance) 余额相对变化率,归一化特征。
double x = time / (D'2024.01.01'-D'2023.01.01') 年周期相位基准,约 31536000 秒。
MathSin(x!=0?2*M_PI*x:0) 非零才算正弦,避免除零脏值进网络。
外汇与贵金属杠杆高,这类 RL 特征工程只是训练端逻辑,实盘信号仍需自行验证胜率与回撤。
targets.Row(target, pos); Result.AssignArray(targets); if(!BLEndpoints.backProp(Result, (CBufferFloat*)NULL)) { PrintFormat("%s -> %d", __FUNCTION__, __LINE__); Stop = true; break; } if(!BLEncoder.backPropGradient((CBufferFloat*)NULL)) { PrintFormat("%s -> %d", __FUNCTION__, __LINE__); Stop = true; break; } bProbs.AssignArray(vector<class="type">float>::Zeros(NForecast)); bProbs.Update((class="type">int)pos, class="num">1); bProbs.BufferWrite(); if(!BLProbability.backProp(GetPointer(bProbs), GetPointer(BLEndpoints))) { PrintFormat("%s -> %d", __FUNCTION__, __LINE__); Stop = true; break; } class=class="str">"cmt">//--- Policy class="type">float PrevBalance = Buffer[tr].States[MathMax(i - class="num">1, class="num">0)].account[class="num">0]; class="type">float PrevEquity = Buffer[tr].States[MathMax(i - class="num">1, class="num">0)].account[class="num">1]; bAccount.Clear(); bAccount.Add((Buffer[tr].States[i].account[class="num">0] - PrevBalance) / PrevBalance); bAccount.Add(Buffer[tr].States[i].account[class="num">1] / PrevBalance); bAccount.Add((Buffer[tr].States[i].account[class="num">1] - PrevEquity) / PrevEquity); bAccount.Add(Buffer[tr].States[i].account[class="num">2]); bAccount.Add(Buffer[tr].States[i].account[class="num">3]); bAccount.Add(Buffer[tr].States[i].account[class="num">4] / PrevBalance); bAccount.Add(Buffer[tr].States[i].account[class="num">5] / PrevBalance); bAccount.Add(Buffer[tr].States[i].account[class="num">6] / PrevBalance); class="type">class="kw">double time = (class="type">class="kw">double)Buffer[tr].States[i].account[class="num">7]; class="type">class="kw">double x = time / (class="type">class="kw">double)(D&class="macro">#x27;class="num">2024.01.class="num">01&class="macro">#x27; - D&class="macro">#x27;class="num">2023.01.class="num">01&class="macro">#x27;); bAccount.Add((class="type">float)MathSin(x != class="num">0 ? class="num">2.0 * M_PI * x : class="num">0)); x = time / (class="type">class="kw">double)PeriodSeconds(PERIOD_MN1); bAccount.Add((class="type">float)MathCos(x != class="num">0 ? class="num">2.0 * M_PI * x : class="num">0)); x = time / (class="type">class="kw">double)PeriodSeconds(PERIOD_W1);
「把账户状态塞进强化学习网络」
这段逻辑干的事很直接:把账户特征按日线周期做正弦编码后写入缓冲区,再依次喂给状态编码器、端点编码器和 Actor 网络。x = time / PeriodSeconds(PERIOD_D1) 意味着时间轴被压缩成以「天」为单位的连续相位,MathSin(2.0*M_PI*x) 给出的是归一化周期信号,避免网络把绝对时间戳当泛化特征。 bAccount.GetIndex() >= 0 才执行 BufferWrite(),是在确认账户向量已成型;任一层 feedForward 返回 false 就打印函数名加行号、置 Stop 并 break,方便你在 MT5 Experts 日志里精准定位是哪一层前向传播崩了。 direct > 0 分支里有个硬阈值:state[4] > 30 且 state[5] > -100 才计算出场参数。tp 用 target[1]/_Point/MaxTP 归一化,sl 取 target[1]/3 与 -target[2] 的最大值再除 _Point 并兜底到 MaxSL/10,仓位则按 risk/(value*sl) 与 0.01 取大。外汇与贵金属杠杆高,这类自动仓位计算仅作信号参考,实盘可能触发超预期回撤。 开 MT5 把这段贴进你的 EA 训练循环,先单步跑 direct>0 分支,看 state[4]、state[5] 在样本上是否真能稳定过阈,再决定是否接实盘。
bAccount.Add((class="type">float)MathSin(x != class="num">0 ? class="num">2.0 * M_PI * x : class="num">0)); x = time / (class="type">class="kw">double)PeriodSeconds(PERIOD_D1); bAccount.Add((class="type">float)MathSin(x != class="num">0 ? class="num">2.0 * M_PI * x : class="num">0)); if(bAccount.GetIndex() >= class="num">0) bAccount.BufferWrite(); class=class="str">"cmt">//--- State embedding if(!StateEncoder.feedForward((CNet *)GetPointer(BLEncoder), -class="num">1, (CBufferFloat*)GetPointer(bAccount))) { PrintFormat("%s -> %d", __FUNCTION__, __LINE__); Stop = true; break; } class=class="str">"cmt">//--- Endpoint embedding if(!EndpointEncoder.feedForward((CNet *)GetPointer(BLEndpoints), -class="num">1, (CNet*)GetPointer(BLProbability))) { PrintFormat("%s -> %d", __FUNCTION__, __LINE__); Stop = true; break; } class=class="str">"cmt">//--- Actor if(!Actor.feedForward((CNet *)GetPointer(StateEncoder), -class="num">1, (CNet*)GetPointer(EndpointEncoder))) { PrintFormat("%s -> %d", __FUNCTION__, __LINE__); Stop = true; break; } if(direct > class="num">0) { if(Buffer[tr].States[i].state[class="num">4] > class="num">30 && Buffer[tr].States[i].state[class="num">5] > -class="num">100 ) { class="type">float tp = class="type">float(target[class="num">1] / _Point / MaxTP); result[class="num">1] = tp; class="type">int sl = class="type">int(MathMax(MathMax(target[class="num">1] / class="num">3, -target[class="num">2]) / _Point, MaxSL / class="num">10)); result[class="num">2] = class="type">float(sl) / MaxSL; result[class="num">0] = class="type">float(MathMax(risk / (value * sl), class="num">0.01)) + FLT_EPSILON; } }
◍ 反向传播与训练进度回显的实现细节
这段逻辑处在强化学习训练循环的反向传播阶段。当状态缓存里 state[4] 低于 70 且 state[5] 低于 100 时,才会计算局部止盈比例 tp(以 -target[2] 换算成点数再除以 MaxTP 归一化),并用 MathMax 约束止损下限不低于 MaxSL/10,避免过小止损在贵金属跳空时直接被打掉。 随后依次对 Actor、StateEncoder、EndpointEncoder 调用 backProp 与 backPropGradient。只要任一环节返回 false,就打印函数名与行号、置 Stop=true 并 break——这种硬退出能保证权重不会被半成品梯度污染。 训练耗时用 GetTickCount 做节流:每超过 500 毫秒才刷新一次 Comment。进度百分比按 (i-state)/(end-state)+iter 除以总迭代次数算,同时把 Actor、Endpoints、Probability 三家网络的近期平均误差用 %15.8f 精度打印出来。在 MT5 里跑时,你能直接看到误差是否从 1e-1 量级往 1e-3 收敛,从而判断该不该提前终止。 外汇与贵金属杠杆高、滑点随机,这类自研训练循环若误差不降反跳,大概率过拟合到历史噪点,实盘使用前务必用样本外数据复核。
else { if(Buffer[tr].States[i].state[class="num">4] < class="num">70 && Buffer[tr].States[i].state[class="num">5] < class="num">100 ) { class="type">float tp = class="type">float((-target[class="num">2]) / _Point / MaxTP); result[class="num">4] = tp; class="type">int sl = class="type">int(MathMax(MathMax((-target[class="num">2]) / class="num">3, target[class="num">1]) / _Point, MaxSL / class="num">10)); result[class="num">5] = class="type">float(sl) / MaxSL; result[class="num">3] = class="type">float(MathMax(risk / (value * sl), class="num">0.01)) + FLT_EPSILON; } } Result.AssignArray(result); if(!Actor.backProp(Result, (CNet *)GetPointer(EndpointEncoder)) || !StateEncoder.backPropGradient(GetPointer(bAccount), (CBufferFloat *)GetPointer(bGradient)) || !EndpointEncoder.backPropGradient((CNet*)GetPointer(BLProbability)) ) { PrintFormat("%s -> %d", __FUNCTION__, __LINE__); Stop = true; break; } if(!BLEncoder.backPropGradient((CBufferFloat*)NULL)) { PrintFormat("%s -> %d", __FUNCTION__, __LINE__); Stop = true; break; } if(GetTickCount() - ticks > class="num">500) { class="type">class="kw">double percent = (class="type">class="kw">double(i - state) / ((end - state)) + iter) * class="num">100.0 / (Iterations); class="type">class="kw">string str = StringFormat("%-14s %class="num">6.2f%% -> Error %class="num">15.8f\n", "Actor", percent, Actor.getRecentAverageError()); str += StringFormat("%-14s %class="num">6.2f%% -> Error %class="num">15.8f\n", "Endpoints", percent, BLEndpoints.getRecentAverageError()); str += StringFormat("%-14s %class="num">6.2f%% -> Error %class="num">15.8f\n", "Probability", percent, BLProbability.getRecentAverageError()); Comment(str); ticks = GetTickCount(); } } } Comment(""); class=class="str">"cmt">//---
用 PrintFormat 把三层网络误差抖出来
在 MT5 的 EA 调试里,直接把每层神经网络的近期平均误差打印到日志,比盲目跑回测更容易定位哪一层在拖后腿。下面这段把 Actor、Endpoints、Probability 三者的 getRecentAverageError() 依次输出,格式固定为左对齐 15 字符名称加 10.7f 精度浮点。 PrintFormat 的第一个 %s 填 __FUNCTION__,第二个 %d 填 __LINE__,这样出错时能直接反查到具体函数与行号。对外汇与贵金属模型来说,三层误差若同时高于 0.05 概率上说明学习率或样本窗口要重调,这类品种波动噪声大、高风险,别拿单次打印当收敛证据。 最后一行 ExpertRemove() 会在打印完立即卸载 EA,适合做一次性诊断而非常驻监控。开 MT5 把这段代码塞进 OnTick 末尾,跑几十根 K 线就能看到三层误差的真实分布。
PrintFormat("%s -> %d -> %-15s %class="num">10.7f", __FUNCTION__, __LINE__, "Actor", Actor.getRecentAverageError()); PrintFormat("%s -> %d -> %-15s %class="num">10.7f", __FUNCTION__, __LINE__, "Endpoints", BLEndpoints.getRecentAverageError()); PrintFormat("%s -> %d -> %-15s %class="num">10.7f", __FUNCTION__, __LINE__, "Probability", BLProbability.getRecentAverageError()); ExpertRemove(); class=class="str">"cmt">//--- }
「EURUSD 上跑出的盈利因子与稀疏交易」
把前面搭好的图卷积层接进 MT5 策略测试器,用 EURUSD H1 的 2023 年前 7 个月数据做训练与测试。之前文章里攒的经验回放缓冲区不用重采,只要把数据文件改名成 BaseLines.bd 就能直接喂给训练流程,省掉了环境交互 EA 重新跑数据的开销。 训练时目标值生成阶段可以反复用同一份训练集调参,不必边训边补数据。但第一版结果并不乐观,测试窗口从 1 个月拉长到 3 个月后才勉强看到正期望。 最终模型在训练集和测试集上都盈利,盈利因子 1.4,且基于 7 个月历史训出的权重在随后至少 3 个月里保持正收益,说明抓到的预测变量可能有一定跨期稳定性。外汇与贵金属属高风险品种,历史稳定性不代表未来延续。 致命短板是交易频次:11 个月只成交 3 笔。对实盘交易者来说,这种稀疏程度基本丧失策略意义,模型方向可能对,但样本量太小,参数或阈值大概率还要重调。
◍ 保守决策是下一步要啃的硬骨头
前面几篇把优化轨迹预测模型的底子打好了,实现思路让训练后的模型能抓住源数据里真正显要的预测因子,训练完挺长时间内跑起来都算稳。 但实测下来有个绕不开的现象:模型做决策偏保守,直接反映在成交次数极少上。我们回测里看到的就是下单频率明显低于传统阈值触发策略。 这种极少成交的保守性,意味着信号过滤可能过狠,漏掉了不少本可捕获的波段。外汇和贵金属市场高波动、高杠杆,模型不敢出手反而容易错过风险收益比合适的窗口。 所以接下来要做的,不是换框架,而是调宽松度——从特征权重衰减系数和决策置信门槛两头松,看成交频次和回撤怎么走。开 MT5 把这两参数拉一下对比曲线,比空谈解释性实在。
随包附带的几支程序
这套 LSTM 多元时间序列预测方案不是只给思路,作者把整套工程文件都打进了 MQL5.zip(871.72 KB)里,直接下下来就能在 MT5 跑。里面分了样本收集、训练、测试三条线:Research.mq5 和 ResearchRealORL.mq5 是两类采集 EA,前者普通采样,后者用 Real-ORL 方法补样本;Study.mq5 管模型训练,Test.mq5 做回测验证。 底层依赖三个库:Trajectory.mqh 定义系统状态结构,NeuroNet.mqh 封装建网类,NeuroNet.cl 是 OpenCL 核函数,显存并行训练就靠它。外汇与贵金属市场高杠杆、滑点随机,拿这些 EA 实盘前务必先在策略测试器用历史数据过一遍。 有一点提醒:压缩包里的代码与文章反映作者个人观点,平台方不对信息准确性及使用后果负责,部分或全文转载均被禁止。你下载后改参调试可以,但别把原文件外发,权归属 MetaQuotes 与作者 Dmitriy Gizlyk。