基于CelebA数据集的Conditional DCGAN训练循环张量匹配错误排查
Conditional DCGAN张量尺寸匹配问题解决方案
问题根源
当前代码混淆了生成器和判别器对条件张量的尺寸需求:
- 生成器需将40维属性标签与噪声在通道维度拼接,噪声形状为
[64, 100, 1, 1],因此条件张量需保持[64, 40, 1, 1]的空间尺寸(与噪声的1x1空间维度一致)。 - 判别器需将条件与真实/生成图片在通道维度拼接,图片形状为
[64, 3, 64, 64],因此条件张量需扩展为[64, 40, 64, 64](与图片的64x64空间维度一致)。
你之前错误地将同一份条件张量重复扩展为64x64尺寸后同时传给生成器和判别器,导致生成器拼接时维度不匹配。
具体修改步骤
1. 训练循环中分离生成器与判别器的条件张量
在每个batch处理时,从原始conditions(形状[64, 40])分别生成两种尺寸的条件张量:
# 原始conditions来自dataloader,形状为[b_size, 40] b_size = real_cpu.size(0) # 生成器用的条件:保持1x1空间维度 gen_conditions = conditions.view(b_size, 40, 1, 1).to(device) # 判别器用的条件:扩展为与图片一致的64x64空间维度 dis_conditions = gen_conditions.repeat(1, 1, image_size, image_size)
2. 修正生成器调用逻辑
生成器调用时传入gen_conditions而非扩展后的条件:
# 替换原错误代码:fake = netG(noise, conditions.to(device)) fake = netG(noise, gen_conditions)
生成器的forward函数无需修改(此时input和label的空间维度均为1x1,可直接在通道维度拼接):
def forward(self, input, label): # input形状[64,100,1,1],label形状[64,40,1,1] input = torch.cat((input, label), 1) # 拼接后形状[64,140,1,1] return self.conv(input)
3. 修正判别器调用逻辑
判别器调用时传入扩展后的dis_conditions:
# 真实图片传入判别器 output = netD(real_cpu.to(device), dis_conditions) # 生成图片传入判别器 output = netD(fake.detach(), dis_conditions)
4. 修正fixed_noise的生成器测试代码
原代码中netG(fixed_noise)未传入条件,需补充固定条件张量:
# 生成固定条件(示例:全0属性或随机属性) fixed_conditions = torch.zeros(fixed_noise.size(0), 40, 1, 1, device=device) with torch.no_grad(): fake = netG(fixed_noise, fixed_conditions).detach().cpu()
5. 移除无效的条件重塑代码
删除你之前尝试的错误重塑逻辑:
# 以下代码直接删除 conditions = conditions.permute(0, 2, 3, 1) conditions = conditions.view(batch_size, 1, 1, 40)
验证修改后张量形状
- 生成器输入:
noise([64,100,1,1])+gen_conditions([64,40,1,1])→ 拼接后[64,140,1,1],符合生成器输入要求。 - 判别器输入:
real_cpu([64,3,64,64])+dis_conditions([64,40,64,64])→ 拼接后[64,43,64,64],符合判别器输入要求。
内容的提问来源于stack exchange,提问作者RaphDaPingu
相关产品推荐
相关产品推荐

