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

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.09.30 09:36:05