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

加载6分类Keras VGG16模型却仅实现二分类的问题排查与解决

问题分析与优化方案

核心问题排查

推理代码中的关键错误

  • 已弃用方法误用:model.predict_classes()在TensorFlow 2.x中已被移除,若模型最后一层用sigmoid激活,该方法会默认以0.5为阈值将输出二值化,直接导致仅输出两类结果。
  • 概率计算错误:np.argmax(model.predict(img), axis=-1)得到的是类别索引,而非概率值,后续乘以100显示百分比完全错误。
  • 重复预测浪费资源:连续两次调用model.predict(img),既降低效率也可能因随机性导致结果不一致。
  • 缺失输入预处理:VGG16训练时要求输入归一化(如像素值缩至[0,1]或匹配ImageNet均值/std),当前代码直接用原始像素值,输入分布与训练时不匹配,严重影响分类精度。
  • 冗余条件判断:6个elif分支逻辑完全重复,可大幅简化。

模型训练阶段可能的隐患

  • 最后一层激活函数错误:多分类任务需用softmax输出类别概率分布,若误用sigmoid,模型会按独立二分类方式输出,无法完成多分类。
  • 损失函数不匹配:多分类应使用sparse_categorical_crossentropy(标签为整数)或categorical_crossentropy(标签为独热编码),若用binary_crossentropy,模型会被训练为二分类任务。
  • 数据不平衡:若训练集中某两类样本占比过高,模型会偏向学习这两类特征,导致其他类别识别失效。

优化后的完整代码

import tensorflow as tf
import numpy as np
import cv2
from tensorflow.keras.models import load_model

# 加载人脸检测器
facedetect = cv2.CascadeClassifier('haarcascade_frontalface_default.xml')

# 初始化摄像头
cap = cv2.VideoCapture(0)
cap.set(3, 640)
cap.set(4, 480)
font = cv2.FONT_HERSHEY_COMPLEX

# 加载模型并重新编译(确保损失函数和激活匹配)
model = load_model('model/FAceRec.h5', compile=False)
# 根据实际训练标签类型选择损失函数,此处假设标签为整数形式
model.compile(optimizer='adam', loss='sparse_categorical_crossentropy', metrics=['accuracy'])

# 类别映射
class_names = ["Karuna", "Suhail", "Uyaam", "Saftan", "Ahmad", "Nikzaad"]

while True:
    success, img_original = cap.read()
    if not success:
        break
    
    # 检测人脸
    faces = facedetect.detectMultiScale(img_original, 1.3, 5)
    for x, y, w, h in faces:
        # 裁剪人脸区域
        crop_img = img_original[y:y+h, x:x+w]
        # 调整尺寸匹配模型输入
        img = cv2.resize(crop_img, (224, 224))
        # VGG16输入归一化(匹配训练时的预处理逻辑)
        img = img / 255.0
        img = np.expand_dims(img, axis=0)  # 增加batch维度
        
        # 单次预测获取概率分布
        predictions = model.predict(img, verbose=0)
        class_index = np.argmax(predictions, axis=-1)[0]
        confidence = np.max(predictions) * 100  # 获取对应类别的置信度
        
        # 绘制人脸框和标签
        cv2.rectangle(img_original, (x, y), (x+w, y+h), (0, 255, 0), 2)
        cv2.rectangle(img_original, (x, y-40), (x+w, y), (0, 255, 0), -2)
        cv2.putText(img_original, class_names[class_index], (x, y-10), font, 0.75, (255, 255, 255), 1, cv2.LINE_AA)
        
        # 显示置信度
        cv2.putText(img_original, f"{confidence:.2f}%", (180, 75), font, 0.75, (255, 0, 0), 2, cv2.LINE_AA)
    
    cv2.imshow("Face Detection", img_original)
    if cv2.waitKey(1) == ord('q'):
        break

cap.release()
cv2.destroyAllWindows()

额外的模型修复建议

  1. 检查模型结构:加载模型后打印结构,确认最后一层输出维度为6,激活函数为softmax:
    print(model.summary())
    # 查看最后一层配置
    print(model.layers[-1].activation)
    print(model.layers[-1].units)
    
  2. 重新训练模型(若结构错误):
    • 确保最后一层定义为:Dense(6, activation='softmax')
    • 损失函数对应:标签为整数用sparse_categorical_crossentropy,独热编码用categorical_crossentropy
    • 训练时加入验证集,监控各类别的分类精度,排查数据不平衡问题

内容的提问来源于stack exchange,提问作者Suhail Razeeth

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.24 21:27:12