资讯动态

基于TensorFlow与Flask的水稻病虫害识别系统实战与踩坑记录

发布时间:2026/8/27 22:09:58 来源:尧图企业网站定制
简介图像识别作为人工智能的核心应用之一近年来在农业植保领域展现出巨大价值。通过深度学习技术对农作物叶片图像进行分类能够快速辅助诊断病虫害降低对人工经验的依赖。在工程落地时通常需要借助开源框架完成模型训练与服务封装其中TensorFlow提供了成熟的训练生态Flask则以轻量灵活的方式快速搭建Web推理接口。两者结合可以构建一套完整的图像分类应用覆盖数据预处理、模型训练、迁移学习调优以及在线预测等环节。这类技术方案不仅适用于田间病虫害监测也可扩展到工业质检、医疗影像等通用图像分类场景。本文以水稻叶片病害识别为例详细记录了从环境配置、数据集构建到模型部署的全过程并总结了版本兼容、数据增广和部署预处理的常见问题为相似项目提供了一份可直接参考的实践指南。 做水稻病虫害识别这个项目算是把这两年踩过的坑一次性踩全了。从TensorFlow版本匹配到Flask部署的坑从数据集标注到模型过拟合每个环节都有值得记录下来东西。这篇文章是把整个项目的完整思路、关键代码和踩坑记录整理出来给正在做类似图像识别项目的朋友一个可以直接参考的路线。先说清楚这个项目是做什么的。它是一个基于深度学习的农作物病虫害识别系统核心流程是用TensorFlow训练一个卷积神经网络模型让它能识别水稻叶片上的常见病害稻瘟病、稻曲病、白叶枯病等然后把这个模型封装成一个Flask Web应用用户上传一张水稻叶片照片网页端就能返回病害类型和置信度。整个项目源码完整从数据预处理到模型部署都有可跑的代码。适合谁来参考做深度学习入门到落地项目的学生、做农业信息化相关开发的工程师、以及想在Web端部署图像识别模型但还没摸清完整链路的人。下面按实际开发顺序来记录从架构设计到环境搭建再到训练和部署最后是问题排查。1. 项目概述与整体架构设计1.1 为什么需要水稻病虫害智能识别传统的水稻病虫害诊断主要靠植保人员肉眼观察一个县级植保站通常只有几名技术人员高峰期根本跑不过来。农户遇到问题拍张照片发到群里问得到的答案往往也是各说各话。深度学习方法在农业植保领域的落地价值就在这里图像识别模型可以在几秒内给出参考诊断结果把专业植保知识以极低成本复制到每个农户的微信里。技术实现上这个项目有三个核心难点需要提前想清楚。第一是数据水稻病害不同生育期的表现差异很大同一病害在不同光照、不同拍摄角度下外观差别明显数据集的覆盖程度直接决定模型上限。第二是模型移动端和实际部署环境的算力有限不能一上来就堆大模型需要在精度和推理速度之间找平衡。第三是部署训练好的模型要变成普通人能用的工具需要在Web端做输入输出封装这部分的坑往往比训练还多。1.2 技术选型为什么是TensorFlow Flask选TensorFlow而不是PyTorch主要出于三个考虑。一是工业部署生态成熟。TensorFlow的SavedModel格式在模型部署上有天然优势配合TensorFlow Serving做后续扩展很方便。如果只是个人项目PyTorch的torchserve也能用但TensorFlow在生产环境的案例积累更久遇到的坑基本都能搜到解决方案。二是Keras API的上手成本低。TensorFlow 2.x的Keras接口非常友好用两三行代码就能搭建一个迁移学习模型这对快速验证方案非常有帮助。PyTorch当然也可以但要写的样板代码多一些。三是Flask框架轻量灵活。相比Django和FastAPIFlask的核心理念是微框架适合做单服务的模型推理接口。FastAPI虽然性能更好但生态相对年轻在Windows环境下的兼容性和文档完整度不如Flask成熟。实际开发中用Flask写一个图片上传、模型推理、结果返回的接口代码量非常少而且flask的调试模式对开发期很有帮助。1.3 系统整体架构与完整源码目录这个项目的整体链路是数据集目录 → 数据预处理脚本 → TensorFlow训练脚本 → SavedModel模型文件 → Flask应用 → 浏览器页面。完整的源码目录结构如下rice_disease_system/ ├── app.py # Flask主应用 ├── model_train.py # 模型训练脚本 ├── data_preprocess.py # 数据预处理与增广脚本 ├── requirements.txt # 项目依赖清单 ├── models/ │ └── rice_disease_model.h5 # 训练好的模型文件 ├── dataset/ │ ├── train/ │ │ ├── rice_blast/ # 稻瘟病 │ │ ├── rice_tungro/ # 东格鲁病 │ │ ├── bacterial_blight/ # 白叶枯病 │ │ └── healthy/ # 健康叶片 │ └── val/ │ ├── rice_blast/ │ ├── rice_tungro/ │ ├── bacterial_blight/ │ └── healthy/ ├── static/ │ ├── uploads/ # 用户上传图片目录 │ └── css/ │ └── style.css └── templates/ └── index.html # 前端页面这套目录结构是实际验证过最顺手的训练代码和Web代码分离模型文件独立存放数据集按类别分文件夹管理和Keras的ImageDataGenerator的目录结构天然匹配。2. 开发环境搭建与依赖配置2.1 Python版本与虚拟环境策略这个项目的开发环境我推荐Python 3.9或3.10不要图新鲜用最新的Python 3.12或3.13。原因很简单TensorFlow对Python版本的支持有明显滞后经常出现Python最新版本无法安装TF的情况。我用的是Python 3.9.13搭配TensorFlow 2.10.0这套组合在Windows 11 RTX 3060环境下非常稳定没有踩到什么版本冲突的坑。如果你的显卡是RTX 30系列以上用TF 2.10搭配CUDA 11.2是性价比最高的选择如果显卡比较新比如RTX 40系列可以考虑TensorFlow 2.12以上版本因为老版本TF对新显卡的支持不太好。虚拟环境必须创建。强烈建议用conda创建独立的虚拟环境避免把系统Python搞乱。我用的是conda命令conda create -n rice_disease python3.9 conda activate rice_disease如果不用conda用venv也完全可以python -m venv rice_env rice_env\Scripts\activate # Windows source rice_env/bin/activate # Linux/Mac虚拟环境里pip安装依赖最稳妥的方式是把pip升级到最新python -m pip install --upgrade pip2.2 TensorFlow与CUDA版本匹配要点TensorFlow版本和CUDA、cuDNN的匹配是新手最容易翻车的地方。装好TensorFlow之后在import时如果报错或者运行训练时提示CUDA相关错误基本都是版本不匹配造成的。不同TensorFlow版本对应的CUDA版本整理如下TensorFlow版本Python版本建议对应CUDA对应cuDNN2.10.03.7-3.1011.28.12.12.03.8-3.1111.88.62.15.03.9-3.1112.28.92.18.03.9-3.1212.39.1上面这张表是参考TensorFlow官方支持列表整理的实操中最省事的方法是安装一个带GPU支持的TensorFlow后用tf.test.is_gpu_available()验证是否能用GPU。如果返回False就去官方对应表查版本。我实际项目的环境配置是Python 3.9.13 TensorFlow 2.10.0 CUDA 11.2 cuDNN 8.1 RTX 3060 12G。这套组合在训练不到2000张图片的数据集时一个epoch只需要几十秒完全够用。要注意CUDA的安装路径里不能有中文和空格否则TensorFlow找不到CUDA库。2.3 Flask与其余依赖的安装细节Flask相关依赖比较简单直接pip安装即可pip install flask3.0.0 pip install pillow10.1.0 pip install numpy1.24.3 pip install tensorflow2.10.0这里有几个容易踩的坑。numpy版本和TensorFlow的兼容性很关键。TensorFlow 2.10.0要求numpy版本不能高于1.26否则会报错。实际我用的是1.24.3稳定运行。pillow库用于处理上传的图片建议不要用太老的版本因为老版本对PNG格式的某些参数处理有问题。安装TensorFlow的时候如果网络不稳定容易下载到一半卡住。可以换用pip镜像源速度会快很多。我这里用的是清华源pip install tensorflow2.10.0 -i https://pypi.tuna.tsinghua.edu.cn/simple常见安装报错速查报错信息原因解决方式Could not find a version that satisfies the requirement tensorflowPython版本太高TensorFlow不支持降到Python 3.9或3.10DLL load failedCUDA版本不匹配或安装路径含中文重新安装匹配的CUDA版本numpy.dtype size changednumpy版本与TF不匹配降低numpy版本推荐pip install numpy1.24.3OOM when allocating tensorGPU显存不足或batch size过大减小batch size或降低图像分辨率3. 数据准备与图像预处理3.1 数据集来源与标注规范水稻病虫害识别任务数据集决定模型的上限。常见类别包括稻瘟病、稻曲病、白叶枯病、纹枯病、胡麻叶斑病等。每类的图片数量建议不少于500张最好在1000张以上类别数建议保持在5-8个之间太多类别容易混淆太少又体现不出项目价值。数据来源有几种方式一是直接从公开数据集下载比如Kaggle上的水稻病害数据集、PlantVillage数据集但PlantVillage里的水稻图片种类比较少二是从农业论文的附件数据里找很多论文会在补充材料里放原始图片三是自己到田里拍这个最花时间但数据质量最好。标注规范上用文件夹名作为类别标签是最省事的方式。比如dataset/train/rice_blast/下的所有图片都视为稻瘟病样本。图片的格式建议统一处理成JPG因为部分PNG图片有透明度通道读入时可能会报错。3.2 数据增广策略与代码实现如果每类的原始图片只有几百张直接训练很容易过拟合。数据增广是深度学习解决小样本问题的经典手段。用Keras的ImageDataGenerator可以一行代码实现增广核心思路是随机对图片做旋转、翻转、平移、亮度调整、缩放、裁剪让模型看到更多样化的数据提升泛化能力。我实际使用的增广参数如下from tensorflow.keras.preprocessing.image import ImageDataGenerator train_datagen ImageDataGenerator( rescale1./255, rotation_range30, # 随机旋转30度 width_shift_range0.2, # 水平平移20% height_shift_range0.2, # 垂直平移20% shear_range0.2, # 剪切变换 zoom_range0.2, # 随机缩放 horizontal_flipTrue, # 水平翻转 brightness_range[0.8, 1.2], # 亮度调整 fill_modenearest ) val_datagen ImageDataGenerator(rescale1./255)这里有个细节验证集和测试集只做归一化不做增广。因为增广的目的是让模型在训练时看到更多变化而验证集需要保留真实的数据分布来评估模型泛化能力。如果把增广也用在了验证集上评估结果会不真实模型看起来很好实际部署就露馅。rescale1./255是将像素值从0-255缩放到0-1这是几乎所有图像分类任务的标准预处理方式。如果不做归一化梯度更新会很慢训练容易不收敛。3.3 数据划分与图片读取管道数据划分直接使用flow_from_directory它会根据文件夹结构自动生成标签train_generator train_datagen.flow_from_directory( dataset/train, target_size(224, 224), # 网络输入尺寸 batch_size32, class_modecategorical ) val_generator val_datagen.flow_from_directory( dataset/val, target_size(224, 224), batch_size32, class_modecategorical )target_size设为224x224是因为它正好是ImageNet预训练模型的默认输入尺寸大部分卷积网络都按这个尺寸设计。如果你用的是其他模型输入尺寸要跟着调整比如EfficientNet的输入是240x240或260x260。class_modecategorical表示使用one-hot编码做多分类输出对应的输出层激活函数必须用softmax。batch_size的选择要根据显存来我用的12G显存batch_size32比较合适显存不够就降到16或8但不要低于8否则梯度更新太频繁训练不稳定。4. 模型搭建与训练调优4.1 迁移学习MobileNetV2作为骨干网络水稻病害识别不是大规模图像分类任务从头训练一个卷积神经网络既费时间又容易过拟合。迁移学习是更明智的选择。核心思路是先用ImageNet数据集上预训练的模型提取通用特征边缘、纹理、形状然后在自己的数据集上微调。我选MobileNetV2作为骨干网络理由有三。第一它参数量小推理速度快适合以后部署到Web服务。第二它在ImageNet上的特征提取能力足够强泛化能力已经验证过。第三Keras内置的MobileNetV2代码非常简洁不需要自己实现网络结构。from tensorflow.keras.applications import MobileNetV2 from tensorflow.keras.models import Model from tensorflow.keras.layers import Dense, GlobalAveragePooling2D, Dropout base_model MobileNetV2( weightsimagenet, include_topFalse, # 去掉顶部分类层 input_shape(224, 224, 3) ) # 冻结预训练权重 base_model.trainable False x base_model.output x GlobalAveragePooling2D()(x) x Dense(128, activationrelu)(x) x Dropout(0.5)(x) predictions Dense(num_classes, activationsoftmax)(x) model Model(inputsbase_model.input, outputspredictions)include_topFalse表示去掉ImageNet分类的1000类输出层换上自己的分类层。GlobalAveragePooling2D替代Flatten大幅减少参数量同时保留空间特征这是经典的微调方案。Dropout设为0.5在数据量不多的情况下有效抑制过拟合这个0.5是经验值太高了欠拟合太低了没效果。这里有个很容易忽略的点base_model.trainable False。冻结预训练层只训练新增的全连接层先让新层学会从特征中做分类等损失下降得差不多了再解冻部分底层做微调精度还能往上提一截。如果一开始就解冻全部层训练预训练权重会被冲坏效果反而更差。4.2 训练参数与回调设置迁移学习分两个阶段训练。第一阶段只用训练新加的分类层优化器用Adam学习率设高一点比如1e-3训练20个epoch。第二阶段解冻MobileNetV2的深层部分再训练20-30个epoch学习率降低到1e-5或5e-6避免破坏已经学好的特征。完整的训练代码如下from tensorflow.keras.optimizers import Adam from tensorflow.keras.callbacks import ModelCheckpoint, EarlyStopping, ReduceLROnPlateau model.compile( optimizerAdam(learning_rate1e-3), losscategorical_crossentropy, metrics[accuracy] ) checkpoint ModelCheckpoint( models/rice_disease_model.h5, monitorval_accuracy, save_best_onlyTrue, verbose1 ) early_stop EarlyStopping( monitorval_loss, patience5, restore_best_weightsTrue ) reduce_lr ReduceLROnPlateau( monitorval_loss, factor0.5, patience3, min_lr1e-7 ) history model.fit( train_generator, steps_per_epochtrain_generator.samples // 32, epochs30, validation_dataval_generator, validation_stepsval_generator.samples // 32, callbacks[checkpoint, early_stop, reduce_lr] )ModelCheckpoint里的save_best_onlyTrue很关键只保存验证集精度最高的模型避免训练后期过拟合导致的模型回退。EarlyStopping设置patience5连续5个epoch验证损失不下降就停止训练能省很多时间。ReduceLROnPlateau当验证损失连续3个epoch不下降时学习率减半让训练在接近收敛时更精细。第二阶段微调的代码解冻部分层降低学习率继续训练。# 解冻base_model的后半部分 base_model.trainable True for layer in base_model.layers[:100]: layer.trainable False model.compile( optimizerAdam(learning_rate5e-6), losscategorical_crossentropy, metrics[accuracy] )4.3 训练过程中的指标监控与调优心得训练过程中要实时关注训练集和验证集的loss变化。如果训练集loss继续下降但验证集loss反而上升说明过拟合了需要扩大数据增广、增加Dropout或者早停。如果训练集和验证集loss都降不下去可能是学习率太大或模型太浅。我训练过程中得到的一组参考指标第一阶段在20个epoch时验证集准确率稳定在90%左右第二阶段微调后能到95%以上。实际使用中这个精度水平基本能满足识别需求。如果数据集比较大比如每类有几千张图可以考虑换成EfficientNet系列精度还会高一些但需要更大显存。另外训练中断了不要慌张有ModelCheckpoint在任何时候都能从保存的模型继续。关键是训练完一定要检查一下模型文件是否能正常加载遇到过有人训练了一晚上发现模型文件没有保存成功那叫一个惨。5. Flask Web应用实现与API设计5.1 模型加载与预测函数封装Flask端的核心逻辑是加载训练好的模型接收用户上传的图片预处理图片传入模型做预测返回识别结果。模型加载建议放在全局避免每次请求都重新加载模型不然卡到怀疑人生。import numpy as np from tensorflow.keras.models import load_model from tensorflow.keras.preprocessing import image from PIL import Image import io MODEL_PATH models/rice_disease_model.h5 model None CLASS_NAMES [稻瘟病, 东格鲁病, 白叶枯病, 健康] def load_global_model(): global model if model is None: model load_model(MODEL_PATH) return model def predict_image(img_bytes): img Image.open(io.BytesIO(img_bytes)).convert(RGB) img img.resize((224, 224)) img_array np.array(img) / 255.0 img_array np.expand_dims(img_array, axis0) model load_global_model() predictions model.predict(img_array, verbose0)[0] top_indices np.argsort(predictions)[::-1][:3] results [] for idx in top_indices: results.append({ disease: CLASS_NAMES[idx], confidence: round(float(predictions[idx]) * 100, 2) }) return results这里有几个容易出错的点。第一图片预处理必须和训练时完全一致。训练时用了rescale1./255预测时也要直接将像素除以255否则模型看到的数据分布和训练时不一致预测结果会莫名其妙地差。第二Image.open().convert(RGB)是必须的用户上传的图片可能是带透明度通道的PNG也可能是RGBA四通道图不转换直接传给模型会报维度错误。第三np.expand_dims是把二维图片扩展到四维批次维度因为Keras的predict需要一批输入即使只有一张图片也要加一个维度。5.2 路由设计与前端交互Flask应用包含两个核心路由/用于返回首页HTML页面/predict用于接收POST请求处理图片上传并返回JSON结果。from flask import Flask, request, jsonify, render_template import os app Flask(__name__) app.config[MAX_CONTENT_LENGTH] 16 * 1024 * 1024 # 限制16MB UPLOAD_FOLDER static/uploads os.makedirs(UPLOAD_FOLDER, exist_okTrue) app.route(/) def index(): return render_template(index.html) app.route(/predict, methods[POST]) def predict(): if file not in request.files: return jsonify({error: 未上传文件}), 400 file request.files[file] if file.filename : return jsonify({error: 文件名为空}), 400 try: img_bytes file.read() results predict_image(img_bytes) # 保存上传的图片用于前端展示 filepath os.path.join(UPLOAD_FOLDER, file.filename) file.save(filepath) return jsonify({ success: True, image_url: / filepath.replace(os.sep, /), results: results }) except Exception as e: return jsonify({error: str(e)}), 500 if __name__ __main__: load_global_model() app.run(host0.0.0.0, port5000, debugFalse)前端页面用简单的HTML JavaScript用户选择图片上传AJAX发送请求展示返回的识别结果列表。核心思路是在fetch请求里带上FormData不用刷新页面就能拿到JSON结果。这个交互体验在本地测试够用了。5.3 性能优化与生产化注意点用app.run()跑起来的Flask服务是开发服务器单线程处理请求并发能力很弱。但作为课程设计、个人项目或小型工具已经够用。如果要部署到生产环境有几个方向可以优化。第一个方向是用多线程模式跑Flaskapp.run(threadedTrue)处理并发请求的能力提升明显。第二个方向是把模型推理放到单独的线程池里避免GPU推理阻塞其他请求。第三个方向是换成生产级WSGI服务器比如Gunicorn配合多个worker或者uWSGI。模型推理的性能瓶颈主要在两处图片读取和模型前向传播。图片读取用PIL的Image.open本身不慢但如果入口图片很大比如手机拍的原图有4000x3000像素读取和resize会比较耗时。建议在路由层就对上传的图片做大小控制先用PIL把图片resize到合理的尺寸再送入网络同时也能减少内存占用。还有个隐藏问题要注意文件名安全。用户上传的文件名可能是中文、包含路径分隔符如果直接拼接保存路径存在路径穿越风险。虽然Flask的secure_filename不是万能的但至少要做一层过滤。6. 常见问题与排查技巧实录6.1 环境与安装阶段遇到的问题这个阶段的问题最常见而且报错信息往往不直观。TensorFlow提示Could not find a version that satisfies the requirement tensorflow基本都是Python版本和TensorFlow版本不对应。比如Python 3.12搭配TensorFlow 2.10就会报这个错解决方式是降Python版本或升TensorFlow版本。提示No module named tensorflow但明明pip install了80%的情况是装错了环境。conda里创建了rice_disease环境但pycharm终端或者VSCode终端跑的是base环境。检查方式是在终端输入which python看当前解释器的路径是不是你虚拟环境里的路径。GPU相关报错Could not load dynamic library cudnn64_8.dll说明cuDNN缺失或版本不匹配。解决方式是安装对应版本的cuDNN并把bin目录加入系统PATH。如果你装了CUDA 11.2但没装cuDNNTensorFlow导入到GPU相关代码时也会报错。6.2 训练过程遇到的坑训练时遇到最典型的问题是loss从一开始就居高不下怎么降都降不动。有一半情况下是数据问题标签和图片对应错了比如某个类别的目录里混了一张其他类别的图片模型就会被带偏。排查方式很简单找一张图片出来看看多抽查几个样本。另一半情况是模型结构问题比如输出层神经元数量不等于类别数或者激活函数用错。多分类任务输出层必须用softmax如果错用sigmoidloss很难收敛到理想值。class_weight也有影响。如果某些病害类的样本特别少模型会偏向预测样本多的类别这时候要给少样本类别更高权重。在fit里设置class_weight或直接用flow_from_directory的class_modecategorical配合自己计算权重。还有一个训练阶段的高频问题OutOfMemoryError。显存溢出处理方式有三种降低batch_size、降低图片分辨率、换显存更大的显卡。在你的数据量不大的情况下用CPU训练也能接受只是慢一些。6.3 Flask部署阶段遇到的坑Flask部署阶段的坑主要集中在模型加载和图片格式问题上。模型加载报错Unable to load weights from the checkpoint file最常见的原因是保存模型用的Keras/TensorFlow版本和加载环境不一致。解决方式有两种一是训练和部署用同一个虚拟环境二是在保存模型时用model.save(xxx.h5)这个格式兼容性最好。TensorFlow 2.6以后还可以保存为.keras格式但这种格式在旧版本里加载不了。预测结果全是同一类别很可能是模型本身有严重过拟合或者数据分布不均衡。可以尝试用之前保存的最优模型文件而不是最后一次训练的模型。ModelCheckpoint保存的best模型通常比最后一轮模型效果更好。上传图片提示TypeError: NoneType object is not callable一般是文件读取方式有问题。用file.read()拿到的二进制字节流可以直接用Image.open(io.BytesIO(...))打开不要用file.save()保存后再去读路径绕了一圈容易出错。Flask启动后提示Address already in use端口被占用了用netstat -ano | findstr 5000找到占用的进程杀掉或换端口。6.4 预测结果不准确时的排查思路如果模型在验证集上准确率很高但实际使用预测结果差首先排查预处理链路。检查上传图片的预处理方式和训练时是否一致尺寸、归一化、通道顺序这三项每一项不一致都会导致预测错误。其次排查类别顺序。训练时flow_from_directory按字母顺序给类别编码比如bacterial_blight是0healthy是1。如果你在Flask端手动写了CLASS_NAMES列表顺序必须和训练时使用的顺序完全一致否则类别对应错位预测结果全是张冠李戴。最后检查图片质量。实际用户上传的图片可能是模糊的、过暗的、手动截图的带边框图片。这些情况在训练数据里很少模型预测不准是正常的。建议在Web端做一个简单的图片质量提示如果清晰度太低就提示用户重新拍摄。7. 后续扩展方向项目的核心功能已经跑通扩展方向可以从识别范围、模型性能和部署形态三个角度考虑。识别范围方面当前只做了叶片病害识别可以扩展到水稻稻穗病虫害、稻飞虱、稻纵卷叶螟等甚至可以把营养元素缺乏的症状也加进来做成一个完整的水稻生育期诊断系统。模型性能方面可以考虑用TensorRT或ONNX Runtime做推理加速。实测下来在同样的硬件上ONNX Runtime比TensorFlow自带的推理速度快20%-30%而且内存占用更低。部署形态方面可以开发一个微信小程序端用Flask只是做后端API前端换成小程序。这样农户使用门槛会更低识别体验也更符合日常习惯。小程序端和Web端的后端接口设计是一样的代码可以复用。模型解释性方面可以集成Grad-CAM热力图可视化让用户看到模型是根据叶片的哪个区域做出的判断。这对植保人员有参考价值也能增加系统可信度。在我实际使用这套系统的时候最大的体会是深度学习项目的难点不在模型结构而在数据质量和部署细节。数据不干净再强大的模型也白搭模型训练好了部署环节如果预处理不一致线上预测效果照样崩。所以新手在做类似项目时一定要在数据清洗和预处理一致性上多花时间这个投入的回报率是最高的。最后分享一个小技巧训练结束后用model.summary()查看一下模型总参数量和层结构把打印结果截图保存下来。答辩、写报告、做项目展示的时候这个信息非常有用能直接证明你确实理解了自己搭的模型。本文还有配套的精品资源点击获取

读完文章,也想定制专属网站?

尧图设计师 24 小时内与您沟通定制方案

免费获取报价