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

训练糖尿病视网膜病变检测GAN模型时遭遇ValueError求助

糖尿病视网膜病变检测模型训练错误解决方案

问题说明

训练糖尿病视网膜病变图像检测模型时触发维度不匹配错误,已确认数据集非空,调整维度后问题仍存在。错误核心提示:标签(target)与模型输出(output)形状不一致,target.shape=(None, 3),output.shape=(None, 5)。

错误信息

Epoch 1/50
Traceback (most recent call last):
  File "C:\Users\asus\OneDrive\Desktop\project\DR-GAN\TrainModel.py", line 65, in <module>
    classifier.fit(X, Y, batch_size=32, epochs=50)
  File "C:\Users\asus\AppData\Roaming\Python\Python312\site-packages\keras\src\utils\traceback_utils.py", line 122, in error_handler
    raise e.with_traceback(filtered_tb) from None
  File "C:\Users\asus\AppData\Roaming\Python\Python312\site-packages\keras\src\backend\tensorflow\nn.py", line 553, in categorical_crossentropy
    raise ValueError(
ValueError: Arguments `target` and `output` must have the same shape. Received: target.shape=(None, 3), output.shape=(None, 5)

训练代码

import numpy as np
import imutils
import sys
import cv2
import os
from tensorflow.keras.utils import to_categorical
from keras.models import model_from_json
from keras.layers import MaxPooling2D
from keras.layers import Dense, Dropout, Activation, Flatten
from keras.layers import Convolution2D
from keras.models import Sequential 

images = []
image_labels  = []
directory = 'dataset'
list_of_files = os.listdir(directory)
index = 0
for file in list_of_files:
    subfiles = os.listdir(directory+'/'+file)
    for sub in subfiles:
        path = directory+'/'+file+'/'+sub
        img = cv2.imread(path)
        #img = cv2.cvtColor(img, cv2.COLOR_BGR2GRAY)
        if img is None:
          print('Wrong path:', path)
        else:
         img = cv2.resize(img, (32,32))
         im2arr = np.array(img)
         im2arr = im2arr.reshape(32,32,3)
         images.append(im2arr)
         image_labels.append(file)
    print(file)    

X = np.asarray(images)
Y = np.asarray(image_labels)
Y = to_categorical(Y)
img = X[20].reshape(32,32,3)
cv2.imshow('ff',cv2.resize(img,(250,250)))
cv2.waitKey(0)
print("shape == "+str(X.shape))
print("shape == "+str(Y.shape))
print(Y)
X = X.astype('float32')
X = X/255

np.save("model/img_data.txt",X)
np.save("model/img_label.txt",Y)

X = np.load('model/img_data.txt.npy')
Y = np.load('model/img_label.txt.npy')
print(Y)
img = X[20].reshape(32,32,3)
cv2.imshow('ff',cv2.resize(img,(250,250)))
cv2.waitKey(0)

classifier = Sequential() #alexnet transfer learning code here
classifier.add(Convolution2D(32, 3, 3, input_shape = (32, 32, 3), activation = 'relu'))
classifier.add(MaxPooling2D((2, 2) , padding='same'))
classifier.add(Convolution2D(32, 3, 3, activation = 'relu'))
classifier.add(MaxPooling2D((2, 2) , padding='same'))
classifier.add(Flatten())
classifier.add(Dense(units = 128, activation = 'relu'))
classifier.add(Dense(units = 5, activation = 'softmax'))
classifier.compile(optimizer = 'adam', loss = 'categorical_crossentropy', metrics = ['accuracy'])
classifier.fit(X, Y, batch_size=32, epochs=50)

解决方案

问题根源

  1. 标签处理错误:直接将字符串类型的文件夹名传入to_categorical,该函数仅支持整数类型的类别索引,导致标签维度识别异常。
  2. 模型输出与类别数量不匹配:数据集实际只有3个类别,但模型最后一层Dense硬编码设置了units=5,导致输出维度与标签维度不一致。

修复步骤

1. 将字符串标签转换为整数索引

在收集标签时,为每个类别文件夹分配唯一整数ID,确保to_categorical能正确生成独热编码:

# 排序类别文件夹,保证索引映射稳定
list_of_files = sorted(os.listdir(directory))
# 建立类别名到整数索引的映射
class_map = {name: idx for idx, name in enumerate(list_of_files)}

2. 匹配模型输出维度与实际类别数量

用数据集的实际类别数量设置模型最后一层的输出单元数,避免硬编码:

num_classes = len(class_map)
classifier.add(Dense(units=num_classes, activation='softmax'))

修改后的完整关键代码片段

images = []
image_labels  = []
directory = 'dataset'
list_of_files = sorted(os.listdir(directory))
class_map = {name: idx for idx, name in enumerate(list_of_files)}

for file in list_of_files:
    subfiles = os.listdir(directory+'/'+file)
    for sub in subfiles:
        path = directory+'/'+file+'/'+sub
        img = cv2.imread(path)
        if img is None:
          print('Wrong path:', path)
        else:
         img = cv2.resize(img, (32,32))
         im2arr = np.array(img)
         im2arr = im2arr.reshape(32,32,3)
         images.append(im2arr)
         # 使用整数索引作为标签
         image_labels.append(class_map[file])
    print(file)    

X = np.asarray(images)
Y = np.asarray(image_labels)
# 转换为独热编码,此时维度为(样本数, 类别数)
Y = to_categorical(Y)
print("Y shape:", Y.shape)  # 验证标签维度

X = X.astype('float32')
X = X/255

# ... 数据保存与加载代码不变 ...

classifier = Sequential()
classifier.add(Convolution2D(32, 3, 3, input_shape = (32, 32, 3), activation = 'relu'))
classifier.add(MaxPooling2D((2, 2) , padding='same'))
classifier.add(Convolution2D(32, 3, 3, activation = 'relu'))
classifier.add(MaxPooling2D((2, 2) , padding='same'))
classifier.add(Flatten())
classifier.add(Dense(units = 128, activation = 'relu'))
num_classes = len(class_map)
# 匹配实际类别数量
classifier.add(Dense(units=num_classes, activation='softmax'))
classifier.compile(optimizer = 'adam', loss = 'categorical_crossentropy', metrics = ['accuracy'])
classifier.fit(X, Y, batch_size=32, epochs=50)

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.24 17:14:52