神经网络变得简单(第 97 部分):搭配 MSFformer 训练模型·进阶篇
(2/3)· 从架构定义到实盘数据评估,手把手把上篇搭好的模块喂进策略测试器
接上篇,我们继续深挖 MSFformer 的落地环节。很多人在自定义模型上卡住,不是算法写不出,而是训练流程和环境状态描述对不齐,导致回测阶段张量维度直接报错。
◍ 策略网络后半段的层定义细节
在 MT5 的强化学习策略构建里,actor 网络从第 4 层到第 7 层直接决定了动作输出的结构。下面这段声明把潜变量层、动作头、VAE 采样层和自由分布层依次挂到 actor 上,任何一层 Add 失败都会释放描述符并返回 false。 第 4、5 层都是 defNeuronBaseOCL 类型,count 分别取 LatentCount 和 2*NActions,激活函数用 SIGMOID,优化器统一 ADAM。第 5 层把输出维度翻倍,是为后续均值方差拆分预留空间。 第 6 层换成 defNeuronVAEOCL,count 等于 NActions,负责把潜向量重参数化为动作分布;第 7 层用 defNeuronFreDFOCL,window 设为 NActions、count 为 1、probability 0.7f,step 以 int(false) 即 0 传入,激活为 None。 在实盘接外汇或贵金属信号前,建议把 probability 从 0.7 调到 0.5 附近观察采样集中度变化,这类品种波动剧烈、杠杆风险高,分布过窄可能错失拐点。
if(!(descr = new CLayerDescription())) class="kw">return false; descr.type = defNeuronBaseOCL; descr.count = LatentCount; descr.activation = SIGMOID; descr.optimization = ADAM; if(!actor.Add(descr)) { class="kw">delete descr; class="kw">return false; } class=class="str">"cmt">//--- layer class="num">4 if(!(descr = new CLayerDescription())) class="kw">return false; descr.type = defNeuronBaseOCL; descr.count = LatentCount; descr.activation = SIGMOID; descr.optimization = ADAM; if(!actor.Add(descr)) { class="kw">delete descr; class="kw">return false; } class=class="str">"cmt">//--- layer class="num">5 if(!(descr = new CLayerDescription())) class="kw">return false; descr.type = defNeuronBaseOCL; descr.count = class="num">2 * NActions; descr.activation = None; descr.optimization = ADAM; if(!actor.Add(descr)) { class="kw">delete descr; class="kw">return false; } class=class="str">"cmt">//--- layer class="num">6 if(!(descr = new CLayerDescription())) class="kw">return false; descr.type = defNeuronVAEOCL; descr.count = NActions; descr.optimization = ADAM; if(!actor.Add(descr)) { class="kw">delete descr; class="kw">return false; } class=class="str">"cmt">//--- layer class="num">7 if(!(descr = new CLayerDescription())) class="kw">return false; descr.type = defNeuronFreDFOCL; descr.window = NActions; descr.count = class="num">1; descr.step = class="type">int(false); descr.probability = class="num">0.7f; descr.activation = None; descr.optimization = ADAM; if(!actor.Add(descr)) { class="kw">delete descr; class="kw">return false; }
「两套 EA 怎么把编码器与参与者训出来」
训练环境状态编码器用 StudyEncoder.mq5,训练参与者政策用 Study.mq5,后者顺带训一个评论者模型。评论者只在训练期引导参与者,实盘部署时不参与推理,这种「为训别人而训自己」的结构容易让人绕晕。 编码器为什么还要去预测指标后续值?多数指标由数字滤波器构成,把原生价格噪声压下去了,序列更平滑、可预测性更强。让编码器顺带拟合指标下一帧,等于用更规整的目标去校准它对价格走势的理解,而不是只盯原始 K 线。 StudyEncoder 初始化时先 LoadTotalBase 载训练集,再尝试 Load 预训练 Enc.nnw。加载失败才走 CreateEncoderDescriptions 建新架构并用随机参数初始化,所以复训是常态,冷启动反而是少数。维度校验块比对 Result.Total() 与 NForecast*BarDescr,防的是「模型文件和当前数据集对不上」,不是防建模型写错常量。 Train 方法里每次迭代按盈利能力给轨迹采样概率(编码器用不上余额持仓,但保留统一框架),前馈后不取预测值,只拿预测偏差对真实后续状态做反向传播,误差最小化即更新。迭代次数由外部参数定,终端手动停或报错才提前断。 Study.mq5 的参与者训练把时间戳编成正弦谐波向量,调用已训编码器前馈拿隐藏状态——注意是隐藏状态不是输出层,因为输出层掺了原始序列统计参数需另做归一化,隐藏态没有这层乖离。评论者吃隐藏态和动作出评估,先用自己的反向传播逼近实际奖励;随后关掉评论者训练模式,只让它把误差梯度传回参与者。 参与者参数从两个方向调:盈利轨迹做监督学习打底;亏损轨迹不扔,靠评论者建模「状态-动作-奖励」关系,把盈利目标抬 1%、亏损降 1% 反传,梯度指示动作该往哪挪。外汇与贵金属杠杆高,这套训练产物仅降低过拟合概率,不承诺样本外表现。 下面这段 OnInit 是编码器 EA 的骨架:载数据、载或建模型、维度校验。
class="type">int OnInit() { class=class="str">"cmt">//--- ResetLastError(); if(!LoadTotalBase()) { PrintFormat("Error of load study data: %d", GetLastError()); class="kw">return INIT_FAILED; } class=class="str">"cmt">//--- load models class="type">float temp; if(!Encoder.Load(FileName + "Enc.nnw", temp, temp, temp, dtStudied, true)) { Print("Create new model"); CArrayObj *encoder = new CArrayObj(); if(!CreateEncoderDescriptions(encoder)) { class="kw">delete encoder; class="kw">return INIT_FAILED; } if(!Encoder.Create(encoder)) { class="kw">delete encoder; class="kw">return INIT_FAILED; } class="kw">delete encoder; } class=class="str">"cmt">//--- Encoder.getResults(Result); if(Result.Total() != NForecast * BarDescr) { PrintFormat("The scope of the Encoder does not match the forecast state count(%d <> %d)", NForecast * BarDescr, Result.Total()); class="kw">return INIT_FAILED; } class=class="str">"cmt">//--- Encoder.GetLayerOutput(class="num">0, Result);
编码器维度校验与训练循环的血泪点
初始化阶段先卡一道硬门槛:若编码器输入规模不等于 HistoryBars 乘 BarDescr,直接 INIT_FAILED 并打出实际值与期望值。这一行能帮你当场抓出状态描述维度配错,不用等跑了几千根 K 线才崩。 随后用 EventChartCustom 往图表抛一个自定义事件 'Init',失败同样退回 INIT_FAILED。MT5 里这类自定义事件若 ChartID 取错或图表已关,GetLastError 会直接暴露问题,开终端看专家日志就能定位。 Train 函数才是耗时大户。用 GetTickCount 记起点,每满 500 毫秒就按 iter/Iterations*100 算一次进度——这意味着若 Iterations 设成 10000,每半秒你至少能读到 5% 级别的反馈,不会黑盒干转。 采样时 i 的算法用了 MathRand 平方再除以 32767 平方,把均匀采样压成偏向 0 的分布,再乘 (Buffer[tr].Total - 2 - NForecast)。若 i<=0 或 state 全零就 iter-- 重抽,等于悄悄丢弃无效轨迹。外汇与贵金属波动跳变多,这种剔除能降低编码器被垃圾状态带偏的概率,但高频品种仍可能因样本稀疏导致过拟合风险。
if(Result.Total() != (HistoryBars * BarDescr)) { PrintFormat("Input size of Encoder doesn&class="macro">#x27;t match state description(%d <> %d)", Result.Total(), (HistoryBars * BarDescr)); class="kw">return INIT_FAILED; } if(!EventChartCustom(ChartID(), class="num">1, class="num">0, class="num">0, "Init")) { PrintFormat("Error of create study event: %d", GetLastError()); class="kw">return INIT_FAILED; } class=class="str">"cmt">//--- class="kw">return(INIT_SUCCEEDED); } 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, state; 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 i = (class="type">int)((MathRand() * MathRand() / MathPow(class="num">32767, class="num">2)) * (Buffer[tr].Total - class="num">2 - NForecast)); if(i <= class="num">0) { iter--; class="kw">continue; } state.Assign(Buffer[tr].States[i].state); if(MathAbs(state).Sum()==class="num">0) { iter--; class="kw">continue; } bState.AssignArray(state); class=class="str">"cmt">//--- State Encoder if(!Encoder.feedForward((CBufferFloat*)GetPointer(bState), class="num">1, false, (CBufferFloat*)NULL)) { PrintFormat("%s -> %d", __FUNCTION__, __LINE__); Stop = true; class="kw">break; } class=class="str">"cmt">//--- Collect target data if(!Result.AssignArray(Buffer[tr].States[i + NForecast].state)) class="kw">continue; if(!Result.Resize(BarDescr * NForecast)) class="kw">continue; if(!Encoder.backProp(Result,(CBufferFloat*)NULL)) { PrintFormat("%s -> %d", __FUNCTION__, __LINE__); Stop = true; class="kw">break; } if(GetTickCount() - ticks > class="num">500) { class="type">class="kw">double percent = class="type">class="kw">double(iter) * class="num">100.0 / (Iterations);
◍ 训练循环里的状态编码与时间特征
在 Train() 函数里,每一轮迭代先从概率轨迹采样,挑出一条轨迹 tr,再用 MathRand() 平方归一化落到 Buffer[tr].Total-2 的范围内取状态索引 i,避开 0 和空状态。这种双重随机的做法让采样偏向分布尾部,训练集覆盖更离散。 时间特征被拆成四条三角函数:以 2023 全年秒数(约 31536000)为周期的 sin、月线 PERIOD_MN1 的 cos、周线 PERIOD_W1 的 sin、日线 PERIOD_D1 的 sin,全部做了非零保护。把它们塞进 bTime 再 BufferWrite(),等于把日历节律编码进网络输入。 状态向量 state 经 bState 中转后送进 Encoder.feedForward(),只跑前向、不训练、无目标缓冲。若返回失败则直接跳过本轮,iter 不自增。外汇与贵金属市场高杠杆、滑点无常,这套编码在实盘重训时可能因价差跳变导致状态矩为 0,需要你自己在 MT5 里打印 state.Sum() 验证。 下面这段是训练主循环尾部与退出日志的原文片段,可见误差打印与 ExpertRemove() 的硬停止逻辑:
class="type">class="kw">string str = StringFormat("%-14s %class="num">6.2f%% -> Error %class="num">15.8f\n", "Encoder", percent, Encoder.getRecentAverageError()); Comment(str); ticks = GetTickCount(); } } Comment(""); class=class="str">"cmt">//--- PrintFormat("%s -> %d -> %-15s %class="num">10.7f", __FUNCTION__, __LINE__, "Encoder", Encoder.getRecentAverageError()); ExpertRemove(); class=class="str">"cmt">//--- } class="type">void Train(class="type">void) { class=class="str">"cmt">//--- vector<class="type">float> probability = GetProbTrajectories(Buffer, class="num">0.9); class=class="str">"cmt">//--- vector<class="type">float> result, target, state; 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 i = (class="type">int)((MathRand() * MathRand() / MathPow(class="num">32767, class="num">2)) * (Buffer[tr].Total - class="num">2)); if(i <= class="num">0) { iter--; class="kw">continue; } state.Assign(Buffer[tr].States[i].state); if(MathAbs(state).Sum()==class="num">0) { iter--; class="kw">continue; } bState.AssignArray(state); bTime.Clear(); 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;); bTime.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); bTime.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); bTime.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); bTime.Add((class="type">float)MathSin(x != class="num">0 ? class="num">2.0 * M_PI * x : class="num">0)); if(bTime.GetIndex() >= class="num">0) bTime.BufferWrite(); class=class="str">"cmt">//--- State Encoder if(!Encoder.feedForward((CBufferFloat*)GetPointer(bState), class="num">1, false,(CBufferFloat*)NULL)) {
「回放缓冲区里的 Critic 与 Actor 训练闭环」
在经验回放的单条轨迹遍历中,代码先取当前步的动作向量喂给 Critic 做前向。若 bActions.GetIndex() >= 0 才写缓冲,否则跳过;任何 feedForward 返回失败都会打印 __FUNCTION__ 与 __LINE__ 并置 Stop=true 跳出,这种防御式退出能避免脏数据继续污染梯度。
Critic 训练目标用两步奖励差构造:result = Buffer[tr].States[i+1].rewards - Buffer[tr].States[i+2].rewards * DiscFactor,折扣因子把未来步奖励折损后从下一步奖励里扣减,再经 backProp 回传。外汇与贵金属行情非平稳,DiscFactor 取 0.9 或 0.95 会明显改变策略对延后奖励的敏感度,建议开 MT5 把该常量改成不同值各跑一轮回测对比。
账户特征构造段把净值、权益变化率等 9 个量压进 bAccount:例如 (Buffer[tr].States[i].account[0] - PrevBalance) / PrevBalance 是余额环比变化率,account[4]/PrevBalance 等三项都做了余额归一。归一后特征再送 Actor 前向;若失败同样断点退出。
这套闭环跑通后,Actor 输出经 Critic 二次前向(TrainMode 切 false)得到动作价值估计。实盘接此逻辑前,先在策略测试器用 2023 年 XAUUSD 的 M5 数据验证缓冲区写入索引是否恒为非负,否则训练可能静默空转。
PrintFormat("%s -> %d", __FUNCTION__, __LINE__); Stop = true; class="kw">break; } class=class="str">"cmt">//--- Critic bActions.AssignArray(Buffer[tr].States[i].action); if(bActions.GetIndex() >= class="num">0) bActions.BufferWrite(); Critic.TrainMode(true); if(!Critic.feedForward((CBufferFloat*)GetPointer(bActions), class="num">1, false, GetPointer(Encoder), LatentLayer)) { PrintFormat("%s -> %d", __FUNCTION__, __LINE__); Stop = true; class="kw">break; } result.Assign(Buffer[tr].States[i + class="num">1].rewards); target.Assign(Buffer[tr].States[i + class="num">2].rewards); result = result - target * DiscFactor; Result.AssignArray(result); if(!Critic.backProp(Result, (CNet *)GetPointer(Encoder), LatentLayer)) { PrintFormat("%s -> %d", __FUNCTION__, __LINE__); Stop = true; class="kw">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); bAccount.AddArray(GetPointer(bTime)); if(bAccount.GetIndex() >= class="num">0) bAccount.BufferWrite(); class=class="str">"cmt">//--- Actor if(!Actor.feedForward((CBufferFloat*)GetPointer(bAccount), class="num">1, false, GetPointer(Encoder), LatentLayer)) { PrintFormat("%s -> %d", __FUNCTION__, __LINE__); Stop = true; class="kw">break; } Critic.TrainMode(false); if(!Critic.feedForward((CNet *)GetPointer(Actor), -class="num">1, (CNet*)GetPointer(Encoder), LatentLayer)) {