You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何修改二分类缺陷检测CNN模型为三分类并实现单张图片测试

二分类CNN调整为三分类缺陷检测方案

1. 模型结构与训练逻辑修改

仅需要修改输出层和损失函数配置,修改后的完整代码如下:

import tensorflow as tf

inputs = tf.keras.Input(shape=(120, 120, 3))
x = tf.keras.layers.Conv2D(filters=16, kernel_size=(3, 3), activation='relu')(inputs)
x = tf.keras.layers.MaxPool2D(pool_size=(2, 2))(x)
x = tf.keras.layers.Conv2D(filters=32, kernel_size=(3, 3), activation='relu')(x)
x = tf.keras.layers.MaxPool2D(pool_size=(2, 2))(x)
x = tf.keras.layers.GlobalAveragePooling2D()(x)
# 输出层修改为3个神经元,激活函数改为softmax
outputs = tf.keras.layers.Dense(3, activation='softmax')(x)

model = tf.keras.Model(inputs=inputs, outputs=outputs)

model.compile(
    optimizer='adam',
    # 损失函数替换:标签为整数编码用sparse_categorical_crossentropy,one-hot编码用categorical_crossentropy
    loss='sparse_categorical_crossentropy',
    metrics=['accuracy']
)

print(model.summary())

# 可选优化:复用二分类预训练权重加快收敛
# model.load_weights('你的二分类模型权重文件路径.h5', by_name=True, skip_mismatch=True)
# 可冻住前几层卷积层,仅训练最后的全连接层
# for layer in model.layers[:-1]:
#     layer.trainable = False

history = model.fit(
    train_data,
    validation_data=val_data,
    epochs=100,
    callbacks=[
        tf.keras.callbacks.EarlyStopping(
            monitor='val_loss',
            patience=3,
            restore_best_weights=True
        )
    ]
)

2. 数据集适配要求

  • 统一三类标签编码规则,例如固定映射为 {'big':0, 'small':1, 'other':2},训练、验证、测试集必须使用完全一致的映射规则
  • 若使用Keras自带的数据集加载工具,将class_names参数设置为['big', 'small', 'other']即可自动完成标签编码
  • 图像预处理规则和二分类阶段保持完全一致,例如归一化比例、尺寸缩放逻辑无需修改

3. 单张图像预测实现

import cv2
import numpy as np

# 与训练阶段完全一致的类别映射
label_map = {0: 'big', 1: 'small', 2: 'other'}

def predict_single_image(img_path, model, target_size=(120, 120)):
    # 读取并预处理图像
    img = cv2.imread(img_path)
    # opencv默认读取为BGR格式,若训练时用RGB格式需要转换
    img = cv2.cvtColor(img, cv2.COLOR_BGR2RGB)
    img = cv2.resize(img, target_size)
    # 与训练时一致的归一化操作,若训练时未除以255可删除该行
    img = img / 255.0
    # 增加batch维度适配模型输入要求
    img_input = np.expand_dims(img, axis=0)
    
    # 推理预测
    pred_prob = model.predict(img_input, verbose=0)[0]
    pred_idx = np.argmax(pred_prob)
    
    return label_map[pred_idx], round(float(pred_prob[pred_idx]), 4)

# 调用示例
# pred_label, pred_conf = predict_single_image("测试图片路径.jpg", 训练好的模型对象)
# print(f"预测类别:{pred_label},置信度:{pred_conf}")

内容的提问来源于stack exchange,提问作者mahshid a.j

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.10.05 00:06:03