如何快速反向解码one-hot编码以获取神经网络图像分类结果
操作步骤
第一步:One-Hot编码反向解码得到类别索引
你的one-hot预测结果每个位置的6维向量中,最大值对应的下标就是所属类别,直接调用对应框架的argmax方法即可,注意指定正确的维度:
- Numpy环境代码:
import numpy as np # 假设pred_onehot是预测结果数组,逐像素分类场景形状为(高, 宽, 6),整图分类场景形状为(样本数,6) pred_classes = np.argmax(pred_onehot, axis=-1)
- PyTorch环境代码:
import torch pred_classes = torch.argmax(pred_onehot, dim=-1)
- TensorFlow环境代码:
import tensorflow as tf pred_classes = tf.argmax(pred_onehot, axis=-1)
运行后得到的pred_classes就是值范围在0-5之间的类别索引数组。
第二步:类别索引映射到灰度等级
你可以根据需求选择均匀映射或者自定义灰度映射规则:
- 均匀映射(6个类别均匀分布在0-255灰度区间):
gray_step = 255 // 5 gray_img = pred_classes * gray_step # 转换为图像兼容的uint8格式 gray_img = gray_img.astype(np.uint8)
- 自定义灰度映射:
# 映射表下标对应类别,值对应目标灰度值,可自行修改 gray_map = [0, 40, 90, 140, 190, 255] gray_img = np.vectorize(lambda x: gray_map[x])(pred_classes).astype(np.uint8)
第三步:展示/保存灰度图
可以用PIL或者OpenCV直接操作生成的灰度数组:
# PIL方式展示、保存 from PIL import Image Image.fromarray(gray_img).show() Image.fromarray(gray_img).save("分类结果灰度图.png") # OpenCV方式展示、保存 import cv2 cv2.imshow("分类结果", gray_img) cv2.waitKey(0) cv2.destroyAllWindows() cv2.imwrite("分类结果灰度图.png", gray_img)
注意事项
- 如果你的one-hot数组类别维度不在最后一位,需要调整
argmax的维度参数,比如形状为(6, 高, 宽)的数组,axis参数要设为0 - 如果是整图分类任务,一个样本对应一个类别,得到类别索引后,生成对应尺寸的纯色灰度图即可
- 必须确保最终生成的灰度数组数据类型为uint8,否则图像展示会出现颜色异常
内容的提问来源于stack exchange,提问作者FRANCESCO VACCARO
相关产品推荐
相关产品推荐

