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 自动选择。
- 用户主要负责导入
X、Y,选对序列轴并调参,固定流程由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条样本,因此该数字主要用于确认流程正常,不代表其他数据也能达到相同结果。

图 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。

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

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

图 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。

图 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. 调参与运行
先固定数据划分,再依次检查序列轴、网络容量和训练节奏。训练集显著优于验证集时,可以减小 convChannels、rnnHidden 或 rnnLayers,增大 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 环境、理解输入格式和体验完整调用流程;正式训练、大数据任务与二次开发建议使用完整版。
五、获取公开版程序
注:公开版适用于 Windows 64位 Python 3.11,最多支持100条样本,最多训练30轮,结果图带试用版水印。
六、获取完整版程序
点击本页面“立即支付”按钮,付款后获取完整版代码下载链接和售后联系方式。付款完成后刷新本页面即可看到下载链接。
(注意支付跳转失败的话,请使用浏览器打开本页面)
七、完整版代码重要更新
- 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:最少需要修改哪些内容?
最少替换 X、Y 和 options['sequenceAxis']。如果样本存在人员、设备或批次关联,再提供 splitLabels。之后根据验证集表现调整网络与训练参数。