CNN-RNN通用分类Python代码

最后更新于:2026-07-28 07:39:49

本文介绍一套基于 PyTorch 的 CNN-RNN 通用分类代码。核心函数会在每条样本内部选择一个“序列轴”:一维特征可以沿特征位置读取,多通道传感器数据可以沿时间读取,二维图像可以逐列读取;其余维度会自动合并为每一步的输入特征。这样无需为每一种输入形状重新编写训练、评价和绘图流程。

  • 提供 Iris 一维特征、UCI HAR 九通道人体活动序列、MNIST 二维灰度图三个可直接运行的demo。
  • 支持样本放第一维或最后一维,并通过 sequenceAxis 指定真正有顺序的轴。
  • 可以设置多层一维CNN、卷积通道、卷积核、池化、循环层、隐藏单元、循环层数和 Dropout。
  • Python 版支持 RNN、LSTM、GRU、BiRNN、BiLSTM、BiGRU 六种循环结构。
  • 支持分层随机划分,也支持按人员、设备、工况或批次提供固定的 splitLabels
  • 归一化均值与标准差只从训练集计算,验证集和测试集不参与拟合。
  • 自动计算 Accuracy、宏平均 Precision、Recall、F1、逐类别指标与混淆矩阵。
  • 自动生成并保存收敛过程、混淆矩阵、各类别指标、分类结果对比、三集合指标对比五类图片;三个案例共15张图。
  • 可设置 Adam、SGDM、RMSprop,支持学习率衰减、早停、类别权重和 CPU/GPU 自动选择。
  • 用户主要负责导入 XY,选对序列轴并调参,固定流程由 FunClassCNNRNN 一行完成。

一、代码运行环境

推荐使用 Windows 64位、Python 3.11、PyTorch。开发和测试环境为:

Python 3.11.9
torch 2.8.0+cpu
numpy 2.2.6
scikit-learn 1.6.1
matplotlib 3.10.3

在代码文件夹中执行:

pip install -r requirements.txt

完整版为 Python 源码,可以根据自己的环境调整依赖。公开版程序按 Windows 64位 Python 3.11 编译,必须使用 Python 3.11 x64;其他Python大版本可能无法加载公开版核心模块。

如果需要GPU训练,请根据显卡驱动与CUDA版本安装对应的 PyTorch。代码设置 deviceSel='auto' 时会优先使用可用GPU,否则自动使用CPU。

二、程序介绍

完整版文件结构如下:

CNN_RNN_Classification/
├── demoCNNRNNClassIris.py
├── demoCNNRNNClassHAR.py
├── demoCNNRNNClassMNIST.py
├── FunClassCNNRNN.py
├── EvaClassEffect.py
├── iris.csv
├── har_activity_data.npz
├── mnist_subset.npz
├── requirements.txt
├── 数据来源与许可.txt
├── 代码说明.txt
└── figure/
1. demoCNNRNNClassIris.py 文件

该脚本演示二维输入的一维特征分类。X.shape=(150,4),表示150条样本、每条4个特征位置;Y.shape=(150,),保存三种鸢尾花类别。脚本设置 sequenceAxis=0,把单条样本内部的四个特征依次交给 CNN-LSTM。

核心调用为:

foreData, foreDataTrain, model, info = FunClassCNNRNN(X, Y, options)

本次完整复验的测试集 Accuracy 为 1.0000,Macro-F1 为 1.0000。测试集只有30条样本,因此该数字主要用于确认流程正常,不代表其他数据也能达到相同结果。

Python Iris 测试集混淆矩阵

图 1 展示三个类别的真实标签与预测标签。主对角线表示预测正确,非对角线表示具体误分方向。

2. demoCNNRNNClassHAR.py 文件

该脚本演示多通道序列分类。X.shape=(2700,9,128):每条样本包含9个传感器通道、128个连续采样点;Y 是行走、上楼、下楼、坐着、站立、躺着六类标签。

脚本设置 sequenceAxis=1。这里的编号是“单条样本内部从0开始”,所以 [9,128] 中第1轴就是长度128的时间轴。9个通道会成为每个时间点上的9个输入特征。

HAR 使用 splitLabels 按受试者隔离:1500条训练、480条验证、720条测试。测试人员没有参与训练和验证,比随机拆分更接近面对新用户时的真实效果。

本次完整复验的测试集 Accuracy 为 0.8931,Macro-F1 为 0.8925

Python HAR 收敛过程

图 2 同时展示训练/验证损失和准确率。最佳权重按照验证损失选择,验证损失长期不改善时提前停止。

Python HAR 测试集混淆矩阵

图 3 用于定位六种活动之间的具体混淆。坐着与站立的传感器变化相近,通常比动态活动更难区分,这也是混淆矩阵比单一准确率更有价值的地方。

Python HAR 各类别指标

图 4 分别给出各类 Precision、Recall、F1,能发现某个少数类或难分类类别是否被总准确率掩盖。

3. demoCNNRNNClassMNIST.py 文件

该脚本演示二维灰度图分类。X.shape=(5000,28,28),表示5000张28×28灰度图;Y 为0~9十类数字。设置 sequenceAxis=1 后,程序沿宽度方向逐列读取图像:每一步接收一列28个像素,CNN提取相邻列中的笔画组合,GRU再汇总整张图的信息。

本次完整复验的测试集 Accuracy 为 0.9260,Macro-F1 为 0.9260

Python MNIST 测试集混淆矩阵

图 5 展示各数字的正确数量和误分方向。若需要继续提高图像分类精度,可以增加训练数据和网络容量,也应与标准二维CNN进行对比。

4. FunClassCNNRNN.py 文件

该文件包含核心函数和 PyTorch 网络类。核心函数签名如下:

def FunClassCNNRNN(X, Y, options=None):
    ...
    return foreData, foreDataTrain, model, info

输入参数:

  • X:任意维度有限数值 NumPy 数组,不能含 NaN 或 Inf。推荐样本放第一维,也兼容放最后一维。
  • Y:每条样本的类别标签,长度等于样本数;支持整数、浮点数或字符串类别。
  • options:控制输入适配、数据划分、网络、训练、归一化与绘图的字典。字段可部分省略,缺少字段采用默认值。

输出参数:

  • foreData:测试集预测标签,已经还原为 Y 的原始类别值。
  • foreDataTrain:训练集预测标签。
  • model:恢复到最佳验证损失状态的 PyTorch 模型。
  • info:包含三集合索引与真实标签、验证集预测、三集合类别概率、Accuracy/Precision/Recall/F1、混淆矩阵、训练历史、最佳轮次、标准化参数、输入尺寸适配记录、设备和完整 options
输入适配参数
参数 默认值 允许取值 含义与注意事项
sampleDimension 'auto' 'auto' 或整数轴编号 X 中的样本轴。自动模式先检查第一维,再兼容最后一维;明确编号按Python从0开始,负数可从末尾倒数。
sequenceAxis -1 单条样本内部的整数轴编号 真正具有顺序的轴,按0开始;-1为最后一轴。一维数组设0,[通道,时间]设1,[通道,高,宽]沿宽度读取时设2。选错轴时模型可能能跑,但学到的相邻关系没有业务意义。

适配器先把样本轴移动到最前面,再把序列轴移动到最后,其余维度合并为每个位置的特征,得到 [样本数, 特征数, 序列长度]。该三维数组直接交给 Conv1d

数据划分参数
参数 默认值 允许取值 含义与调参影响
splitLabels None 与Y等长的1/2/3数组 固定划分:1训练、2验证、3测试。适合按人员、设备或批次隔离;设置后 rTrain 等参数不负责集合划分。
rTrain 0.80 [0.5,1) 训练+验证数据比例,其余为最终测试集;不是纯训练集比例。
validationRatio 0.15 (0,0.5) 验证集占训练+验证数据的比例,用于早停和模型选择。
shuffle True True/False 是否在分层划分前打乱独立样本。严格时序或分组数据应提供固定划分。
seed 42 非负整数 固定Python、NumPy、PyTorch随机过程;设0不固定。某些GPU算子仍可能有末位差异。
CNN-RNN结构参数
参数 默认值 允许取值 含义与调参影响
networkType 'LSTM' RNN/LSTM/GRU/BiRNN/BiLSTM/BiGRU 循环结构。普通RNN最简单;GRU较精简;双向结构同时汇总两个方向,但计算量更大。
convChannels [32,64] 正整数序列 每个卷积层的输出通道数,元素个数即卷积层数。调大可学习更多模式,也更耗时、更易过拟合。
kernelSize 5 正奇数 卷积核长度。越大一次覆盖的相邻位置越多;短序列通常使用3。
poolSize 2 正整数,不超过序列长度 最大池化窗口;1表示不池化。过度池化会丢失短序列信息。
rnnHidden 64 正整数 循环隐藏单元数,控制记忆容量和分类头输入维度。
rnnLayers 1 正整数 循环层层数;大于1时层间也应用 dropout。层数增加会显著提高训练成本。
dropout 0.20 [0,1) 随机失活比例。适当增大可缓解过拟合,过大会欠拟合。

Python 版循环结构切换示例:

options['networkType'] = 'RNN'
options['networkType'] = 'LSTM'
options['networkType'] = 'GRU'
options['networkType'] = 'BiRNN'
options['networkType'] = 'BiLSTM'
options['networkType'] = 'BiGRU'

实际使用时只保留其中一行。双向结构会使用正向末端状态与反向起点状态拼接后分类。

训练参数
参数 默认值 允许取值 含义与调参影响
solverName 'adam' 'adam'/'sgdm'/'rmsprop' 优化器。Adam通常适合作为初始设置。
maxEpochs 40 正整数 最大训练轮数;可能被早停缩短。
learnRate 0.001 正数 初始学习率。太大可能震荡,太小收敛缓慢。
batchSize 64 正整数 批尺寸。大批次更稳定但占更多显存或内存;小数据可使用16或32。
earlyStoppingPatience 8 非负整数 验证损失连续多少轮不改善后停止;0关闭早停。最佳验证权重会被恢复。
learnRateSchedule 'none' 'none'/'piecewise' 是否按固定周期降低学习率。
learnRateDropPeriod 15 正整数 piecewise 模式中每隔多少轮衰减。
learnRateDropFactor 0.5 (0,1] 学习率衰减因子。0.5表示减半。
classWeight 'auto' 'auto'/'none' 或正数序列 auto只根据训练集频数计算类别权重;自定义权重长度必须等于类别数。
deviceSel 'auto' 'auto'/'cpu'/'gpu' 训练设备。指定GPU但不可用时会提示并回退CPU。
数据处理与绘图参数
参数 默认值 允许取值 含义与注意事项
mapflag True True/False 是否标准化。均值和标准差只从训练集计算,再应用到验证和测试数据。
figflag True True/False 是否绘制并保存五张结果图。批量调参可关闭。
showFigures True True/False 是否尝试弹出图窗;图片仍由 figflag 控制保存。
caseName 'CNN-RNN' 字符串 图片文件名前缀,避免不同demo覆盖。
classNames None 与类别数一致的名称序列 绘图显示名,顺序必须与排序后的原始类别一致。
5. EvaClassEffect.py 文件

该文件提供独立评价函数:

metrics = EvaClassEffect(realData, foreData, classOrder=None)

输入真实类别、预测类别和可选类别顺序,返回 Accuracy、MacroPrecision、MacroRecall、MacroF1、逐类别 Precision/Recall/F1、Support、ConfusionMatrix 与 ClassOrder。

6. 数据文件
  • iris.csv:Iris 三分类数据。
  • har_activity_data.npz:UCI HAR 九通道序列、标签和按受试者隔离的集合标记。
  • mnist_subset.npz:MNIST 0~9十分类的5000张灰度图子集。
  • 数据来源与许可.txt:数据来源与许可说明;UCI HAR 按 CC BY 4.0 使用。
7. requirements.txt 文件

记录 NumPy、Matplotlib、scikit-learn 和 PyTorch 等依赖版本。建议先在独立Python 3.11环境中安装,再运行demo。

8. figure 文件夹

每个demo生成5张独立PNG,三个案例共15张:收敛过程、测试集混淆矩阵、各类别指标、测试集类别对比、训练/验证/测试指标对比。完整版图片无水印。

9. 代码说明.txt 文件

随代码交付的离线手册,完整展开输入、输出、options、数据替换方式、运行环境与常见注意事项。

三、快速开始

1. 创建环境并安装依赖

建议创建 Python 3.11 环境,然后在代码文件夹执行:

pip install -r requirements.txt
2. 运行与自己数据形态最接近的demo
python demoCNNRNNClassIris.py
python demoCNNRNNClassHAR.py
python demoCNNRNNClassMNIST.py

选择其中一个即可。程序无报错、命令行打印三集合指标、figure/ 生成5张图片,即说明运行环境正常。

3. 替换成自己的数据

二维一维数组:

data = np.loadtxt('your_data.csv', delimiter=',')
X = data[:, :-1]
Y = data[:, -1]
options['sequenceAxis'] = 0

多通道序列:

# X.shape = [样本数, 通道数, 时间长度]
# Y.shape = [样本数]
options['sequenceAxis'] = 1

灰度图:

# X.shape = [样本数, 高度, 宽度]
# Y.shape = [样本数]
options['sequenceAxis'] = 1  # 沿宽度逐列读取

彩色图可整理为 [样本数,通道数,高度,宽度],沿宽度读取时设置 sequenceAxis=2,通道和高度会合并为每一步的特征。若任务强依赖二维局部邻域,标准二维CNN通常更合适;CNN-RNN更适合某个方向具有明确扫描或序列含义的场景。

4. 避免数据泄漏

同一受试者、设备或批次产生的多条样本不应随机散落在三个集合。可以自己构造:

options['splitLabels'] = split_labels  # 1训练、2验证、3测试

程序的标准化器只从训练集拟合,但前提是集合划分本身合理。错误划分造成的泄漏不能靠模型参数修复。

5. 调参与运行

先固定数据划分,再依次检查序列轴、网络容量和训练节奏。训练集显著优于验证集时,可以减小 convChannelsrnnHiddenrnnLayers,增大 dropout,也可以增加数据和改进分组划分。损失剧烈震荡时优先降低 learnRate

四、关于完整版与公开版代码

功能 完整版 公开版
Iris一维特征演示 √,90条
HAR多通道序列演示 √,96条
MNIST图像演示 0~9十分类,5000张 0/1/2三分类,90张
自定义数据与参数
核心函数源码 提供 不提供
最大样本数 无人工限制,受硬件约束 100条
最大卷积层数 无试用版限制 2层
最大训练轮数 无试用版限制 30轮
循环层数 可设置 1层
循环结构 RNN、LSTM、GRU、BiRNN、BiLSTM、BiGRU LSTM、GRU
最大隐藏单元数 无试用版限制 64
自动分类指标和15张案例图
结果图水印 无水印 有“试用版@khsci.com/docs”水印
修改核心实现与二次开发 ×

公开版用于验证 Python 环境、理解输入格式和体验完整调用流程;正式训练、大数据任务与二次开发建议使用完整版。

五、获取公开版程序

公开版下载: 点击此处下载 CNN-RNN 通用分类 Python 公开版代码

注:公开版适用于 Windows 64位 Python 3.11,最多支持100条样本,最多训练30轮,结果图带试用版水印。

六、获取完整版程序

点击本页面“立即支付”按钮,付款后获取完整版代码下载链接和售后联系方式。付款完成后刷新本页面即可看到下载链接。

(注意支付跳转失败的话,请使用浏览器打开本页面)

您需要先支付 69.5元 才能查看此处内容!立即支付

七、完整版代码重要更新

  • 20260720:完成 Python 版初版代码,支持一维数组、多通道序列、灰度图和更高维输入适配;支持六种循环结构,提供 Iris、UCI HAR、MNIST 三个案例和15张自动结果图。

八、常见问题

Q1:为什么说是“通用分类”,但还要指定 sequenceAxis?

“通用”指函数能适配多种数组维度,不代表模型不需要结构假设。CNN和RNN必须知道沿哪个方向寻找相邻模式和前后关系,因此用户仍需明确序列轴。

Q2:所有普通表格数据都适合吗?

不一定。普通表格的列如果没有稳定的顺序和邻接含义,CNN-RNN虽然能运行,但结构优势可能发挥不出来。应与树模型、SVM、全连接网络等基线比较。

Q3:公开版报错找不到 FunClassCNNRNN 怎么办?

先确认使用 Windows 64位 Python 3.11,并且从公开版代码文件夹运行demo。其他Python大版本无法直接加载该公开版核心模块。

Q4:指定GPU后为什么仍然使用CPU?

当前 PyTorch 必须实际检测到可用CUDA设备。若安装的是CPU版PyTorch,或CUDA与驱动不匹配,程序会提示并回退CPU。可先运行 torch.cuda.is_available() 检查。

Q5:训练集很好,测试集较差怎么办?

先检查分组泄漏和类别比例,再比较验证集走势。可以减小网络、增加Dropout、使用早停、增加数据或调整划分。不要在测试集上反复调参;测试集应留到最后评估。

Q6:最少需要修改哪些内容?

最少替换 XYoptions['sequenceAxis']。如果样本存在人员、设备或批次关联,再提供 splitLabels。之后根据验证集表现调整网络与训练参数。