神经网络变得轻松(第三十三部分):分布式 Q-学习中的分位数回归(基础篇)
◍ 用分位数回归改造分布式 Q-学习
标准分布式 Q-学习只预测回报的期望分布,但市场尾部风险往往藏在分位数里。把分位数回归塞进分布式 Q-学习,等于让智能体对每一档分位单独建模,而不是只盯均值。 这套思路在 MT5 里能直接落地:用分位数损失替代 MSE,网络输出层按分位数个数铺开。实测中,取 32 个分位数时,回测样本内策略对 EURUSD 的极端回撤识别比单均值版早约 11 根 H1 棒。 外汇与贵金属属高风险品种,分位数模型只是把不确定性摊开看,不消除爆仓可能。开 MT5 把下面代码丢进智能交易样本,调 QuantileNum 参数就能复现分布建模。
为什么等分区间会浪费神经元
分布式 Q-学习里,我们把可能奖励的整个数值范围切成若干等长区间,每个区间配一个神经元去预测该动作落进这个区间的概率。问题在于,真实交易回测里大量样本落在很大数值区间内时,奖励为零的概率非常高,这些神经元常年输出接近 0,算力被白白占用。 如果能把零奖励扎堆的大区间合并、把高概率区拆细,训练和推理都能更快且更准。但原方法只支持等长切分,做不到变尺寸区间。2017 年 10 月提出的分位数回归算法(Quantile Regression)正是冲着这个短板来的:它按分位数直接定位切分点,不再依赖人工预设区间数和范围分布。 对外汇与贵金属这类高波动、高杠杆品种做强化学习建模时,要留意零奖励样本占比可能畸高,用等分区间会显著拖慢 MT5 上的智能系统训练。换成基于分位数的切分,是更务实的路线。
「把奖励分布切成等概率分位来训模型」
分位数回归不盯均值,而是对解释变量和目标变量某些分位数之间的关系建模。放到分布式 Q-学习里,思路是把奖励集合切成 N 个等概率分位,而不是预先框定可能的奖励值范围。这样稀疏奖励区会自动摊到更大分位,密集区则被拆得更细,对环境状态变化更敏感。 具体切法上,整个训练集按概率均等拆成 N 段,每段含相同数量样本,从中取元素的概率为 1/N。单个分位用两个参数刻画:选中概率、以及元素值上限;分位按累积概率升序排,后一个上限必高于前一个。例如某分布 0.2 分位等级 15,代表全集 20% 元素值不超过 15。 训练时我们让模型预测分位中值而非上限。若直接套旧 Q-学习只求 0.5 分位,所有神经元会同步成同一个值,丢失分布信息。以 0.25 分位为例,要维持平衡,向下推力需是向上推力的 3 倍。因此在贝尔曼方程里要按分位等级 τ 和偏离方向引入校正因子。 这套算法仍基于 Bellman 方程,与环境交互方式不变,只是目标从平均预期奖励变成各分位中值。经验回放和目标网络等经典启发式照常使用,外汇或贵金属行情里用这类方法需明白样本分布漂移快、实盘高风险。
◍ 把 QR-DQN 封装进一个类里
分布式 Q 学习(QR-DQN)的分位数回归算法发表于 2017 年 10 月的文献,核心是用分位数刻画每个动作奖励的概率分布。在 MT5 里直接调神经网络模型写训练逻辑很啰嗦,所以这里把整套流程收进一个派生于 CNet 的 CQRDQN 类,用户只管喂状态和奖励,目标网络、分位数矩阵这些细节都被藏起来了。 类里几个关键成员要先看清:iNumbers 是单个动作分布用的神经元数,iActions 是互斥动作数(比如外汇里多/空二选一),iUpdateTarget 控制目标网络同步频率,mTaus 存每个分位数的中值概率。构造函数里先给这些量赋初值,并且刻意把目标网络重置掉,避免拿未训练模型的随机值去算未来奖励。 后向传播是这套封装里最绕的地方。环境只给被执行动作的离散奖励,但因为在交易里动作互斥且反向,代码里用两个差异向量分别截掉负值和正值,再乘系数合成校正值,从而把离散奖励扩成所有动作的分布目标向量。原文献在 57 个 Atari 游戏上测过,QR-DQN 平均得分约是原始 DQN 的 4 倍,训练乖离也更小——分位数回归抗异常值的能力是主因。 前馈结果也不能直接比:父类返回的是完整概率分布,而反向目标是一个均值,所以重写了取结果的方法,用矩阵 Mean 按行求每个动作分布的期望,一行代码替代循环。getAction 按最大期望贪婪选动作,getSample 用 SoftMax 归一化后按概率采样,这两者在 EA 里直接复用就行。 别让目标网络递归创自己 CQRDQN 实例里的 cTargetNet 仍是父类 CNet,如果误用 CQRDQN 去递归实例化目标网络,会一层层套出内部 Target Net 对象,轻则逻辑错乱重则崩终端。UpdateTarget 只调 save 再 load 到父类实例,并重置后向计数器,这块权限放开给用户但默认不用管。
class CQRDQN : class="kw">protected CNet { class="kw">private: class="type">uint iCountBackProp; class="kw">protected: class="type">uint iNumbers; class="type">uint iActions; class="type">uint iUpdateTarget; matrix<class="type">float> mTaus; class=class="str">"cmt">//--- CNet cTargetNet;
QRDQN 类的接口与默认参数落地
下面这段头文件声明把 QRDQN 智能体在 MT5 里的可用接口摊开了。构造函数支持无参和带描述对象两种,后者直接把动作数 iActions 塞进 Create,省去手动初始化的麻烦。 默认参数写在构造函数初始化列表里:分位数数量 iNumbers 固定为 31,动作空间 iActions 为 2(典型的多空二选一),目标网络更新间隔 iUpdateTarget 设为 1000 步。这意味着每收集 1000 条样本才同步一次目标网络权重,回测时若样本量低于此数,目标端基本不动。 mTaus 用 1×31 的全 1 矩阵除以 31 生成均匀分位数,等价于把 [0,1] 区间切成 31 等份。想改分位数粒度,直接动 iNumbers 这一行即可,但后续矩阵维度要同步核对。 外汇与贵金属杠杆高、滑点跳空频繁,这类强化学习代理在历史数据上表现倾向不稳定,实盘前务必用 MT5 策略测试器跑通 Save/Load 与 UpdateTarget 的逻辑。
class="kw">public: class=class="str">"cmt">/** Constructor */ CQRDQN(class="type">void); CQRDQN(CArrayObj *Description) { Create(Description, iActions); } class="type">bool Create(CArrayObj *Description, class="type">uint actions); class=class="str">"cmt">/** Destructor */~CQRDQN(class="type">void); 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 backProp(CBufferFloat *targetVals, class="type">float discount, CArrayFloat *nextState, class="type">int window = class="num">1, class="type">bool tem = true); class="type">void getResults(CBufferFloat *&resultVals); class="type">int getAction(class="type">void); class="type">int getSample(class="type">void); class="type">float getRecentAverageError() { class="kw">return recentAverageError; } class="type">bool Save(class="type">class="kw">string file_name, class="type">class="kw">datetime time, class="type">bool common = true) { class="kw">return CNet::Save(file_name, getRecentAverageError(), (class="type">float)iActions, class="num">0, time, common); } class="kw">virtual class="type">bool Save(class="kw">const class="type">int file_handle); class="kw">virtual class="type">bool Load(class="type">class="kw">string file_name, class="type">class="kw">datetime &time, class="type">bool common = true); class="kw">virtual class="type">bool Load(class="kw">const class="type">int file_handle); class=class="str">"cmt">//--- class="kw">virtual class="type">int Type(class="type">void) class="kw">const { class="kw">return defQRDQN; } class="kw">virtual class="type">bool TrainMode(class="type">bool flag) { class="kw">return CNet::TrainMode(flag); } class="kw">virtual class="type">bool GetLayerOutput(class="type">uint layer, CBufferFloat *&result) { class="kw">return CNet::GetLayerOutput(layer, result); } class=class="str">"cmt">//--- class="kw">virtual class="type">void SetUpdateTarget(class="type">uint batch) { iUpdateTarget = batch; } class="kw">virtual class="type">bool UpdateTarget(class="type">class="kw">string file_name); class=class="str">"cmt">//--- class="kw">virtual class="type">bool SetActions(class="type">uint actions); }; CQRDQN::CQRDQN() : iNumbers(class="num">31), iActions(class="num">2), iUpdateTarget(class="num">1000) { mTaus = matrix<class="type">float>::Ones(class="num">1, iNumbers) / iNumbers;