基于Matlab的GRU多输入单输出回归预测实战
2026/9/17 0:36:50 网站建设 项目流程

简介:面向计算机、电子信息工程、数学等专业学生及机器学习初学者的GRU多输入单输出回归预测完整Matlab方案,可直接用于课程设计、期末大作业与毕业设计。整套代码采用参数化编程思路,模型结构清晰、关键参数易调,并附带细致注释,便于在Matlab2023b及以上环境中运行、修改与迁移到自己的数据场景。压缩包共8个文件,以Matlab源程序(.m)、网络结构与结果可视化图(.png)、样本数据文件(.csv、.mat)以及结果输出文本(.txt)为主,整体大小为263KB,下载与部署较为轻量。实现中内置多维度误差统计,自动输出MAE、MAPE、MSE、RMSE与R2等回归评价指标,便于从不同角度验证模型表现。目前已有73人学习下载,适合希望快速掌握门控循环单元建模流程并完成回归预测实验的读者参考。

1. 为什么是GRU:多输入单输出回归预测的实战选择

拿到一份day.csv,几十行数据,七八个特征,要预测其中一个连续值。很多人第一反应是上BP神经网络或者XGBoost,但如果在Matlab里想快速出一个可解释、可调参的回归模型,GRU门控循环单元往往比LSTM更省心。GRU只有两个门,参数量大约是LSTM的75%,训练快且不容易过拟合。

这套源码用GRU.m和calc_error.m两个脚本,配day.csv数据集,完整实现多输入单输出回归预测。运行环境要求Matlab2023b及以上,输出MAE、MAPE、MSE、RMSE、R2五个指标。计算机、电子信息工程、数学专业的课程设计、期末大作业、毕业设计拿来就能跑。

这里“多输入单输出”不是把多个独立样本拼成一个序列,而是按时间窗口切分:用前几个时刻的多个特征预测下一时刻的目标值。这个前提后面所有代码都围绕它展开。

2. GRU门控机制与Matlab深度学习网络搭建

2.1 更新门和重置门如何减少参数

GRU的核心是两个门:更新门和重置门。更新门决定上一时刻的隐状态保留多少,重置门决定当前候选隐状态对历史信息的依赖程度。和LSTM相比,省略了独立的记忆单元和输出门,参数量更少。在数据量只有几千条的场景下,参数越少越不容易过拟合,调参空间也更大。

Matlab的深度学习工具箱从R2021b开始完整支持gruLayer,所以2023b跑起来没有任何障碍。比Python的PyTorch写法直观,不需要手动定义状态流,底层自动完成反向传播。使用gruLayer时,最关键的是OutputMode:回归任务用'last',只取最后一个时间步的隐状态;如果要做序列到序列,才用'sequence'。

2.2 用Matlab深度学习工具箱组装GRU网络

构建网络层的代码非常短。把多输入特征放在sequenceInputLayer里,中间接一个GRU层,再经过全连接和回归输出层即可。

% 参数化输入 numFeatures = 7; % 输入特征维度,需要根据day.csv实际列数修改 numHiddenUnits = 64; numResponses = 1; % 构建GRU回归网络 layers = [ sequenceInputLayer(numFeatures, 'Normalization', 'zscore') gruLayer(numHiddenUnits, 'OutputMode', 'last') fullyConnectedLayer(numResponses) regressionLayer ];

sequenceInputLayer的Normalization参数设置成zscore,省去手动对输入数据做标准化。gruLayer第二参数必须写成名值对,OutputMode设为last表示多对一预测。fullyConnectedLayer输出维度为1,对应单输出回归。regressionLayer计算均方误差损失。

再看训练选项。Matlab的trainingOptions直接支持Adam和验证集早停,不需要自己写循环。

options = trainingOptions('adam', ... 'MaxEpochs', 150, ... 'MiniBatchSize', 32, ... 'InitialLearnRate', 0.01, ... 'ValidationData', {XVal, YVal}, ... 'ValidationFrequency', 20, ... 'Plots', 'training-progress', ... 'Verbose', true);

这里ValidationData使用验证集而不是测试集,是为了在训练过程中看泛化曲线。ValidationFrequency表示每20个迭代评估一次验证损失。Plots可以打开训练进度图,对调试很有用。MiniBatchSize如果设得过大,小数据集上梯度更新不平稳。

2.3 训练选项参数表

参数名示例作用调节建议
MaxEpochs150完整遍历训练集的次数先设100~200,看验证损失再增减
MiniBatchSize32每次迭代使用的样本数小数据用16~64,过大容易过拟合
InitialLearnRate0.01初始学习率0.001~0.01之间优先尝试
ValidationFrequency20每隔多少轮计算验证损失通常设为总迭代数的1/10左右
ValidationPatience10验证损失连续几次不下降就早停设10~20,防过拟合

这个表里的参数几乎可以原样套用到其他GRU回归任务中。经常有人把ValidationData设成测试集,这是错的,因为会泄漏信息到训练过程,最后测试指标虚高。

3. day.csv加载、清洗与时间窗口序列构造

3.1 读取CSV并识别输入输出列

day.csv放在GRU.zip里,本身是一张带日期的表格。第一列可能是日期,后续多列是可用的输入特征,最后一列是预测目标。用readtable读取后,需要先观察数据格式。

data = readtable('day.csv'); disp(head(data, 5)); disp(data.Properties.VariableNames);

运行后确认每一列的数据类型。如果日期列是datetime类型,直接把它排除在特征之外;如果某些特征列里有NaN,用fillmissing线性插值填充。常见做法是:data = fillmissing(data, 'linear');,但要保证插值方向按行顺序,避免破坏时间序列的因果性。

确定特征列和目标列的通用写法:

varNames = data.Properties.VariableNames; featureCols = varNames(2:end-1); % 假设第一列是日期,最后一列是目标 targetCol = varNames{end}; features = table2array(data(:, featureCols)); target = table2array(data(:, targetCol)); % 计算特征维度 numFeatures = size(features, 2);

这里没有硬编码列号,后面换数据集时只要保持“第一列日期、最后一列目标”的结构,就可以直接跑。如果day.csv没有日期列,只要把featureCols的范围改成1:end-1就行。

3.2 时间窗口切分函数

GRU要求输入是序列,所以要把原始表格转成“样本×时间步×特征”的三维数组。下面这个函数是这套源码的灵魂:

function [X, Y] = makeSequences(features, target, steps) n = size(features, 1); numFeatures = size(features, 2); X = zeros(n - steps, steps, numFeatures); Y = zeros(n - steps, 1); for i = 1:n - steps X(i, :, :) = features(i:i+steps-1, :); Y(i, :) = target(i+steps, :); end end

调用方式[X, Y] = makeSequences(features, target, 7)表示使用前7天的全部特征预测第8天的目标值。X的第一维是样本数,第二维是时间步,第三维是特征数,这正是trainNetwork对序列输入的要求格式。Y的维度是样本数×1,对应单输出。steps如果太大,样本数量会变少;太小则学不到周期规律,一般取3、7、14试验。

参数含义建议范围
steps时间窗口长度3、7、14
numFeatures输入特征维度day.csv实际列数
target目标列向量单输出列

3.3 训练集/验证集划分和归一化

时间序列回归不能直接随机抽train-test split,否则会破坏时间顺序。按时间顺序取前80%作为训练,剩余20%作为测试,再从训练集中取最后10%作为验证集。

numSamples = size(X, 1); trainEnd = floor(numSamples * 0.8); valEnd = trainEnd + floor(numSamples * 0.1); Xtr = X(1:trainEnd, :, :); Ytr = Y(1:trainEnd, :); XVal = X(trainEnd+1:valEnd, :, :); YVal = Y(trainEnd+1:valEnd, :); XTest = X(valEnd+1:end, :, :); YTest = Y(valEnd+1:end, :);

验证集夹在训练集和测试集之间,用来早停。归一化最好用zscore,但需要先计算训练集的均值和标准差,再应用到验证集和测试集,避免信息泄漏。如果用了sequenceInputLayer的'Normalization','zscore',特征部分可以不手动标准化;但目标值Y还需要手动处理,因为网络输出层没有做逆变换。

muY = mean(Ytr); stdY = std(Ytr); YtrNorm = (Ytr - muY) / stdY; % 预测完成后要反归一化 Ypred = Ypred * stdY + muY;

很多人只归一化特征不归一化目标,会让回归损失在数量级上失衡。目标值范围很大时,训练初期loss会很大,导致学习率不好选。

4. 训练GRU并输出多指标评价

4.1 从GRU.m看训练主流程

GRU.m把整个流程串起来。核心顺序是读表、构序列、划分、训练、预测、反归一化、计算误差、画图。在GRU.zip里看到的1.png、2.png、3.png就是训练过程图、预测对比图、误差分布图。主流程代码框架如下:

% GRU.m 核心流程 data = readtable('day.csv'); mat = table2array(data(:, 2:end)); % 去掉日期列 features = mat(:, 1:end-1); target = mat(:, end); [X, Y] = makeSequences(features, target, 7); % 划分训练/验证/测试 % ... 同第3章,省略 % 定义网络层和训练选项 % layers = [...] % options = trainingOptions(...) % 训练 net = trainNetwork(Xtr, YtrNorm, layers, options); % 测试预测 Ypred = predict(net, XTest); % 反归一化 Ypred = Ypred * stdY + muY; % 计算并保存结果 [mae, mape, mse, rmse, r2] = calc_error(YTest, Ypred);

trainNetwork第一个参数是三维数组,第二个是归一化后的目标向量。predict输出维度是样本数×1。注意预测之前不要对测试集目标做任何变换,因为YTest是原始值。

4.2 calc_error.m中的评价指标计算

calc_error.m是整个源码里最值得抄的段落。它用五句话算出五个指标,公式和顺序都照顾到了。

function [mae, mape, mse, rmse, r2] = calc_error(ytest, ypred) e = ytest - ypred; mae = mean(abs(e)); mape = mean(abs(e ./ ytest)) * 100; mse = mean(e .^ 2); rmse = sqrt(mse); ssres = sum(e .^ 2); sstot = sum((ytest - mean(ytest)) .^ 2); r2 = 1 - ssres / sstot; end

MAE是绝对误差的均值,单位与原始数据一致。MAPE用百分比表示,适合向业务方汇报。MSE给大误差更高惩罚,RMSE是MSE的开平方,恢复量纲后更容易解释。R2等于1表示完美拟合,0表示模型等于直接用均值,负值说明模型比均值基线还差。

调用时保持顺序一致:

[mae, mape, mse, rmse, r2] = calc_error(YTest, Ypred); fprintf('MAE=%.4f MAPE=%.2f%% MSE=%.4f RMSE=%.4f R2=%.4f\n', ... mae, mape, mse, rmse, r2);

4.3 指标解读与多输入单输出常见误区

结果.txt里会有一行五个指标输出。这里给出参考判读标准,注意不是绝对标准:

指标取值范围好模型参考说明
MAE0~∞越小越好平均绝对误差,看量纲
MAPE0~100%<10% 较好相对误差,对接近0的目标敏感
MSE0~∞越小越好大误差惩罚强
RMSE0~∞与MAE接近则稳定比MAE大说明存在离群误差
R2-∞~1>0.8可用反映模型解释方差的比例

经常有人把MAPE计算成mean(abs(ypred - ytest)) ./ ytest的两倍误差,其实关键在于除的是真实值。另外,如果ytest里有0或接近0的值,MAPE会爆炸,这时建议改用SMAPE或者直接去掉零值样本。在多输入单输出场景里,还要检查Ypred的排序是否和YTest对应:时间序列预测一旦做过随机打乱,两条曲线就错位了,R2会变成负的。上面的代码严格按时间顺序划分,就不会出这个问题。

5. GRU超参数调优与结果诊断

5.1 先调学习率还是先调隐含单元数

GRU调参顺序不是从隐含单元开始,而是先从学习率入手。初始化学习率过大,loss曲线振荡;过小,训练几十个epoch还在原地不动。我的做法是固定GRU隐含单元数到32,先试0.001、0.005、0.01三档,用验证集loss看哪个最稳。

然后再调隐含单元数。GRU状态维度太小,欠拟合;太大,在小样本上会记忆噪声。电力和气象数据集通常20~100之间就够用。补充一句,如果发现验证集loss下降但测试集指标差,问题出在归一化或数据泄漏,而不是隐含单元数。

优先级参数建议值观察点
1InitialLearnRate0.001~0.01loss曲线是否振荡
2numHiddenUnits20~100R2是否达到瓶颈
3steps3~21样本数量变化

5.2 时间步长对预测精度的影响

时间窗口steps是序列模型特有的超参数。之前用7天的窗口预测第8天,但day.csv如果带有明显周趋势,7会很好;如果是月度周期,可能要试14或者30。窗口增大,样本数减少,所以不是越大越好。用以下代码快速扫描不同窗口长度下的R2:

stepList = [3 7 14 21]; for s = stepList [X, Y] = makeSequences(features, target, s); % 这里省略划分、归一化和训练,直接复用主循环 fprintf('steps=%d R2=%.4f\n', s, r2); end

5.3 用验证集早停防止过拟合

Matlab的trainingOptions里ValidationPatience就是为这个设计的。当验证损失连续多次不下降时,训练自动停止,返回当前最优模型。设置代码如下:

options = trainingOptions('adam', ... 'MaxEpochs', 300, ... 'ValidationData', {XVal, YValNorm}, ... 'ValidationFrequency', 20, ... 'ValidationPatience', 15, ... 'OutputNetwork', 'best-validation');

OutputNetwork设为best-validation,训练结束后取验证损失最小的网络,而不是最后一个epoch的网络。这个参数是2023b环境下的推荐设置。很多入门代码忘了这个细节,导致后续预测用的是过拟合后的权重。

6. 把GRU封装成可复用的Matlab函数

6.1 参数化训练入口

GRU.m是一次性脚本,验证好参数后可以封装成函数,输入数据路径和超参数,返回指标。这样做的好处是换一份csv,不需要改训练代码。

function metrics = trainGruPredictor(dataFile, steps, numHidden, lr) data = readtable(dataFile); mat = table2array(data(:, 2:end)); [X, Y] = makeSequences(mat(:, 1:end-1), mat(:, end), steps); % 按顺序划分、归一化 % 组装网络 % trainNetwork % 返回 metrics 结构体 end

这样后面批量调参时,只要写一个for循环调用trainGruPredictor即可。比如尝试3组窗口、3组隐含单元,一次跑完。

6.2 用结构体批量测试超参数

configs(1) = struct('steps', 7, 'hidden', 32, 'lr', 0.01); configs(2) = struct('steps', 14, 'hidden', 64, 'lr', 0.005); for i = 1:numel(configs) m = trainGruPredictor('day.csv', configs(i).steps, ... configs(i).hidden, configs(i).lr); fprintf('config %d: R2=%.4f RMSE=%.4f\n', i, m.r2, m.rmse); end

结构体数组比cell数组直观,字段名就能看出参数含义。参数化编程是这套源码的一个优点,改steps、hidden、lr都在入口处完成,不需要去中间代码里找魔法数字。未来即使Codex能像执行Python一样操作Matlab任务,参数化的函数接口仍然是批量实验的基础。把这个函数存成trainGruPredictor.m,连同makeSequences.m和calc_error.m,就是一套可移植的GRU回归工具链。

本文还有配套的精品资源,点击获取

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询