十年匠心定制 · 商业建站与技术教学双线并行 咨询热线:400-886-1026 service@lmnt.cn
ARTICLE DETAIL

资讯详情

深耕网站建设与运营推广的一线实战洞察。

TensorFlow图像分类完整代码模板:从CNN到迁移学习与Transformer实战

TensorFlow图像分类完整代码模板:从CNN到迁移学习与Transformer实战 1. 模板不是背代码而是建立一套可复用的图像分类流水线先说个现象。很多初学者一上来就盯着 LeNet、ResNet 的论文复现把网络结构抄了一遍又一遍但真到自己的数据集上数据加载写不清楚、训练日志看不懂、模型保存加载一塌糊涂。折腾半个月准确率上不去也不知道该调哪里。我这些年用 TensorFlow 做图像分类项目从 Kaggle 比赛到工业质检再到遥感影像识别最深的体会是真正卡住进度的从来不是模型结构本身而是模型之外的整套流程。数据怎么组织、训练怎么配置、日志怎么记录、模型怎么保存和部署这些环节只要有一个不顺手整个项目就会被拖住。所以这篇文章不是给你讲某个花哨的网络结构而是给出一套我实测过很多次的 TensorFlow 图像分类完整代码模板从环境搭建到训练评估到推理部署每一个环节都有可以直接抄的代码同时把每个关键选择背后的原因讲清楚。你用这套模板去套自己的数据集只需要改数据路径和几个超参数就能跑起来一个完整的图像分类项目。这套模板适合谁刚入门深度学习、被各种教程碎片信息搞晕的新手已经跑通过 Mnist 但不知道怎么迁移到自己数据集上的同学以及工作中需要快速验证一个图像分类想法、不想从零重写流程的工程师。我不保证你用这套模板能拿到 SOTA但我可以保证你会拥有一套结构清晰、能改能调、出了问题知道去哪查的工程化代码。先说清楚这套模板基于 TensorFlow 2.x 的 Keras 高层 API。Keras 的 Sequential 和 Model 子类化接口足够简洁同时保留了底层灵活性是平衡开发效率和功能覆盖面的最佳选择。后面所有代码我都尽量完整给出你复制到自己项目里改改就能用。2. 开工前的环境准备别让环境问题浪费一下午2.1 TensorFlow 安装的环境选型图像分类项目第一步是装环境这一关能卡住不少人。装 TensorFlow 之前先想清楚一个问题你是在个人电脑上学习调试还是要跑正经规模的训练个人学习调试CPU 版本完全够用。很多人一上来就追求 GPU 版本结果驱动、CUDA、cuDNN 版本对不上折腾两三天装不好TensorFlow 代码一行没写。我的建议是先装 CPU 版本把流程跑通代码没问题了再考虑 GPU 加速。图像分类的 MNIST、CIFAR-10 这类小数据集CPU 训练虽然慢一点但完全能接受。真正要上 GPU用 Anaconda 创建独立环境是最稳妥的方案好处是不会把系统 Python 环境搞乱。安装命令很简单conda create -n tf python3.10 conda activate tf pip install tensorflow安装完成后用一段小代码验证环境是否正常import tensorflow as tf print(tf.__version__) print(tf.config.list_physical_devices(GPU))这里有个经验TensorFlow 2.10 是最后一个在 Windows 上原生支持 GPU 的版本之后的版本在 Windows 上要 GPU 得用 WSL2。如果你在 Windows 上开发建议直接装 2.10如果你用 Linux装最新版没问题。另外别用 conda 装 tensorflow-gpu直接把tensorflow装上自带 GPU 支持。2.2 用 Anaconda 管理依赖的实操建议用 Anaconda 不是为了装 TensorFlow 本身而是为了管理项目依赖。我见过太多人把 TensorFlow、PyTorch、各种包全部装进 base 环境最后版本冲突到怀疑人生。推荐的做法是为每个项目建独立环境。除了 TensorFlow图像分类项目还需要这些基础依赖pip install numpy matplotlib scikit-learn pandas pillownumpy 做数组运算matplotlib 画训练曲线和可视化结果scikit-learn 用来算分类报告和混淆矩阵pandas 偶尔处理标签文件pillow 处理图像读取。这些库在图像分类项目中出现的频率极高提前装好省得后面一次次补。环境的可复现性也很重要。项目跑通了以后用pip freeze requirements.txt把依赖导出来换机器或者团队协作时直接pip install -r requirements.txt就能恢复环境。这个习惯能救你于水火之中。3. 数据准备图像分类项目里最容易被低估的环节3.1 数据集目录组织的推荐方案很多教程用 Mnist、CIFAR 这种内置数据集数据自动下载、标签自动对应根本接触不到真实世界的数据处理。但实际项目里你拿到的数据往往是一堆图片放在一个文件夹里标签可能要自己整理。最通用的数据组织方式是按类别分文件夹data/ train/ cat/ cat_001.jpg cat_002.jpg dog/ dog_001.jpg val/ cat/ cat_001.jpg dog/ dog_001.jpg test/ ...这种组织方式的优势在于TensorFlow 的image_dataset_from_directory可以直接读取文件夹名作为类别标签不需要额外写 CSV 或者其他格式的标注文件。从零整理数据时只需要写一个简单的脚本按类别把图片放到对应文件夹里。按类别分文件夹虽然简单但它隐含了一个假设单个文件夹里的图片都是同一个类别。这就引出一个关键问题——数据质量必须提前检查。我有一次做花卉分类发现某个类别的文件夹里混了几张明显是杂草的图片后来排查准确率上不去的原因数据混杂是重要因素。3.2 数据增强让模型更鲁棒图像分类模型动辄几十万、上百万参数而真实项目里的数据往往只有几千张直接训练很容易过拟合。数据增强是解决数据量不足的最有效手段之一没有之一。TensorFlow Keras 里可以用预处理层直接实现数据增强不用在数据管线里单独操作。增强操作在训练时随机执行验证和测试时自动关闭这个机制非常关键。data_augmentation tf.keras.Sequential([ tf.keras.layers.RandomFlip(horizontal), tf.keras.layers.RandomRotation(0.1), tf.keras.layers.RandomZoom(0.1), ])这里选择的增强策略是有讲究的。水平翻转对大多数图像分类任务都安全因为物体在水平方向翻转后还是同一类但垂直翻转就要小心比如识别猫和狗倒过来的猫图可能会让模型学到奇怪的特征。随机旋转的幅度也不能太大旋转 10 度的花还是花旋转 90 度的花可能看起来就像另一类了。一个实用的经验是先从小的增强幅度开始观察训练集和验证集准确率的差距。如果训练集准确率远高于验证集说明过拟合可以加大增强力度如果两者都低说明欠拟合应该先增强模型容量而不是继续加增强。3.3 Dataset API 的高效加载与预处理TensorFlow 的tf.data.DatasetAPI 是官方推荐的高性能数据管线方案。它能把数据读取、预处理、增强、批处理这些操作串成一条流水线还能自动并行处理。用文件夹数据训练最直接的方式是image_dataset_from_directorytrain_ds tf.keras.utils.image_dataset_from_directory( data/train, image_size(224, 224), batch_size32, label_modecategorical, validation_split0.2, subsettraining, seed123 ) val_ds tf.keras.utils.image_dataset_from_directory( data/train, image_size(224, 224), batch_size32, label_modecategorical, validation_split0.2, subsetvalidation, seed123 )这里label_modecategorical生成的是一组 one-hot 编码的标签向量适配后面用 Softmax 做多分类输出的情况。如果类别特别多也可以考虑label_modesparse配合 SparseCategoricalCrossentropy 损失函数内存占用更小。注意validation_split参数需要seed指定随机种子才能保证训练集和验证集是同一套数据拆分出来的。数据流水线构建完成后还需要做两个优化.cache()把预处理后的数据缓存在内存或磁盘上.prefetch()让数据加载和模型训练并行这两个操作能让训练速度有明显提升。train_ds train_ds.map(lambda x, y: (data_augmentation(x), y)).prefetch(tf.data.AUTOTUNE) val_ds val_ds.map(lambda x, y: (x / 255.0, y)).prefetch(tf.data.AUTOTUNE)数据归一化被我放在了这里而不是模型内部是因为image_dataset_from_directory默认输出的像素值是 0 到 255 的整数需要先转成 0 到 1 的浮点数才能喂给模型。用 Lambda 层在模型内部做归一化也不是不行但放在数据管线的语义更清晰也方便调试。4. 模型构建从 CNN 基础到迁移学习到 Transformer4.1 自己搭一个 CNN 基线模型图像分类的模型选择我的习惯永远是先搭一个简单的 CNN 基线模型跑通整个流程再考虑换更强的模型。直接一上来就用 EfficientNet 或者 ViT一旦效果不好你很难判断是模型问题、数据问题还是训练配置问题。而基线模型足以帮你验证数据管线和训练流程是否正常。一个经典的 CNN 基线模型长这样def build_cnn_model(input_shape(224, 224, 3), num_classes10): model tf.keras.Sequential([ tf.keras.layers.Input(shapeinput_shape), tf.keras.layers.Conv2D(32, (3, 3), activationrelu, paddingsame), tf.keras.layers.MaxPooling2D((2, 2)), tf.keras.layers.Conv2D(64, (3, 3), activationrelu, paddingsame), tf.keras.layers.MaxPooling2D((2, 2)), tf.keras.layers.Conv2D(128, (3, 3), activationrelu, paddingsame), tf.keras.layers.MaxPooling2D((2, 2)), tf.keras.layers.GlobalAveragePooling2D(), tf.keras.layers.Dropout(0.5), tf.keras.layers.Dense(num_classes, activationsoftmax) ]) return model这个模型用三个卷积块提取特征每个卷积块后接最大池化降低特征图尺寸最后用全局平均池化把特征压缩成一维向量。这里有两个细节值得说用GlobalAveragePooling2D而不是Flatten好处是参数量大幅减少同时能减轻过拟合。全连接层前加Dropout(0.5)这是最常用的正则化手段训练时随机丢弃一半神经元让模型不过分依赖某个特征。4.2 迁移学习快速提升精度CNN 基线模型在自己搭的时候能跑通但精度天花板很低。如果数据量有限——比如每类只有几千张图——从头训练一个大模型极容易过拟合。这时候工业界最成熟的做法是迁移学习。迁移学习的核心原理很简单一个在大规模数据集比如 ImageNet上预训练好的模型已经学会了通用的视觉特征——边缘、纹理、形状、物体部件。这些特征是通用的可以迁移到你的目标任务上。你只需要在预训练模型后面接上自己的分类头然后微调。TensorFlow Keras 里用预训练模型做迁移学习非常方便base_model tf.keras.applications.ResNet50( weightsimagenet, include_topFalse, input_shape(224, 224, 3) ) base_model.trainable False model tf.keras.Sequential([ base_model, tf.keras.layers.GlobalAveragePooling2D(), tf.keras.layers.Dropout(0.3), tf.keras.layers.Dense(num_classes, activationsoftmax) ])这里把base_model.trainable设为 False意思是冻结预训练模型的参数只训练后面新加的分类层。为什么要先冻结因为预训练模型的参数已经在 ImageNet 上训练好了如果一开始就让它们跟着你的小数据集更新很容易把学好的特征破坏掉这叫灾难性遗忘。先用冻结的方式训几个 epoch等新的分类头收敛得差不多了再解冻部分底层做微调。微调时用很小的学习率通常在 1e-5 到 1e-5 之间让预训练参数缓慢适应你的数据分布。整个策略类似先学走路再学跑步。选择哪个预训练模型我的经验是ResNet50 是最稳妥的默认选择训练速度适中、精度可靠、资源占用不算夸张。EfficientNet 系列在精度和效率的权衡上更优但训练时对图片尺寸和缩放比较敏感。MobileNetV3 适合部署到移动端或边缘设备。如果你的数据有上百万张或者有充足的 GPU 资源再考虑 ViT 变体。4.3 Transformer 架构与传统 CNNs 的取舍ViTVision Transformer最近两年很火它把 NLP 里的 Transformer 架构搬到了图像上做法是把图片切成固定大小的 patch每个 patch 线性映射成 token然后交给 Transformer Encoder 处理。热词里提到 Transformer 图像分类这里多说几句。ViT 的优势在于它能建模图像中远距离像素之间的关系这是 CNN 的卷积核受限于局部感受野而做不到的。但这不意味着 ViT 在所有场景下都碾压 CNN。ViT 的问题也很明显它需要大量数据才能训好。在 ImageNet 这种大规模数据集上 ViT 确实能超过 CNN但在小数据集上效果反而不如 ResNet 这类有归纳偏置的模型。在小数据集上做迁移学习TensorFlow 官方提供的ViTBase16等预训练权重也可以直接用但微调时的细节比 CNN 多不少。我的建议是默认场景先用 ResNet50 跑迁移学习这个方案在大多数实际项目里已经够用。只有当你确认数据量足够大、或者特征确实需要全局建模时再考虑 ViT。技术选型的关键不是哪个模型更新而是哪个模型在数据、算力、精度三者的约束下综合最优。5. 训练配置与完整流程参数背后的推演逻辑5.1 模型编译时如何选损失函数和优化器编译模型是训练前的最后一步这一步的选择直接决定优化过程能否收敛。对于多分类图像识别任务标准的搭配是Softmax 输出层 分类交叉熵损失 Adam 优化器。model.compile( optimizertf.keras.optimizers.Adam(learning_rate1e-4), losstf.keras.losses.CategoricalCrossentropy(), metrics[accuracy] )交叉熵损失函数衡量的是预测概率分布和真实标签分布之间的差距它的特点是梯度在模型完全错误时大、接近正确时小很适合分类任务的训练。如果你的标签用的是sparse模式而不是 one-hot损失函数要换成SparseCategoricalCrossentropy否则会报错。学习率从 1e-4 起步是迁移学习常用的策略。从头训练 CNN 基线模型时可以用 1e-3因为所有参数都是随机初始化的需要更大的步长快速收敛。而迁移学习时预训练参数已经比较接近最优点学习率太大一步就跨过头了反而破坏已有特征。5.2 回调函数与训练日志监控训练模型时要实时看指标不能训练完了才知道结果。Keras 的回调机制提供了训练过程中执行额外操作的钩子我最常用的三个是 ModelCheckpoint、EarlyStopping、ReduceLROnPlateau。callbacks [ tf.keras.callbacks.ModelCheckpoint( best_model.keras, monitorval_accuracy, save_best_onlyTrue, modemax ), tf.keras.callbacks.EarlyStopping( monitorval_loss, patience10, restore_best_weightsTrue ), tf.keras.callbacks.ReduceLROnPlateau( monitorval_loss, factor0.5, patience3, min_lr1e-6 ) ]ModelCheckpoint 的职责是在训练过程中自动保存表现最好的模型monitorval_accuracy表示监控验证集准确率save_best_onlyTrue保证磁盘里永远只留最优的那一版模型。EarlyStopping 在验证损失连续patience个 epoch 没有下降时提前终止训练既能省时间又能防止后期过拟合。ReduceLROnPlateau 在验证损失陷入平台期时把学习率减半帮助模型跨过局部最优。这三个回调配合起来的效果等于有一个细心的助手在训练过程中帮你盯指标、自动做决策。5.3 训练执行参数怎么定训练执行代码本身很简单history model.fit( train_ds, validation_dataval_ds, epochs50, callbackscallbacks )epochs 我一般设一个较大的值比如 50配合 EarlyStopping 让它自动在合适的时候停下来。batch_size在数据加载阶段已经设置好了这里不用重复。如果你的数据集特别大还可以设置steps_per_epoch和validation_steps来限制每个 epoch 跑多少步但默认情况下不需要动。训练完成后用history.history能拿到每一个 epoch 的损失和准确率曲线把训练集和验证集的两条曲线画在一起能直观地判断模型是否过拟合——训练准确率持续上升、验证准确率停滞或下降就是过拟合的典型信号。6. 模型评估与推理训练完只是开始6.1 测试集评估与混淆矩阵训练完成后一定要在独立的测试集上做最终评估不能只用验证集结果说话。验证集在训练过程中已经参与了 EarlyStopping 等决策有信息泄露的风险测试集才能反映模型在没见过数据上的真实表现。test_ds tf.keras.utils.image_dataset_from_directory( data/test, image_size(224, 224), batch_size32, label_modecategorical, shuffleFalse ) test_loss, test_acc model.evaluate(test_ds) print(fTest accuracy: {test_acc:.4f})整体准确率只能说明模型的平均表现掩盖了类别之间的差异。分类报告和混淆矩阵能告诉你模型在哪些类别上表现好、在哪些类别上容易混淆。数据集类别不均衡时这一点尤其关键——有可能整体准确率 95%但某个稀有类别几乎全错。import numpy as np from sklearn.metrics import classification_report, confusion_matrix y_true np.concatenate([y.numpy().argmax(axis1) for x, y in test_ds]) y_pred np.concatenate([model.predict(x).argmax(axis1) for x, y in test_ds]) print(classification_report(y_true, y_pred, target_namesclass_names)) print(confusion_matrix(y_true, y_pred))6.2 保存、加载与单张图片推理ModelCheckpoint保存的是训练过程中最好的模型权重和结构但在部署场景下我更推荐把完整的模型保存为一个独立文件这样加载时不需要重新定义模型结构model.save(final_model.h5)加载模型并做单张图片推理代码也很直接loaded_model tf.keras.models.load_model(final_model.h5) img tf.keras.utils.load_img(test_single.jpg, target_size(224, 224)) img_array tf.keras.utils.img_to_array(img) img_array tf.expand_dims(img_array, 0) img_array / 255.0 predictions loaded_model.predict(img_array) predicted_class np.argmax(predictions[0]) confidence np.max(predictions[0])关于保存格式多说一句TensorFlow 2.x 里默认推荐的是.keras格式它把模型结构和权重打包在一个文件里跨版本兼容性更好。.h5格式是旧版 HDF5 格式虽然也能用但新项目建议直接用.keras。7. 完整代码模板一套可以直接复制修改的工程这一节把前面所有环节整合成一个完整的训练脚本你可以保存为train.py只需要修改开头的配置参数就能直接跑。import tensorflow as tf import numpy as np from sklearn.metrics import classification_report, confusion_matrix # ---------- 配置参数 ---------- TRAIN_DIR data/train TEST_DIR data/test IMG_SIZE (224, 224) BATCH_SIZE 32 EPOCHS 50 NUM_CLASSES 10 # 根据类别数修改 LEARNING_RATE 1e-4 MODEL_SAVE_PATH best_model.keras # ---------- 数据加载 ---------- train_ds tf.keras.utils.image_dataset_from_directory( TRAIN_DIR, image_sizeIMG_SIZE, batch_sizeBATCH_SIZE, label_modecategorical, validation_split0.2, subsettraining, seed123 ) val_ds tf.keras.utils.image_dataset_from_directory( TRAIN_DIR, image_sizeIMG_SIZE, batch_sizeBATCH_SIZE, label_modecategorical, validation_split0.2, subsetvalidation, seed123 ) test_ds tf.keras.utils.image_dataset_from_directory( TEST_DIR, image_sizeIMG_SIZE, batch_sizeBATCH_SIZE, label_modecategorical, shuffleFalse ) class_names train_ds.class_names # ---------- 数据增强 ---------- data_augmentation tf.keras.Sequential([ tf.keras.layers.RandomFlip(horizontal), tf.keras.layers.RandomRotation(0.1), tf.keras.layers.RandomZoom(0.1), ]) AUTOTUNE tf.data.AUTOTUNE train_ds train_ds.map( lambda x, y: (data_augmentation(x), y), num_parallel_callsAUTOTUNE ).prefetch(AUTOTUNE) val_ds val_ds.map( lambda x, y: (x / 255.0, y), num_parallel_callsAUTOTUNE ).prefetch(AUTOTUNE) test_ds test_ds.map( lambda x, y: (x / 255.0, y), num_parallel_callsAUTOTUNE ).prefetch(AUTOTUNE) # ---------- 构建迁移学习模型 ---------- base_model tf.keras.applications.ResNet50( weightsimagenet, include_topFalse, input_shape(*IMG_SIZE, 3) ) base_model.trainable False model tf.keras.Sequential([ base_model, tf.keras.layers.GlobalAveragePooling2D(), tf.keras.layers.Dropout(0.3), tf.keras.layers.Dense(NUM_CLASSES, activationsoftmax) ]) model.compile( optimizertf.keras.optimizers.Adam(learning_rateLEARNING_RATE), losstf.keras.losses.CategoricalCrossentropy(), metrics[accuracy] ) model.summary() # ---------- 回调 ---------- callbacks [ tf.keras.callbacks.ModelCheckpoint( MODEL_SAVE_PATH, monitorval_accuracy, save_best_onlyTrue, modemax ), tf.keras.callbacks.EarlyStopping( monitorval_loss, patience10, restore_best_weightsTrue ), tf.keras.callbacks.ReduceLROnPlateau( monitorval_loss, factor0.5, patience3, min_lr1e-6 ) ] # ---------- 训练 ---------- history model.fit( train_ds, validation_dataval_ds, epochsEPOCHS, callbackscallbacks ) # ---------- 评估 ---------- best_model tf.keras.models.load_model(MODEL_SAVE_PATH) test_loss, test_acc best_model.evaluate(test_ds) print(fTest accuracy: {test_acc:.4f}) y_true np.concatenate([y.numpy().argmax(axis1) for x, y in test_ds]) y_pred np.concatenate([best_model.predict(x).argmax(axis1) for x, y in test_ds]) print(classification_report(y_true, y_pred, target_namesclass_names)) print(confusion_matrix(y_true, y_pred))这套模板的核心设计原则是配置和逻辑分离。所有可能要改的参数集中在文件开头的常量区块里后面没有任何硬编码的魔法数字。换数据集时改目录路径、图片尺寸、类别数和学习率基本就可以了。8. 踩坑自查表我实际项目中遇到的问题8.1 过拟合训练集准确率接近 100%验证集却卡在 70% 上不去。这是图像分类最经典的问题解决方案按优先级排序增加数据增强强度这是最直接的手段然后加大 Dropout 比例从 0.3 提到 0.5 甚至 0.6如果数据量允许解冻更多预训练层做微调让模型用更多参数拟合任务。8.2 数据类别不均衡某个类别的图片数量是其他类别的好几倍时模型会偏向数量多的类别。解决思路有三个层面一是数据层面对少数类做过采样或者数据增强加倍二是损失函数层面给少数类更高的权重三是评估指标层面别只看准确率多关注精确率、召回率和 F1 分数。8.3 训练显存不足ResourceExhaustedError是我见过最多的报错。最直接的解法是减小BATCH_SIZE从 32 减到 16 甚至 8。如果模型还是太大考虑换小一点的预训练模型比如 ResNet50 换成 MobileNetV3。另外确认一下是不是有其他程序占用显存用nvidia-smi可以查看。8.4 图像预处理不一致训练时用了数据增强验证测试时只做了归一化这个不一致是对的。容易出错的是忘了对测试集做归一化或者归一化的方法不一致——训练时x / 255.0测试时忘了除模型输出就会全部乱掉。常见的问题包括忘记归一化、归一化尺度不一致、训练和测试图片尺寸不一致。8.5 从 CSV 或者 DataFrame 加载数据如果你的标签数据存在于 CSV 而不是文件夹名里image_dataset_from_directory就不适用了。此时应该用tf.data.Dataset.from_tensor_slices配合自定义的加载函数def load_image(image_path, label): image tf.io.read_file(image_path) image tf.image.decode_jpeg(image, channels3) image tf.image.resize(image, IMG_SIZE) image / 255.0 return image, label dataset tf.data.Dataset.from_tensor_slices((image_paths, labels)) dataset dataset.map(load_image, num_parallel_callsAUTOTUNE).batch(BATCH_SIZE)9. 模板还能怎么扩展从单标签到多标签到目标检测这套模板做的是单标签多分类也就是一张图只属于一个类别。但实际项目里经常会遇到更复杂的需求一张图同时包含多个标签比如一幅画面里既有猫又有狗或者不仅要分类还要定位出物体位置那就变成目标检测问题。多标签分类的改动不复杂最后一层改成 Sigmoid 而不是 Softmax损失函数换成 BinaryCrossentropy评估指标换成精确率和召回率。目标检测则要从头换一套框架TensorFlow 官方的 Object Detection API 或者 Keras CV 的 bounding box 模块都值得看看。另外在部署环节如果要把模型放到服务端做在线推理推荐把模型转换成 TensorFlow Lite 格式.tflite或者用 TensorFlow Serving 做标准化的模型服务。转换代码很简单converter tf.lite.TFLiteConverter.from_keras_model(best_model) tflite_model converter.convert() with open(model.tflite, wb) as f: f.write(tflite_model)这套模板不是我凭空设计的而是从一个个实际项目的坑里刨出来的。我的个人体会是模板的价值不在于让你少写代码而在于让你少做决策。每一步都已经有相对成熟的默认选择你只需要在真正需要偏离的地方花心思把精力集中在数据质量和模型调优上而不是浪费在环境配置和代码报错上。最后再分享一个小技巧训练完成后用model.predict随机挑几张测试图片把预测概率分布打印出来。你会发现即使最终分类正确模型在相近类别上的置信度也会给出非常有价值的信息——比如一个猫的样本在豹子上也有不低的得分说明训练数据里可能存在相似度极高的样本这往往是进一步改进模型的重要线索。
返回列表