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

使用自定义数据集替代MNIST运行DCGAN时遇类型错误求助

解决DCGAN自定义数据集加载的TypeError问题

错误原因

你用flow_from_directory得到的是DirectoryIterator迭代器,而原代码针对的是MNIST加载出的numpy数组,直接对迭代器做/255这类数值运算自然会报错。

修复步骤

1. 正确加载自定义数据集并预处理

替换原MNIST加载代码,把归一化逻辑整合到数据生成器中,避免直接对迭代器做数值运算:

from tensorflow.keras.preprocessing.image import ImageDataGenerator

# 初始化数据生成器,同步完成像素值归一化(适配DCGAN的tanh输出,缩至[-1,1]区间)
datagen = ImageDataGenerator(rescale=1./127.5 - 1)

# 从目录加载数据,需保证目录结构为:根目录/类别子目录/图片文件
train_generator = datagen.flow_from_directory(
    './your_dataset_dir',  # 替换为你的数据集根目录路径
    target_size=(28, 28),  # 匹配原DCGAN的输入尺寸,可根据你的图片修改
    batch_size=32,
    color_mode='grayscale',  # 彩色图改为'rgb'
    class_mode=None  # GAN无需标签,设为None
)

2. 修改训练循环的数据集读取逻辑

原代码直接操作numpy数组,现在要改为从迭代器中逐批次获取数据:
找到原代码的训练循环部分,替换为以下逻辑:

for epoch in range(epochs):
    # 遍历数据集的所有批次
    for batch in train_generator:
        # 生成随机噪声
        noise = np.random.normal(0, 1, (batch.shape[0], latent_dim))
        
        # 生成假样本
        gen_imgs = generator.predict(noise)
        
        # 拼接真实样本与假样本
        X = np.concatenate((batch, gen_imgs))
        
        # 创建判别器标签(加入少量噪声做标签平滑)
        y = np.concatenate((np.ones((batch.shape[0], 1)), np.zeros((batch.shape[0], 1))))
        y += 0.05 * np.random.random(y.shape)
        
        # 训练判别器
        d_loss = discriminator.train_on_batch(X, y)
        
        # 训练生成器
        noise = np.random.normal(0, 1, (batch.shape[0], latent_dim))
        y_gen = np.ones((batch.shape[0], 1))
        g_loss = combined.train_on_batch(noise, y_gen)
        
        # 打印训练进度(可选)
        print(f"{epoch} [D loss: {d_loss}] [G loss: {g_loss}]")
        
        # 当batch_index归0时,说明已遍历完一轮数据集,跳出循环
        if train_generator.batch_index == 0:
            break

3. 适配图片尺寸与通道数

如果你的自定义图片不是28x28灰度图,需同步修改生成器和判别器的输入输出维度:

  • 彩色图需将判别器输入层改为Input(shape=(img_rows, img_cols, 3))
  • 生成器输出层改为Conv2D(3, kernel_size=3, activation='tanh', padding='same')

关键注意事项

  • 数据集目录必须遵循根目录/类别子目录/图片的结构,即使GAN不需要类别标签,flow_from_directory也要求该格式
  • 归一化逻辑要与原DCGAN一致:用rescale=1./127.5 -1将像素值从0-255缩至[-1,1],匹配生成器的tanh输出
  • 训练循环中需判断batch_index ==0,确保每个epoch只遍历一次数据集,避免无限循环

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.15 12:40:37