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
相关产品推荐
相关产品推荐

