news 2026/9/16 11:36:15

用TensorFlow实现Iris分类神经网络:从数据预处理到checkpoint保存

作者头像

张小明

前端开发工程师

1.2k 24
文章封面图
用TensorFlow实现Iris分类神经网络:从数据预处理到checkpoint保存

简介:面向机器学习初学者的课程设计参考,围绕Iris鸢尾花数据集,演示从零搭建神经网络、训练分类模型的完整流程。资源包共25个文件,压缩后仅75KB,核心包含Python训练脚本main.py、Iris数据集iris.csv、TensorFlow模型文件(meta、index、checkpoint、data)以及TensorBoard训练日志。模型文件覆盖20%、50%、70%等不同训练进度,可直观查看训练中期与最终的参数状态。附带的VS Code配置(settings.json、launch.json、code-workspace)便于一键复现调试环境。已有281人学习。通过学习这份项目,读者可以掌握数据加载与预处理、网络结构定义、模型训练与保存、可视化监控等关键环节,并借助已有模型权重理解checkpoint的恢复机制。整体结构清晰,适合作为神经网络入门实验或同类大作业的参考。

1. 把 Iris 分类做成神经网络课程设计,值不值

Iris 鸢尾花分类是课程设计里最容易被低估的一道题:150 条样本、4 个特征、3 个类别,拿 sklearn 的 LogisticRegression 三行代码就能跑到 95% 以上准确率。但题目一旦限定必须用神经网络,事情就完全不同——从 pandas 读 iris.csv 到 one-hot 编码,从隐藏层宽度到学习率,从 Saver 保存 checkpoint 到 TensorBoard 画损失曲线,每一步都有参数要调、有坑要踩。

这份 Practicum-master 工程把整条链路跑通,在 5%、20%、50%、70% 四个训练进度分别留下检查点,适合赶课程设计大作业的同学,也适合想搞懂 checkpoint 文件族怎么用的工程师。下面从数据处理开始,按训练一个能交差的 Iris 分类模型的顺序逐段拆解。

2. 数据预处理:从 iris.csv 到神经网络的输入张量

2.1 先看原始数据:150 条样本的统计特征

iris.csv 是经典 Iris 数据集的原始导出,150 行,每行对应一株鸢尾花,包含花萼长度、花萼宽度、花瓣长度、花瓣宽度四个连续特征,以及一个品种标签。三个类别各 50 条:setosa、versicolor、virginica。样本量小、特征维度低,意味着模型很容易过拟合,训练时要盯住验证集,不能只看训练集准确率。

列名含义取值范围(近似)
sepal_length花萼长度4.3 ~ 7.9
sepal_width花萼宽度2.0 ~ 4.4
petal_length花瓣长度1.0 ~ 6.9
petal_width花瓣宽度0.1 ~ 2.5

pandas 读进来后第一件事是确认表头和分布,我一般会先打印columns而不是直接取列名,因为网上下载的版本表头常有差异:

import pandas as pd df = pd.read_csv('iris.csv') print(df.columns.tolist()) print(df.head()) print(df.describe()) print(df['species'].value_counts())

describe()输出的均值、标准差直接决定后面归一化的参数,value_counts()用来确认三类样本是否均衡。这份数据是 50/50/50,不需要做类别权重;如果拿到的是不平衡版本,要额外考虑加权损失或者过采样。实际下载的 csv 表头可能有大小写差异,比如SepalLengthCmsepal_length,先确认列名再取特征列,能省掉后面不少对齐的麻烦。

2.2 Z-score 归一化:为什么小数据集更要做

决策树、随机森林对特征尺度不敏感,神经网络不行。输入特征进入第一层就要和权重做点积,花萼长度范围约 4.3~7.9,花瓣长度范围约 1.0~6.9,量纲虽然接近,梯度更新时仍会被数值较大的特征主导;如果某个特征范围是 0~1000,模型几乎收敛不了。课程设计里最常见的处理是 Z-score 归一化,把每个特征变成均值 0、标准差 1。

from sklearn.model_selection import train_test_split feature_cols = ['sepal_length', 'sepal_width', 'petal_length', 'petal_width'] X = df[feature_cols].values y = df['species'].values X_train, X_test, y_train, y_test = train_test_split( X, y, test_size=0.2, random_state=42) # 归一化参数只在训练集上拟合,测试集复用同一套参数,避免数据泄漏 mean = X_train.mean(axis=0) std = X_train.std(axis=0) X_train = (X_train - mean) / std X_test = (X_test - mean) / std print(X_train.shape, X_test.shape)

test_size=0.2表示测试集 30 条、训练集 120 条,对课程设计够用;random_state=42固定随机种子,保证多次运行得到同样的划分,答辩复现时不会出偏差。这里的关键点是归一化参数只从训练集计算,测试集用同一套 mean/std 变换,否则测试集信息提前混入训练流程,属于数据泄漏,最后评估的数字会虚高。

2.3 one-hot 编码:标签顺序要和输出层对齐

Iris 的品种是互斥类别,不能直接编码成 0、1、2 当回归目标,网络会学出「类别 1 介于类别 0 和类别 2 之间」的错误先验。课程设计里最常用的处理是 one-hot:setosa 变成 [1,0,0],versicolor 变成 [0,1,0],virginica 变成 [0,0,1]。

import numpy as np species_list = ['setosa', 'versicolor', 'virginica'] y_train_idx = np.array([species_list.index(s) for s in y_train]) y_test_idx = np.array([species_list.index(s) for s in y_test]) y_train_hot = np.eye(3)[y_train_idx] y_test_hot = np.eye(3)[y_test_idx] print(y_train_hot[:3])

np.eye(3)生成 3×3 单位阵,按索引取行就得到 one-hot 向量。species_list的顺序必须和后面输出层 3 个神经元的含义一一对应,一旦错位,训练过程 loss 会掉,但分类结果永远错位。TensorFlow 里也可以不手动 one-hot,直接用整数标签配合sparse_softmax_cross_entropy_with_logits,但显式 one-hot 的好处是能直接拿tf.argmax(pred, 1)和标签做可视化对比,答辩画图时更方便。

3. 网络结构设计:4-8-3 前馈网络为什么够用

很多人看到神经网络就想到卷积神经网络,但 CNN 是为图像这类有空间局部性的数据设计的;Iris 是 4 维表格特征,用前馈神经网络(也就是常说的 BP 神经网络)就够。同一份数据,sklearn 的 SVM 也能跑到 97%,但那种方式把梯度下降和反向传播全封装了,课程设计的评分点恰恰在这些细节里。

3.1 从输入到输出的维度推演

输入层 4 个神经元对应四个特征,输出层 3 个神经元对应三个类别,中间夹一个隐藏层。隐藏层宽度是第一个要定的超参数:太窄拟合能力不够,准确率停在 90% 上下;太宽在 150 条样本上必然过拟合,训练集 100% 而测试集只有 80%。选 8 是经验做法:输入 4、输出 3,取中间值两倍左右,同时 8 是 2 的幂,向量化计算友好。想要更复杂可以扩到 4-16-8-3 两层隐藏层,但 Iris 用单隐藏层 4-8-3 已经能稳定到 95% 以上。

import tensorflow as tf inputs = tf.placeholder(tf.float32, [None, 4], name='inputs') labels = tf.placeholder(tf.float32, [None, 3], name='labels') w1 = tf.Variable(tf.random_normal([4, 8], stddev=0.1), name='w1') b1 = tf.Variable(tf.zeros([8]), name='b1') h1 = tf.nn.relu(tf.matmul(inputs, w1) + b1) w2 = tf.Variable(tf.random_normal([8, 3], stddev=0.1), name='w2') b2 = tf.Variable(tf.zeros([3]), name='b2') logits = tf.matmul(h1, w2) + b2 pred = tf.nn.softmax(logits, name='pred')

占位符第一个维度是None,表示批次大小可变,训练和推理共用一套图。stddev=0.1控制初始化方差,配合 ReLU 避免激活值过大或过小;偏置初始化为 0 是常规做法。name参数不是装饰,后面从 checkpoint 恢复时靠它定位张量,我习惯把inputspred这种关键节点都命名,命名不规范会导致get_tensor_by_name找不到节点。

3.2 输出层为什么用 softmax 交叉熵而不是均方误差

分类任务的输出层接 softmax,把 3 个 logits 变成和为 1 的概率分布。损失用交叉熵而不是均方误差,原因在梯度:MSE 对 softmax 输出求导,在概率接近 0 或 1 的区域梯度趋近于零,训练会卡住;交叉熵的梯度是pred - label,错误越严重梯度越大,收敛更快且不会饱和。

# logits 直接传入,不要在外部先做 softmax loss = tf.reduce_mean( tf.nn.softmax_cross_entropy_with_logits_v2(logits=logits, labels=labels))

这个函数在内部完成 softmax 和交叉熵的联合计算,并用 log-sum-exp 技巧保证数值稳定。如果自己先算softmax(logits)再算交叉熵,log(0)会出现 nan。reduce_mean对 batch 内所有样本的损失取平均,让 loss 的量级与 batch size 无关,换 batch size 时不需要重新调学习率。

3.3 优化器与学习率:先跑通再调优

课程设计阶段不需要在优化器上花太多精力,Adam 在 Iris 这种小数据集上几乎不需要调参:

train_op = tf.train.AdamOptimizer(learning_rate=0.01).minimize(loss)
优化器学习率建议在 Iris 上的表现说明
SGD0.01~0.1收敛慢,可能要 500 epoch配合 momentum 更稳
Adam0.001~0.01200 epoch 内稳定收敛默认 beta1/beta2 即可
RMSProp0.001~0.01居中课程设计里少见

学习率 0.01 是这份工程的取值,实际跑的时候如果 loss 震荡,降到 0.001;如果 50 个 epoch 还不掉,检查归一化而不是调学习率。批次大小 16,也就是每个 epoch 用 120 条训练样本滚动 8 次更新,比全量梯度更新快,比单样本更新稳。

4. 训练与检查点:把训练过程存成可复现的证据

4.1 checkpoint 文件族拆解:meta、index、data 各管什么

先看这份工程解压后的目录结构:

Practicum-master/ ├── main.py ├── iris.csv ├── Practicum.code-workspace ├── .vscode/ ├── model/ │ ├── checkpoint │ ├── model.ckpt.meta │ ├── model.ckpt.index │ └── model.ckpt.data-00000-of-00001 ├── tensorboard/ ├── 5%Training/ ├── 20%Training/ ├── 50%Training/ └── 70%Training/

model/下四个文件是一套完整的 TensorFlow 1.x Saver 产物,各自职责如下:

文件内容恢复时作用
model.ckpt.meta计算图结构 MetaGraphDef不需要重写网络定义就能重建图
model.ckpt.index张量名到文件偏移的映射定位每个变量在 data 里的位置
model.ckpt.data-00000-of-00001变量数值,分片文件权重、偏置的实际值
checkpoint记录最新一次保存的路径latest_checkpoint()读取它

>import os saver = tf.train.Saver(max_to_keep=5) EPOCHS = 200 BATCH_SIZE = 16 # key 是训练进度,value 是保存目录 save_points = {0.05: '5%Training', 0.20: '20%Training', 0.50: '50%Training', 0.70: '70%Training'} with tf.Session() as sess: sess.run(tf.global_variables_initializer()) for epoch in range(EPOCHS): idx = np.random.permutation(len(X_train)) for i in range(0, len(X_train), BATCH_SIZE): batch_idx = idx[i:i + BATCH_SIZE] _, l = sess.run([train_op, loss], feed_dict={ inputs: X_train[batch_idx], labels: y_train_hot[batch_idx]}) progress = (epoch + 1) / EPOCHS for ratio, folder in save_points.items(): if progress >= ratio and not os.path.exists(folder): saver.save(sess, os.path.join(folder, 'model.ckpt')) print('saved at epoch', epoch + 1, '->', folder) saver.save(sess, 'model/model.ckpt')

max_to_keep=5限制磁盘上最多保留 5 份检查点,覆盖 4 个中间进度加最后一份。用os.path.exists(folder)判断是否已保存,避免同一个 epoch 重复触发多个保存条件,也让断点续跑时不会覆盖已有进度。saver.save不传global_step,文件名就是朴素的model.ckpt.meta形式,和目录清单里看到的一致;如果传了global_step,文件名会变成model.ckpt-73.meta这种带步数的形式。

4.3 TensorBoard 可视化:损失曲线是答辩材料

在定义图的同时挂上标量汇总,训练循环里每个 epoch 结束追加一次:

tf.summary.scalar('loss', loss) tf.summary.scalar('acc', acc) merged = tf.summary.merge_all() summary_writer = tf.summary.FileWriter('tensorboard', sess.graph) # 在 4.2 的训练循环内,每个 epoch 结束后追加: summary = sess.run(merged, feed_dict={inputs: X_train, labels: y_train_hot}) summary_writer.add_summary(summary, epoch)

启动可视化:

tensorboard --logdir=./tensorboard --port=6006

sess.graph会把整张计算图写进 events 文件,浏览器里能看到每个 op 的输入输出和变量形状;add_summary的第二个参数是 step,横轴对应 epoch。答辩时打开 TensorBoard 展示损失曲线,比口述「模型收敛了」有说服力得多。曲线平滑下降说明学习率合适,锯齿剧烈说明学习率偏大,可以降到 0.003 再试。

4.4 训练阶段最常踩的三个坑

第一个坑是 loss 变成 nan。绝大多数情况不是网络写错,而是喂进去的数据没有归一化,或者学习率过大,特征数值一大,logits 溢出到 float32 的表示范围。先检查 feed_dict 里数据有没有经过 2.2 的变换。

第二个坑是准确率卡在 33%(随机水平)附近。优先怀疑标签映射错位:species_list的顺序、csv 里species列的种类、输出层 3 个神经元三者必须一致。打印y_train_hot[:3]np.argmax(pred, 1)[:3]对比是最快的排查方式。

第三个坑是训练集 100%、测试集只有 80% 的过拟合。Iris 只有 150 条样本,隐藏层超过 16 个神经元就会开始记样本。把隐藏层缩回 8,或者加 dropout(keep_prob 0.8),都能缓解。课程设计要求不高的话,单隐藏层 4-8-3 加 200 个 epoch 是最稳的组合。

提示:测试集只在最终评估时碰一次。调参过程中反复看测试集准确率,相当于用测试集训练,最后的数字会虚高。

5. 恢复检查点做分类评估:把模型从磁盘请回来

5.1 用 meta 图恢复模型并跑推理

训练完拿到的是四个 checkpoint 文件,交付时不会带训练环境。恢复有两种方式:重跑网络定义再 restore,要求变量名完全一致;或用 meta 文件重建图,不依赖原始代码。跨机器复现时第二种更实用:

tf.reset_default_graph() saver = tf.train.import_meta_graph('model/model.ckpt.meta') graph = tf.get_default_graph() inputs = graph.get_tensor_by_name('inputs:0') pred = graph.get_tensor_by_name('pred:0') with tf.Session() as sess: saver.restore(sess, tf.train.latest_checkpoint('model/')) prob = sess.run(pred, feed_dict={inputs: X_test}) y_pred = np.argmax(prob, axis=1)

import_meta_graph只读图结构,restore 才把 data 文件里的权重加载进来。latest_checkpoint会读 checkpoint 文件里记录的路径,不用手写文件名。get_tensor_by_name找的inputs:0pred:0对应 3.1 的 name 参数,命名不一致会在这里报 KeyError。

5.2 准确率之外:混淆矩阵与每类精确率/召回率

分类评估不能只报准确率。setosa 与另两类线性可分,准确率很容易到 95%,真正要看的是 versicolor 和 virginica 的重叠区域。直接出混淆矩阵和分类报告:

from sklearn.metrics import confusion_matrix, classification_report y_true = np.argmax(y_test_hot, axis=1) print(confusion_matrix(y_true, y_pred)) print(classification_report(y_true, y_pred, target_names=species_list))

混淆矩阵对角线是正确分类数,非对角线是错分样本。classification_report给出每类精确率、召回率、F1。典型结果是 setosa 100%,versicolor 与 virginica 互相错分 1~2 条。测试集只有 30 条,单次划分波动大,严谨做法是 5 折交叉验证取均值。

5.3 一个答辩加分的技巧:验证损失早停

从训练集再切 20% 当验证集,每个 epoch 结束后跑一次验证集 loss,与训练 loss 画在同一张 TensorBoard 里。验证集 loss 连续 10 个 epoch 不降反升,就 restore 回滚到上一次保存的检查点,Iris 的过拟合拐点通常在第 80~120 epoch 之间。触发早停后注意一件事:不要先执行tf.global_variables_initializer()再 restore,初始化会把刚恢复的权重清空;正确做法是 restore 之后直接前向推理,输出验证集的分类报告。把「训练集降、验证集升」的双曲线截图放进报告,比单贴一个准确率更能说明你理解过拟合。

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

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

熟练掌握SpringCloud流行技术栈,如Nacos、Seata、Zookeeper、Dubbo、OpenFeign、GateWay、Sentinel、SkyWalking和Discovery。熟悉

Nacos 注册中心原理nacos1.x 1.提供方调用rest接口发起服务注册,注册自己的ip:端口 2.消费方每隔10秒拉取服务提供方注册列表 3.消费方提供方都需要和nacos每隔5秒进行心跳检测,15秒没有心跳,标志为不健康,30秒没有心跳标志位下线…

作者头像 李华
网站建设 2026/9/16 11:35:52

微群人脉微信小程序源码部署与LNMP环境搭建实战

简介:微群人脉微信小程序源码是一套基于微信生态的社群运营与流量裂变系统,适合开发者、站长及寻求私域裂变工具的个人或团队部署使用。新版针对旧版痛点做了关键优化,用户登录后直接进入微群界面,不再被广告页拦截,有…

作者头像 李华
网站建设 2026/9/16 11:34:40

工业自动化中EtherCAT实现跨品牌设备高精度同步控制

1. 项目概述:工业自动化中的异构设备集成挑战在工业自动化产线升级项目中,我们经常遇到不同品牌设备协同工作的技术难题。最近完成的某汽车零部件生产线改造中,就涉及到用基恩士KV8000系列PLC控制松下A6系列伺服电机的典型场景。这种日系PLC与…

作者头像 李华
网站建设 2026/9/16 11:33:31

AI协同PCB设计:从自然语言到量产Gerber的工程实践

1. 这不是科幻预告片,是硬件工程师晨会的真实议题“GPT-6 都能自己画 PCB 了”——这句话最近在几个硬件工程师群和EDA工具论坛里反复刷屏,语气里混着调侃、焦虑,还有点将信将疑的试探。我上周在苏州一家做工业传感器的公司做技术交流&#x…

作者头像 李华
网站建设 2026/9/16 11:32:26

低压直流伺服驱动器怎么选?电压电流、编码器与通讯协议全解析

我做运动控制集成这些年,收到最多的咨询就是:低压直流伺服驱动器到底怎么选?电机功率、驱动器电流、通讯接口、编码器协议,每一关都有人踩坑。最常见的情况是,设备都装好了,上电一跑才发现,要么…

作者头像 李华