基于MNIST训练GAN时Discriminator层收到2个输入张量报错问题
你遇到的报错核心原因是构造判别器训练集的代码逻辑错误,导致传入train_on_batch的输入参数不是单个样本数组,而是包含两个样本数组的二元组,判别器检测到两个输入张量与自己单输入的定义不符,触发报错。
具体错误位置与原因
训练函数中如下代码行存在逻辑错误:
X_train_dis, y_train_dis = shuffle(X_neg_train_dis, X_pos_train_dis), shuffle(y_neg_train_dis, y_pos_train_dis)
sklearn的shuffle函数支持传入多个数组同步打乱,传入n个数组就会返回n个独立打乱后的数组。你这行代码的实际运行结果为:
- 第一个
shuffle(X_neg_train_dis, X_pos_train_dis)返回两个数组:打乱后的假样本集、打乱后的真样本集 - 第二个
shuffle(y_neg_train_dis, y_pos_train_dis)返回两个数组:打乱后的假标签集、打乱后的真标签集 - 左侧仅两个变量接收右侧的四个返回值,最终
X_train_dis实际是(假样本集, 真样本集)组成的二元组,y_train_dis是(假标签集, 真标签集)组成的二元组 - 二元组传入
train_on_batch时,会被模型识别为两个独立的输入张量,因此触发“期望1个输入,收到2个输入”的报错。
修复方法
直接替换原构造判别器训练集的代码,先拼接正负样本与对应标签,再做同步shuffle即可,保证样本与标签的对应关系不会错乱:
# 拼接正负样本 X_train_dis = np.concatenate([X_neg_train_dis, X_pos_train_dis], axis=0) # 拼接正负标签 y_train_dis = np.concatenate([y_neg_train_dis, y_pos_train_dis], axis=0) # 同步打乱样本和标签,保证对应关系 X_train_dis, y_train_dis = shuffle(X_train_dis, y_train_dis)
补充优化建议:你当前生成器最后输出层用了ReLU激活,MNIST数据集通常归一化到[-1,1]或[0,1]区间,建议将生成器最后全连接层的激活改为tanh或sigmoid,避免生成样本取值范围不符合数据集分布,影响训练效果。
内容的提问来源于stack exchange,提问作者logankilpatrick
相关产品推荐
相关产品推荐

