Keras多分类模型ROC AUC评估遇形状不兼容错误求助
解决Keras多分类模型中Shapes (None,1)和(None,4)不兼容问题
错误核心原因
你使用了CategoricalCrossentropy损失函数,它要求真实标签y_true是独热编码格式(形状为(样本数, 类别数)),但你的训练/测试标签是整数索引格式(形状为(样本数,1)),两者形状不匹配,直接触发形状兼容错误。同时自定义指标和回调逻辑也需要对应调整,避免后续问题。
修复方案(二选一即可)
方案1:将整数标签转为独热编码
对y_train和y_test执行独热编码转换,匹配CategoricalCrossentropy的要求:
import tensorflow as tf # 你的任务是4分类,对应num_classes=4 num_classes = 4 y_train = tf.keras.utils.to_categorical(y_train, num_classes=num_classes) y_test = tf.keras.utils.to_categorical(y_test, num_classes=num_classes)
对应调整自定义指标
需要先把独热编码的y_true转为整数标签,再和模型输出的预测类别对比:
class MulticlassTruePositives(tf.keras.metrics.Metric): def __init__(self, name='multiclass_true_positives', **kwargs): super(MulticlassTruePositives, self).__init__(name=name, **kwargs) self.true_positives = self.add_weight(name='tp', initializer='zeros') def update_state(self, y_true, y_pred, sample_weight=None): # 独热标签转整数索引 y_true = tf.argmax(y_true, axis=1) y_pred = tf.argmax(y_pred, axis=1) values = tf.cast(y_true, 'int32') == tf.cast(y_pred, 'int32') values = tf.cast(values, 'float32') if sample_weight is not None: sample_weight = tf.cast(sample_weight, 'float32') values = tf.multiply(values, sample_weight) self.true_positives.assign_add(tf.reduce_sum(values)) def result(self): return self.true_positives def reset_states(self): self.true_positives.assign(0.)
对应调整回调函数
混淆矩阵需要整数标签,因此要把独热编码的y_true转换后再传入:
class PerformanceVisualizationCallback(tf.keras.callbacks.Callback): def __init__(self, model, test_data, image_dir): super().__init__() self.model = model self.test_data = test_data os.makedirs(image_dir, exist_ok=True) self.image_dir = image_dir def on_epoch_end(self, epoch, logs={}): y_pred = np.asarray(self.model.predict(self.test_data[0])) y_true = self.test_data[1] y_pred_class = np.argmax(y_pred, axis=1) y_true_class = np.argmax(y_true, axis=1) # 独热转整数 # 保存混淆矩阵 fig, ax = plt.subplots(figsize=(16,12)) plot_confusion_matrix(y_true_class, y_pred_class, ax=ax) fig.savefig(os.path.join(self.image_dir, f'confusion_matrix_epoch_{epoch}')) plt.close(fig) # 保存ROC曲线(多分类ROC一般接收独热标签) fig, ax = plt.subplots(figsize=(16,12)) plot_roc(y_true, y_pred, ax=ax) fig.savefig(os.path.join(self.image_dir, f'roc_curve_epoch_{epoch}')) plt.close(fig)
方案2:改用SparseCategoricalCrossentropy损失
如果不想修改标签格式,直接替换损失函数为SparseCategoricalCrossentropy,它专门适配整数索引标签:
hypermodel.compile(optimizer='sgd', loss=tf.keras.losses.SparseCategoricalCrossentropy(), metrics=[tf.keras.metrics.AUC(), MulticlassTruePositives()])
对应调整自定义指标
简化形状处理逻辑,直接对比整数标签和预测类别:
class MulticlassTruePositives(tf.keras.metrics.Metric): def __init__(self, name='multiclass_true_positives', **kwargs): super(MulticlassTruePositives, self).__init__(name=name, **kwargs) self.true_positives = self.add_weight(name='tp', initializer='zeros') def update_state(self, y_true, y_pred, sample_weight=None): y_pred = tf.argmax(y_pred, axis=1) # 直接对比,自动兼容(None,1)和(None,)形状 values = tf.cast(y_true, 'int32') == tf.cast(y_pred, 'int32') values = tf.cast(values, 'float32') if sample_weight is not None: sample_weight = tf.cast(sample_weight, 'float32') values = tf.multiply(values, sample_weight) self.true_positives.assign_add(tf.reduce_sum(values)) def result(self): return self.true_positives def reset_states(self): self.true_positives.assign(0.)
对应调整回调函数
仅需处理整数标签的维度压缩,ROC曲线需要时可临时转独热编码:
class PerformanceVisualizationCallback(tf.keras.callbacks.Callback): def __init__(self, model, test_data, image_dir): super().__init__() self.model = model self.test_data = test_data os.makedirs(image_dir, exist_ok=True) self.image_dir = image_dir def on_epoch_end(self, epoch, logs={}): y_pred = np.asarray(self.model.predict(self.test_data[0])) y_true = self.test_data[1] y_pred_class = np.argmax(y_pred, axis=1) y_true_class = y_true.squeeze() # 把(None,1)转为(None,) # 保存混淆矩阵 fig, ax = plt.subplots(figsize=(16,12)) plot_confusion_matrix(y_true_class, y_pred_class, ax=ax) fig.savefig(os.path.join(self.image_dir, f'confusion_matrix_epoch_{epoch}')) plt.close(fig) # 保存ROC曲线:如果plot_roc要求独热标签,临时转换 y_true_onehot = tf.keras.utils.to_categorical(y_true_class, num_classes=4) fig, ax = plt.subplots(figsize=(16,12)) plot_roc(y_true_onehot, y_pred, ax=ax) fig.savefig(os.path.join(self.image_dir, f'roc_curve_epoch_{epoch}')) plt.close(fig)
内容的提问来源于stack exchange,提问作者melolilili
相关产品推荐
相关产品推荐

