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

Python新手构建ROC曲线遇多类报错,请求代码修正

问题分析与修正方案

看起来你在使用自编码器构建ROC曲线时,混淆了自编码器的输出特性和ROC曲线的输入要求,这是新手很容易踩的坑~让我一步步帮你理清问题并修正代码:

核心错误原因

你的自编码器是图像重构模型,输出decoded_imgs的形状是(样本数, 256, 256, 3),和输入图像维度一致。但ROC曲线要求的是每个样本对应一个数值型得分(比如分类概率、重构误差),而你直接对decoded_imgs做argmax(axis=1)得到的是每个像素维度的索引,完全不符合ROC的输入要求,这就是你遇到各种形状不匹配错误的根源。

正确的解决思路

自编码器做ROC曲线(通常用于异常检测或分类任务),应该用每个样本的重构误差作为得分:

  • 正常样本的重构误差小,异常样本的重构误差大(异常检测场景)
  • 或者用误差来区分不同类别的样本(分类场景)

具体步骤:

  1. 计算每个测试样本的重构误差(比如MSE或MAE,把三维图像压缩成一个数值)
  2. 确保真实标签Y_test是一维数组(每个样本对应一个类别/异常标签)
  3. 用这两个一维数组调用roc_curve

修正后的完整代码

替换你原代码中ROC曲线相关的部分,完整修正代码如下:

import keras
import numpy as np
from keras.datasets import mnist
from get_dataset import get_dataset
from stack import keras_model
X_train, X_test, Y_train, Y_test = get_dataset()
from keras.layers import Input, Conv2D, MaxPooling2D, UpSampling2D, Dense
from keras.models import Model
input_img = Input(shape=(256, 256, 3))
x = Conv2D(32, (3, 3), activation='relu', padding='same')(input_img)
x = MaxPooling2D((2, 2), padding='same')(x)
x = Conv2D(64, (3, 3), activation='relu', padding='same')(x)
x = MaxPooling2D((2, 2), padding='same')(x)
x = Conv2D(64, (3, 3), activation='relu', padding='same')(x)
encoded = MaxPooling2D((2, 2), padding='same')(x)
x = Conv2D(64, (3, 3), activation='relu', padding='same')(encoded)
x = UpSampling2D((2, 2))(x)
x = Conv2D(64, (3, 3), activation='relu', padding='same')(x)
x = UpSampling2D((2, 2))(x)
x = Conv2D(32, (3, 3), activation='relu', padding='same')(x)
x = UpSampling2D((2, 2))(x)
decoded = Conv2D(3, (3, 3), activation='sigmoid', padding='same')(x)
autoencoder = Model(input_img, decoded)
autoencoder.compile(optimizer='rmsprop', loss='mae',metrics=['mse', 'accuracy'])
from keras.callbacks import ModelCheckpoint, TensorBoard
checkpoints = []
from keras.preprocessing.image import ImageDataGenerator
generated_data = ImageDataGenerator(featurewise_center=False, samplewise_center=False, featurewise_std_normalization=False, samplewise_std_normalization=False, zca_whitening=False, rotation_range=0, width_shift_range=0.1, height_shift_range=0.1, horizontal_flip = True, vertical_flip = False)
generated_data.fit(X_train)
epochs = 1
batch_size = 5
autoencoder.fit_generator(generated_data.flow(X_train, X_train, batch_size=batch_size), steps_per_epoch=X_train.shape[0]/batch_size, epochs=epochs, validation_data=(X_test, X_test), callbacks=[TensorBoard(log_dir='/tmp/autoencoder')])
autoencoder.fit(X_train, X_train, batch_size=batch_size, epochs=epochs, validation_data=(X_test, X_test), shuffle=True, callbacks=[TensorBoard(log_dir='/tmp/auti')])
decoded_imgs = autoencoder.predict(X_test)
from sklearn.metrics import roc_curve

# ------------------- 修正部分开始 -------------------
# 计算每个样本的重构MSE误差:将三维图像压缩为单个数值
reconstruction_errors = np.mean(np.square(X_test - decoded_imgs), axis=(1,2,3))

# 确保Y_test是一维标签数组(如果是one-hot编码,用argmax转成类别索引)
y_true = Y_test.argmax(axis=1)

# 现在调用roc_curve,两个输入都是一维数组,形状匹配
fpr_keras, tpr_keras, thresholds_keras = roc_curve(y_true, reconstruction_errors)
# ------------------- 修正部分结束 -------------------

# 可选:打印验证形状是否匹配
print(f"y_true形状: {y_true.shape}")
print(f"重构误差形状: {reconstruction_errors.shape}")

额外说明

  • 如果你的任务是异常检测,Y_test应该是0(正常)/1(异常)的标签,重构误差越大代表越可能是异常,这时候ROC曲线可以有效区分正常和异常样本。
  • 如果是多分类任务,你可能需要针对每个类别做One-vs-Rest的ROC曲线,或者考虑改用带分类头的自编码器(在编码层后加Dense分类层)。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.11 08:51:46