多标签CNN模型预测输出数字数组,如何转换为对应标签?
多标签CNN模型预测结果转标签问题
我是神经网络新手,训练了一个多标签CNN模型,用model.predict(image)对新图像预测时,得到的是数字数组而非对应标签,输出如下:
1/1 [==============================] - 0s 294ms/step array([[1.2423882e-07, 4.8644133e-10, 6.2801077e-11, ..., 5.9059038e-09, 7.9583238e-08, 1.8839316e-07], [6.7338499e-07, 3.7959146e-11, 6.3526181e-12, ..., 6.6384481e-10, 6.5803110e-09, 7.2412973e-09], [3.4079651e-06, 1.3292389e-12, 1.1482074e-09, ..., 2.3795290e-09, 4.0437489e-09, 2.5665909e-09], ..., [4.5694168e-09, 5.8126518e-14, 7.4569753e-18, ..., 9.1445762e-16, 8.6637960e-15, 1.7640524e-11], [7.3903612e-09, 1.1566522e-13, 6.5436875e-17, ..., 7.0952320e-15, 8.3692771e-14, 1.2283501e-10], [7.1919551e-08, 1.9803466e-14, 2.5265929e-16, ..., 5.3995962e-13, 5.4146050e-12, 3.9835553e-09]], dtype=float32)
我使用的预测代码如下:
from PIL import Image import numpy as np from skimage import transform def load(filename): np_image = Image.open(filename) np_image = np.array(np_image).astype('float32')/255 np_image = transform.resize(np_image, (256, 256, 3)) np_image = np.expand_dims(np_image, axis=0) return np_image image = load('/content/1000_IM-0003-1001.dcm.png') loaded_2.predict(image)
希望能将预测结果转换为对应的标签,而非数字数组。
模型架构
Model: "sequential_1" _________________________________________________________________ Layer (type) Output Shape Param # ================================================================= conv2d_4 (Conv2D) (None, 196, 196, 16) 1216 batch_normalization_4 (BatchNormalization) (None, 196, 196, 16) 64 max_pooling2d_4 (MaxPooling2D) (None, 98, 98, 16) 0 dropout_6 (Dropout) (None, 98, 98, 16) 0 conv2d_5 (Conv2D) (None, 94, 94, 32) 12832 max_pooling2d_5 (MaxPooling2D) (None, 47, 47, 32) 0 batch_normalization_5 (BatchNormalization) (None, 47, 47, 32) 128 dropout_7 (Dropout) (None, 47, 47, 32) 0 conv2d_6 (Conv2D) (None, 43, 43, 64) 51264 max_pooling2d_6 (MaxPooling2D) (None, 21, 21, 64) 0 batch_normalization_6 (BatchNormalization) (None, 21, 21, 64) 256 dropout_8 (Dropout) (None, 21, 21, 64) 0 conv2d_7 (Conv2D) (None, 17, 17, 64) 102464 max_pooling2d_7 (MaxPooling2D) (None, 8, 8, 64) 0 batch_normalization_7 (BatchNormalization) (None, 8, 8, 64) 256 dropout_9 (Dropout) (None, 8, 8, 64) 0 flatten_1 (Flatten) (None, 4096) 0 dense_3 (Dense) (None, 128) 524416 dropout_10 (Dropout) (None, 128) 0 dense_4 (Dense) (None, 64) 8256 dropout_11 (Dropout) (None, 64) 0 dense_5 (Dense) (None, 524) 34060 ================================================================= Total params: 735,212 Trainable params: 734,860 Non-trainable params: 352
标签信息
标签采用one-hot编码,形状如下:
y.shape (4312, 524)
解决方案
1. 理解预测输出含义
模型最后一层是Dense(524),未指定激活函数时默认用线性激活,输出的是原始logit值。多标签分类任务通常需要用sigmoid激活将输出转为0-1之间的概率值,再通过阈值判断标签是否激活。
如果训练时已用sigmoid激活,输出直接是概率值;否则先做转换:
predictions = loaded_2.predict(image) probabilities = tf.nn.sigmoid(predictions).numpy() # 转为概率值
2. 阈值筛选标签
多标签分类中每个标签独立判断,设定阈值(比如0.5),筛选出概率大于阈值的标签索引:
threshold = 0.5 predicted_indices = np.where(probabilities[0] > threshold)[0] # 取单张图片的预测结果
3. 映射到真实标签名称
你需要有一个标签索引与名称的映射列表(顺序必须和one-hot编码一致),通过索引获取对应标签:
# 替换为你训练时使用的真实标签名称列表 label_names = ["标签1", "标签2", ..., "标签524"] predicted_labels = [label_names[idx] for idx in predicted_indices] print("预测标签:", predicted_labels)
完整示例代码
import tensorflow as tf import numpy as np from PIL import Image from skimage import transform def load(filename): np_image = Image.open(filename) np_image = np.array(np_image).astype('float32')/255 np_image = transform.resize(np_image, (256, 256, 3)) np_image = np.expand_dims(np_image, axis=0) return np_image image = load('/content/1000_IM-0003-1001.dcm.png') # 预测并转换为标签 predictions = loaded_2.predict(image) # 若模型最后一层无sigmoid激活,启用下面一行 probabilities = tf.nn.sigmoid(predictions).numpy() threshold = 0.5 predicted_indices = np.where(probabilities[0] > threshold)[0] # 替换为你的真实标签名称列表 label_names = ["标签1", "标签2", ..., "标签524"] predicted_labels = [label_names[i] for i in predicted_indices] print("预测的标签:", predicted_labels)
注意事项
- 阈值可根据任务调整:想减少误判就提高阈值(如0.7),想召回更多标签就降低阈值(如0.3)。
- 必须保证
label_names的顺序和训练时one-hot编码的顺序完全一致,否则会出现标签映射错误。 - 若训练时模型最后一层已用
sigmoid激活,无需再调用tf.nn.sigmoid,直接用predictions作为概率值即可。
内容的提问来源于stack exchange,提问作者Saad Khattak
相关产品推荐
相关产品推荐

