Keras构建GAN判别器报错预期1个输入却收到3个张量如何解决?
错误原因
- 重复指定输入形状:Sequential模型仅第一层需要指定
input_shape,你给判别器的第二、第三层Conv2D也加了input_shape = img_shape参数,而经过第一层卷积后特征图形状已经和原始输入图片不同,强行指定输入形状会导致形状匹配逻辑错乱,将输入形状元组的三个元素识别为三个独立的输入张量。 - Dense层语法错误:
Dense(1,Activation('sigmoid'))是错误写法,你把Activation层实例传递给了Dense层第二个位置参数(本应为use_bias参数的位置),参数传递错乱导致层的输入输出逻辑异常。 - 代码缩进错误:第三层
Conv2D之后的BatchNormalization、LeakyReLU等层缩进和前面的model.add不对齐,导致这些层没有被正确加入判别器模型。 - 缺失生成器定义:你直接调用了
generator(z_dim)但没有提前定义生成器函数,且如果生成器输出形状和判别器输入形状(32,32,3)不匹配,拼接GAN时也会报错。 - 命名冲突:函数名和实例化后的变量名完全重名,后续如果需要重新调用构造函数会找不到原函数。
修复方案
修正后的完整可运行代码如下:
from tensorflow.keras.models import Sequential from tensorflow.keras.layers import Conv2D, LeakyReLU, BatchNormalization, Flatten, Dense, Reshape, UpSampling2D from tensorflow.keras.optimizers import Adam img_rows = 32 img_cols = 32 channels = 3 img_shape = (img_rows, img_cols, channels) z_dim = 150 # 修正函数名,避免和实例变量重名 def build_discriminator(img_shape): model = Sequential() # 仅第一层指定input_shape model.add(Conv2D(32, kernel_size=3, strides=2, input_shape=img_shape, padding='same')) model.add(LeakyReLU(alpha=0.01)) # 移除多余的input_shape参数 model.add(Conv2D(64, kernel_size=3, strides=2, padding='same')) model.add(BatchNormalization()) model.add(LeakyReLU(alpha=0.01)) # 移除多余的input_shape参数,修正缩进 model.add(Conv2D(128, kernel_size=3, strides=2, padding='same')) model.add(BatchNormalization()) model.add(LeakyReLU(alpha=0.01)) model.add(Flatten()) # 修正Dense层激活函数写法 model.add(Dense(1, activation='sigmoid')) return model # 补充生成器定义,确保输出形状为(32,32,3) def build_generator(z_dim): model = Sequential() model.add(Dense(128 * 8 * 8, input_dim=z_dim)) model.add(LeakyReLU(alpha=0.01)) model.add(Reshape((8, 8, 128))) model.add(UpSampling2D()) model.add(Conv2D(128, kernel_size=3, padding='same')) model.add(BatchNormalization()) model.add(LeakyReLU(alpha=0.01)) model.add(UpSampling2D()) model.add(Conv2D(64, kernel_size=3, padding='same')) model.add(BatchNormalization()) model.add(LeakyReLU(alpha=0.01)) model.add(Conv2D(3, kernel_size=3, padding='same', activation='tanh')) return model def build_gan(generator, discriminator): model = Sequential() model.add(generator) model.add(discriminator) return model discriminator = build_discriminator(img_shape) discriminator.compile(loss='binary_crossentropy', optimizer=Adam(), metrics=['accuracy']) discriminator.trainable = False generator = build_generator(z_dim) gan = build_gan(generator, discriminator) gan.compile(loss='binary_crossentropy', optimizer=Adam())
修复要点说明:
- 调整函数命名,避免和实例变量重名
- 移除第二、三层卷积多余的
input_shape参数,统一缩进格式 - 修正Dense层激活函数的传递方式
- 补充符合输入输出形状要求的生成器实现,保证生成器输出和判别器输入形状匹配
内容的提问来源于stack exchange,提问作者DellFan0914
相关产品推荐
相关产品推荐

