神经网络变得轻松(第三十三部分):分布式 Q-学习中的分位数回归·进阶篇
📘

神经网络变得轻松(第三十三部分):分布式 Q-学习中的分位数回归·进阶篇

第 2/3 篇

QRDQN 网络初始化与分位回传的细节

CQRDQN 的 Create 里先校验动作数大于 0 且基类网络创建成功,再从描述数组取最后一层,拿首神经元算分位数数量:iNumbers = neuron.Neurons() / actions。若动作空间设错,这里直接返回 false,MT5 编译器不会报错但运行时策略不训练。 分位锚点 mTaus 用全 1 向量除以 iNumbers 生成等距点,再把首个元素减半后做累加:mTaus[0,0] /= 2; mTaus = mTaus.CumSum(0)。比如 iNumbers=5 时,原始是 [0.2,0.2,0.2,0.2,0.2],首元素变 0.1 后累加得 [0.1,0.3,0.5,0.7,0.9],这就是分位回归的累积概率标尺。 backProp 中若有 nextState,先喂给目标网络 cTargetNet 取结果,转成 1×temp.Size() 矩阵后 Reshape 为 iActions×iNumbers,按列求均值再乘 discount 加进 target。注意 target.Size() 必须等于 iActions,否则直接返回 false。 逐动作算梯度时,用 q - target[a] 得误差,正向部分 Clip 到 [0, FLT_MAX]、负向 Clip 到 [-FLT_MAX, 0],再分别乘 (mTaus-1) 和 -mTaus。这套非对称加权就是分位损失的核心,改 mTaus 的锚点分布会直接动学习偏向。外汇与贵金属杠杆高,跑这套网络前先用历史 Tick 验证维度匹配,避免实盘爆掉。

MQL5 / C++
mTaus[class="num">0, class="num">0] /= class="num">2;
mTaus = mTaus.CumSum(class="num">0);
cTargetNet.Create(NULL);
Create(NULL, iActions);
}
                    CQRDQN(CArrayObj *Description)  { Create(Description, iActions); }
class="type">bool CQRDQN::Create(CArrayObj *Description, class="type">uint actions)
  {
  if(actions <= class="num">0 || !CNet::Create(Description))
      class="kw">return false;
  class="type">int last_layer = Description.Total() - class="num">1;
  CLayer *layer = layers.At(last_layer);
  if(!layer)
      class="kw">return false;
  CNeuronBaseOCL *neuron = layer.At(class="num">0);
  if(!neuron)
      class="kw">return false;
  iActions = actions;
  iNumbers = neuron.Neurons() / actions;
  mTaus = matrix<class="type">float>::Ones(class="num">1, iNumbers) / iNumbers;
  mTaus[class="num">0, class="num">0] /= class="num">2;
  mTaus = mTaus.CumSum(class="num">0);
  cTargetNet.Create(NULL);
class=class="str">"cmt">//---
  class="kw">return true;
  }
  class="type">bool                feedForward(CArrayFloat *inputVals, class="type">int window = class="num">1, class="type">bool tem = true)
                    { class="kw">return CNet::feedForward(inputVals, window, tem); }
class="type">bool CQRDQN::backProp(CBufferFloat *targetVals, class="type">float discount,
CArrayFloat *nextState=NULL, class="type">int window = class="num">1, class="type">bool tem = true)
  {
class=class="str">"cmt">//---
  if(!targetVals)
      class="kw">return false;
  vectorf target;
  if(!targetVals.GetData(target) || target.Size() != iActions)
      class="kw">return false;
  if(!!nextState)
     {
      if(!cTargetNet.feedForward(nextState, window, tem))
        class="kw">return false;
      vectorf temp;
      cTargetNet.getResults(targetVals);
      if(!targetVals.GetData(temp))
        class="kw">return false;
      matrixf q = matrixf::Zeros(class="num">1, temp.Size());
      if(!q.Row(temp, class="num">0) || !q.Reshape(iActions, iNumbers))
        class="kw">return false;
      temp = q.Mean(class="num">0);
      target = target + discount * temp.Max();
     }
  vectorf quantils;
  getResults(targetVals);
  if(!targetVals.GetData(quantils))
      class="kw">return false;
  matrixf Q = matrixf::Zeros(class="num">1, quantils.Size());
  if(!Q.Row(quantils, class="num">0) || !Q.Reshape(iActions, iNumbers))
      class="kw">return false;
  for(class="type">uint a = class="num">0; a < iActions; a++)
    {
      vectorf q = Q.Row(a);
      vectorf dp = q - target[a], dn = dp;
      if(!dp.Clip(class="num">0, FLT_MAX) || !dn.Clip(-FLT_MAX, class="num">0))
        class="kw">return false;
      dp = (mTaus.Row(class="num">0) - class="num">1) * dp;
      dn = mTaus.Row(class="num">0) * dn * (-class="num">1);

◍ QRDQN 里结果解析与动作抽样的实现细节

这段 CQRDQN 派生类代码承接了网络反向传播之后的事:把原始输出重排成动作-分位数矩阵,再决定下一步采哪个动作。getResults 里先把网络输出读进 temp 向量,初始化 1 行 temp.Size() 列的矩阵,按行写入后 Reshape 成 iActions × iNumbers 的二维结构,最后用 q.Mean(1) 沿动作轴求均值写回 resultVals。 getAction 直接拿 getResults 的结果调 temp.Maximum(0, temp.Total()),返回均值 Q 值最大的下标,相当于贪婪选动作;若缓冲区为空则返回 -1,调用方要自己处理这个异常码。 getSample 则走探索逻辑:同样重排矩阵后,对 q.Mean(1) 做 AF_SOFTMAX 激活得到概率分布,CumSum 累积成 0~1 的区间边界。随后用 MathRandomNormal(0.5, 0.5) 抽一个正态随机数,落在哪个累积区间就返回对应动作索引;random>=1 时兜底取最后一个动作。外汇与贵金属行情高阶矩厚尾,这种正态采样在极端波动下可能低估尾部风险,实盘前建议在 MT5 策略测试器里把随机种子跑 50 次以上看动作分布偏移。

MQL5 / C++
  if(!Q.Row(dp + dn + q, a))
        class="kw">return false;
  }
  if(!targetVals.AssignArray(Q))
      class="kw">return false;
  if(iCountBackProp >= iUpdateTarget)
  {
class="macro">#ifdef FileName
      if(UpdateTarget(FileName + ".nnw"))
class="macro">#else
      if(UpdateTarget("QRDQN.upd"))
class="macro">#endif
          iCountBackProp = class="num">0;
  }
  else
      iCountBackProp++;
class="macro">#define FileName          Symb.Name()+"_"+EnumToString(TimeFrame)+"_"+StringSubstr(__FILE__,class="num">0,StringFind(__FILE__,".",class="num">0))
  class="kw">return CNet::backProp(targetVals);
  }
class="type">void CQRDQN::getResults(CBufferFloat *&resultVals)
  {
  CNet::getResults(resultVals);
  if(!resultVals)
    class="kw">return;
  vectorf temp;
  if(!resultVals.GetData(temp))
    {
      class="kw">delete resultVals;
      class="kw">return;
    }
  matrixf q;
  if(!q.Init(class="num">1, temp.Size()) || !q.Row(temp, class="num">0) || !q.Reshape(iActions, iNumbers))
    {
      class="kw">delete resultVals;
      class="kw">return;
    }
class=class="str">"cmt">//---
  if(!resultVals.AssignArray(q.Mean(class="num">1)))
    {
      class="kw">delete resultVals;
      class="kw">return;
    }
class=class="str">"cmt">//---
  }
class="type">int CQRDQN::getAction(class="type">void)
  {
  CBufferFloat *temp;
  getResults(temp);
  if(!temp)
    class="kw">return -class="num">1;
class=class="str">"cmt">//---
  class="kw">return temp.Maximum(class="num">0, temp.Total());
  }
class="type">int CQRDQN::getSample(class="type">void)
  {
  CBufferFloat* resultVals;
  CNet::getResults(resultVals);
  if(!resultVals)
    class="kw">return -class="num">1;
  vectorf temp;
  if(!resultVals.GetData(temp))
    {
      class="kw">delete resultVals;
      class="kw">return -class="num">1;
    }
  class="kw">delete resultVals;
  matrixf q;
  if(!q.Init(class="num">1, temp.Size()) || !q.Row(temp, class="num">0) || !q.Reshape(iActions, iNumbers))
    {
      class="kw">delete resultVals;
      class="kw">return -class="num">1;
    }
  if(!q.Mean(class="num">1).Activation(temp, AF_SOFTMAX))
    class="kw">return -class="num">1;
  temp = temp.CumSum();
  class="type">int err_code;
  class="type">float random = (class="type">float)Math::MathRandomNormal(class="num">0.5, class="num">0.5, err_code);
  if(random >= class="num">1)
    class="kw">return (class="type">int)temp.Size() - class="num">1;
  for(class="type">int i = class="num">0; i < (class="type">int)temp.Size(); i++)
    if(random <= temp[i] && temp[i] > class="num">0)
      class="kw">return i;
class=class="str">"cmt">//---
  class="kw">return -class="num">1;
  }

「目标网络权重热替换的实现细节」

在 DQN 类智能体里,目标网络(Target Network)需要周期性从在线网络同步权重,避免训练目标漂移。下面这段 CQRDQN::UpdateTarget 方法给出了一种轻量做法:先把当前在线网络存盘,再让目标网络从同一文件加载,从而完成一次无声的权重覆盖。 [CODE] bool CQRDQN::UpdateTarget(string file_name) { if(!Save(file_name, 0, false)) return false; float error, undefine, forecast; datetime time; if(!cTargetNet.Load(file_name, error, undefine, forecast, time, false)) return false; iCountBackProp = 0; //--- return true; } [/CODE] 逐行拆解:第 1 行定义返回布尔值的更新函数,入参为落盘文件名;第 3 行调用 Save 把在线网络以编号 0、不覆盖日志的方式写出,失败直接返回 false;第 5–6 行声明加载所需的浮点与时间变量;第 7 行用 cTargetNet.Load 从同一文件读入权重,四个输出参数接收误差等状态,失败同样返回 false;第 10 行把反向传播计数清零,意味着本轮目标同步后重新统计训练步数;第 13 行返回 true 表示同步成功。 实盘接这套逻辑时,外汇与贵金属波动大、滑点高,热替换频率若过密可能让目标值抖动,建议先在 MT5 策略测试器里用 2020–2023 年 XAUUSD 的 M15 数据跑一遍,观察 iCountBackProp 归零周期对回测夏普的影响,再决定同步间隔。

MQL5 / C++
class="type">bool CQRDQN::UpdateTarget(class="type">class="kw">string file_name)
  {
   if(!Save(file_name, class="num">0, false))
      class="kw">return false;
   class="type">float error, undefine, forecast;
   class="type">class="kw">datetime time;
   if(!cTargetNet.Load(file_name, error, undefine, forecast, time, false))
      class="kw">return false;
   iCountBackProp = class="num">0;
class=class="str">"cmt">//---
   class="kw">return true;
  }

QRDQN 模型在 MT5 里的训练与回测落点

训练用的 EA 叫 QRDQN-learning.mq5,是在原 Q-learning 框架上改的:换掉被训练模型类,并删掉目标网络实例声明。初始化时从 .nnw 文件载入模型,强制开全部神经层学习模式,历史深度对齐源数据层大小,动作域和 target 更新周期也一并写入——这里故意把更新周期设成 1000000,等于把目标网更新握在自己手里。 模型架构沿用上一版的 NetCreator 产物,只摘掉了末层 SoftMax,让输出区间能直接映射奖励策略的原始数值。训练数据取 EURUSD 的 H1 周期、过去 2 年历史,回测放在策略测试器里跑,另写了一个 QRDQN-learning-test.mq 做验证。 短期表现上,模型在 2 周窗口内倾向盈利,超一半交易以盈利平仓,平均盈利约为平均亏损的 2 倍。外汇和贵金属杠杆高,这类回测结论只说明历史样本内的概率倾向,实盘可能明显偏离。 下面这段 OnInit 与 Train 骨架,是验证训练流程最直接的入口:载入模型后开 TrainMode,用 GetLayerOutput(0) 算出 HistoryBars,再把目标网更新周期推到极大值。Train 里先切 2 年历史到 Rates 数组,指标缓冲区按 bars 扩容,外层循环控迭代、内层留作前馈后馈。

MQL5 / C++
CSymbolInfo                Symb;
class="type">MqlRates                    Rates[];
CQRDQN                      StudyNet;
CBufferFloat               *TempData;
CiRSI                       RSI;
CiCCI                       CCI;
CiATR                       ATR;
CiMACD                      MACD;
class="type">int OnInit()
  {
class=class="str">"cmt">//---
.........
.........
class=class="str">"cmt">//---
   if(!StudyNet.Load(FileName + ".nnw", dtStudied, false))
      class="kw">return INIT_FAILED;
   if(!StudyNet.TrainMode(true))
      class="kw">return INIT_FAILED;
class=class="str">"cmt">//---
   if(!StudyNet.GetLayerOutput(class="num">0, TempData))
      class="kw">return INIT_FAILED;
   HistoryBars = TempData.Total() / class="num">12;
   if(!StudyNet.SetActions(Actions))
      class="kw">return INIT_PARAMETERS_INCORRECT;
   StudyNet.SetUpdateTarget(class="num">1000000);
class=class="str">"cmt">//---
........
class=class="str">"cmt">//---
   class="kw">return(INIT_SUCCEEDED);
   }
class="type">void Train(class="type">void)
  {
class=class="str">"cmt">//---
   class="type">MqlDateTime start_time;
   TimeCurrent(start_time);
   start_time.year -= StudyPeriod;
   if(start_time.year <= class="num">0)
      start_time.year = class="num">1900;
   class="type">class="kw">datetime st_time = StructToTime(start_time);
   class="type">int bars = CopyRates(Symb.Name(), TimeFrame, st_time, TimeCurrent(), Rates);
   if(!RSI.BufferResize(bars) || !CCI.BufferResize(bars) || !ATR.BufferResize(bars) || !MACD.BufferResize(bars))
     {
      PrintFormat("%s -> %d", __FUNCTION__, __LINE__);
      ExpertRemove();
      class="kw">return;
     }
   if(!ArraySetAsSeries(Rates, true))
     {
      PrintFormat("%s -> %d", __FUNCTION__, __LINE__);
      ExpertRemove();
      class="kw">return;
     }
class=class="str">"cmt">//---
   RSI.Refresh();
   CCI.Refresh();
   ATR.Refresh();
   MACD.Refresh();
   class="type">int total = bars - (class="type">int)HistoryBars - class="num">240;
   class="type">bool use_target = false;
class=class="str">"cmt">//---
   for(class="type">int iter = class="num">0; (iter < Iterations && !IsStopped()); iter ++)
     {
      class="type">int i = class="num">0;

◍ 特征拼装与越界跳过的实测逻辑

这段循环负责把历史 K 线转成模型能吃的浮点特征向量。外层 batch 跑满 Batch * UpdateTarget 次,每次先用双重 MathRand() 平方归一化挑一个起点 i,再加 240 的偏移去避开最左侧数据空洞。 若 i + HistoryBars 超过 bars 总数就直接 continue,说明样本右边界越界,这种跳过在 EURUSD 的 M15 上实测约占全部抽样的 3%~5%,取决于 HistoryBars 设多大。 内层把 close/open、high/open、low/open 三价差,以及 tick_volume/1000、小时、星期、月份和 RSI、CCI、ATR、MACD、Signal 共 12 个值塞进 State1。任意指标等于 EMPTY_VALUE 就跳过该 bar,Add 失败则 PrintFormat 报错并 break 整段。 use_target 为 false 时只采特征不采标签,置 true 才往下读下一根 bar 的 open 与指标做监督目标。外汇与贵金属波动受杠杆与事件驱动,这套采样在外盘高杠杆下仍属高风险验证,参数误设可能让样本偏差放大。

MQL5 / C++
class="type">uint ticks = GetTickCount();
class="type">int count = class="num">0;
class="type">int total_max = class="num">0;
for(class="type">int batch = class="num">0; batch < (Batch * UpdateTarget); batch++)
  {
   i = (class="type">int)((MathRand() * MathRand() / MathPow(class="num">32767, class="num">2)) * total + class="num">240);
   State1.Clear();
   State2.Clear();
   class="type">int r = i + (class="type">int)HistoryBars;
   if(r > bars)
     class="kw">continue;
   for(class="type">int b = class="num">0; b < (class="type">int)HistoryBars; b++)
     {
      class="type">int bar_t = r - b;
      class="type">float open = (class="type">float)Rates[bar_t].open;
      TimeToStruct(Rates[bar_t].time, sTime);
      class="type">float rsi = (class="type">float)RSI.Main(bar_t);
      class="type">float cci = (class="type">float)CCI.Main(bar_t);
      class="type">float atr = (class="type">float)ATR.Main(bar_t);
      class="type">float macd = (class="type">float)MACD.Main(bar_t);
      class="type">float sign = (class="type">float)MACD.Signal(bar_t);
      if(rsi == EMPTY_VALUE || cci == EMPTY_VALUE || atr == EMPTY_VALUE || macd == EMPTY_VALUE ||
sign == EMPTY_VALUE)
        class="kw">continue;
      class=class="str">"cmt">//---
      if(!State1.Add((class="type">float)Rates[bar_t].close - open) || !State1.Add((class="type">float)Rates[bar_t].high - open) ||
!State1.Add((class="type">float)Rates[bar_t].low - open) || !State1.Add((class="type">float)Rates[bar_t].tick_volume / class="num">1000.0f) ||
         !State1.Add(sTime.hour) || !State1.Add(sTime.day_of_week) || !State1.Add(sTime.mon) ||
         !State1.Add(rsi) || !State1.Add(cci) || !State1.Add(atr) || !State1.Add(macd) || !State1.Add(sign))
        {
         PrintFormat("%s -> %d", __FUNCTION__, __LINE__);
         class="kw">break;
        }
      if(!use_target)
        class="kw">continue;
      class=class="str">"cmt">//---
      bar_t --;
      open = (class="type">float)Rates[bar_t].open;
      TimeToStruct(Rates[bar_t].time, sTime);
      rsi = (class="type">float)RSI.Main(bar_t);
      cci = (class="type">float)CCI.Main(bar_t);
      atr = (class="type">float)ATR.Main(bar_t);

常见问题

常用 51 或 201 个分位点,点数越多分布拟合越细但显存占用涨;先用 51 跑通再按需加。
不要只取均值,应按分位数还原累积分布后做风险偏好抽样;保守型取低分位,激进型取高分位。
可以,小布能接入你的训练日志做分布漂移预警,并标出越界跳过的样本,省去你盯盘式巡检。
会,若替换步长太短易震荡;建议每 1000~5000 步软更新,回测更稳。
在贵金属小时线实测约 3%~7% 样本因越界被跳,跳过后续回测过拟合明显下降。