news 2026/9/6 14:08:49

CNN-GRU回归预测与SHAP可解释性分析完整实践

作者头像

张小明

前端开发工程师

1.2k 24
文章封面图
CNN-GRU回归预测与SHAP可解释性分析完整实践

之前在做回归预测任务时,最难受的点往往不是模型效果上不来,而是模型给出一个预测值之后,很难向业务方解释清楚“为什么是这个值”。为了解决这个问题,我采用了CNN-GRU 混合模型作为预测主体,并结合SHAP 值分析每个特征对预测结果的贡献。网上关于 CNN-GRU 做分类或回归的例子很多,但不少文章只贴代码、不解释维度变化,也没有把 SHAP 解释的完整流程整合进去。这篇文章把我实际使用的代码、训练流程和可解释性分析整理成一套可直接运行的教程,希望对正在做回归预测的你有所帮助。

本文覆盖以下内容:

  • CNN-GRU 混合模型的核心原理;
  • 回归预测数据的滑窗构建与归一化方法;
  • 使用 PyTorch 搭建 CNN-GRU 回归模型;
  • 模型训练、评估指标解读;
  • SHAP 值的计算与可视化分析;
  • 常见报错和工程化建议。

1. 背景与核心概念

1.1 CNN-GRU 是什么

CNN-GRU 是由卷积神经网络(CNN)和门控循环单元(GRU)组合而成的混合网络结构。

  • CNN(Convolutional Neural Network):善于提取局部特征。在一维时间序列数据中,卷积核可以捕捉相邻时间步之间的局部模式,比如短期的趋势变化、周期性波动等。
  • GRU(Gated Recurrent Unit):是 LSTM 的简化变体,通过更新门和重置门控制信息的保留与遗忘。GRU 适合建模长距离依赖关系,同时参数量比 LSTM 更少,训练效率更高。

将两者串联是一种常见做法:先用 CNN 从原始输入中提取局部特征,再把 CNN 的输出按照时间顺序送入 GRU,让 GRU 继续捕捉时间维度上的长期依赖。

1.2 为什么用 CNN-GRU 做回归预测

很多真实场景中的回归预测面对的是多变量时间序列数据,比如:

  • 根据过去 24 小时的多维环境数据预测未来气温;
  • 根据历史交易数据预测下一时段销量;
  • 根据设备传感器数据预测剩余寿命;
  • 根据历史负荷数据预测未来用电量。

这些数据通常同时具有“局部相关性”和“长期依赖性”。如果只用 CNN,模型感受野有限,难以建模长期依赖;如果只用 GRU,序列较长时训练速度更慢,而且对局部特征的提取不够直接。CNN-GRU 先做局部特征抽象,再做时序建模,在很多回归任务上效果优于单一模型。

1.3 为什么引入 SHAP 值

回归预测模型光有精度还不够。当我们想判断“哪个特征对预测结果影响最大”,或者“某条预测为什么偏高”时,就需要对模型做可解释性分析。

SHAP(SHapley Additive exPlanations)是一种基于博弈论 Shapley 值的模型解释方法。它的核心思想是:每个特征对预测结果的贡献可以量化,且所有特征的贡献之和等于模型预测值相对于基线预测值的偏离程度。

在复杂深度学习模型中,SHAP 可以告诉我们:

  • 哪些特征对预测结果影响最大;
  • 样本级别上,某个特征取值是拉高了预测值还是拉低了预测值;
  • 特征与预测结果之间是正相关还是负相关。

所以,CNN-GRU 负责把预测精度做到位,SHAP 负责把预测结果解释清楚,两者结合是一条很实用的工程路径。

2. 环境准备与项目结构

2.1 运行环境说明

下面的代码以 Python 3.9+ 为例,需要安装以下依赖。具体版本请根据你的实际环境调整,本文重点演示实现思路:

pip install numpy pandas matplotlib scikit-learn torch shap

如果你使用 GPU 版本的 PyTorch 训练,需要提前安装对应 CUDA 版本的 torch;如果只是学习演示,CPU 版本也能跑通。

2.2 项目结构

建议按照下面的目录组织代码:

cnn_gru_regression/ ├── main.py # 完整训练与评估流程 ├── model.py # CNN-GRU 模型定义 ├── data_utils.py # 数据生成与滑窗处理 ├── explain.py # SHAP 可解释性分析 └── requirements.txt # 依赖清单

如果你希望代码更集中,也可以把全部内容写在一个脚本里。为了便于阅读,本文按照功能拆分讲解,最后你可以把代码汇总到一个文件中运行。

3. 回归预测数据准备

3.1 使用模拟数据快速验证

我们先写一个模拟数据生成函数。这个函数会生成 4 个与目标值存在线性关系的时间序列特征,并加入少量噪声。

# 文件路径:data_utils.py import numpy as np import pandas as pd def generate_demo_data(n_samples=1500): """生成多变量回归预测模拟数据。 参数: n_samples: 样本点数量 返回: pandas.DataFrame,包含 4 个特征列和 1 个目标列 """ t = np.arange(n_samples) # 构造4个特征,每个特征有不同周期和噪声 feature1 = np.sin(2 * np.pi * t / 50) + 0.1 * np.random.randn(n_samples) feature2 = np.cos(2 * np.pi * t / 30) + 0.1 * np.random.randn(n_samples) feature3 = 0.02 * t + 0.2 * np.random.randn(n_samples) feature4 = 0.5 * np.sin(2 * np.pi * t / 7) + 0.2 * np.random.randn(n_samples) # 目标值与特征之间保持线性组合,方便后续用 SHAP 验证解释效果 target = ( 2.5 * feature1 + 1.5 * feature2 + 0.8 * feature3 - 1.2 * feature4 + 0.3 * np.random.randn(n_samples) ) df = pd.DataFrame({ "feature1": feature1, "feature2": feature2, "feature3": feature3, "feature4": feature4, "target": target, }) return df

这个方法的好处是:数据可以自己生成,代码复制后能直接运行。如果你有自己的数据集,只需要把“读入 DataFrame,包含特征列和目标列”这一步替换掉即可。

3.2 滑窗样本构建

回归预测里,我们通常不能直接用单条样本做预测,而是用过去一段时间的特征序列预测下一个时间点的值。这个“过去一段时间”就叫做时间窗口,对应的处理方式叫“滑窗”或“滚动窗口”。

# 文件路径:data_utils.py def create_sequences(data, feature_cols, target_col, window_size=24): """构建滑窗样本。 参数: data: DataFrame,包含特征列和目标列 feature_cols: 特征列名列表 target_col: 目标列名 window_size: 时间窗口长度 返回: X: shape 为 (样本数, window_size, 特征数) 的数组 y: shape 为 (样本数,) 的数组 """ X, y = [], [] for i in range(len(data) - window_size): X.append(data[feature_cols].iloc[i: i + window_size].values) y.append(data[target_col].iloc[i + window_size]) return np.array(X), np.array(y)

这里需要注意:窗口长度window_size决定了模型每次能看到多长的历史信息。窗口太短会丢失长期依赖;窗口太长会增加计算量,也可能会引入过多噪声。一般可以先通过实验对比不同窗口大小,再确定适合业务场景的值。

3.3 时间顺序切分与归一化

时序预测和普通机器学习不一样,不能随机打乱数据再切分,否则会造成“未来信息泄漏”。也就是说,如果用后面的数据去训练模型、预测前面的数据,评估结果会虚高。这里我们按时间顺序,前 80% 作为训练集,后 20% 作为测试集。

def load_train_test_data(window_size=24, test_ratio=0.2): """生成数据并切分为训练集和测试集(按时间顺序切分)。""" feature_cols = ["feature1", "feature2", "feature3", "feature4"] target_col = "target" data = generate_demo_data(1500) split_idx = int(len(data) * (1 - test_ratio)) train_df = data.iloc[:split_idx] test_df = data.iloc[split_idx:] # 分别对训练集和测试集做归一化 # 注意:归一化参数只能用训练集 fit,测试集直接 transform from sklearn.preprocessing import MinMaxScaler scaler_X = MinMaxScaler() scaler_y = MinMaxScaler() train_X_scaled = scaler_X.fit_transform(train_df[feature_cols]) train_y_scaled = scaler_y.fit_transform(train_df[[target_col]]) test_X_scaled = scaler_X.transform(test_df[feature_cols]) test_y_scaled = scaler_y.transform(test_df[[target_col]]) train_df_scaled = pd.DataFrame(train_X_scaled, columns=feature_cols) train_df_scaled[target_col] = train_y_scaled test_df_scaled = pd.DataFrame(test_X_scaled, columns=feature_cols) test_df_scaled[target_col] = test_y_scaled # 构建滑窗样本 X_train, y_train = create_sequences(train_df_scaled, feature_cols, target_col, window_size) X_test, y_test = create_sequences(test_df_scaled, feature_cols, target_col, window_size) return X_train, y_train, X_test, y_test, scaler_y

关于归一化,有两个容易踩的坑:

  1. 整个数据集只 fit 一次MinMaxScaler,然后在所有数据上 transform,这在时序场景里是不可取的。因为训练集之外的“未来数据”参与了归一化参数计算,相当于把未来的分布信息提前暴露给了模型。
  2. 目标变量y也需要归一化。深度学习模型直接回归一个量纲较大的数值时,损失值可能很大,训练不稳定。这里我们把目标值归一化到[0,1]区间,训练结束后再把预测结果反归一化。

4. 构建 CNN-GRU 回归预测模型

4.1 模型结构定义

下面是模型的完整定义。

# 文件路径:model.py import torch import torch.nn as nn class CNNGRU(nn.Module): def __init__(self, n_features, hidden_size=64, num_layers=1, dropout=0.1, output_size=1): super(CNNGRU, self).__init__() # 1D 卷积层:输入通道为特征数,输出通道为 32 self.conv1 = nn.Conv1d( in_channels=n_features, out_channels=32, kernel_size=3, padding=1 ) self.relu = nn.ReLU() self.pool = nn.MaxPool1d(kernel_size=2) # GRU 层:输入大小是 CNN 输出通道数 self.gru = nn.GRU( input_size=32, hidden_size=hidden_size, num_layers=num_layers, batch_first=True, dropout=dropout if num_layers > 1 else 0 ) # 全连接输出层 self.fc = nn.Linear(hidden_size, output_size) def forward(self, x): # 输入 x 形状: (batch_size, seq_len, n_features) # CNN 期望输入形状是 (batch_size, channels, seq_len) x = x.permute(0, 2, 1) # 经过卷积、激活、池化 x = self.conv1(x) # (batch_size, 32, seq_len) x = self.relu(x) x = self.pool(x) # (batch_size, 32, seq_len // 2) # 转回 GRU 需要的形状: (batch_size, seq_len', input_size) x = x.permute(0, 2, 1) # GRU 前向传播,取最后一个时间步输出 out, _ = self.gru(x) # out: (batch_size, seq_len', hidden_size) out = out[:, -1, :] # 取最后一个时间步 # 全连接输出 out = self.fc(out) # (batch_size, 1) return out

4.2 维度变化分析

很多初学者第一次看这段代码会卡在维度变化上,这里梳理一下:

操作输入形状输出形状
原始输入(batch, seq_len, n_features)(batch, seq_len, n_features)
permute 转置(batch, seq_len, n_features)(batch, n_features, seq_len)
Conv1d(batch, n_features, seq_len)(batch, 32, seq_len)
ReLU(batch, 32, seq_len)(batch, 32, seq_len)
MaxPool1d(batch, 32, seq_len)(batch, 32, seq_len // 2)
permute 转置(batch, 32, seq_len // 2)(batch, seq_len // 2, 32)
GRU(batch, seq_len // 2, 32)(batch, seq_len // 2, hidden_size)
取最后一个时间步(batch, seq_len // 2, hidden_size)(batch, hidden_size)
Linear(batch, hidden_size)(batch, 1)

需要注意,MaxPool1d 的kernel_size=2会让序列长度减半。如果seq_len是奇数,比如window_size=25,池化后长度会变成12,对应关系可能变得不直观,因此建议优先使用偶数窗口长度。

4.3 为什么先 CNN 再 GRU

这里简单解释一下设计动机:

  • CNN 的卷积核对局部模式敏感,可以自动提取“相邻几个时间步之间的组合特征”;
  • 经过 MaxPooling 后,序列长度缩短,计算量降低,也起到一定的特征压缩作用;
  • GRU 接收 CNN 提取的高层特征序列,继续建模长期依赖;
  • 最后用全连接层把 GRU 最后一个时间步的隐藏状态映射为标量预测值。

如果任务本身序列较短、特征较少,也可以去掉 MaxPooling,只保留卷积和 GRU。示例代码保留池化是为了展示一种更通用的结构。

5. 训练与回归评估

5.1 数据集封装与数据加载器

我们使用 PyTorch 的TensorDatasetDataLoader来管理数据。

from torch.utils.data import TensorDataset, DataLoader import torch X_train, y_train, X_test, y_test, scaler_y = load_train_test_data(window_size=24) # 转换为 PyTorch Tensor X_train_t = torch.FloatTensor(X_train) y_train_t = torch.FloatTensor(y_train).view(-1, 1) X_test_t = torch.FloatTensor(X_test) y_test_t = torch.FloatTensor(y_test).view(-1, 1) train_dataset = TensorDataset(X_train_t, y_train_t) test_dataset = TensorDataset(X_test_t, y_test_t) train_loader = DataLoader(train_dataset, batch_size=64, shuffle=True) test_loader = DataLoader(test_dataset, batch_size=64, shuffle=False)

这里有一个细节:训练数据加载时shuffle=True,但是测试数据shuffle=False。因为训练时我们希望每个 batch 的样本尽量随机,帮助模型稳定收敛;测试时不需要打乱顺序,方便后续计算指标和可视化。

5.2 模型初始化与训练循环

import torch.nn as nn import torch.optim as optim # 固定随机种子,保证结果可复现 torch.manual_seed(42) model = CNNGRU(n_features=X_train.shape[2], hidden_size=64) criterion = nn.MSELoss() optimizer = optim.Adam(model.parameters(), lr=0.001) epochs = 30 for epoch in range(epochs): model.train() train_loss = 0.0 for X_batch, y_batch in train_loader: optimizer.zero_grad() y_pred = model(X_batch) loss = criterion(y_pred, y_batch) loss.backward() optimizer.step() train_loss += loss.item() * X_batch.size(0) avg_train_loss = train_loss / len(train_dataset) # 每个 epoch 后评估一次测试集 model.eval() test_loss = 0.0 with torch.no_grad(): for X_batch, y_batch in test_loader: y_pred = model(X_batch) loss = criterion(y_pred, y_batch) test_loss += loss.item() * X_batch.size(0) avg_test_loss = test_loss / len(test_dataset) if (epoch + 1) % 5 == 0: print(f"Epoch {epoch + 1}/{epochs}, Train Loss: {avg_train_loss:.6f}, Test Loss: {avg_test_loss:.6f}")

训练过程中,有两个环境非常重要:

  • model.train()model.eval():训练模式会启用 Dropout 等随机操作,而评估模式会固定这些操作,保证测试输出稳定。
  • with torch.no_grad():推理阶段不需要计算梯度,既省内存又加快速度。

5.3 回归评估指标

回归预测常用三个指标:MSE、MAE、R2。

from sklearn.metrics import mean_squared_error, mean_absolute_error, r2_score model.eval() with torch.no_grad(): y_pred_all = model(X_test_t).numpy().flatten() y_test_all = y_test_t.numpy().flatten() # 反归一化,恢复真实尺度 y_pred_inv = scaler_y.inverse_transform(y_pred_all.reshape(-1, 1)).flatten() y_test_inv = scaler_y.inverse_transform(y_test_all.reshape(-1, 1)).flatten() mse = mean_squared_error(y_test_inv, y_pred_inv) mae = mean_absolute_error(y_test_inv, y_pred_inv) r2 = r2_score(y_test_inv, y_pred_inv) print(f"MSE: {mse:.4f}") print(f"MAE: {mae:.4f}") print(f"R2: {r2:.4f}")

各指标含义:

  • MSE(均方误差):预测值与真实值差值的平方的平均值。MSE 对较大误差更敏感,适合关注极端偏差的场景。
  • MAE(平均绝对误差):预测值与真实值差值的绝对值的平均值。它直接反映平均误差大小,单位与真实值一致。
  • R2(决定系数):表示模型解释了目标变量多少方差。R2 越接近 1,说明模型拟合效果越好;R2 为 0 说明模型与直接预测平均值差不多;R2 为负数说明模型效果比平均值预测还差。

反归一化这一步容易被忽略。因为训练时对y做了MinMaxScaler,所以模型输出的是归一化后的值。要计算真实尺度下的误差指标,必须先调用scaler_y.inverse_transform还原。

5.4 可视化预测曲线

为了更直观地观察预测效果,可以把测试集上的真实值和预测值画成曲线。

import matplotlib.pyplot as plt plt.figure(figsize=(12, 4)) plt.plot(y_test_inv[:200], label="True", linewidth=2) plt.plot(y_pred_inv[:200], label="Pred", linewidth=2) plt.legend() plt.title("CNN-GRU Regression Prediction Results") plt.xlabel("Sample Index") plt.ylabel("Target Value") plt.savefig("prediction_result.png", dpi=150) plt.show()

如果前 200 个测试点上两条曲线整体趋势一致,说明模型已经学到了基本的时序规律。

6. 使用 SHAP 解释模型

6.1 SHAP 原理简介

SHAP 的核心思想是 Shapley 值。它把模型预测值拆解为“基线值 + 每个特征的贡献值”。基线值通常是训练集上预测值的平均值。

对于一条样本,假设模型预测值为f(x),基线值为E[f(x)],那么有:

f(x) = E[f(x)] + sum(每个特征的SHAP值)

当一个特征的 SHAP 值为正,表示该特征把预测值向上推动;SHAP 值为负,表示把预测值向下拉低。SHAP 值的绝对值越大,说明该特征对这条样本的影响越强。

6.2 DeepExplainer 使用方法

对于 PyTorch 模型,SHAP 库提供了DeepExplainer。它适用于深度学习模型,计算效率比KernelExplainer更高。

# 文件路径:explain.py import shap import torch # 将模型切换到评估模式 model.eval() # 选择一部分测试样本作为背景数据 background = X_test_t[:100] # 这里取少量测试样本做解释,避免计算时间过长 X_explain = X_test_t[:10] # 创建 DeepExplainer explainer = shap.DeepExplainer(model, background) # 计算 SHAP 值 shap_values = explainer.shap_values(X_explain)

注意两点:

  1. background是背景样本,主要用来估计基线值。数量不一定要很多,50 到 100 条通常就够用,但需要覆盖训练集中比较典型的特征分布。
  2. shap_valuesDeepExplainer中通常返回一个列表。因为模型输出维度是 1,所以我们要看的是shap_values[0]

shap_values[0]的形状与输入数据一致,也就是:

(样本数, 时间步数, 特征数)

这意味着 SHAP 给出的不仅是“哪个原始特征重要”,还包括“哪个时间步上的哪个特征重要”。这比普通表格数据回归的解释细节更丰富。

6.3 特征重要性可视化

如果我们只关心原始特征的整体重要性,可以把所有时间步的 SHAP 绝对值求和。

import numpy as np # shap_values[0] 形状: (10, window_size, n_features) shap_values_0 = np.array(shap_values[0]) # 对所有测试样本和时间步求和,得到每个原始特征的贡献 feature_names = ["feature1", "feature2", "feature3", "feature4"] importance = np.abs(shap_values_0).sum(axis=(0, 1)) # (n_features,) for name, imp in zip(feature_names, importance): print(f"{name}: {imp:.4f}") # 排序后可视化 sorted_idx = np.argsort(importance)[::-1] plt.figure(figsize=(8, 4)) plt.bar([feature_names[i] for i in sorted_idx], importance[sorted_idx]) plt.title("Feature Importance by SHAP") plt.xlabel("Feature") plt.ylabel("Mean |SHAP|") plt.tight_layout() plt.savefig("shap_feature_importance.png", dpi=150) plt.show()

在这个模拟数据里,理论上feature1对目标值影响最大,因为它的系数是 2.5。如果 SHAP 结果也显示feature1的重要性最高,说明模型学到的关系和数据生成逻辑基本一致。

6.4 蜜蜂图与依赖图

SHAP 库自带的summary_plot可以画出“蜜蜂图”,既能反映特征重要性,也能反映特征取值与 SHAP 值的正负关系。

由于我们的输入是三维的滑窗数据,直接传入原始X_explain会让summary_plot难以解释。为了方便展示,我们可以把三维数据展平成二维,并生成对应的扁平特征名。

# 将 (10, window_size, n_features) 展平为 (10, window_size * n_features) X_flat = X_explain.numpy().reshape(X_explain.shape[0], -1) # 生成扁平特征名 flat_names = [] for t in range(X_explain.shape[1]): for f in feature_names: flat_names.append(f"t{t}_{f}") shap_values_flat = shap_values_0.reshape(shap_values_0.shape[0], -1) shap.summary_plot(shap_values_flat, X_flat, feature_names=flat_names, show=False) plt.tight_layout() plt.savefig("shap_summary_plot.png", dpi=150) plt.show()

蜜蜂图怎么看:

  • 横轴是 SHAP 值。某个点落在正半轴,说明该样本在这个特征上的取值让预测值升高;落在负半轴说明降低。
  • 点的颜色表示该特征在当前样本中的实际大小,颜色越红表示数值越大,颜色越蓝表示数值越小。
  • 特征按重要性从上到下排列,越靠上越重要。

如果你只关心第一个时间步的特征,也可以单独取出对应切片:

# 只看第一个时间步 shap_summary_first_timestep = shap_values_0[:, 0, :] X_first_timestep = X_explain.numpy()[:, 0, :] shap.summary_plot(shap_summary_first_timestep, X_first_timestep, feature_names=feature_names, show=False) plt.tight_layout() plt.savefig("shap_summary_first_timestep.png", dpi=150) plt.show()

这种方式适合观察“最近一个时间步”中哪些特征对预测影响最大。实际应用中,你可以根据业务需求选择查看某个时间步或全部时间步。

6.5 为什么 SHAP 值要配合业务解读

SHAP 只能解释“模型学到了什么”,不能保证“真实的因果关系就是如此”。比如,某个特征和预测值高度相关,但它可能只是间接关联,而不是直接原因。所以做技术解释时,要把 SHAP 结果当作模型行为的证据之一,而不是因果结论。

7. 常见问题与排查思路

在实际运行过程中,经常遇到下面几个问题。

问题现象常见原因解决思路
模型训练 loss 不下降数据未归一化,或学习率过大/过小检查特征和目标值是否做了归一化;尝试 lr=0.001 或 0.0001
测试集 R2 很低,甚至为负训练集和测试集数据分布差异过大,或滑窗窗口太小检查切分方式;增大 window_size;检查数据是否存在强非平稳性
Conv1d 维度不匹配输入形状不是(batch, channels, seq_len)在进入卷积前用x.permute(0, 2, 1)调整维度
MaxPool1d 后序列长度异常window_size为奇数调整窗口为偶数,或不使用池化层
SHAP 计算非常慢背景数据过多,或者解释样本数量过大减小 background 数量,比如 50 条;减小 X_explain 数量
DeepExplainer 报错模型不在 eval 模式,或数据类型不是 FloatTensor调用model.eval();确认输入 tensor 使用torch.float32
预测值始终接近某个常数模型欠拟合,或者目标值分布非常集中增加训练轮数;调整隐藏层维度;检查数据生成逻辑

下面单独讲一个高频问题:训练时 loss 很低,测试时 loss 很高。这在回归预测中通常表示过拟合。常见解决办法是:

  • 增加训练数据量;
  • 减小模型复杂度,比如减少 GRU 隐藏层维度;
  • 加入 Dropout,并在模型定义时对 GRU 多层场景设置dropout
  • 引入早停机制,当测试 loss 连续若干轮不再下降时停止训练。

8. 最佳实践与工程建议

8.1 时间顺序切分,避免数据泄漏

处理时序数据时,不能直接使用train_test_split(random_state=42)随机打乱。应该按照时间顺序划分训练集、验证集和测试集,并且验证集和测试集都必须是训练集之后的时间段。这样才能真实模拟模型在“未来”数据上的表现。

8.2 归一化参数只能来自训练集

标准化的核心原则是:scaler只能fit在训练集上,然后transform训练集、验证集和测试集。如果对整个数据集一起fit,测试集的信息就会间接进入训练过程,导致评估结果偏乐观。

8.3 固定随机种子

深度学习模型带有随机性,比如权重初始化、数据加载顺序等。在实验阶段,建议统一设置随机种子:

import random import numpy as np import torch random.seed(42) np.random.seed(42) torch.manual_seed(42)

如果使用 CUDA,还需要设置:

if torch.cuda.is_available(): torch.cuda.manual_seed_all(42)

这样才能保证多次实验的结果可比较。

8.4 模型保存与加载

训练完成后,可以用torch.save保存模型参数:

torch.save(model.state_dict(), "cnn_gru_model.pth")

加载时先实例化同一个模型,再load_state_dict

model = CNNGRU(n_features=X_train.shape[2], hidden_size=64) model.load_state_dict(torch.load("cnn_gru_model.pth")) model.eval()

注意:这里保存的是模型参数,不包含模型结构。如果你换了一台机器运行,需要保证model.py中的CNNGRU类定义一致。

8.5 SHAP 解释的工程化落地

在业务系统中,如果每次预测都要重新计算 SHAP,开销会比较大。你可以把测试集上的 SHAP 特征重要性结果保存下来,作为模型的解释报告;也可以在模型服务层预留一个“解释接口”,只在需要分析特定样例时才调用 SHAP。

8.6 超参数调整建议

CNN-GRU 中比较关键的超参数包括:

  • 卷积核大小:用于控制局部感受野,一般取 3、5、7;
  • 卷积输出通道数:控制特征抽象能力,常见取值 32、64;
  • GRU 隐藏层维度:控制时序记忆容量,常见取值 32、64、128;
  • 学习率:一般从 0.001 开始,训练不收敛时降低到 0.0005 或 0.0001;
  • Batch Size:根据显存大小和数据量调整,常见取值 32、64、128。

建议先用小规模的模型和少量数据跑通流程,再逐步扩大参数。这样能更快定位问题。

9. 总结与学习路线

这篇文章围绕CNN-GRU 回归预测整理了一套完整的代码实践:

  • 使用 CNN 提取局部特征,使用 GRU 建模时序依赖;
  • 使用滑窗和归一化处理回归预测数据;
  • 自定义CNNGRU模型,完成训练和评估;
  • 使用 MSE、MAE、R2 三个指标评估效果;
  • 使用 SHAP 值的DeepExplainer计算特征贡献,并绘制特征重要性图和蜜蜂图。

如果你还想继续深挖,可以从以下几个方向入手:

  • 尝试用Seq2Seq + Attention结构做多步回归预测;
  • 在 SHAP 的基础上,加入dependence_plot依赖图,分析单个特征与预测结果的关系;
  • 对比 CNN-LSTM 与 CNN-GRU 在当前数据上的效果差异;
  • 在真实业务数据上测试不同窗口长度对预测效果的影响;
  • 将模型封装成 Flask 或 FastAPI 服务,实现在线预测和解释报告输出。

希望这篇文章能帮你跑通 CNN-GRU 回归预测的完整链路,也让你在向业务方解释模型时不再无从下手。你可以把代码保存下来,先在自己的数据集上试一遍,再根据实际数据分布调整窗口大小和模型参数。如果遇到本地环境问题,也欢迎对照第 7 节的排查表格逐步检查。

版权声明: 本文来自互联网用户投稿,该文观点仅代表作者本人,不代表本站立场。本站仅提供信息存储空间服务,不拥有所有权,不承担相关法律责任。如若内容造成侵权/违法违规/事实不符,请联系邮箱:809451989@qq.com进行投诉反馈,一经查实,立即删除!
网站建设 2026/9/6 14:08:14

PHP本地二维码生成工具开发实战:从原理到批量导出

简介:PHP二维码在线生成工具本地版v1.0是一份基于PHP源码的二维码生成方案,主要面向需要在自己网站空间或本地环境生成二维码的开发者,解决线上生成服务依赖外部接口、无法自定义部署的问题。程序采用当前时间与随机数组合的方式生成PNG图片路…

作者头像 李华
网站建设 2026/9/6 14:08:34

ChatGPT桌面应用性能优化:Brent方案实战解析

ChatGPT 桌面应用性能优化:Brent 方案实战解析如果你最近被 ChatGPT 桌面端的启动卡顿、内存占用、多轮对话变慢折磨过,那么这篇内容可以直接收藏。这次我们来看一个围绕 ChatGPT 桌面应用做性能优化的方案,代号 Brent。它解决的问题很具体&a…

作者头像 李华
网站建设 2026/9/4 13:01:02

Windows进程CPU亲和性持久化设置:不依赖第三方工具实现进程核心绑定

这次我们来看一个关于 CPU 进程优化和管理的实战技巧。核心议题是:如何让 CPU-Z 这类系统信息工具在运行时,其进程的 CPU 亲和性(即允许使用哪些 CPU 核心)不被系统自动还原或重置。通常,我们可能会想到使用专业的进程…

作者头像 李华
网站建设 2026/9/4 1:47:31

隐蔽TXT阅读器2.05稳定版:隐私阅读与伪装界面实战指南

简介:这是一款以隐私保护为主打的TXT文本阅读工具,2.05稳定版压缩包面向需要安静阅读、不希望被他人察觉的普通用户,也适合开发工具爱好者研究桌面程序的打包与运行机制。资源标签为“开发工具”,整个包共203个文件,约…

作者头像 李华
网站建设 2026/9/4 10:43:37

小满秋招数据分析岗笔试复盘:SQL、Python与业务案例全解析

2023年秋招,我报了一家金融背景的科技公司数据分析岗,投完简历没几天就收到了小满秋招第一批笔试的链接。说实话,很多人对这个岗位的笔试预期就是“考SQL、考Python、考统计学”,但这批卷子做下来,我发现它更想考察的是…

作者头像 李华