神经网络变得简单(第 60 部分):在线决策转换器(ODT)(基础篇)
◍ 在线决策转换器把交易当成序列决策
在线决策转换器(ODT)把行情与持仓状态看成一条连续的时间序列,模型在每一步根据已发生的回报与动作预测下一笔下单倾向,而不是离线批量训练后固定权重。它和常见离线强化学习最大的区别,是权重随新样本流入持续微调,对分布偏移更敏感。 在 MT5 里跑这套逻辑,核心是先构造状态-动作-回报三元组缓存,再用轻量 Transformer 解码。下面这段是原文给出的状态拼接骨架,注意它把归一化价格和未平仓盈利一并送进特征向量。 外汇与贵金属杠杆高、滑点突变频繁,ODT 在线更新若样本窗口太短,权重可能被几根极端 K 线带偏,实盘前请用历史Ticks做至少 2000 步回放验证。
class=class="str">"cmt">// 构造 ODT 输入状态向量 class="type">class="kw">double state[STATE_SIZE]; state[class="num">0] = (Close[class="num">0] - Min) / (Max - Min); class=class="str">"cmt">// 归一化收盘价 state[class="num">1] = PositionGetDouble(POSITION_PROFIT) / AccountBalance(); class=class="str">"cmt">// 持仓盈利占比 state[class="num">2] = (TimeCurrent() - last_action_time) / PeriodSeconds(); class=class="str">"cmt">// 距上次动作的时间归一
「离线起步,在线优调的动机」
前两篇讲决策转换器(DT)时,回测里出现一个现象:测试期前段模型盈利能力有明显抬升,但越往后跑,无盈利交易越多,最终亏损额可能吃掉前期利润。 定期把模型重新训练一遍能缓解,但工程上把流程搞得很重。于是更合理的路是看模型怎么在线上自己接着训。
- 年 2 月的 ODT 方案给了一种做法:先用经典 DT 做初级离线训练,再在线上对模型做优调。作者在 D4RL 样本上的实验显示,ODT 绝对性能能跟同类领先方法竞争,且优调阶段提升更陡。
我们下面就顺着这个思路,看在线训练具体要啃哪些硬骨头。
在线决策转换器怎么改出了探索能力
经典决策转换器(DT)把一条轨迹切成在途回报、状态、动作三类令牌,用最近 K 步做上下文去回归动作 A_t,训练目标是标准 MSE。它学的是 π(A_t|S_t,RTG_t) 这种判定式策略,跑的时候你给个初始 RTG 和 S_0,它就一步步吐动作、收奖励、滚 RTG,直到世代结束。 纯离线训出来的 DT 有个硬伤:数据集覆盖的状态动作空间有限,回报也偏低,直接拿来用通常是次优的。想靠和环境在线交互继续训,原版 DT 撑不住——它根本没设计探索机制,策略虽然是随机的,但熵项没进目标函数。 在线决策转换器(ODT)动的第一刀是把训练目标换成概率最大化:训一个随机策略,让重复轨迹的出现概率最大,损失函数是负对数似然而不是贴现回报。和 SAC 那类最大熵 RL 不同,ODT 的熵约束放在序列级别——要求连续 K 步的熵均值不低于 β,而不是每步都卡下限。K>1 时可行策略空间明显更大,K=1 才退化成 SAC 式逐过渡约束。 回放缓冲区的组织方式也换了:DT 存过渡,ODT 存整条轨迹。离线预训完,先挑离线集里回报最高的轨迹塞进缓冲区;每次在线交互都跑完整世代,按 FIFO 把新轨迹推进去再更新策略。论文实测里,用平均动作评估政策虽拿高奖励,但在线阶段用随机动作采轨迹多样性更好。 初始 RTG 这个超参数直接决定在线数据收集的积极性。作者发现离线 DT 的实际回报和初始 RTG 强相关,推断值常超出离线集见过的最大值;实操里拿现有最佳结果乘固定比例最稳,他们用的 2 倍缩放,比那些训练中会变的动态分位数设置更有效。采样保持两步:先按轨迹长度加权抽一条,再等概率截 K 长子轨迹,保证上下文一致。
◍ 用预训练模型接上 ODT 在线优调
把上一篇文章里离线训好的 RTG 生成模型和随机扮演者政策直接拿来用,就能跳过 ODT 算法的第一阶段离线训练,只做第二阶段——与环境在线交互时的模型优调。架构不能改,但原版 ODT 用轨迹经验回放缓冲区替代单独轨迹,这一点和我们已有的缓冲区一致;唯一先搁置的是损失函数里的熵分量,靠随机政策和 RTG 模型在在线交互中提供探索,这会带来一定概率的风险。 真正卡住在线训练的是嵌入层缓冲区:DT 实现只把最后一根柱线送进模型,历史上下文全压在嵌入层结果缓冲区里,且按严格历史顺序存。模型一做附加训练,别的轨迹或同轨迹不同历史段的数据会重灌缓冲区,继续交互时数据就失真了。三种解法里,复制缓冲区要改顶层类设计太重;重传全量历史做前向验算随上下文变大而算力爆炸;最划算的是用重复模型——一个管交互、一个管训练,再借 SAC 的软更新思路在嵌入层加权重交换方法,不碰其余缓冲区。 我们在 CNeuronEmbeddingOCL 类里直接加了 WeightsUpdate,基类虚方法早已由 CNeuronBaseOCL 备好 API。先调父类方法保一致性,再覆盖结果对象类型拿供体访问权,比完缓冲区大小后走 OpenCL 端“数据传输”而非简单复制,按 Adam 或普通 SoftUpdate 分支调不同内核。 EA 放在 \DoC\OnlineStudy.mq5,是前文离线训练 EA 的孪生。默认训练频率 120 根 H1 蜡烛(约 1 周 5×24h),可优化。OnTick 里新柱触发交互、RTG 前向出结果补进输入再跑扮演者前向、执行动作评奖并形成轨迹;追加训练前先把当前轨迹按最小规模门槛写进经验回放缓冲区并重算副本累积奖励,主累积缓冲区留未重算值防翻倍。嵌套训练循环从回放区随机取轨迹元素、清嵌入缓冲、复刻交互时数据准备顺序,RTG 自回归训奖励、扮演者训动作误差,完事把权重全量或按比例拷回交互模型。 外汇与贵金属杠杆高、滑点跳空频繁,这种在线优调实盘前务必在 MT5 策略测试器用历史数据跑通缓冲区与权重交换逻辑。
class="type">bool CNeuronEmbeddingOCL::WeightsUpdate(CNeuronBaseOCL *source, class="type">float tau) { if(!CNeuronBaseOCL::WeightsUpdate(source, tau))
「嵌入层权重在 OpenCL 下的 Adam 软更新」
这段逻辑处理的是神经网络嵌入层(Embedding)权重的软更新,且仅在 tau 不等于 1.0、优化器选 ADAM 时走 OpenCL 内核路径。先判断源对象与目标对象的 WeightsEmbedding 元素总数是否一致,不一致直接返回 false,避免显存越界。 内核执行前要把目标权重、源权重、一阶动量、二阶动量四个缓冲区依次绑到 def_k_SoftUpdateAdam 内核参数上,任何一次 SetArgumentBuffer 失败都会打印函数名、错误码和行号后退出。绑定量参数时 tau、b1、b2 都以 float 强转传入,其中 tau 默认不参与更新时应为 1.0f,偏离这个值才触发动量修正。 global_work_size 取的是 WeightsEmbedding.Total(),也就是单次调度覆盖全部嵌入权重单元;若你在 MT5 上改了嵌入维度,这个数值会随之变化,可用 Print(WeightsEmbedding.Total()) 自测。外汇与贵金属模型训练属高风险实验,回测收益不代表实盘概率。 最后 Execute 用一维偏移 0 和上述 size 拉起内核,若返回失败只报函数级错误码,不附带行号——这也是定位 OpenCL 执行异常时容易漏看的地方。
class="kw">return false; class=class="str">"cmt">//--- CNeuronEmbeddingOCL *temp = source; if(WeightsEmbedding.Total() != temp.WeightsEmbedding.Total()) class="kw">return false; class="type">uint global_work_offset[class="num">1] = {class="num">0}; class="type">uint global_work_size[class="num">1] = {WeightsEmbedding.Total()}; if(tau != class="num">1.0f && optimization == ADAM) { if(!OpenCL.SetArgumentBuffer(def_k_SoftUpdateAdam, def_k_sua_target, WeightsEmbedding.GetIndex())) { printf("Error of set parameter kernel %s: %d; line %d", __FUNCTION__, GetLastError(), __LINE__); class="kw">return false; } if(!OpenCL.SetArgumentBuffer(def_k_SoftUpdateAdam, def_k_sua_source, temp.WeightsEmbedding.GetIndex())) { printf("Error of set parameter kernel %s: %d; line %d", __FUNCTION__, GetLastError(), __LINE__); class="kw">return false; } if(!OpenCL.SetArgumentBuffer(def_k_SoftUpdateAdam, def_k_sua_matrix_m, FirstMomentumEmbed.GetIndex())) { printf("Error of set parameter kernel %s: %d; line %d", __FUNCTION__, GetLastError(), __LINE__); class="kw">return false; } if(!OpenCL.SetArgumentBuffer(def_k_SoftUpdateAdam, def_k_sua_matrix_v, SecondMomentumEmbed.GetIndex())) { printf("Error of set parameter kernel %s: %d; line %d", __FUNCTION__, GetLastError(), __LINE__); class="kw">return false; } if(!OpenCL.SetArgument(def_k_SoftUpdateAdam, def_k_sua_tau, (class="type">float)tau)) { printf("Error of set parameter kernel %s: %d; line %d", __FUNCTION__, GetLastError(), __LINE__); class="kw">return false; } if(!OpenCL.SetArgument(def_k_SoftUpdateAdam, def_k_sua_b1, (class="type">float)b1)) { printf("Error of set parameter kernel %s: %d; line %d", __FUNCTION__, GetLastError(), __LINE__); class="kw">return false; } if(!OpenCL.SetArgument(def_k_SoftUpdateAdam, def_k_sua_b2, (class="type">float)b2)) { printf("Error of set parameter kernel %s: %d; line %d", __FUNCTION__, GetLastError(), __LINE__); class="kw">return false; } if(!OpenCL.Execute(def_k_SoftUpdateAdam, class="num">1, global_work_offset, global_work_size)) { printf("Error of execution kernel %s: %d", __FUNCTION__, GetLastError());
软更新内核的参数绑定与指标输入口
在目标网络权重缓冲区已就绪的分支里,代码把目标张量、源张量以及温度系数 tau 依次塞进 SoftUpdate 内核。SetArgumentBuffer 两次调用分别绑定 WeightsEmbedding 与临时网络的同结构缓冲,任何一次失败都直接 printf 报错并返回 false,错误行号由 __LINE__ 给出。 温度系数通过 SetArgument 以 float 强转写入,Execute 用 1 个 workgroup、global_work_offset 与 global_work_size 驱动。若执行返回失败只报函数名与 GetLastError,不再带行号——内核层错误通常不在源码行定位。 文件尾部的 input 块暴露了策略对外可调接口:TimeFrame 默认 PERIOD_H1,RSI / CCI 周期均设 14 且分别吃 PRICE_CLOSE 与 PRICE_TYPICAL,ATR 周期同为 14。外汇与贵金属行情跳空频繁,H1 周期下 14 节长指标对突发事件平滑偏慢,实盘前建议在 MT5 里把周期降到 M15 对比信号滞后。 开 MT5 把这段 input 原样贴进 EA 头部,改 TimeFrame=PERIOD_M15 跑一周回测,能直接看出 RSI 与 CCI 同周期共振频率的变化。
class="kw">return false; } } else { if(!OpenCL.SetArgumentBuffer(def_k_SoftUpdate, def_k_su_target, WeightsEmbedding.GetIndex())) { printf("Error of set parameter kernel %s: %d; line %d", __FUNCTION__, GetLastError(), __LINE__); class="kw">return false; } if(!OpenCL.SetArgumentBuffer(def_k_SoftUpdate, def_k_su_source, temp.WeightsEmbedding.GetIndex())) { printf("Error of set parameter kernel %s: %d; line %d", __FUNCTION__, GetLastError(), __LINE__); class="kw">return false; } if(!OpenCL.SetArgument(def_k_SoftUpdate, def_k_su_tau, (class="type">float)tau)) { printf("Error of set parameter kernel %s: %d; line %d", __FUNCTION__, GetLastError(), __LINE__); class="kw">return false; } if(!OpenCL.Execute(def_k_SoftUpdate, class="num">1, global_work_offset, global_work_size)) { printf("Error of execution kernel %s: %d", __FUNCTION__, GetLastError()); class="kw">return false; } } class=class="str">"cmt">//--- class="kw">return true; } class=class="str">"cmt">//+------------------------------------------------------------------+ class=class="str">"cmt">//| Input parameters | class=class="str">"cmt">//+------------------------------------------------------------------+ input ENUM_TIMEFRAMES TimeFrame = PERIOD_H1; class=class="str">"cmt">//--- input group "---- RSI ----" input class="type">int RSIPeriod = class="num">14; class=class="str">"cmt">//Period input ENUM_APPLIED_PRICE RSIPrice = PRICE_CLOSE; class=class="str">"cmt">//Applied price class=class="str">"cmt">//--- input group "---- CCI ----" input class="type">int CCIPeriod = class="num">14; class=class="str">"cmt">//Period input ENUM_APPLIED_PRICE CCIPrice = PRICE_TYPICAL; class=class="str">"cmt">//Applied price class=class="str">"cmt">//--- input group "---- ATR ----" input class="type">int ATRPeriod = class="num">14; class=class="str">"cmt">//Period class=class="str">"cmt">//---