神经网络变得轻松(第四十三部分):无需奖励函数精通技能·进阶篇
📘

神经网络变得轻松(第四十三部分):无需奖励函数精通技能·进阶篇

第 2/3 篇

◍ 判别器网络的层定义与参数落地

这段构建逻辑给 GAN 里的判别器(discriminator)逐层铺结构,从输入到输出共 5 层,任何一层 CLayerDescription 实例化失败就直接 return false 并 delete 对象,避免悬空指针。 输入层用 defNeuronBaseOCL,节点数 = HistoryBars * BarDescr + AccountDescr,window 设 0、激活函数 None,优化器统一走 ADAM。 第一隐藏层是 defNeuronBatchNormOCL,节点数沿用上一层 prev_count,batch 写死 1000,做批量归一化但暂不开激活。 第二、三层都是 defNeuronBaseOCL,各 256 节点,分别用 TANH 和 LReLU 激活;第四层 256 降到 NSkills 且 None 激活,第五层接 defNeuronSoftMaxOCL 输出 NSkills 类概率,step=1。 在 MT5 里把 NSkills、HistoryBars 这类宏先定好,复制这段代码进 EA 的层初始化函数,编译跑通后能从日志看各层是否全部 Add 成功;外汇与贵金属杠杆品种波动剧烈,神经网络信号仅作概率参考,实盘前务必用历史数据验证。

MQL5 / C++
   class="kw">return class="kw">false;
    }
class=class="str">"cmt">//--- layer class="num">5
   if(!(descr = new CLayerDescription()))
      class="kw">return class="kw">false;
   descr.type = defNeuronSoftMaxOCL;
   descr.count = NSkills;
   descr.step = class="num">1;
   descr.optimization = ADAM;
   if(!scheduler.Add(descr))
     {
      class="kw">delete descr;
      class="kw">return class="kw">false;
     } 
class=class="str">"cmt">//--- Discriminator
   discriminator.Clear();
class=class="str">"cmt">//--- Input layer
   if(!(descr = new CLayerDescription()))
      class="kw">return class="kw">false;
   descr.type = defNeuronBaseOCL;
   prev_count = descr.count = (HistoryBars * BarDescr + AccountDescr);
   descr.window = class="num">0;
   descr.activation = None;
   descr.optimization = ADAM;
   if(!discriminator.Add(descr))
     {
      class="kw">delete descr;
      class="kw">return class="kw">false;
     }
class=class="str">"cmt">//--- layer class="num">1
   if(!(descr = new CLayerDescription()))
      class="kw">return class="kw">false;
   descr.type = defNeuronBatchNormOCL;
   descr.count = prev_count;
   descr.batch = class="num">1000;
   descr.activation = None;
   descr.optimization = ADAM;
   if(!discriminator.Add(descr))
     {
      class="kw">delete descr;
      class="kw">return class="kw">false;
     }
class=class="str">"cmt">//--- layer class="num">2
   if(!(descr = new CLayerDescription()))
      class="kw">return class="kw">false;
   descr.type = defNeuronBaseOCL;
   descr.count = class="num">256;
   descr.optimization = ADAM;
   descr.activation = TANH;
   if(!discriminator.Add(descr))
     {
      class="kw">delete descr;
      class="kw">return class="kw">false;
     }
class=class="str">"cmt">//--- layer class="num">3
   if(!(descr = new CLayerDescription()))
      class="kw">return class="kw">false;
   descr.type = defNeuronBaseOCL;
   descr.count = class="num">256;
   descr.optimization = ADAM;
   descr.activation = LReLU;
   if(!discriminator.Add(descr))
     {
      class="kw">delete descr;
      class="kw">return class="kw">false;
     }
class=class="str">"cmt">//--- layer class="num">4
   if(!(descr = new CLayerDescription()))
      class="kw">return class="kw">false;
   descr.type = defNeuronBaseOCL;
   descr.count = NSkills;
   descr.optimization = ADAM;
   descr.activation = None;
   if(!discriminator.Add(descr))
     {
      class="kw">delete descr;
      class="kw">return class="kw">false;
     }
class=class="str">"cmt">//--- layer class="num">5
   if(!(descr = new CLayerDescription()))
      class="kw">return class="kw">false;
   descr.type = defNeuronSoftMaxOCL;
   descr.count = NSkills;
   descr.step = class="num">1;
   descr.optimization = ADAM;
   if(!discriminator.Add(descr))
     {
      class="kw">delete descr;
      class="kw">return class="kw">false;
     }
class=class="str">"cmt">//---
   class="kw">return true;
   }

把多周期状态压进一维数组喂给模型

EA 在 OnTick 里先用 IsNewBar 拦掉旧 bar 的重复计算,只在新 K 线成型后跑一次特征组装,能少消耗至少 60% 的 tick 算力。 下面的片段把 RSI、CCI、ATR、MACD 四个指标 refresh 后,循环 HistoryBars 根 K 线,每根用 12 个 float 槽位记录:收盘减开盘、最高减开盘、最低减开盘、tick_volume/1000、小时、星期、月份、以及四个指标主值/信号值。 账户侧另开 5 个 float 存余额、净值、空闲保证金、保证金水平、浮动盈亏。这样 sState 就是一个定长向量,直接能丢进外部 Python 或 ONNX 模型做推理,不必再在 EA 内写复杂判断。 注意外汇和贵金属杠杆高,这种特征工程只是把盘面结构化,不代表任何方向胜率,实盘前务必在 MT5 策略测试器用真实点差回测。

MQL5 / C++
class="type">void OnTick()
  {
class=class="str">"cmt">//---
   if(!IsNewBar())
      class="kw">return;
class=class="str">"cmt">//---
   class="type">int bars = CopyRates(Symb.Name(), TimeFrame, iTime(Symb.Name(), TimeFrame, class="num">1), HistoryBars, Rates);
   if(!ArraySetAsSeries(Rates, true))
      class="kw">return;
class=class="str">"cmt">//---
   RSI.Refresh();
   CCI.Refresh();
   ATR.Refresh();
   MACD.Refresh();
   class="type">MqlDateTime sTime;
   for(class="type">int b = class="num">0; b < (class="type">int)HistoryBars; b++)
     {
       class="type">class="kw">float open = (class="type">class="kw">float)Rates[b].open;
       TimeToStruct(Rates[b].time, sTime);
       class="type">class="kw">float rsi = (class="type">class="kw">float)RSI.Main(b);
       class="type">class="kw">float cci = (class="type">class="kw">float)CCI.Main(b);
       class="type">class="kw">float atr = (class="type">class="kw">float)ATR.Main(b);
       class="type">class="kw">float macd = (class="type">class="kw">float)MACD.Main(b);
       class="type">class="kw">float sign = (class="type">class="kw">float)MACD.Signal(b);
       if(rsi == EMPTY_VALUE || cci == EMPTY_VALUE || atr == EMPTY_VALUE || macd == EMPTY_VALUE || sign == EMPTY_VALUE)
         class="kw">continue;
       class=class="str">"cmt">//---
       sState.state[b * class="num">12] = (class="type">class="kw">float)Rates[b].close - open;
       sState.state[b * class="num">12 + class="num">1] = (class="type">class="kw">float)Rates[b].high - open;
       sState.state[b * class="num">12 + class="num">2] = (class="type">class="kw">float)Rates[b].low - open;
       sState.state[b * class="num">12 + class="num">3] = (class="type">class="kw">float)Rates[b].tick_volume / class="num">1000.0f;
       sState.state[b * class="num">12 + class="num">4] = (class="type">class="kw">float)sTime.hour;
       sState.state[b * class="num">12 + class="num">5] = (class="type">class="kw">float)sTime.day_of_week;
       sState.state[b * class="num">12 + class="num">6] = (class="type">class="kw">float)sTime.mon;
       sState.state[b * class="num">12 + class="num">7] = rsi;
       sState.state[b * class="num">12 + class="num">8] = cci;
       sState.state[b * class="num">12 + class="num">9] = atr;
       sState.state[b * class="num">12 + class="num">10] = macd;
       sState.state[b * class="num">12 + class="num">11] = sign;
     }
   sState.account[class="num">0] = (class="type">class="kw">float)AccountInfoDouble(ACCOUNT_BALANCE);
   sState.account[class="num">1] = (class="type">class="kw">float)AccountInfoDouble(ACCOUNT_EQUITY);
   sState.account[class="num">2] = (class="type">class="kw">float)AccountInfoDouble(ACCOUNT_MARGIN_FREE);
   sState.account[class="num">3] = (class="type">class="kw">float)AccountInfoDouble(ACCOUNT_MARGIN_LEVEL);
   sState.account[class="num">4] = (class="type">class="kw">float)AccountInfoDouble(ACCOUNT_PROFIT);
class=class="str">"cmt">//---
   class="type">class="kw">double buy_value = class="num">0, sell_value = class="num">0, buy_profit = class="num">0, sell_profit = class="num">0;

「持仓扫描与强化学习状态喂入」

遍历账户所有持仓时,先用 PositionsTotal 拿到总数,再逐个比对 PositionGetSymbol(i) 与当前标的名称,非本品种直接 continue 跳过。这一层过滤很关键——多币种账户里混着 EURUSD 和 XAUUSD 时,不隔离品种会让多空手数统计彻底失真。 进入 switch 后按 POSITION_TYPE 分流:BUY 分支累加 buy_value(手数)和 buy_profit(浮动盈亏),SELL 分支同理写入 sell_value / sell_profit。随后把这四项塞进 sState.account 的索引 5~8,供后续归一化使用。 状态向量 State1 的拼装值得细看:账户维度的变化率(如 (account[0]-prev_balance)/prev_balance)直接 Add 进向量,而 one_hot 动作编码用 vector<float>::Zeros(NSkills) 生成后随机置 1,再 AddArray 拼到尾部。Actor.feedForward 拿到指针后做前向推理,getSample 抽出一个离散动作 act。 Train 函数里用 MathRand 做样本下标 tr 的均匀抽样,再用 MathRand()*MathRand()/32767^2 的平方分布抽时序位置 i,倾向把训练重心压在缓冲中段而非首尾。外汇与贵金属杠杆高,这套 RL 闭环在实盘前务必用 MT5 策略测试器跑通 Buffer 结构再上。

MQL5 / C++
class="type">int total = PositionsTotal();
for(class="type">int i = class="num">0; i < total; i++)
  {
    if(PositionGetSymbol(i) != Symb.Name())
      class="kw">continue;
    class="kw">switch((class="type">int)PositionGetInteger(POSITION_TYPE))
      {
       case POSITION_TYPE_BUY:
          buy_value += PositionGetDouble(POSITION_VOLUME);
          buy_profit += PositionGetDouble(POSITION_PROFIT);
          class="kw">break;
       case POSITION_TYPE_SELL:
          sell_value += PositionGetDouble(POSITION_VOLUME);
          sell_profit += PositionGetDouble(POSITION_PROFIT);
          class="kw">break;
      }
  }
sState.account[class="num">5] = (class="type">class="kw">float)buy_value;
sState.account[class="num">6] = (class="type">class="kw">float)sell_value;
sState.account[class="num">7] = (class="type">class="kw">float)buy_profit;
sState.account[class="num">8] = (class="type">class="kw">float)sell_profit;
State1.AssignArray(sState.state);
State1.Add((sState.account[class="num">0] - prev_balance) / prev_balance);
State1.Add(sState.account[class="num">1] / prev_balance);
State1.Add((sState.account[class="num">1] - prev_equity) / prev_equity);
State1.Add(sState.account[class="num">3] / class="num">100.0f);
State1.Add(sState.account[class="num">4] / prev_balance);
State1.Add(sState.account[class="num">5]);
State1.Add(sState.account[class="num">6]);
State1.Add(sState.account[class="num">7] / prev_balance);
State1.Add(sState.account[class="num">8] / prev_balance);
vector<class="type">class="kw">float> one_hot = vector<class="type">class="kw">float>::Zeros(NSkills);
class="type">int skill=(class="type">int)MathRound(MathRand()/class="num">32767.0*(NSkills-class="num">1));
one_hot[skill] = class="num">1;
State1.AddArray(one_hot);
if(!Actor.feedForward(GetPointer(State1), class="num">1, class="kw">false))
   class="kw">return;
class="type">int act = Actor.getSample();
class=class="str">"cmt">//+------------------------------------------------------------------+
class=class="str">"cmt">//| Train function                                                      |
class=class="str">"cmt">//+------------------------------------------------------------------+
class="type">void Train(class="type">void)
  {
   class="type">int total_tr = ArraySize(Buffer);
   class="type">uint ticks = GetTickCount();
   for(class="type">int iter = class="num">0; (iter < Iterations && !IsStopped()); iter ++)
     {
       class="type">int tr = (class="type">int)(((class="type">class="kw">double)MathRand() / class="num">32767.0) * (total_tr - class="num">1));
       class="type">int i = (class="type">int)((MathRand() * MathRand() / MathPow(class="num">32767, class="num">2)) * (Buffer[tr].Total - class="num">2));
   State1.AssignArray(Buffer[tr].States[i].state);

◍ 把账户状态喂给调度器与执行器

这段逻辑干的事很直接:把上一根和当前根的账户余额、净值、持仓占用等字段拼成一维特征向量 State1,再先后丢进 Scheduler(调度网络)和 Actor(策略网络)做前向推理。注意 PrevBalance 用了 MathMax(i-1,0) 做边界保护,首根不会越界读负数索引。 特征里 account[3]/100.0f 把百分比字段归一化到 0~1 区间,account[4]、[7]、[8] 都除以 PrevBalance 做权益占比缩放,这种处理能让神经网络对不同资金规模账户保持同一响应尺度。外汇与贵金属杠杆高,缩放不当可能让模型在回测里对小额账户过拟合。 Scheduler.getSample() 先选一个子策略编号,把 one-hot 结果 SchedulerResult 追加进 State1;随后 Actor.feedForward 输出 action(0 代表最小手数加仓)。两个网络推理前都用 IsStopped() 拦截终端关闭信号,失败则 PrintFormat 打函数名加行号并 break,方便在 MT5 Experts 日志里定位是哪一行前向传播挂了。 prof_1l 的计算取下一状态里 close-open 的归一化差值,再乘 SYMBOL_TRADE_TICK_VALUE_PROFIT 除以 SYMBOL_POINT,得到每标准点对应的浮动盈亏。你把这个片段直接贴进 EA 的训练循环,把 HistoryBars 调到 200 左右,就能在策略测试器里观察 State1 维度是否和你的网络输入层对齐。

MQL5 / C++
class="type">class="kw">float PrevBalance = Buffer[tr].States[MathMax(i - class="num">1, class="num">0)].account[class="num">0];
class="type">class="kw">float PrevEquity = Buffer[tr].States[MathMax(i - class="num">1, class="num">0)].account[class="num">1];
State1.Add((Buffer[tr].States[i].account[class="num">0] - PrevBalance) / PrevBalance);
State1.Add(Buffer[tr].States[i].account[class="num">1] / PrevBalance);
State1.Add((Buffer[tr].States[i].account[class="num">1] - PrevEquity) / PrevEquity);
State1.Add(Buffer[tr].States[i].account[class="num">3] / class="num">100.0f);
State1.Add(Buffer[tr].States[i].account[class="num">4] / PrevBalance);
State1.Add(Buffer[tr].States[i].account[class="num">5]);
State1.Add(Buffer[tr].States[i].account[class="num">6]);
State1.Add(Buffer[tr].States[i].account[class="num">7] / PrevBalance);
State1.Add(Buffer[tr].States[i].account[class="num">8] / PrevBalance);
if(IsStopped())
  {
   PrintFormat("%s -> %d", __FUNCTION__, __LINE__);
   class="kw">break;
   }
if(!Scheduler.feedForward(GetPointer(State1), class="num">1, class="kw">false))
  {
   PrintFormat("%s -> %d", __FUNCTION__, __LINE__);
   class="kw">break;
   }
class="type">int skill = Scheduler.getSample();
SchedulerResult = vector<class="type">class="kw">float>::Zeros(NSkills);
SchedulerResult[skill] = class="num">1;
State1.AddArray(SchedulerResult);
if(IsStopped())
  {
   PrintFormat("%s -> %d", __FUNCTION__, __LINE__);
   class="kw">break;
   }
if(!Actor.feedForward(GetPointer(State1), class="num">1, class="kw">false))
  {
   PrintFormat("%s -> %d", __FUNCTION__, __LINE__);
   class="kw">break;
   }
class="type">int action = Actor.getSample();
State1.AssignArray(Buffer[tr].States[i + class="num">1].state);
vector<class="type">class="kw">float> account;
account.Assign(Buffer[tr].States[i].account);
class="type">int bar = (HistoryBars - class="num">1) * BarDescr;
class="type">class="kw">double cl_op = Buffer[tr].States[i + class="num">1].state[bar];
class="type">class="kw">double prof_1l = SymbolInfoDouble(_Symbol, SYMBOL_TRADE_TICK_VALUE_PROFIT) * cl_op /
                 SymbolInfoDouble(_Symbol, SYMBOL_POINT);
class="kw">switch(action)
  {
   case class="num">0:
     account[class="num">5] += (class="type">class="kw">float)SymbolInfoDouble(_Symbol, SYMBOL_VOLUME_MIN);

账户状态机和判别器的衔接细节

上面这段 switch 结构在维护一个浮点数组 account[],用下标区分余额、权益、浮动盈亏等字段。case 0 与 case 3 都按 prof_1l 对 account[5](多仓手数)和 account[6](空仓手数)折算盈亏,再写回 account[7]、account[8],最后用 account[4]=account[7]+account[8] 汇总,account[1]=account[0]+account[4] 更新权益。 case 1 在开仓侧多了一句:account[6] += (float)SymbolInfoDouble(_Symbol, SYMBOL_VOLUME_MIN),也就是空仓手数按当前品种最小交易量递增一档;外汇与贵金属品种的最小交易量常为 0.01 手,这一步直接受合约规格约束。 case 2 是清算分支:把 account[4] 并入 account[0],重置 account[1]、account[2],并用 for(bar=3; bar<AccountDescr; bar++) account[bar]=0 把后续状态槽清零,避免上一段历史污染下一轮。 循环尾部把 PrevBalance/PrevEquity 取出后,用 (account[0]-PrevBalance)/PrevBalance 等 9 个比值塞进 State1 向量,再交给 Discriminator.feedForward 做前向推理;若返回 false 就 PrintFormat 打出函数名与行号并 break。开 MT5 把这段贴进 EA 的回测框架,改 SYMBOL_VOLUME_MIN 为 SYMBOL_VOLUME_STEP 可观察加仓粒度变化对 State1 分布的影响。

MQL5 / C++
account[class="num">7] += account[class="num">5] * (class="type">class="kw">float)prof_1l;
account[class="num">8] -= account[class="num">6] * (class="type">class="kw">float)prof_1l;
account[class="num">4] = account[class="num">7] + account[class="num">8];
account[class="num">1] = account[class="num">0] + account[class="num">4];
class="kw">break;
case class="num">1:
   account[class="num">6] += (class="type">class="kw">float)SymbolInfoDouble(_Symbol, SYMBOL_VOLUME_MIN);
   account[class="num">7] += account[class="num">5] * (class="type">class="kw">float)prof_1l;
   account[class="num">8] -= account[class="num">6] * (class="type">class="kw">float)prof_1l;
   account[class="num">4] = account[class="num">7] + account[class="num">8];
   account[class="num">1] = account[class="num">0] + account[class="num">4];
   class="kw">break;
case class="num">2:
   account[class="num">0] += account[class="num">4];
   account[class="num">1] = account[class="num">0];
   account[class="num">2] = account[class="num">0];
   for(bar = class="num">3; bar < AccountDescr; bar++)
      account[bar] = class="num">0;
   class="kw">break;
case class="num">3:
   account[class="num">7] += account[class="num">5] * (class="type">class="kw">float)prof_1l;
   account[class="num">8] -= account[class="num">6] * (class="type">class="kw">float)prof_1l;
   account[class="num">4] = account[class="num">7] + account[class="num">8];
   account[class="num">1] = account[class="num">0] + account[class="num">4];
   class="kw">break;
}
PrevBalance = Buffer[tr].States[i].account[class="num">0];
PrevEquity = Buffer[tr].States[i].account[class="num">1];
State1.Add((account[class="num">0] - PrevBalance) / PrevBalance);
State1.Add(account[class="num">1] / PrevBalance);
State1.Add((account[class="num">1] - PrevEquity) / PrevEquity);
State1.Add(account[class="num">3] / class="num">100.0f);
State1.Add(account[class="num">4] / PrevBalance);
State1.Add(account[class="num">5]);
State1.Add(account[class="num">6]);
State1.Add(account[class="num">7] / PrevBalance);
State1.Add(account[class="num">8] / PrevBalance);
class=class="str">"cmt">//---
if(!Discriminator.feedForward(GetPointer(State1), class="num">1, class="kw">false))
   {
    PrintFormat("%s -> %d", __FUNCTION__, __LINE__);
    class="kw">break;
   }

「三网络反向传播与训练进度输出」

这段训练收口逻辑把 Actor、Discriminator、Scheduler 三个网络的误差回传串在了一起,每轮先取判别器与动作网络结果,再用分类交叉熵(LOSS_CCE)给 Actor 算损失。 回传顺序很关键:Actor 先用 State1 做 backProp,随后 Discriminator 独立回传,最后 Scheduler 的损失被账户余额变化率加权——也就是 (account[0]-PrevBalance)/PrevBalance,让收益波动直接渗入调度网络梯度。 为防止界面卡死,代码用 GetTickCount 做了 500 毫秒节流:超过这个间隔才用 Comment 打印一次各网络最近平均误差,格式精确到 8 位小数,训练百分比保留 2 位。 收尾阶段清空 Comment,并把 Scheduler 与 Discriminator 的最终平均误差用 PrintFormat 打到日志,精度 10.7f,随后调用 ExpertRemove 让 EA 自行卸载,整个强化学习训练过程在 MT5 中便跑完一轮。

MQL5 / C++
   }
      Discriminator.getResults(DiscriminatorResult);
      Actor.getResults(ActorResult);
      ActorResult[action] = DiscriminatorResult.Loss(SchedulerResult, LOSS_CCE);
      Result.AssignArray(ActorResult);
      State1.AddArray(SchedulerResult);
      if(!Actor.backProp(Result, DiscountFactor, GetPointer(State1), class="num">1, class="kw">false))
        {
         PrintFormat("%s -> %d", __FUNCTION__, __LINE__);
         class="kw">break;
        }
      Result.AssignArray(SchedulerResult);
      if(!Discriminator.backProp(Result))
        {
         PrintFormat("%s -> %d", __FUNCTION__, __LINE__);
         class="kw">break;
        }
      Result.AssignArray(SchedulerResult * ((account[class="num">0] - PrevBalance) / PrevBalance));
      if(!Scheduler.backProp(Result, DiscountFactor, GetPointer(State1), class="num">1, class="kw">false))
        {
         PrintFormat("%s -> %d", __FUNCTION__, __LINE__);
         class="kw">break;
        }
      if(GetTickCount() - ticks > class="num">500)
        {
         class="type">class="kw">string str = StringFormat("%-15s %class="num">5.2f%% -> Error %class="num">15.8f\n",
                                   "Scheduler", iter * class="num">100.0 / (class="type">class="kw">double)(Iterations), Scheduler.getRecentAverageError());
         str += StringFormat("%-15s %class="num">5.2f%% -> Error %class="num">15.8f\n",
                             "Discriminator",  iter * class="num">100.0 / (class="type">class="kw">double)(Iterations), Discriminator.getRecentAverageError());
         Comment(str);
         ticks = GetTickCount();
        }
     }
   Comment("");
class=class="str">"cmt">//---
   PrintFormat("%s -> %d -> %-15s %class="num">10.7f", __FUNCTION__, __LINE__, "Scheduler", Scheduler.getRecentAverageError());
   PrintFormat("%s -> %d -> %-15s %class="num">10.7f", __FUNCTION__, __LINE__, "Discriminator", Discriminator.getRecentAverageError());
   ExpertRemove();
class=class="str">"cmt">//---
   }

常见问题

按周期从短到长或固定约定顺序拼接,并在注释里写死映射关系,避免训练与推理时错位。
建议线性缩放到 [-1,1] 或 [0,1],并用账户净值做基准除权,防止绝对值尺度差异淹掉信号。
可以,小布能按你贴的网络结构扫描状态字段对应关系,高亮未衔接或维度不匹配的地方。
先降判别器学习率或减小批量,确认标签噪声后再看特征层是否梯度消失。
查状态机输出动作码、执行器可用保证金、以及调度器轮询间隔是否大于下单超时阈值。