SparseCategoricalCrossentropy形状不匹配报错,其支持的输入形状是什么
SparseCategoricalCrossentropy 输入要求与问题修复
官方输入形状要求
SparseCategoricalCrossentropy 是为非one-hot编码的类别标签设计的损失函数,输入形状规则如下:
y_pred:形状为[batch_size, 类别总数],最后一维存储每个类别的预测概率(如果设置from_logits=True则存储未激活的原始logits值)y_true:形状为[batch_size],每个元素为对应样本的真实类别整数索引,不需要做one-hot编码。如果你的标签已经是one-hot格式,应该使用普通CategoricalCrossentropy损失函数。
报错的核心原因
你当前传入的y_true是形状为(1,1000)的one-hot编码结果,不符合Sparse版本损失函数对标签的形状要求,因此触发维度不匹配报错。
两种修复方案
方案1:继续使用SparseCategoricalCrossentropy,调整标签格式
直接传入类别索引即可,无需做one-hot编码:
import keras.backend as K import numpy as np import tensorflow as tf full_model = tf.keras.applications.MobileNetV2( input_shape=(224,224,3), alpha=1.0, include_top=True, weights="imagenet", input_tensor=None, pooling=None, classes=1000, classifier_activation="softmax", ) func = K.function(full_model.layers[1].input, full_model.layers[155].output) conv_output = func([processed_image]) y_pred = np.single(conv_output) # 仅修改y_true定义即可,直接传类别索引282,形状为(1,)匹配batch_size=1 y_true = np.array([282]) scce = tf.keras.losses.SparseCategoricalCrossentropy() print(scce(y_true, y_pred).numpy())
方案2:保留one-hot标签,更换损失函数
如果需要保持现有one-hot标签格式,把损失函数替换为普通分类交叉熵即可:
# 原有y_true定义不变 y_true = np.zeros(1000).reshape(1,1000) y_true[0][282] = 1 # 换用普通CategoricalCrossentropy cce = tf.keras.losses.CategoricalCrossentropy() print(cce(y_true, y_pred).numpy())
内容的提问来源于stack exchange,提问作者Ricardo
相关产品推荐
相关产品推荐

