神经网络变得轻松(第三十三部分):分布式 Q-学习中的分位数回归·进阶篇
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 验证维度匹配,避免实盘爆掉。
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 次以上看动作分布偏移。
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 归零周期对回测夏普的影响,再决定同步间隔。
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 扩容,外层循环控迭代、内层留作前馈后馈。
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 与指标做监督目标。外汇与贵金属波动受杠杆与事件驱动,这套采样在外盘高杠杆下仍属高风险验证,参数误设可能让样本偏差放大。
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);