PyTorch条件GAN运行报错:张量维度不匹配(4维与1维)
PyTorch条件GAN维度不匹配错误修复及256x256图像生成参考资源
问题描述
首次实现PyTorch条件GAN,原常规GAN运行正常,修改后执行fit函数触发RuntimeError,报错信息:Tensors must have same number of dimensions: got 4 and 1,错误出现在判别器forward方法的torch.cat拼接图像张量与标签张量步骤。关键代码及报错堆栈如下:
关键实现代码
import torch.nn as nn class discriminator(nn.Module): def __init__(self): super(discriminator, self).__init__() self.main = nn.Sequential( nn.Conv2d(3, 32, kernel_size=3, stride=2, padding=1, bias=False), nn.BatchNorm2d(32), nn.LeakyReLU(0.1, inplace=True), nn.Conv2d(32, 64, kernel_size=4, stride=2, padding=1, bias=False), nn.BatchNorm2d(64), nn.LeakyReLU(0.1, inplace=True), nn.Conv2d(64, 128, kernel_size=4, stride=2, padding=1, bias=False), nn.BatchNorm2d(128), nn.LeakyReLU(0.1, inplace=True), nn.Conv2d(128, 256, kernel_size=4, stride=2, padding=1, bias=False), nn.BatchNorm2d(256), nn.LeakyReLU(0.1, inplace=True), nn.Conv2d(256, 512, kernel_size=4, stride=2, padding=1, bias=False), nn.BatchNorm2d(512), nn.LeakyReLU(0.1, inplace=True), nn.Conv2d(512, 1024, kernel_size=4, stride=2, padding=1, bias=False), nn.BatchNorm2d(1024), nn.LeakyReLU(0.1, inplace=True), nn.Conv2d(1024, 1, kernel_size=4, stride=1, padding=0, bias=False), nn.Flatten(), nn.Sigmoid() ) def forward(self, x, labels): x = torch.cat((x, labels), dim=1) return self.main(x) discriminator = discriminator() discriminator = to_device(discriminator,device) class generator(nn.Module): def __init__(self, latent_dim): super(generator, self).__init__() self.main = nn.Sequential( nn.ConvTranspose2d(latent_dim, 1024, kernel_size=4, stride=1, padding=0, bias=False), nn.BatchNorm2d(1024), nn.LeakyReLU(0.1, inplace=True), nn.ConvTranspose2d(1024, 512, kernel_size=4, stride=2, padding=1, bias=False), nn.BatchNorm2d(512), nn.LeakyReLU(0.1, inplace=True), nn.ConvTranspose2d(512, 256, kernel_size=4, stride=2, padding=1, bias=False), nn.BatchNorm2d(256), nn.LeakyReLU(0.1, inplace=True), nn.ConvTranspose2d(256, 128, kernel_size=4, stride=2, padding=1, bias=False), nn.BatchNorm2d(128), nn.LeakyReLU(0.1, inplace=True), nn.ConvTranspose2d(128, 64, kernel_size=4, stride=2, padding=1, bias=False), nn.BatchNorm2d(64), nn.LeakyReLU(0.1, inplace=True), nn.ConvTranspose2d(64, 32, kernel_size=4, stride=2, padding=1, bias=False), nn.BatchNorm2d(32), nn.LeakyReLU(0.1, inplace=True), nn.ConvTranspose2d(32, 3, kernel_size=4, stride=2, padding=1, bias=False), nn.Tanh() ) def forward(self, z, labels): x = torch.cat((z, labels), dim=1) return self.main(x) generator = generator(latent_sz) generator = to_device(generator,device) def train_discriminator(real_images, real_labels, opt_d): opt_d.zero_grad() real_preds = discriminator(real_images, real_labels) real_targets = torch.ones(real_images.size(0), 1, device=device) real_loss = F.binary_cross_entropy(real_preds, real_targets) real_score = torch.mean(real_preds).item() latent = torch.randn(batch_size, latent_sz, 1, 1, device=device) fake_labels = torch.randint(0, num_classes, (batch_size,), device=device) fake_images = generator(latent, fake_labels) fake_preds = discriminator(fake_images, fake_labels) fake_targets = torch.zeros(fake_images.size(0), 1, device=device) fake_loss = F.binary_cross_entropy(fake_preds, fake_targets) fake_score = torch.mean(fake_preds).item() loss = fake_loss + real_loss loss.backward() opt_d.step() return loss.item(), real_score, fake_score def train_generator(opt_g): opt_g.zero_grad() latent = torch.randn(batch_size, latent_sz, 1, 1, device=device) fake_labels = torch.randint(0, num_classes, (batch_size,), device=device) fake_images = generator(latent, fake_labels) targets = torch.ones(batch_size, 1, device=device) score = discriminator(fake_images, fake_labels) loss = F.binary_cross_entropy(score, targets) loss.backward() opt_g.step() return loss.item() def fit(epochs, lr, start_idx=1): loss_d = [] loss_g = [] real_scores = [] fake_scores = [] opt_d = torch.optim.Adam(discriminator.parameters(), lr=lr, betas=(0.5, 0.999)) opt_g = torch.optim.Adam(generator.parameters(), lr=lr, betas=(0.5, 0.999)) for epoch in range(epochs): for img, labels in tqdm(train_loader): img = img.to(device) labels = labels.to(device) loss, real_score, fake_score = train_discriminator(img, labels, opt_d) lossg = train_generator(opt_g) loss_d.append(loss) loss_g.append(lossg) real_scores.append(real_score) fake_scores.append(fake_score) print("Epoch [{}/{}], loss_g: {:.4f}, loss_d: {:.4f}, real_score: {:.4f}, fake_score: {:.4f}, memory_usage: {:.4f}".format( epoch + 1, epochs, loss, lossg, real_score, fake_score, psutil.virtual_memory()[2])) save_samples(epoch + start_idx, fixed_latent, fixed_labels, show=False) return loss_g, loss_d, real_scores, fake_scores lr = 5e-4 epochs = 20 history = fit(epochs,lr)
报错堆栈信息
--------------------------------------------------------------------------- RuntimeError Traceback (most recent call last) /scratch/ipykernel_151601/1116731599.py in <module> 1 lr = 5e-4 2 epochs = 20 ----> 3 history = fit(epochs,lr) /scratch/ipykernel_151601/2486237605.py in fit(epochs, lr, start_idx) 13 labels = labels.to(device) 14 ---> 15 loss, real_score, fake_score = train_discriminator(img, labels, opt_d) 16 lossg = train_generator(opt_g) 17 /scratch/ipykernel_151601/2842949976.py in train_discriminator(real_images, real_labels, opt_d) 2 opt_d.zero_grad() 3 ----> 4 real_preds = discriminator(real_images, real_labels) 5 real_targets = torch.ones(real_images.size(0), 1, device=device) 6 ~/anaconda3/envs/Ashank/lib/python3.7/site-packages/torch/nn/modules/module.py in _call_impl(self, *input, **kwargs) 1192 if not (self._backward_hooks or self._forward_hooks or self._forward_pre_hooks or _global_backward_hooks 1193 or _global_forward_hooks or _global_forward_pre_hooks): -> 1194 return forward_call(*input, **kwargs) 1195 # Do not call functions when jit is used 1196 full_backward_hooks, non_full_backward_hooks = [], [] /scratch/ipykernel_151601/3323902491.py in forward(self, x, labels) 34 35 def forward(self, x, labels): ---> 36 x = torch.cat((x, labels), dim=1) 37 return self.main(x) 38 RuntimeError: Tensors must have same number of dimensions: got 4 and 1
错误修复方案
错误核心是张量维度不匹配:图像张量是4维([batch_size, channel, height, width]),而输入的标签是1维([batch_size]),无法直接在通道维度拼接。需要将标签转换为与图像/潜在向量匹配的维度,并转为one-hot编码(离散类别需转为向量形式)。
1. 判别器修复
修改判别器初始化
将第一个卷积层的输入通道数改为3 + num_classes(原图像3通道 + 类别数通道):
# 替换判别器__init__中的第一个Conv2d nn.Conv2d(3 + num_classes, 32, kernel_size=3, stride=2, padding=1, bias=False),
修改判别器forward方法
把标签转为one-hot编码,并扩展空间维度至与图像一致:
import torch.nn.functional as F def forward(self, x, labels): # 将标签转为one-hot编码,shape变为[batch_size, num_classes] labels_onehot = F.one_hot(labels, num_classes=num_classes).float() # 扩展为4维:[batch_size, num_classes, 1, 1] labels_onehot = labels_onehot.unsqueeze(2).unsqueeze(3) # 重复空间维度,匹配图像的height和width labels_onehot = labels_onehot.repeat(1, 1, x.size(2), x.size(3)) # 在通道维度拼接图像和标签 x = torch.cat((x, labels_onehot), dim=1) return self.main(x)
2. 生成器修复
修改生成器初始化
将第一个转置卷积层的输入通道数改为latent_dim + num_classes(潜在向量维度 + 类别数):
# 替换生成器__init__中的第一个ConvTranspose2d nn.ConvTranspose2d(latent_dim + num_classes, 1024, kernel_size=4, stride=1, padding=0, bias=False),
修改生成器forward方法
把标签转为one-hot编码,并扩展为与潜在向量一致的4维:
import torch.nn.functional as F def forward(self, z, labels): # 将标签转为one-hot编码并扩展为4维:[batch_size, num_classes, 1, 1] labels_onehot = F.one_hot(labels, num_classes=num_classes).float().unsqueeze(2).unsqueeze(3) # 在通道维度拼接潜在向量和标签 x = torch.cat((z, labels_onehot), dim=1) return self.main(x)
3. 补充说明
- 需提前定义
num_classes为数据集的类别总数 - 如果数据加载器返回的标签已经是one-hot编码,可跳过
F.one_hot步骤
256x256图像生成的条件GAN优化思路
针对256x256高分辨率图像生成,可从以下方向优化模型:
- 网络结构:用残差块替换部分卷积层,提升模型表达能力,缓解梯度消失问题
- 归一化策略:判别器改用
InstanceNorm2d替代BatchNorm2d,小批量训练时稳定性更强;生成器保留BatchNorm2d - 损失函数:替换二元交叉熵为WGAN-GP(带梯度惩罚的Wasserstein GAN),训练更稳定,收敛效果更好
- 训练策略:采用交替训练(判别器训练1-5次,生成器训练1次);加入学习率衰减,后期逐步降低学习率
- 数据增强:对训练图像进行随机裁剪、翻转、颜色抖动等操作,提升生成图像的多样性和泛化能力
内容的提问来源于stack exchange,提问作者Ashank Shah
相关产品推荐
相关产品推荐

