神经网络变得轻松(第二十五部分):实践迁移学习·进阶篇
(2/3)· 同样的问题、平等的初始条件,借来的编码器到底能省多少训练成本?
◍ 初始化与训练触发的底层约束
这段逻辑集中在 EA 的 OnInit、OnDeinit 与自定义图表事件里,先把神经网络和四类指标句柄一次性建好,再决定何时拉起训练。 网络层输出若取不到,直接返 INIT_FAILED;TempData 总量必须能被 12 整除,HistoryBars 由此算出,而 getResults 后 Total() 不等于 3 会判为参数错误——这三个断言少一个,EA 都起不来。 RSI、CCI、ATR、MACD 四个 Create 调用全失败即退,说明指标句柄是后续推理的硬依赖,少一个品种周期组合都会让初始化崩。 EventChartCustom 把训练起始时间塞进图表事件:用 100*recentAverageSmoothingFactor 再乘 1 或 10(dForecast>=70 取 1,否则 10)做偏移,这种写法让高预测值反而更早触发回测。 OnChartEvent 只认 id==1001,收到就调 Train(lparam);Train 里把当前时间逐年回退 StudyPeriod 年,year<=0 则兜底 1900,dtStudied 取传入值与回退时间的最大值,训练窗口由此锁定。 别把初始化当成走过场:在 MT5 里把 StudyPeriod 改成 5 和 20,看 OnInit 后 HistoryBars 与 dtStudied 的差值,能直接验证回测样本长度是否被悄悄腰斩。外汇与贵金属杠杆高,这类自动训练 EA 若参数错配,可能在不经意间用极小样本过拟合。
class="kw">return INIT_PARAMETERS_INCORRECT; } if(!Net.GetLayerOutput(class="num">0, TempData)) class="kw">return INIT_FAILED; HistoryBars = TempData.Total() / class="num">12; Net.getResults(TempData); if(TempData.Total() != class="num">3) class="kw">return INIT_PARAMETERS_INCORRECT; if(!Symb.Name(_Symbol)) class="kw">return INIT_FAILED; Symb.Refresh(); if(!RSI.Create(Symb.Name(), TimeFrame, RSIPeriod, RSIPrice)) class="kw">return INIT_FAILED; if(!CCI.Create(Symb.Name(), TimeFrame, CCIPeriod, CCIPrice)) class="kw">return INIT_FAILED; if(!ATR.Create(Symb.Name(), TimeFrame, ATRPeriod)) class="kw">return INIT_FAILED; if(!MACD.Create(Symb.Name(), TimeFrame, FastPeriod, SlowPeriod, SignalPeriod, MACDPrice)) class="kw">return INIT_FAILED; bEventStudy = EventChartCustom(ChartID(), class="num">1, (class="type">long)MathMax(class="num">0, MathMin(iTime(Symb.Name(), PERIOD_CURRENT, (class="type">int)(class="num">100 * Net.recentAverageSmoothingFactor * (dForecast >= class="num">70 ? class="num">1 : class="num">10))), dtStudied)), class="num">0, "Init"); class=class="str">"cmt">//--- class="kw">return(INIT_SUCCEEDED); } class=class="str">"cmt">//+------------------------------------------------------------------+ class=class="str">"cmt">//| Expert deinitialization function | class=class="str">"cmt">//+------------------------------------------------------------------+ class="type">void OnDeinit(class="kw">const class="type">int reason) { class=class="str">"cmt">//--- if(CheckPointer(TempData) != POINTER_INVALID) class="kw">delete TempData; } class=class="str">"cmt">//+------------------------------------------------------------------+ class=class="str">"cmt">//| ChartEvent function | class=class="str">"cmt">//+------------------------------------------------------------------+ class="type">void OnChartEvent(class="kw">const class="type">int id, class="kw">const class="type">long &lparam, class="kw">const class="type">class="kw">double &dparam, class="kw">const class="type">class="kw">string &sparam) { class=class="str">"cmt">//--- if(id == class="num">1001) Train(lparam); } class="type">void Train(class="type">class="kw">datetime StartTrainBar = class="num">0) { class="type">int count = class="num">0; 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); dtStudied = MathMax(StartTrainBar, st_time); class="type">class="kw">ulong last_tick = class="num">0; class="type">class="kw">double prev_er = DBL_MAX;
「多指标特征拼装与回测缓冲清洗」
这段逻辑干的事很直接:把 RSI、CCI、ATR、MACD 四个指标连同 OHLC 与时间结构塞进一个临时特征容器,供后续模型或统计使用。注意它先对四个指标缓冲区做 BufferResize,只要有一个失败就 ExpertRemove 退出,避免拿空数组去算。 复制行情用 CopyRates 从 st_time 拉到 TimeCurrent,ArraySetAsSeries 设成时间倒序,这样 Rates[0] 是最新 bar。四个指标都 Refresh(OBJ_ALL_PERIODS) 强制刷新全周期缓存,否则可能读到上一 tick 的滞后值。 真正拼特征时,total 被压到 bars - max(HistoryBars,0) - 300,等于人为留了 300 根 bar 的边界余量,防止索引越界。每根采样窗取 HistoryBars 根历史,把 close-open、high-open、low-open、tick_volume/1000、小时、星期、月份、以及 4 个指标值共 12 维塞进 TempData。 若 TempData.Total() 小于 HistoryBars*12,说明中间有 EMPTY_VALUE 被跳过、特征没凑齐,直接 continue 放弃该样本。外汇与贵金属波动受消息面干扰大,这类特征缺失样本若强行保留,模型倾向学到噪声,实盘高风险。 让小布替你跑这套 把 HistoryBars 从默认改到 50,观察 TempData 丢弃率:若超过 20% 样本因 EMPTY_VALUE 被 continue,说明指标周期设得比可用历史还长,应先降指标周期而非硬加 bars。
class="type">class="kw">datetime bar_time = class="num">0; class="type">bool stop = IsStopped(); 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; } RSI.Refresh(OBJ_ALL_PERIODS); CCI.Refresh(OBJ_ALL_PERIODS); ATR.Refresh(OBJ_ALL_PERIODS); MACD.Refresh(OBJ_ALL_PERIODS); class="type">MqlDateTime sTime; class="type">int total = (class="type">int)(bars - MathMax(HistoryBars, class="num">0) - class="num">300); do { prev_er = dError; stop = IsStopped(); for(class="type">int it = total; it > class="num">1 && !stop; t--) { TempData.Clear(); class="type">int i = it + class="num">299; 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(!TempData.Add((class="type">float)Rates[bar_t].close - open) || !TempData.Add((class="type">float)Rates[bar_t].high - open) || !TempData.Add((class="type">float)Rates[bar_t].low - open) || !TempData.Add((class="type">float)Rates[bar_t].tick_volume / class="num">1000.0f) || !TempData.Add(sTime.hour) || !TempData.Add(sTime.day_of_week) || !TempData.Add(sTime.mon) || !TempData.Add(rsi) || !TempData.Add(cci) || !TempData.Add(atr) || !TempData.Add(macd) || !TempData.Add(sign)) class="kw">break; } if(TempData.Total() < (class="type">int)HistoryBars * class="num">12) class="kw">continue;
神经网络输出的概率归一与信号判定
这段逻辑紧接前一步的训练数据准备,把网络前向推理的结果做 softmax 归一,再按最大激活节点映射成交易信号。外汇与贵金属市场高杠杆、高波动,以下仅作技术验证参考,信号倾向不代表确定性方向。 先跑前向传播并取回输出:Net.feedForward(TempData, 12, true) 用 12 个隐含节点推理,Net.getResults(TempData) 把结果写回 TempData。随后对 3 个输出节点做指数化并累加,再用除法把三项压成和为 1 的概率分布——这是标准 softmax,避免输出值量级失衡。 归一后通过 TempData.Maximum(0,3) 找最强节点:返回 1 表示偏多,返回 2 表示偏空,其余归为未定义(Undefine)。case 1 里若节点1与节点2不相等则取节点1值作为 dPrevSignal,case 2 直接取负节点2值,default 置 0。 每 250 毫秒(GetTickCount64 差值 >= 250)用 Comment 把代次、误差、未定义概率、预测概率及买/卖/未定义三项数值打印到图表,方便肉眼盯训练过程。循环末尾清 TempData,按相邻三根 K 线高低点构造新标签:中间柱最高且两侧更低判 sell,最低且两侧更高判 buy,否则未定义,再喂给 Net.backProp 做反向传播。
Net.feedForward(TempData, class="num">12, true); Net.getResults(TempData); class="type">float sum = class="num">0; for(class="type">int res = class="num">0; res < class="num">3; res++) { class="type">float temp = exp(TempData.At(res)); sum += temp; TempData.Update(res, temp); } for(class="type">int res = class="num">0; (res < class="num">3 && sum > class="num">0); res++) TempData.Update(res, TempData.At(res) / sum); class=class="str">"cmt">//--- class="kw">switch(TempData.Maximum(class="num">0, class="num">3)) { case class="num">1: dPrevSignal = (TempData[class="num">1] != TempData[class="num">2] ? TempData[class="num">1] : class="num">0); class="kw">break; case class="num">2: dPrevSignal = -TempData[class="num">2]; class="kw">break; class="kw">default: dPrevSignal = class="num">0; class="kw">break; } if((GetTickCount64() - last_tick) >= class="num">250) { class="type">class="kw">string s = StringFormat("Study -> Era %d -> %.2f -> Undefine %.2f%% foracast %.2f%%\n %d of %d -> %.2f%% \nError %.2f\n%s -> %.2f ->> Buy %.5f - Sell %.5f - Undef %.5f", count, dError, dUndefine, dForecast, total - it - class="num">1, total, (class="type">class="kw">double)(total - it - class="num">1.0) / (total) * class="num">100, Net.getRecentAverageError(), EnumToString(DoubleToSignal(dPrevSignal)), dPrevSignal, TempData[class="num">1], TempData[class="num">2], TempData[class="num">0]); Comment(s); last_tick = GetTickCount64(); } stop = IsStopped(); if(!stop) { TempData.Clear(); class="type">bool sell = (Rates[i - class="num">1].high <= Rates[i].high && Rates[i + class="num">1].high < Rates[i].high); class="type">bool buy = (Rates[i - class="num">1].low >= Rates[i].low && Rates[i + class="num">1].low > Rates[i].low); TempData.Add(!(buy || sell)); TempData.Add(buy); TempData.Add(sell); Net.backProp(TempData); ENUM_SIGNAL signal = DoubleToSignal(dPrevSignal); if(signal != Undefine) {
◍ 信号打分与历史回放的数据装配
这段逻辑在做两件事:一是根据实时信号对看涨/看跌倾向值 dForecast 做指数式平滑修正,二是把过去 300 根候选 bar 逐根重构成神经网络可读取的特征向量。 当 signal 与 buy/sell 标记吻合时,dForecast 按 (100 - dForecast) / Net.recentAverageSmoothingFactor 向上逼近;否则按对称方式衰减。dUndefine 则在无明确信号时反向累积,用来衡量“看不清”的概率权重。外汇与贵金属波动受消息扰动大,这类打分只反映历史样本的统计倾向,实盘须警惕滑点与跳空风险。 回放循环里,i 从 0 到 299 遍历,r = i + HistoryBars 定位终点 bar,超出总 bars 就跳过。每根回溯 bar 抽取 open 差值、高低点偏移、tick_volume/1000、以及小时/星期/月份等时间结构,再拼上 RSI、CCI、ATR、MACD 主线与信号线,共 12 维特征。 只要任一指标返回 EMPTY_VALUE,或 TempData 总数不足 HistoryBars * 12,该样本直接丢弃;达标后才送 Net.feedForward(..., 12, true) 做前向推理。打开 MT5 把 HistoryBars 调到 20 以上,能直观看到特征矩阵密度对推理耗时的拖累。
if((signal == Sell && sell) || (signal == Buy && buy)) dForecast += (class="num">100 - dForecast) / Net.recentAverageSmoothingFactor; else dForecast -= dForecast / Net.recentAverageSmoothingFactor; dUndefine -= dUndefine / Net.recentAverageSmoothingFactor; } else { if(!(buy || sell)) dUndefine += (class="num">100 - dUndefine) / Net.recentAverageSmoothingFactor; } } count++; for(class="type">int i = class="num">0; i < class="num">300; i++) { TempData.Clear(); class="type">int r = i + (class="type">int)HistoryBars; if(r > bars) class="kw">continue; class=class="str">"cmt">//--- 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(!TempData.Add((class="type">float)Rates[bar_t].close - open) || !TempData.Add((class="type">float)Rates[bar_t].high - open) || !TempData.Add((class="type">float)Rates[bar_t].low - open) || !TempData.Add((class="type">float)Rates[bar_t].tick_volume / class="num">1000.0f) || !TempData.Add(sTime.hour) || !TempData.Add(sTime.day_of_week) || !TempData.Add(sTime.mon) || !TempData.Add(rsi) || !TempData.Add(cci) || !TempData.Add(atr) || !TempData.Add(macd) || !TempData.Add(sign)) class="kw">break; } if(TempData.Total() < (class="type">int)HistoryBars * class="num">12) class="kw">continue; Net.feedForward(TempData, class="num">12, true);
「softmax 归一与信号落图的收尾实现」
这段逻辑紧跟前文的网络推理,先把 Net.getResults 拿到的原始输出做 softmax 处理。循环里对 3 个输出节点取 exp 并累加,再用总和归一,保证三个概率之和为 1,后续比较才有意义。 归一后通过 TempData.Maximum(0,3) 找最大索引:索引 1 代表偏多信号,索引 2 代表偏空信号,两者概率不等时才给 dPrevSignal 赋值,相等则视为无信号。外汇与贵金属波动剧烈,这种概率判定只代表模型倾向,不保证方向。 若 DoubleToSignal 解析为 Undefine 就删掉对应 K 线时间的对象,否则在 high/low 区间画出信号标记。训练没被中断时,把平均误差和预测值写进 .nnw 与 .csv,printf 打出类似 "Era 120 -> error 0.83 % forecast 0.45" 的日志,方便你直接开 MT5 看回测轨迹。 最外层的 do-while 以 dError<0.01 且误差改善不足 0.01 为停止条件,命中后清 Comment 并 ExpertRemove 结束 EA。你可以把 0.01 这个阈值调大来加速收敛,但预测精度会相应下降。
Net.getResults(TempData); class=class="str">"cmt">//--- class="type">float sum = class="num">0; for(class="type">int res = class="num">0; res < class="num">3; res++) { class="type">float temp = exp(TempData.At(res)); sum += temp; TempData.Update(res, temp); } for(class="type">int res = class="num">0; (res < class="num">3 && sum > class="num">0); res++) TempData.Update(res, TempData.At(res) / sum); class=class="str">"cmt">//--- class="kw">switch(TempData.Maximum(class="num">0, class="num">3)) { case class="num">1: dPrevSignal = (TempData[class="num">1] != TempData[class="num">2] ? TempData[class="num">1] : class="num">0); class="kw">break; case class="num">2: dPrevSignal = (TempData[class="num">1] != TempData[class="num">2] ? -TempData[class="num">2] : class="num">0); class="kw">break; class="kw">default: dPrevSignal = class="num">0; class="kw">break; } if(DoubleToSignal(dPrevSignal) == Undefine) DeleteObject(Rates[i].time); else DrawObject(Rates[i].time, dPrevSignal, Rates[i].high, Rates[i].low); } if(!stop) { dError = Net.getRecentAverageError(); Net.Save(FileName + ".nnw", dError, dUndefine, dForecast, Rates[class="num">0].time, false); printf("Era %d -> error %.2f %% forecast %.2f", count, dError, dForecast); class="type">int h = FileOpen(FileName + ".csv", FILE_READ | FILE_WRITE | FILE_CSV); if(h != INVALID_HANDLE) { FileSeek(h, class="num">0, SEEK_END); FileWrite(h, eta, count, dError, dUndefine, dForecast); FileFlush(h); FileClose(h); } } } class="kw">while(!(dError < class="num">0.01 && (prev_er - dError) < class="num">0.01) && !stop); Comment(""); ExpertRemove(); }