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

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 
---&gt; 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):
-&gt; 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):
---&gt; 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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.14 17:41:01