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

TensorFlow报InvalidArgumentError x==y 多分类自定义损失形状不匹配如何修复

问题根因
  • 输入形状传递引入了多余维度:你定义的输入shape是(1,5),输入数据集X的形状为(6,1,5),经过dense1层后输出形状为(6,1,10),三个分类头输出的形状均为(6,1,4),拼接后的整体输出形状为(6,1,12)。你在损失函数中直接用y_pred[:, 4*i:4*(i+1)]索引时,默认第二个维度是分类输出维度,但实际第二个维度是多余的空维度,导致索引出来的张量形状错乱。
  • 损失函数的reshape操作破坏了样本对应关系:你将三个分类头的损失拼接后reshape为(-1,1),会得到形状为(18,1)的张量,但Keras要求损失的第一维度必须和样本数(此处为6)对齐,因此触发了形状不匹配断言。
  • 不必要的reshape操作不符合sparse交叉熵的输入要求:sparse_categorical_crossentropy的真实标签输入不需要额外加最后一维,预测值也不需要额外插入空维度,多余的reshape反而会导致形状校验失败。
修复代码

修复后的自定义模型

import tensorflow as tf
from tensorflow import keras
from tensorflow.keras import layers
import tensorflow.keras.backend as K

class Model(keras.Model):
    def __init__(self):
        super(Model, self).__init__()
        self.dense1 = layers.Dense(10, input_shape=(1, 5), activation="relu")
        self.u = layers.Dense(4, activation="softmax")
        self.c = layers.Dense(4, activation="softmax")
        self.k = layers.Dense(4, activation="softmax")
        self.outputs = layers.Concatenate()

    def call(self, inputs):
        x = tf.convert_to_tensor(inputs)
        x = self.dense1(x)
        # 去掉多余的空维度
        x = tf.squeeze(x, axis=1)
        ls = []
        u = self.u(x)
        ls.append(u)
        c = self.c(x)
        ls.append(c)
        k = self.k(x)
        ls.append(k)
        return self.outputs(ls)

    def process(self, observations):
        action_probs = self.predict_on_batch(observations)
        return action_probs

修复后的自定义损失函数

def custom_cross_entropy(y_true, y_pred):
    total_loss = 0.0
    for i in range(3):
        # 直接取对应位置的标签和预测值,无需多余reshape
        y_true_head = y_true[:, i]
        y_pred_head = y_pred[:, 4*i : 4*(i+1)]
        total_loss += tf.keras.losses.sparse_categorical_crossentropy(y_true_head, y_pred_head, from_logits=False)
    # 返回每个样本的平均损失,形状为(None,)和样本数对齐
    return total_loss / 3.0

测试运行代码

# 模拟数据
X = [[[1, 2, 3, 4, 5]], [[1, 2, 3, 4, 5]], [[1, 2, 3, 4, 5]], [[1, 2, 3, 4, 5]], [[1, 2, 3, 4, 5]], [[1, 2, 3, 4, 5]]]
y_t = [[0, 1, 3], [0, 2, 1], [2, 1, 1], [1, 2, 0], [1, 2, 3], [1, 2, 1]]

model = Model()
model.compile(loss=custom_cross_entropy, optimizer='adam', metrics=['accuracy'])
model.fit(X, y_t, epochs=5)
补充说明

如果需要保留多输出的结构,也可以直接让模型返回三个分类头的输出列表,Keras原生支持多输出损失配置,不需要手动拼接输出和写自定义损失,实现更简洁也更不容易出错。

内容的提问来源于stack exchange,提问作者Stat_prob_001

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.09.26 15:24:01