神经网络变得轻松(第二十六部分):强化学习·进阶篇
「用交叉熵给代理者找路:有限状态里的试错收敛」
交叉熵方法本质就是结构化试错。它要求环境状态数和代理者行动数都有限,且训练区间有限,符合马尔可夫过程前提。代理者会在环境里跑多轮完整验算,每步状态对应一个行动——随机或按初始政策指派,轮数由架构师定为超参数。 每轮验算从头跑到尾,我们记下状态、行动、总奖励。之后挑总奖励前 20%~50% 的验算作为精英样本,用它们更新政策,再重复跑、再挑、再更新,直到收益不再增长或达到预期。外汇与贵金属市场高风险,这种收敛只是概率倾向,不是确定性路径。 落地到 MT5,状态有限性是个硬骨头。原文用 k-means 把市场形态聚成 500 类,相当于把连续行情压成有限状态空间,足够演示算法。代理者行动数明确(示例设 3),但状态必须靠聚类截断。 代码里外部参数直接暴露了算法骨架:Samples=100 每轮验算次数,Percentile=70 选前 70% 精英,Actions=3 行动空间,Clusters=500 状态聚类数。调这三个数就能改变探索密度与收敛速度。 别把正态当圣经 示例里奖励被硬编码为正确行动 +1、其余 -1,且故意排除行动对后续状态的影响,只为跑通技术。实盘每步都要回环境取真实状态,不可能预知 target 向量。
class="macro">#include "..\Unsupervised\K-means\kmeans.mqh" class="macro">#include <Trade\SymbolInfo.mqh> class="macro">#include <Indicators\Oscilators.mqh> class="kw">input class="type">int StudyPeriod = class="num">15; class=class="str">"cmt">//Study period, years class="kw">input class="type">uint HistoryBars = class="num">20; class=class="str">"cmt">//Depth of history class="kw">input class="type">int Clusters = class="num">500; class=class="str">"cmt">//Clusters ENUM_TIMEFRAMES TimeFrame = PERIOD_CURRENT; class=class="str">"cmt">//--- class="kw">input class="type">int Samples = class="num">100; class="kw">input class="type">int Percentile = class="num">70; class="type">int Actions = class="num">3; class=class="str">"cmt">//--- class="kw">input group "---- RSI ----" class="kw">input class="type">int RSIPeriod = class="num">14; class=class="str">"cmt">//Period class="kw">input ENUM_APPLIED_PRICE RSIPrice = PRICE_CLOSE; class=class="str">"cmt">//Applied price class=class="str">"cmt">//--- class="kw">input group "---- CCI ----"
◍ 指标句柄与K线对象的初始化链路
这段初始化代码把后续聚类要用到的四类技术指标和符号对象一次性挂到当前图表上。CCI 取 14 周期典型价、ATR 同样 14 周期、MACD 用 12/26/9 加收盘价——这些都是 MT5 标准类的默认惯性参数,改一个数就能换一种波动灵敏度。 OnInit 里先 new 出 CSymbolInfo 并刷新品种属性,再依次构建 CiRSI、CiCCI、CiATR、CiMACD 四个指标实例,每个都做了指针有效性和 Create 成功的双重校验,任一失败直接返回 INIT_FAILED,避免后面算聚类时读到空句柄。 最后 new 了 CKmeans 对象并投了一个 ChartCustom 事件(ID=1,文本"Init"),相当于给面板发个“就绪”信号。外汇与贵金属杠杆高,这类多指标共振策略在极端跳空时可能失效,实盘前务必在 MT5 策略测试器跑通初始化。
class="kw">input class="type">int CCIPeriod = class="num">14; class=class="str">"cmt">//Period class="kw">input ENUM_APPLIED_PRICE CCIPrice = PRICE_TYPICAL; class=class="str">"cmt">//Applied price class=class="str">"cmt">//--- class="kw">input group "---- ATR ----" class="kw">input class="type">int ATRPeriod = class="num">14; class=class="str">"cmt">//Period class=class="str">"cmt">//--- class="kw">input group "---- MACD ----" class="kw">input class="type">int FastPeriod = class="num">12; class=class="str">"cmt">//Fast class="kw">input class="type">int SlowPeriod = class="num">26; class=class="str">"cmt">//Slow class="kw">input class="type">int SignalPeriod = class="num">9; class=class="str">"cmt">//Signal class="kw">input ENUM_APPLIED_PRICE MACDPrice = PRICE_CLOSE; class=class="str">"cmt">//Applied price class="type">int OnInit() { class=class="str">"cmt">//--- Symb = new CSymbolInfo(); if(CheckPointer(Symb) == POINTER_INVALID || !Symb.Name(_Symbol)) class="kw">return INIT_FAILED; Symb.Refresh(); class=class="str">"cmt">//--- RSI = new CiRSI(); if(CheckPointer(RSI) == POINTER_INVALID || !RSI.Create(Symb.Name(), TimeFrame, RSIPeriod, RSIPrice)) class="kw">return INIT_FAILED; class=class="str">"cmt">//--- CCI = new CiCCI(); if(CheckPointer(CCI) == POINTER_INVALID || !CCI.Create(Symb.Name(), TimeFrame, CCIPeriod, CCIPrice)) class="kw">return INIT_FAILED; class=class="str">"cmt">//--- ATR = new CiATR(); if(CheckPointer(ATR) == POINTER_INVALID || !ATR.Create(Symb.Name(), TimeFrame, ATRPeriod)) class="kw">return INIT_FAILED; class=class="str">"cmt">//--- MACD = new CiMACD(); if(CheckPointer(MACD) == POINTER_INVALID || !MACD.Create(Symb.Name(), TimeFrame, FastPeriod, SlowPeriod, SignalPeriod, MACDPrice)) class="kw">return INIT_FAILED; class=class="str">"cmt">//--- Kmeans = new CKmeans(); if(CheckPointer(Kmeans) == POINTER_INVALID) class="kw">return INIT_FAILED; class=class="str">"cmt">//--- class="type">bool bEventStudy = EventChartCustom(ChartID(), class="num">1, class="num">0, class="num">0, "Init"); class=class="str">"cmt">//--- class="kw">return(INIT_SUCCEEDED); } class="type">void OnDeinit(class="kw">const class="type">int reason) { class=class="str">"cmt">//--- if(CheckPointer(Symb) != POINTER_INVALID)
训练前的指针清理与无监督加载
在 EA 析构阶段,必须逐个用 CheckPointer 判断指标与聚类对象是否仍占用内存,再执行 delete。Symb、RSI、CCI、ATR、MACD、Kmeans 任一未释放都会拖慢 MT5 终端,尤其在切换周期时可能触发内存泄漏告警。 Train 函数先通过 OpenCLCreate(cl_unsupervised) 拿到无监督计算上下文;若返回 POINTER_INVALID 直接 ExpertRemove 退出,避免空指针往下跑。Kmeans.SetOpenCL 失败同样清场走人,这一步决定了后续聚类能否走 GPU 加速。 回测窗口由 StudyPeriod 控制:把当前时间年份减掉 StudyPeriod,若跌到 0 以下则兜底设为 1900,再 StructToTime 转成起始 datetime。CopyRates 按此区间抓取 Symb.Name() 与 TimeFrame 的报价,返回 bars 总数,紧接着给四个指标 BufferResize(bars),任一失败即终止 EA。 模型文件按聚类数命名,如 kmeans_5.net;用 FILE_READ|FILE_BIN 打开后先比对 FileReadInteger(handl) 与 Kmeans.Type(),类型不符立即退出。读入后 total = bars - HistoryBars - 480,数据矩阵按 total*8*HistoryBars 与 total*3 两档扩容,外汇与贵金属品种在此类无监督训练上波动剧烈,实盘前请在模拟盘验证内存占用。
class="kw">delete Symb; class=class="str">"cmt">//--- if(CheckPointer(RSI) != POINTER_INVALID) class="kw">delete RSI; class=class="str">"cmt">//--- if(CheckPointer(CCI) != POINTER_INVALID) class="kw">delete CCI; class=class="str">"cmt">//--- if(CheckPointer(ATR) != POINTER_INVALID) class="kw">delete ATR; class=class="str">"cmt">//--- if(CheckPointer(MACD) != POINTER_INVALID) class="kw">delete MACD; class=class="str">"cmt">//--- if(CheckPointer(Kmeans) != POINTER_INVALID) class="kw">delete Kmeans; class=class="str">"cmt">//--- } class="type">void Train(class="type">void) { COpenCLMy *opencl = OpenCLCreate(cl_unsupervised); if(CheckPointer(opencl) == POINTER_INVALID) { ExpertRemove(); class="kw">return; } if(!Kmeans.SetOpenCL(opencl)) { class="kw">delete opencl; ExpertRemove(); class="kw">return; } 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=class="str">"cmt">//--- 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)) { ExpertRemove(); class="kw">return; } if(!ArraySetAsSeries(Rates, true)) { ExpertRemove(); class="kw">return; } class=class="str">"cmt">//--- RSI.Refresh(); CCI.Refresh(); ATR.Refresh(); MACD.Refresh(); class="type">int handl = FileOpen(StringFormat("kmeans_%d.net", Clusters), FILE_READ | FILE_BIN); if(handl == INVALID_HANDLE) { ExpertRemove(); class="kw">return; } if(FileReadInteger(handl) != Kmeans.Type()) { ExpertRemove(); class="kw">return; } class="type">bool result = Kmeans.Load(handl); FileClose(handl); if(!result) { ExpertRemove(); class="kw">return; } class="type">int total = bars - (class="type">int)HistoryBars - class="num">480; class="type">class="kw">double data[], fractals[]; if(ArrayResize(data, total * class="num">8 * HistoryBars) <= class="num">0 || ArrayResize(fractals, total * class="num">3) <= class="num">0) { ExpertRemove(); class="kw">return; }
「把K线特征塞进强化学习状态矩阵」
这段逻辑干的事很直白:把每个样本窗口的八维特征(含开高低收偏离、RSI、CCI、ATR、MACD双线)平铺进一维 data 数组,每个样本占 HistoryBars×8 个槽位,起始偏移额外加 480 根作为预热。 外层循环用 i 遍历 total 个样本,每次先 Comment 打印进度「Create data: i of total」,再嵌套 b 循环搬 HistoryBars 根 bar 的数据;bar 索引 = i+b+480,shift 按 i*HistoryBars+b 乘 8 定位。注意第 7 维写的是 MACD 信号线,不是主线重复。 分形标记单独存 fractals 数组,每个样本占 3 槽:上分形看左右两根 high 是否更低、下分形看 low 是否更高,第三槽用 (fractals[shift]+fractals[shift])==0 判定「无分形」——这里原代码疑似笔误,本该是 up+down 相加。 IsStopped 触发就 ExpertRemove 并 return,避免终端关闭时还硬写。之后把 data、fractals 包成 CBufferFloat 交给 Kmeans.SoftMax 做归一。 环境向量 env 按簇数切分取最大值偏移,目标向量 target 从 fractals 同法提取;policy 矩阵先用 1/Actions 均匀填充,供后续 Q 学习更新。外汇与贵金属行情跳空频繁,480 根预热在跳空周可能引入偏移,建议开 MT5 用 EURUSD 的 M15 实测一次再调 HistoryBars。
for(class="type">int i = class="num">0; (i < total && !IsStopped()); i++) { Comment(StringFormat("Create data: %d of %d", i, total)); for(class="type">int b = class="num">0; b < (class="type">int)HistoryBars; b++) { class="type">int bar = i + b + class="num">480; class="type">int shift = (i * (class="type">int)HistoryBars + b) * class="num">8; class="type">class="kw">double open = Rates[bar] .open; data[shift] = open - Rates[bar].low; data[shift + class="num">1] = Rates[bar].high - open; data[shift + class="num">2] = Rates[bar].close - open; data[shift + class="num">3] = RSI.GetData(MAIN_LINE, bar); data[shift + class="num">4] = CCI.GetData(MAIN_LINE, bar); data[shift + class="num">5] = ATR.GetData(MAIN_LINE, bar); data[shift + class="num">6] = MACD.GetData(MAIN_LINE, bar); data[shift + class="num">7] = MACD.GetData(SIGNAL_LINE, bar); } class="type">int shift = i * class="num">3; class="type">int bar = i + class="num">480; fractals[shift] = (class="type">int)(Rates[bar - class="num">1].high <= Rates[bar].high && Rates[bar + class="num">1].high < Rates[bar].high); fractals[shift + class="num">1] = (class="type">int)(Rates[bar - class="num">1].low >= Rates[bar].low && Rates[bar + class="num">1].low > Rates[bar].low); fractals[shift + class="num">2] = (class="type">int)((fractals[shift] + fractals[shift]) == class="num">0); } if(IsStopped()) { ExpertRemove(); class="kw">return; } CBufferFloat *Data = new CBufferFloat(); if(CheckPointer(Data) == POINTER_INVALID || !Data.AssignArray(data)) class="kw">return; CBufferFloat *Fractals = new CBufferFloat(); if(CheckPointer(Fractals) == POINTER_INVALID || !Fractals.AssignArray(fractals)) class="kw">return; class=class="str">"cmt">//--- ResetLastError(); Data = Kmeans.SoftMax(Data); vector env = vector::Zeros(Data.Total() / Clusters); vector target = vector::Zeros(env.Size()); matrix states = matrix::Zeros(Samples, env.Size()); matrix actions = matrix::Zeros(Samples, env.Size()); vector CumRewards = vector::Zeros(Samples); class="type">class="kw">double average = class="num">1.0 / Actions; matrix policy = matrix::Full(Clusters, Actions, average); for(class="type">class="kw">ulong state = class="num">0; state < env.Size(); state++) { class="type">class="kw">ulong shift = state * Clusters; env[state] = (class="type">class="kw">double)(Data.Maximum((class="type">int)shift, Clusters) - shift); shift = state * Actions; target[state] = Fractals.Maximum((class="type">int)shift, Actions) - shift; }