神经网络变得轻松(第五十二部分):研究乐观情绪和分布校正·进阶篇
(2/3)· 当经验回放越攒越多,扮演者动作却离缓冲区越来越远,这篇讲怎么用分布校正和乐观情绪模型兜住训练效率
经验回放缓冲区塞得越满,模型见过的环境样本越杂,但扮演者更新越多,它当下的动作分布和缓冲区里老样本的距离就越拉越大。多数人在软性扮演者-评论者里只盯着评论者的悲观下限,没察觉动作同质化已经把探索空间压扁了。
SAC-DICE 初始化与训练入口的容错逻辑
在 MT5 里跑强化学习代理,第一道关卡是 OpenCL 上下文。若 opencl 对象没成功建出 context,后面所有神经网络前向传播都会失效,代码直接打印 "Don't opened OpenCL context" 并返回 false,这一步卡不住就会在显卡上静默算错。
Critic 与函数网络(zeta、nu)的创建也必须成对成功。任意一张网 Create 返回 false,就用 PrintFormat 带 GetLastError() 把错误码甩出来——实盘前先在策略测试器看日志,错误码非 0 即代表显存或结构参数有问题。
目标网络(TargetCritic1/2、TargetNu)建好后,用 WeightsUpdate(GetPointer(...), 1.0) 把在线网络权重整份拷过去,系数 1.0 表示硬拷贝而非滑动平均;随后 fLambda 初始化为 1e-5f,这是 DICE 框架里的正则项起点,调大可能让策略偏保守。
Study 方法开头就做维度校验:Actions.Total() 必须等于 ActionsLogProbab.Size(),否则直接 false 退出。之后对 NextState 做 feedForward 并喂给两个 TargetCritic 跑前向,任一失败则返回 false——这意味着你的状态缓冲区和网络层索引 iLatentLayer 必须对齐,否则回测时每根 K 线都白算。
Print("Don&class="macro">#x27;t opened OpenCL context"); class="kw">return false; } if(!cCritic1.Create(critic) || !cCritic2.Create(critic)) { PrintFormat("Error of create Critic: %d", GetLastError()); class="kw">return false; } if(!cZeta.Create(zeta) || !cNu.Create(nu)) { PrintFormat("Error of create function nets: %d", GetLastError()); class="kw">return false; } class=class="str">"cmt">//--- if(!cTargetCritic1.Create(critic) || !cTargetCritic2.Create(critic) || !cTargetNu.Create(nu)) { PrintFormat("Error of create target models: %d", GetLastError()); class="kw">return false; } cActorExploer.SetOpenCL(opencl); cCritic1.SetOpenCL(opencl); cCritic2.SetOpenCL(opencl); cZeta.SetOpenCL(opencl); cNu.SetOpenCL(opencl); cTargetCritic1.SetOpenCL(opencl); cTargetCritic2.SetOpenCL(opencl); cTargetNu.SetOpenCL(opencl); if(!cTargetCritic1.WeightsUpdate(GetPointer(cCritic1), class="num">1.0) || !cTargetCritic2.WeightsUpdate(GetPointer(cCritic2), class="num">1.0) || !cTargetNu.WeightsUpdate(GetPointer(cNu), class="num">1.0)) { PrintFormat("Error of update target models: %d", GetLastError()); class="kw">return false; } fLambda = class="num">1.0e-5f; fLambda_m = class="num">0; fLambda_v = class="num">0; fZeta = class="num">0; iLatentLayer = latent_layer; class=class="str">"cmt">//--- class="kw">return true; } class="type">bool CNet_SAC_DICE::Study(CArrayFloat *State, CArrayFloat *SecondInput, CBufferFloat *Actions, vector<class="type">float> &ActionsLogProbab, CBufferFloat *NextState, CBufferFloat *NextSecondInput, class="type">float reward, class="type">float discount, class="type">float tau) { class=class="str">"cmt">//--- if(!Actions || Actions.Total()!=ActionsLogProbab.Size()) class="kw">return false; if(!CNet::feedForward(NextState, class="num">1, false, NextSecondInput)) class="kw">return false; if(!cTargetCritic1.feedForward(GetPointer(this), iLatentLayer, GetPointer(this), layers.Total() - class="num">1) || !cTargetCritic2.feedForward(GetPointer(this), iLatentLayer, GetPointer(this), layers.Total() - class="num">1))
「Dual-Critic 里的拉格朗日乘子怎么动」
这段逻辑跑在强化学习交易智能体的训练回路里,用双评论家(nu / target_nu)加一个约束网络 zeta 来近似对偶变量。bellman_residuals 由下一状态价值、折扣因子 discount、策略概率比 policy_ratio 与即时 reward 拼出,公式写成 next_nu * discount * policy_ratio - nu + policy_ratio * reward,任何一项为 NaN 都会让后续梯度崩掉。 zeta_loss 与 nu_loss 都带 L2 正则项(各自除以 2 或 2.0f),其中 fLambda 是对偶乘子,靠 Adam 风格更新:fLambda_m、fLambda_v 分别用 b1、b2 做一阶与二阶矩滑动平均,步长 lr 乘上 m 再除以 sqrt(v)(v 为 0 时兜底成 1.0f)。在 MT5 里把 b1 设 0.9、b2 设 0.999 是常见起点,但外汇与贵金属波动下可能需调小 lr 防震荡。 更新 nu 梯度时,代码从最后一层取出 CNeuronBaseOCL 的 gradient buffer,把 nu_grad = nu_loss * (zeta * bellman_residuals / MathAbs(bellman_residuals) + nu) 写回再反向传播;这里用 MathAbs 做符号函数近似,若 bellman_residuals 恰为 0 会除到 0,实盘前最好在本地用 EURUSD 5 分钟样本跑一次确认无 inf。
if(!cTargetNu.feedForward(GetPointer(this), iLatentLayer, GetPointer(this), layers.Total() - class="num">1)) class="kw">return false; if(!CNet::feedForward(State, class="num">1, false, SecondInput)) class="kw">return false; CBufferFloat *output = ((CNeuronBaseOCL*)((CLayer*)layers.At(layers.Total() - class="num">1)).At(class="num">0)).getOutput(); output.AssignArray(Actions); output.BufferWrite(); if(!cNu.feedForward(GetPointer(this), iLatentLayer, GetPointer(this))) class="kw">return false; if(!cZeta.feedForward(GetPointer(this), iLatentLayer, GetPointer(this))) class="kw">return false; vector<class="type">float> nu, next_nu, zeta, ones; cNu.getResults(nu); cTargetNu.getResults(next_nu); cZeta.getResults(zeta); ones = vector<class="type">float>::Ones(zeta.Size()); vector<class="type">float> log_prob = GetLogProbability(output); class="type">float policy_ratio = MathExp((log_prob - ActionsLogProbab).Sum()); vector<class="type">float> bellman_residuals = next_nu * discount * policy_ratio - nu + policy_ratio * reward; vector<class="type">float> zeta_loss = zeta * (MathAbs(bellman_residuals) - fLambda) * (-class="num">1) + MathPow(zeta, class="num">2.0f) / class="num">2; vector<class="type">float> nu_loss = zeta * MathAbs(bellman_residuals) + MathPow(nu, class="num">2.0f) / class="num">2.0f; class="type">float lambda_los = fLambda * (ones - zeta).Sum(); class=class="str">"cmt">//--- update lambda class="type">float grad_lambda = (ones - zeta).Sum() * (-lambda_los); fLambda_m = b1 * fLambda_m + (class="num">1 - b1) * grad_lambda; fLambda_v = b2 * fLambda_v + (class="num">1 - b2) * MathPow(grad_lambda, class="num">2); fLambda += lr * fLambda_m / (fLambda_v != class="num">0.0f ? MathSqrt(fLambda_v) : class="num">1.0f); class=class="str">"cmt">//--- CBufferFloat temp; temp.BufferInit(MathMax(Actions.Total(), SecondInput.Total()), class="num">0); temp.BufferCreate(opencl); class=class="str">"cmt">//--- update nu class="type">int last_layer = cNu.layers.Total() - class="num">1; CLayer *layer = cNu.layers.At(last_layer); if(!layer) class="kw">return false; CNeuronBaseOCL *neuron = layer.At(class="num">0); if(!neuron) class="kw">return false; CBufferFloat *buffer = neuron.getGradient(); if(!buffer) class="kw">return false; vector<class="type">float> nu_grad = nu_loss * (zeta * bellman_residuals / MathAbs(bellman_residuals) + nu); if(!buffer.AssignArray(nu_grad) || !buffer.BufferWrite()) class="kw">return false; if(!cNu.backPropGradient(output, GetPointer(temp))) class="kw">return false; class=class="str">"cmt">//--- update zeta last_layer = cZeta.layers.Total() - class="num">1; layer = cZeta.layers.At(last_layer); if(!layer) class="kw">return false; neuron = layer.At(class="num">0); if(!neuron)
◍ 分布式critic的梯度回传与ζ缩放
| 这段逻辑跑在强化学习智能体的更新步里:先取神经元梯度缓冲,任一环节拿不到指针就直接 return false,避免脏数据写进权重。zeta_grad 按 zeta_loss * (zeta - | bellman_residuals | + fLambda) * (-1) 构造,负号代表沿损失减小的反方向推。 |
|---|
| 前馈阶段把状态送进两个 critic 网络(cCritic1 / cCritic2),任一 feedForward 失败即退出。fZeta 用 0.9/0.1 的 EMA 平滑绝对值,初始为 0 时直接取 | zeta[0] | ,后续滚动更新。 |
|---|
| zeta[0] 被立方根归一化: | zeta[0] | ^(1/3) / (10 * fZeta^(1/3)),把量纲压到接近 1 的区间,防止 critic 损失被极端残差带飞。目标值取双 critic 的较小者减 LogProbMultiplier * log_prob.Sum(),这是 SAC 类算法的标准 min-Q 操作。 |
|---|
critic1 的 loss 写成 zeta[0] * (Q - target)^2,fLoss1 同样用 0.999/0.001 的滑动均方根跟踪。梯度 grad = loss * 2 * zeta[0] * (target - result[0]),写进最后一层神经元缓冲后触发 backPropGradient。critic2 对称处理,但代码里 fLoss2 的 EMA 误用了 fLoss1 的幂次——在 MT5 里跑这套时建议改成 MathPow(fLoss2, 2.0f),否则第二个网络的损失尺度会失真。外汇与贵金属行情跳空频繁,这类 RL 估值网络在高波动时段可能给出偏离较大的 Q 值,实盘前务必用历史 tick 回测确认。
class="kw">return false; buffer = neuron.getGradient(); if(!buffer) class="kw">return false; vector<class="type">float> zeta_grad = zeta_loss * (zeta - MathAbs(bellman_residuals) + fLambda) * (-class="num">1); if(!buffer.AssignArray(zeta_grad) || !buffer.BufferWrite()) class="kw">return false; if(!cZeta.backPropGradient(output, GetPointer(temp))) class="kw">return false; class=class="str">"cmt">//--- feed forward critics if(!cCritic1.feedForward(GetPointer(this), iLatentLayer, output) || !cCritic2.feedForward(GetPointer(this), iLatentLayer, output)) class="kw">return false; vector<class="type">float> result; if(fZeta == class="num">0) fZeta = MathAbs(zeta[class="num">0]); else fZeta = class="num">0.9f * fZeta + class="num">0.1f * MathAbs(zeta[class="num">0]); zeta[class="num">0] = MathPow(MathAbs(zeta[class="num">0]), class="num">1.0f / class="num">3.0f) / (class="num">10.0f * MathPow(fZeta, class="num">1.0f / class="num">3.0f)); cTargetCritic1.getResults(result); class="type">float target = result[class="num">0]; cTargetCritic2.getResults(result); target = reward + discount * (MathMin(result[class="num">0], target) - LogProbMultiplier * log_prob.Sum()); class=class="str">"cmt">//--- update critic1 cCritic1.getResults(result); class="type">float loss = zeta[class="num">0] * MathPow(result[class="num">0] - target, class="num">2.0f); if(fLoss1 == class="num">0) fLoss1 = MathSqrt(loss); else fLoss1 = MathSqrt(class="num">0.999f * MathPow(fLoss1, class="num">2.0f) + class="num">0.001f * loss); class="type">float grad = loss * class="num">2 * zeta[class="num">0] * (target - result[class="num">0]); last_layer = cCritic1.layers.Total() - class="num">1; layer = cCritic1.layers.At(last_layer); if(!layer) class="kw">return false; neuron = layer.At(class="num">0); if(!neuron) class="kw">return false; buffer = neuron.getGradient(); if(!buffer) class="kw">return false; if(!buffer.Update(class="num">0, grad) || !buffer.BufferWrite()) class="kw">return false; if(!cCritic1.backPropGradient(output, GetPointer(temp)) || !backPropGradient(SecondInput, GetPointer(temp), iLatentLayer)) class="kw">return false; class=class="str">"cmt">//--- update critic2 cCritic2.getResults(result); loss = zeta[class="num">0] * MathPow(result[class="num">0] - target, class="num">2.0f); if(fLoss2 == class="num">0) fLoss2 = MathSqrt(loss); else fLoss2 = MathSqrt(class="num">0.999f * MathPow(fLoss1, class="num">2.0f) + class="num">0.001f * loss); grad = loss * class="num">2 * zeta[class="num">0] * (target - result[class="num">0]); last_layer = cCritic2.layers.Total() - class="num">1; layer = cCritic2.layers.At(last_layer);
SAC-DICE 的梯度回传与目标网络更新
这段实现把双评论家(cCritic1/2)的梯度先取出来,再做策略层面的反向传播。任何一层指针为空或缓冲区写入失败都直接 return false,说明在 MT5 的 OpenCL 张量层里,空指针防护比模型结构本身更容易让训练静默中断。 均值与方差的融合写得很直白:两个评论家各出 result[0],mean 取平均,var 取绝对值折半。目标值用 zeta[0]*(mean - 2.5f*var + discount*log_prob.Sum()*LogProbMultiplier) + result[0] 计算,其中 2.5f 这个惩罚系数偏大,倾向压制高方差动作——在外汇 15M 级别上可能让 agent 过早保守。 探索策略分支里 target 改成 mean + 2.0f*var,符号反转即鼓励不确定性,用于离线条目采集。最后三行 WeightsUpdate 用 tau 做软更新,若返回 false 会 PrintFormat 报错码,开 MT5 跑时建议把 Experts 日志打开看 Error of update target models 是否刷屏。 保存方法以 .set 二进制落盘,common 参数控制是否进公共目录;handle 为 INVALID_HANDLE 即放弃,没有异常抛出,调用方必须自己判断返回值。
if(!layer) class="kw">return false; neuron = layer.At(class="num">0); if(!neuron) class="kw">return false; buffer = neuron.getGradient(); if(!buffer) class="kw">return false; if(!buffer.Update(class="num">0, grad) || !buffer.BufferWrite()) class="kw">return false; if(!cCritic2.backPropGradient(output, GetPointer(temp)) || !backPropGradient(SecondInput, GetPointer(temp), iLatentLayer)) class="kw">return false; class=class="str">"cmt">//--- update policy cCritic1.getResults(result); class="type">float mean = result[class="num">0]; class="type">float var = result[class="num">0]; cCritic2.getResults(result); mean += result[class="num">0]; var -= result[class="num">0]; mean /= class="num">2.0f; var = MathAbs(var) / class="num">2.0f; target = zeta[class="num">0] * (mean - class="num">2.5f * var + discount * log_prob.Sum() * LogProbMultiplier) + result[class="num">0]; CBufferFloat bTarget; bTarget.Add(target); cCritic2.TrainMode(false); if(!cCritic2.backProp(GetPointer(bTarget), GetPointer(this)) || !backPropGradient(SecondInput, GetPointer(temp))) { cCritic2.TrainMode(true); class="kw">return false; } class=class="str">"cmt">//--- update exploration policy if(!cActorExploer.feedForward(State, class="num">1, false, SecondInput)) { cCritic2.TrainMode(true); class="kw">return false; } output = ((CNeuronBaseOCL*)((CLayer*)cActorExploer.layers.At(layers.Total() - class="num">1)).At(class="num">0)).getOutput(); output.AssignArray(Actions); output.BufferWrite(); cActorExploer.GetLogProbs(log_prob); target = zeta[class="num">0] * (mean + class="num">2.0f * var + discount * log_prob.Sum() * LogProbMultiplier) + result[class="num">0]; bTarget.Update(class="num">0, target); if(!cCritic2.backProp(GetPointer(bTarget), GetPointer(cActorExploer)) || !cActorExploer.backPropGradient(SecondInput, GetPointer(temp))) { cCritic2.TrainMode(true); class="kw">return false; } cCritic2.TrainMode(true); if(!cTargetCritic1.WeightsUpdate(GetPointer(cCritic1), tau) || !cTargetCritic2.WeightsUpdate(GetPointer(cCritic2), tau) || !cTargetNu.WeightsUpdate(GetPointer(cNu), tau)) { PrintFormat("Error of update target models: %d", GetLastError()); class="kw">return false; } class=class="str">"cmt">//--- class="kw">return true; } class="type">bool CNet_SAC_DICE::Save(class="type">class="kw">string file_name, class="type">bool common = true) { if(file_name == NULL) class="kw">return false; class="type">int handle = FileOpen(file_name + ".set", (common ? FILE_COMMON : class="num">0) | FILE_BIN | FILE_WRITE); if(handle == INVALID_HANDLE)
「SAC-DICE 模型的存档与读档落点」
这段逻辑属于 SAC-DICE 强化学习框架的持久化部分:训练中途把超参与各子网络权重落盘,下次启动直接 Load 恢复,避免从零收敛。Save 里先写 .set 二进制配置,再逐个调子网络对象的 Save 方法,文件名后缀区分角色——Act 是 actor,Crt1/Crt2 是双 critic,Zeta 和 Nu 对应 DICE 的约束网络。 .set 文件用 FileWriteFloat 存了 3 个 float(fLambda、fLambda_m、fLambda_v)加 1 个 int(iLatentLayer),任一写入字节数不足 sizeof 就 return false,说明对磁盘写入完整性是零容忍的。FileFlush 紧跟在配置写完后调用,确保缓冲区进盘再关句柄。 Load 侧对称:先开 .set 读配置,每次 FileReadFloat 之后都查 FileIsEnding,读到尾直接 false,防止截断文件让后续网络维度错配。子网络 Load 用临时 float temp 和 datetime dt 接住不需要的返回值字段,common 参数控制走 FILE_COMMON 共享目录还是本地 sandbox。 在 MT5 里跑这套,重点验证两件事:一是 .set 和 .nnw 是否同目录且 common 标志一致;二是改了 iLatentLayer 后旧存档读进来维度对不上会静默失败。外汇与贵金属行情下用这类模型,过拟合与实时分布偏移风险偏高,复盘务必用样本外数据。
class="kw">return false; if(FileWriteFloat(handle, fLambda) < class="kw">sizeof(fLambda) || FileWriteFloat(handle, fLambda_m) < class="kw">sizeof(fLambda_m) || FileWriteFloat(handle, fLambda_v) < class="kw">sizeof(fLambda_v) || FileWriteInteger(handle, iLatentLayer) < class="kw">sizeof(iLatentLayer)) class="kw">return false; FileFlush(handle); FileClose(handle); if(!CNet::Save(file_name + "Act.nnw", class="num">0, class="num">0, class="num">0, TimeCurrent(), common)) class="kw">return false; class=class="str">"cmt">//--- if(!cActorExploer.Save(file_name + "ActExp.nnw", class="num">0, class="num">0, class="num">0, TimeCurrent(), common)) class="kw">return false; class=class="str">"cmt">//--- if(!cTargetCritic1.Save(file_name + "Crt1.nnw", fLoss1, class="num">0, class="num">0, TimeCurrent(), common)) class="kw">return false; class=class="str">"cmt">//--- if(!cTargetCritic2.Save(file_name + "Crt2.nnw", fLoss2, class="num">0, class="num">0, TimeCurrent(), common)) class="kw">return false; class=class="str">"cmt">//--- if(!cZeta.Save(file_name + "Zeta.nnw", class="num">0, class="num">0, class="num">0, TimeCurrent(), common)) class="kw">return false; class=class="str">"cmt">//--- if(!cTargetNu.Save(file_name + "Nu.nnw", class="num">0, class="num">0, class="num">0, TimeCurrent(), common)) class="kw">return false; class=class="str">"cmt">//--- class="kw">return true; } class="type">bool CNet_SAC_DICE::Load(class="type">class="kw">string file_name, class="type">bool common = true) { if(file_name == NULL) class="kw">return false; class=class="str">"cmt">//--- class="type">int handle = FileOpen(file_name + ".set", (common ? FILE_COMMON : class="num">0) | FILE_BIN | FILE_READ); if(handle == INVALID_HANDLE) class="kw">return false; if(FileIsEnding(handle)) class="kw">return false; fLambda = FileReadFloat(handle); if(FileIsEnding(handle)) class="kw">return false; fLambda_m = FileReadFloat(handle); if(FileIsEnding(handle)) class="kw">return false; fLambda_v = FileReadFloat(handle); if(FileIsEnding(handle)) class="kw">return false; iLatentLayer = FileReadInteger(handle);; FileClose(handle); class=class="str">"cmt">//--- class="type">float temp; class="type">class="kw">datetime dt; if(!CNet::Load(file_name + "Act.nnw", temp, temp, temp, dt, common)) class="kw">return false; class=class="str">"cmt">//---
◍ 模型权重装载与状态结构落地
强化学习智能体在 MT5 里跑起来之前,得先把离线训好的网络权重文件逐个读进内存。下面这段Load逻辑分别加载 Actor 探索网络、两套 Critic 及其目标网络、Zeta 与 Nu 网络,只要任意一个 .nnw 读取失败就直接返回 false 中断初始化。 if(!cActorExploer.Load(file_name + "ActExp.nnw", temp, temp, temp, dt, common)) return false; //---
| if(!cCritic1.Load(file_name + "Crt1.nnw", fLoss1, temp, temp, dt, common) |
|---|
!cTargetCritic1.Load(file_name + "Crt1.nnw", temp, temp, temp, dt, common)) return false; //---
| if(!cCritic2.Load(file_name + "Crt2.nnw", fLoss2, temp, temp, dt, common) |
|---|
!cTargetCritic2.Load(file_name + "Crt2.nnw", temp, temp, temp, dt, common)) return false; //--- if(!cZeta.Load(file_name + "Zeta.nnw", temp, temp, temp, dt, common)) return false; //---
| if(!cNu.Load(file_name + "Nu.nnw", temp, temp, temp, dt, common) |
|---|
!cTargetNu.Load(file_name + "Nu.nnw", temp, temp, temp, dt, common)) return false; 权重读完后统一调 SetOpenCL 把计算丢给显卡,8 个网络对象一个不落。SAC-DICE 这类算法靠双 Critic 和目标网络做稳定性约束,缺一个目标网就可能让训练偏移,实盘外汇或贵金属推理时误差会被杠杆放大,属高风险操作。 状态用 SState 结构体封装:state 长度固定为 HistoryBars * BarDescr,account 比 AccountDescr 少 4 个浮点,action 与 log_prob 各占 NActions。重载等号时用 ArrayCopy 做深拷贝,避免回放缓冲区里多条轨迹互相污染。 OnInit 里先 ResetLastError 再调 LoadTotalBase,失败就 PrintFormat 打出错误码并返回 INIT_FAILED;随后 Net.Load(FileName, true) 装载总模型。开 MT5 把对应 .nnw 丢进终端数据目录,缺文件时日志会直接暴露加载断点,方便你定位是哪一环没训完。
if(!cActorExploer.Load(file_name + "ActExp.nnw", temp, temp, temp, dt, common)) class="kw">return false; class=class="str">"cmt">//--- if(!cCritic1.Load(file_name + "Crt1.nnw", fLoss1, temp, temp, dt, common) || !cTargetCritic1.Load(file_name + "Crt1.nnw", temp, temp, temp, dt, common)) class="kw">return false; class=class="str">"cmt">//--- if(!cCritic2.Load(file_name + "Crt2.nnw", fLoss2, temp, temp, dt, common) || !cTargetCritic2.Load(file_name + "Crt2.nnw", temp, temp, temp, dt, common)) class="kw">return false; class=class="str">"cmt">//--- if(!cZeta.Load(file_name + "Zeta.nnw", temp, temp, temp, dt, common)) class="kw">return false; class=class="str">"cmt">//--- if(!cNu.Load(file_name + "Nu.nnw", temp, temp, temp, dt, common) || !cTargetNu.Load(file_name + "Nu.nnw", temp, temp, temp, dt, common)) class="kw">return false; cActorExploer.SetOpenCL(opencl); cCritic1.SetOpenCL(opencl); cCritic2.SetOpenCL(opencl); cZeta.SetOpenCL(opencl); cNu.SetOpenCL(opencl); cTargetCritic1.SetOpenCL(opencl); cTargetCritic2.SetOpenCL(opencl); cTargetNu.SetOpenCL(opencl); class=class="str">"cmt">//--- class="kw">return true; } class="kw">struct SState { class="type">float state[HistoryBars * BarDescr]; class="type">float account[AccountDescr - class="num">4]; class="type">float action[NActions]; class="type">float log_prob[NActions]; class=class="str">"cmt">//--- SState(class="type">void); class=class="str">"cmt">//--- class="type">bool Save(class="type">int file_handle); class="type">bool Load(class="type">int file_handle); class=class="str">"cmt">//--- overloading class="type">void class="kw">operator=(const SState &obj) { ArrayCopy(state, obj.state); ArrayCopy(account, obj.account); ArrayCopy(action, obj.action); ArrayCopy(log_prob, obj.log_prob); } }; class="macro">#include "Net_SAC_DICE.mqh" STrajectory Buffer[]; CNet_SAC_DICE Net; class=class="str">"cmt">//--- class="type">float dError; class="type">class="kw">datetime dtStudied; class=class="str">"cmt">//--- CBufferFloat bState; CBufferFloat bAccount; CBufferFloat bActions; CBufferFloat bNextState; CBufferFloat bNextAccount; class="type">int OnInit() { class=class="str">"cmt">//--- ResetLastError(); if(!LoadTotalBase()) { PrintFormat("Error of load study data: %d", GetLastError()); class="kw">return INIT_FAILED; } class=class="str">"cmt">//--- load models if(!Net.Load(FileName, true)) {