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

VGG16网络调用predict_classes报错,求非序列模型分类预测方法

解决VGG16自定义模型无predict_classes方法的问题

嘿,这个坑我之前踩过!你遇到的问题本质是:predict_classes()这个方法是旧版Keras中Sequential模型专属的快捷方法,而你自定义的VGG16应该是用Functional API构建的Model对象,这类对象并没有内置这个属性,所以会报AttributeError。

不用慌,我们可以手动实现和predict_classes()完全一样的效果,分两种场景处理:

多分类场景(输出层是Softmax)

如果你的VGG16是做多分类(比如ImageNet的1000类),输出层用了Softmax激活,那只需要先通过predict()得到每个类别的概率,再用np.argmax()取概率最大的类别索引即可:

import numpy as np

# 获取每个样本的类别概率分布
pred_probs = custom_vgg_model.predict(x)
# 取概率最大的索引作为预测类别
pred_classes = np.argmax(pred_probs, axis=1)

这里的axis=1表示沿着样本的类别维度取最大值,得到的pred_classes和Sequential模型predict_classes(x)返回的结果完全一致。

二分类场景(输出层是Sigmoid)

如果是二分类任务,输出层用了Sigmoid激活,那可以通过阈值(通常是0.5)来判断类别:

import numpy as np

pred_probs = custom_vgg_model.predict(x)
# 大于0.5为类别1,否则为类别0,flatten()把结果转成一维数组
pred_classes = (pred_probs > 0.5).astype(int).flatten()

其实这么手动处理反而更灵活——你可以自己调整阈值、查看概率分布,比依赖封装好的方法更可控。

内容的提问来源于stack exchange,提问作者Khan W.

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.27 09:56:07