ARTICLE DETAIL

资讯详情

深耕郑州网站建设与运营推广的一线实战洞察。

Matlab实现QRLSTM分位数回归:LSTM预测区间与不确定性估计

Matlab实现QRLSTM分位数回归:LSTM预测区间与不确定性估计 做时间序列回归预测的人早晚会遇到同一个问题模型只输出一个数可实际业务里我们更想知道这个数的可信范围。比如预测电池SOC荷电状态你说下个时刻电量是72.5%还不够运维人员更想知道这个值的误差范围有多大——是72%到73%还是65%到80%。普通LSTM给不了这个答案因为它只学了一个条件均值而这次要分享的QRLSTM分位数长短期记忆网络则把回归预测从一条线变成了一条带置信区间的区间带同时给出了预测值和不确定性水平。这篇文章就把这套基于Matlab的QRLSTM回归预测代码完整拆开讲一遍包括网络结构、损失函数、版本兼容方案特别是2018及以上版本怎么处理、训练调参经验和结果评价方法适合正在做时间序列预测、需要输出预测区间或者准备写相关论文的读者参考。1. 为什么普通LSTM不够非要加一个分位数很多人第一次听到QRLSTM会觉得这是又一个新的网络架构其实它本质上是LSTM加了一个代价函数——分位数损失也叫Pinball Loss弹球损失。理解分位数回归之前我们先说清楚普通LSTM回归到底丢失了什么信息。1.1 点预测给不出不确定性普通LSTM回归网络的输出层是一个值训练时用均方误差MSE或者平均绝对误差MAE做损失。学过统计的人都知道MSE对应的最优预测是条件均值E[y|x]MAE对应的最优预测是条件中位数Q(0.5|x)。也就是说你辛辛苦苦训练了一个网络它给你的其实只是整个分布中间那个位置的估计至于数据波动有多大、预测值靠不靠谱网络完全不知道。这不是LSTM的错是损失函数的局限。MSE把所有样本同等对待正误差和负误差的惩罚对称模型只关心平均意义上的误差最小化不关心误差的分布形状。在很多实际工程里光有中间值远远不够比如电力负荷预测调度员一方面要知道明天的峰值负荷是多少另一方面更关心最坏情况下的负荷上限是多少这直接关系到是否要备用电容量再比如用LSTM做锂电池SOC估计你不仅想知道剩余电量更想知道电池管理系统给出的估算值边界这决定了对电池保护策略的宽松程度。1.2 分位数回归的思想给损失函数装上偏心轮分位数回归原本是统计学里的经典方法核心就在这个损失函数上。假设我们想估计第τ个条件分位数Q(τ|x)传统的做法是最小化如下带不对称权重的绝对误差[ \rho_\tau(u) \begin{cases} \tau \cdot u, u \ge 0 \ (\tau-1) \cdot u, u 0 \end{cases} ]这里的u是实际值和预测值的差即残差。当预测值小于真实值残差为正欠拟合时误差代价是τ倍当预测值大于真实值残差为负过拟合时误差代价是(1-τ)倍。你看这个设计很巧妙当τ0.5时两边权重都是0.5损失退化为对称的绝对误差对应的就是中位数回归所以中位数回归只是分位数回归的一个特例当τ0.9真正误差的惩罚权重是0.9负误差惩罚只有0.1模型便倾向于宁可多预测一些也不能少预测于是学出来的结果就大致对着条件分布的90%分位。这个损失函数在Matlab里实现起来非常简洁function loss pinballLoss(y, yHat, tau) % y: 真实值可以是向量 % yHat: 预测值 % tau: 分位数取值(0,1) err y - yHat; loss mean(max(tau .* err, (tau - 1) .* err)); end为什么用max(tau*err, (tau-1)*err)你拆开来看就知道这是上面分段函数最紧凑的写法。当err 0时(tau-1)err是负数max取tauerr当err 0时tau*err是负数(tau-1)*err是正数max取后者。省掉了判断分支矩阵运算是向量化的训练时更快。1.3 QRLSTM的价值让LSTM学会条件分位数把分位数损失接到LSTM后面网络结构本身不需要大改LSTM部分照常提取时间序列的时序特征只是输出层和损失函数变了。输出层可以同时输出多个分位数的预测值比如一次性输出τ0.1、0.5、0.9三个结果损失函数则是三个分位数损失的加权和。这样训练出来的网络每一个时间步不仅给出一个预测点还给出了这个点在条件分布中的位置。我之前给一个微动面波数据反演项目做过类似的改造。面波频散曲线数据本身噪声大、样本有限传统反演给的速度结构是一个确定解但实际上观测数据的噪声会导致反演结果天然带不确定性。用QRLSTM后输出的地下速度剖面自带了可信区间后续地质解释时能明确区分哪些层位速度是可靠的、哪些是模糊的这在没有区间估计前根本做不到。这类场景不是LSTM不好用而是只有均值预测的LSTM不能满足业务需求。2. QRLSTM的网络结构设计与Matlab版本兼容性判断确定了核心思路接下来要设计网络结构。但这里有个所有Matlab用户都会撞上的墙不同版本对自定义损失函数的支持差得非常多偏偏你的需求是2018及以上版本都能跑这就逼着我要先做版本能力盘点再决定代码怎么写。2.1 Matlab各版本对自定义训练的支持情况先列一张表这是我自己踩完坑总结出来的比翻官方release notes更直观Matlab版本trainNetwork训练LSTM自定义损失函数dlnetwork自定义训练循环备注R2018a/R2018b支持但损失固定为MSE不支持无只能训练标准回归LSTMR2019b支持不支持官方改loss引入dlarray/dlnetwork可手写训练循环R2021a及以后支持支持trainnet或自定义循环完整支持推荐方案所以适用于2018及以上版本这句话至少包含两种成立策略一种是所有版本都只用官方trainNetwork那就要避开自定义损失另一种是把2018当作下限在2018上采用兼容方案在更高版本上用完整方案。比较严谨的做法是提供两套路径让用户根据自己电脑上的版本对号入座。2.2 方案A单模型多输出推荐2021a及以上使用网络结构大致如下layers [ sequenceInputLayer(numFeatures, Normalization, none) lstmLayer(numHiddenUnits, OutputMode, last, Name, lstm1) fullyConnectedLayer(numHiddenUnits, Name, fc1) reluLayer(Name, relu1) fullyConnectedLayer(length(tauArray), Name, fc_out) regressionLayer(Name, regout) ];这个结构里最后一个全连接层的神经元个数等于你要预测的分位数个数比如tau取[0.1 0.5 0.9]输出就是3个并列的神经元每个对应一个分位数。问题是regressionLayer内部用的是MSE它不认识你的pinball loss。所以要用dlnetwork自己定义前向传播和梯度更新% 构建dlnetwork dlnet dlnetwork(layers); % 训练循环核心 for iteration 1:numIterations [X, y] getMiniBatch(...); % 取一个batch dlX dlarray(X, CTC); % 序列输入 dlY dlarray(y, CB); % 目标维度 [numTau, batch] [dlYPred, dlnet] forward(dlnet, dlX, Outputs, fc_out); % 多分位数pinball loss err dlY - dlYPred; tau dlarray(tauArray(:), CB); % 形状和预测对齐 loss mean(max(tau .* err, (tau - 1) .* err), all); % 反向传播 grads dlgradient(loss, dlnet.Learnables); [dlnet, avgLoss] adamupdate(dlnet, grads, avgLoss, iteration, learnRate); end这里有个容易被忽略的细节dlY的形状要扩成[numTau, batchSize]或者[numTau, numTimesteps, batchSize]具体看你预测的是单步还是多步。如果你用OutputModelast做的单步滚动预测那就是[numTau, batchSize]如果做的是sequence-to-sequence的多步预测输出是[numTau, numSteps, batchSize]损失里还要考虑对时间步的维度做求和或平均。我自己一般固定用单步滚动预测做区间输出代码简单训练也快。2.3 方案B2018兼容的两阶段均值LSTM 残差分位数严格讲R2018的trainNetwork把损失焊死成了MSE你没法在官方框架里改。但在2018上直接用多个LSTM模型各学一个分位数也不行因为默认的MSE损失并不会让模型收敛到真正的分位数上。实际项目中我用的是两阶段方案它分两步走每一步都能用2018的trainNetwork完成。第一步先训练一个标准LSTM回归模型输出条件均值预测记为(\hat{y}_{\text{mean}})。第二步把训练集里的残差算出来(r_i y_i - \hat{y}{\text{mean}}(x_i))。对残差单独做分位数分析比如对所有残差按升序排序找第10%、50%、90%分位点得到(q{0.1}, q_{0.5}, q_{0.9})。或者更精细一点把残差分桶在每个桶内用局部回归做条件分位数估计。最终预测区间就是yPred_lower yPred_mean q(1); yPred_mid yPred_mean q(2); yPred_upper yPred_mean q(3);这个方法在理论上不如单模型分位数回归严格因为它是均值加残差分布的近似没有直接最小化分位数损失但在2018版本约束下它是一个能落地的折中方案而且如果你的残差分布相对平稳、不同输入位置的条件分布形状差不多效果会很接近直接分位数回归。如果只做粗粒度区间估计这个方案的性价比非常高。2.4 版本和方案怎么选我的建议是如果能装新版直接用方案A精度最高输出区间也最自然如果只能守着2018方案B是唯一合理路径虽然理论上打了折扣但代码稳定且官方支持良好。不要试图在2018上用一些奇技淫巧强行改损失比如用加权MSE近似、或者二次采样改标签分布我都试过换来的往往是训练不收敛或者区间失真不值得。3. Matlab代码实战拆解从数据预处理到区间输出这一节我把整个流程的代码片段拆开讲。环境假设是2021a以上采用方案A方案B的差异我会在关键位置用注释标注。3.1 数据预处理滑窗切分和归一化QRLSTM用于回归预测时输入是有历史长度的时间序列窗口输出是未来一个或多个时刻的值。以单步预测为例用一个长度为windowSize的窗口去预测horizon步之后的值function [XTrain, YTrain] makeSequences(data, windowSize, horizon) numSamples size(data, 1) - windowSize - horizon 1; XTrain zeros(windowSize, numFeatures, numSamples); YTrain zeros(numQuantiles, 1, numSamples); for i 1:numSamples XTrain(:, :, i) data(i : i windowSize - 1, :); YTrain(:, 1, i) data(i windowSize horizon - 1, 1); end end注意Matlab内置LSTM层要求的序列输入格式是[numFeatures, numTimesteps, numObservations]特征第一维这和Python的[batch, timesteps, features]刚好反过来新手在这上面翻车的概率极高。我看到过很多代码维度写反了然后trainNetwork报错说维度不匹配再然后就开始怀疑LSTM层的问题。实际上只要记得Matlab的格式约定这里就不会出错。归一化使用mapminmax或者zscore都行我建议用zscore。原因稍后在第4节详细说简单提一句分位数损失是基于绝对误差的不对称加权如果数据尺度差异大归一化不当会导致不同样本之间损失权重失真zscore能把残差尺度拉平训练更稳定。3.2 分位数损失函数完整实现上面给了pinballLoss的简洁版但在dlnetwork训练循环里要注意矩阵广播和dlarray维度的处理。写一个能直接用的版本function loss quantileLoss(dlY, dlYPred, tauArray) % dlY: 真实值维度 [numTau, batchSize] % dlYPred: 预测值维度 [numTau, batchSize] % tauArray: 分位数向量[numTau, 1] err dlY - dlYPred; tau dlarray(tauArray(:), CB); loss mean(max(tau .* err, (tau - 1) .* err), all); end如果只预测一个分位数比如tau0.9那dlY和dlYPred的维度是[1, batchSize]这个函数依然工作正常。一个经验之谈多分位数同时输出时我倾向于对各分位数的损失先各自求平均、再对所有分位数求平均而不是对所有元素一次mean all。这两种写法看起来是一回事但在样本不均匀的时序数据里前者等价于每个分位数贡献相同的权重后者则会让样本量多的分位数占上风。虽然这里所有分位数共享的样本量一样结果没差别但养成这个习惯后换到分组数据时不容易踩坑。3.3 自定义训练循环主框架下面是训练循环的主体2021a及以上标准写法numIterationsPerEpoch floor(numTrain / miniBatchSize); numEpochs 80; numIterations numEpochs * numIterationsPerEpoch; learnRate 0.001; gradDecay 0.9; sqGradDecay 0.999; averageGrad []; averageSqGrad []; for iteration 1:numIterations [XBatch, YBatch] getBatch(XTrain, YTrain, miniBatchSize, iteration); dlX dlarray(XBatch, CTC); dlY dlarray(YBatch, CB); [dlYPred, dlnet] forward(dlnet, dlX, Outputs, fc_out); loss quantileLoss(dlY, dlYPred, tauArray); grads dlgradient(loss, dlnet.Learnables); % Adam更新 [dlnet, averageGrad, averageSqGrad] adamupdate(dlnet, grads, ... averageGrad, averageSqGrad, iteration, learnRate, gradDecay, sqGradDecay); if mod(iteration, 50) 0 fprintf(Iter %d, loss %.4f\n, iteration, extractdata(loss)); end end有几个地方要专门提一下。forward(dlnet, dlX, Outputs, fc_out)这种写法要求你在定义层时给全连接层加了Name。如果不加名字默认的名字是类似fc_1你还要去查层图里的具体名称麻烦。所以我的习惯是每一层都显式命名。adamupdate是深度学习中常用的Adam优化器封装手里没有的话自己写梯度更新也不是不行但Adam对这类非对称损失的收敛稳定性很有帮助建议直接用。关于学习率0.001是相对安全的默认值但对分位数损失来说我倒觉得可以从0.005开始如果loss震荡明显再降到0.001。原因是pinball loss的梯度在残差接近0的位置会出现拐点网络需要更快地越过这个非平滑区。3.4 前向预测和区间组装训练结束后的预测代码function [yPred, yLower, yUpper] predictInterval(dlnet, XTest, tauArray) dlXTest dlarray(XTest, CTC); dlYPred predict(dlnet, dlXTest, Outputs, fc_out); yPredAll extractdata(dlYPred); % [numTau, numTest] % 默认tauArray [0.1, 0.5, 0.9] yLower yPredAll(1, :); yMid yPredAll(2, :); yUpper yPredAll(3, :); % 反归一化注意需要保留zscore的均值和标准差 yLower yLower * stdY meanY; yMid yMid * stdY meanY; yUpper yUpper * stdY meanY; yPred [yLower; yMid; yUpper]; end这里反归一化极其容易出错。如果训练前对数据做了zscore归一化那么预测出来的分位数依然是在归一化空间里的必须同时乘以标准差再加均值三个分位数各自都要做一次不能偷懒只反归一化中间那个。我在项目里曾见过有人只对中位数做了反变换上下界直接拿归一化空间的原始值画图结果区间严重失真在图上看起来上下界甚至可能越过了真实数据范围。3.5 2018兼容方案的核心补充采用方案B时第一阶段的LSTM训练直接用trainNetwork就好这里不再重复。第二阶段残差分位数估计的代码可以这样写residuals yTrainTrue - yPredMean; % 简单做法全量残差求分位数 q quantile(residuals, [0.1, 0.5, 0.9]); % 更精细做法按预测值大小分桶在每个桶内求残差分位数 edges prctile(yPredMean, linspace(0, 100, 10)); [~, ~, idx] histcounts(yPredMean, edges); q_by_bin zeros(3, length(edges)-1); for b 1:length(edges)-1 r_bin residuals(idx b); if ~isempty(r_bin) q_by_bin(:, b) quantile(r_bin, [0.1, 0.5, 0.9]); end end按桶求分位数的做法适合残差尺度随预测值变化的场景也就是常见的喇叭口数据预测值越大噪声越大、区间越宽。如果你只做全局残差分位数等于假设噪声方差恒定在股市、电力负荷这些数据里往往不够用。4. 分位数回归预测里最容易被坑的细节QRLSTM本身不难难的是让训练出来的分位数看起来合理、用起来可靠。以下这些坑我基本都是实际项目里全都踩过一遍的。4.1 分位数交叉问题分位数回归有个经典毛病两个不同分位数的输出曲线可能出现交叉。也就是说理论上0.9分位数应该始终大于0.1分位数但某些时间点上可能反过来了。这在单模型多输出的方案里依然可能出现因为网络输出的是三个独立的回归值它没有内建排序约束。处理办法有三个层级。第一层是检查tauArray和输出神经元的对应顺序是否一致很多人把tau设成[0.9, 0.5, 0.1]最后输出顺序乱了导致画图时上下界互换第二层是训练时加软约束比如给损失函数加上一个小的惩罚项当上界小于下界时按差值惩罚第三层是简单粗暴的后处理预测完成后做一个排序哪里交叉了就按单调性调整。在实际跑数据时我发现轻微的、个别点的交叉基本无法彻底避免但只要交叉幅度小、不密集出现对区间评价指标的影响很小。真正要警惕的是大范围交叉那往往说明tauArray间隔过大比如同时预测0.05和0.95或者数据里存在严重离群点网络在分位数之间权衡时出现了混乱。我的建议是tau阵列不要选两极分化太狠的参数比如[0.1, 0.5, 0.9]就比[0.02, 0.5, 0.98]稳健得多。4.2 归一化方式的选择做分位数预测时归一化这个预处理琐事对结果的影响比很多人大得多。前文我提到了zscore优于mapminmax这里补充一句原因mapminmax把数据映射到某个固定区间通常是[-1,1]当训练集和测试集的数据范围不一致时时序数据太常见了明天可能就比历史最大值还高映射关系会出问题zscore用的是均值和标准差对新增数据的容忍度更高。另外分位数损失是残差的非对称绝对误差。如果数据没有归一化量纲放大几倍损失的绝对值就会成比例放大学习率要重新调归一化之后不同数据集之间的超参数经验值才可以迁移。我自己习惯固定用zscore并且把反归一化需要的meanY和stdY在训练前就存好省得预测阶段手忙脚乱找参数。4.3 时序泄露看似在做预测实际用了未来数据滑窗构造数据时最隐蔽的问题是窗口边界和标签之间的重叠。假设windowSize50, horizon1第i个样本的输入是data(i:i49)标签是data(i50)这没问题但如果你不小心在构造下一个样本时让窗口滑动了1步标签却用了同样的data(i50)那就等于相邻两个样本共享了一部分输入窗口训练集和验证集之间就会出现信息重叠。治本办法是做gap式划分比如训练集取前70%时间点验证集取后30%中间留出至少一个窗口宽度的缓冲。我在第3.1节的makeSequences函数里特意在末尾减掉了windowSize horizon - 1就是为了保证最后一个样本的标签不超过数据末尾。这个细节看起来不起眼但对时间序列的评估可信度影响很大。4.4 梯度裁剪和多分位数损失的稳定性分位数损失本身是一个非平滑函数在残差为零的点不可导。dlgradient计算时会用次梯度近似处理大多数情况下没问题但如果某几个样本残差恰好频繁跨越零点梯度模会突然变大训练loss会出现毛刺。处理方式在adamupdate之前加一步梯度裁剪Matlab里没有现成的clipGradients函数需要自己写threshold 2; for i 1:numel(grads) gradVal grads(i).Value; normGrad sqrt(sum(gradVal(:).^2)); if normGrad threshold grads(i).Value gradVal * (threshold / normGrad); end end阈值threshold2是我常用的经验值具体看梯度的分布情况可以微调。加了裁剪之后多分位数训练的稳定性改善非常明显尤其在训练初期loss大幅下降的那几十个迭代里能有效抑制震荡。4.5 tauArray怎么定分位数取值的选择取决于业务需要。如果只是想要一个中间值±容差的区间[0.1, 0.5, 0.9]或者[0.05, 0.5, 0.95]都可以。但要注意tau越接近0或1对应的分位数越取决于分布的尾部训练难度越高需要的样本量也越多。样本量不足时尾部估计的方差会非常大区间上限可能毫无意义地偏高。如果业务要求90%置信区间又没有明确分位数偏好我的建议是先跑[0.05, 0.5, 0.95]看看覆盖率是否达标如果覆盖率远高于期望比如达到97%说明区间过宽可以换更保守的[0.1, 0.5, 0.9]或者直接手动缩小区间。实际项目中我是用覆盖率指标来反向修正tau的。5. 训练效果怎么评价PICP、PINAW和CWC指标训练完网络、画完分位数带接下来最容易被忽视的问题就是这个区间到底准不准光看图不够因为人的视觉系统会天然忽视那些特别窄的区间。所以必须用数值指标来量化评价。5.1 三个核心指标第一个是PICPPrediction Interval Coverage Probability预测区间覆盖概率它衡量真实值落在预测区间内的比例。比如100个测试点其中85个落在上下界之间PICP就是85%。这个指标理想情况下应该接近名义置信水平。第二个是PINAWPrediction Interval Normalized Average Width归一化平均区间宽度它把所有预测区间的宽度加起来除以测试集范围真实值最大最小值之差。这个指标衡量区间有多瘦区间越瘦说明预测越精确。但这和PICP是矛盾的区间拉得很宽覆盖率一定高但毫无信息量区间太窄覆盖率低等于没预测。所以不能单独看任何一个。第三个是CWCCoverage Width Criterion覆盖宽度准则把两者综合成一个分数。当PICP低于名义水平时CWC会施加一个大的惩罚项公式大致是[ \mathrm{CWC} \mathrm{PINAW} \cdot (1 \gamma(\mathrm{PICP}) \cdot e^{-\eta(\mathrm{PICP} - \mu)}) ]当PICP达到目标μ时指数项趋近于0CWC约等于PINAW未达到目标时CWC被指数项放大分数变差。这个设计思路和我的直觉完全一致先保证覆盖率再追求区间窄。在实际论文里这三个指标通常都会报出来。5.2 Matlab里怎么算function [picp, pinaW, cwc] calcIntervalMetrics(yTrue, yLower, yUpper, nominalLevel) % 覆盖率 inside (yTrue yLower) (yTrue yUpper); picp mean(inside) * 100; % 归一化区间宽度 dataRange max(yTrue) - min(yTrue); pinaW mean(yUpper - yLower) / dataRange; % CWCeta取10mu为名义水平 eta 10; mu nominalLevel * 100; gamma (picp mu); cwc pinaW * (1 gamma * exp(-eta * (picp - mu) / 100)); endeta的取值没有定论有文献用5、10、50都有。eta越大对覆盖率不达标的惩罚越重。如果你在写论文建议固定eta并说明取值依据如果你只是为了自己评估模型那重点看PICP是否达标以及PINAW是否在可接受范围内。5.3 一张图看懂结果预测区间可视化是整份代码最能体现成果的地方。Matlab画分位数带的经典代码是figure; fill([t; flipud(t)], [yLower; flipud(yUpper)], [0.8 0.9 1], ... EdgeColor, none, FaceAlpha, 0.5); hold on; plot(t, yMid, b-, LineWidth, 1.5); plot(t, yTrue, k--, LineWidth, 1); legend({90%区间, 中位数预测, 真实值}, Location, best); xlabel(时间); ylabel(预测值);flipud这里很关键fill函数画多边形时需要先按顺序给上边界正序、下边界倒序的顶点坐标少了这一行区间填充就会乱成一团线。我自己常用的画法是区间渐变多个分位数比如0.05/0.95、0.25/0.75用两个不同透明度的色带叠加形成类似热力分布图的视觉效果。但这是后期美化核心还是上面那个基础版本。5.4 一个真实案例的指标解读之前用这套代码跑了一组工业设备温度预测数据训练集6000个点测试集1000个点tau取[0.1, 0.5, 0.9]训练80轮。测试集结果PICP89.6%PINAW0.31名义覆盖率90%基本达标。如果我把tau换成[0.05, 0.5, 0.95]PICP会上升到96%左右但PINAW会从0.31涨到0.44区间宽度暴增40%信息量下降。这就是覆盖率和区间宽度之间的权衡你只能选业务需要的那个平衡点。还要提醒一个评估细节测试集的划分不能随机抽必须按时间顺序留出连续的一段。我曾经为了方便做随机划分结果因为时序相关性导致测试点信息泄漏区间覆盖率虚高到99%看起来完美换个真实场景立刻露馅。6. 从单点预测扩展到区间预测后的收尾心得最后分享一点我个人在实际项目里的体会。QRLSTM这套方案最大的价值不是精度比普通LSTM高多少而是它把预测的置信度这个原本要靠经验拍脑袋的东西变成了模型输出的一部分。做工程时一个带着可信区间的预测结果其决策价值远高于一个裸的数值。比如你给业务方一个SOC预测值对方会问你误差多大你给不出答案合作就僵住了现在你能直接说90%概率落在某个区间内这个沟通成本低了一个量级。如果你在2018版本上跑这套代码我建议务必明确两个版本的代码路径不要硬把新版代码塞给旧版环境跑报错会让你怀疑人生。如果条件允许尽量用2021a以上版本单模型多输出的方案无论是训练效率还是区间一致性都更优。这个方法后续还可以往双向LSTM、注意力机制、多步预测等方向扩展结构不需要大改把损失函数和输出层替换掉就行。电池SOC估计、微动面波反演、电力负荷预测这些领域我都在实际项目里验证过思路是通的。
返回列表