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

PyTorch搭建GAN自定义batch size及DataLoader维度报错咨询

修复方案

1. 修改判别器前向传播逻辑

把原判别器全局展平的逻辑,改为仅展平样本维度,保留batch维度:

# 原代码
def forward(self, x):
    x = x.flatten()
    x = self.fc(x)
    return x

# 修改后代码
def forward(self, x):
    # start_dim=1 表示从第1维(即样本维度)开始展平,第0维batch维度保留
    x = x.flatten(start_dim=1)
    x = self.fc(x)
    return x

修改后输入为(batch_size,15,50)的批量数据时,展平后的维度为(batch_size,750),可以完美匹配全连接层的输入维度要求。

2. 修正训练循环逻辑

需要调整标签维度、生成器输入噪声维度,补充生成器梯度清零逻辑,修改后代码如下:

for epoch in range(2000):
    for i, series in enumerate(dataloader):
        batch_size = series.shape[0]
        # 动态生成对应batch大小的标签,不写死固定值适配最后一个不足batch的情况
        valid = torch.ones(batch_size, 1, requires_grad=False).to("cuda")
        fake = torch.zeros(batch_size, 1, requires_grad=False).to("cuda")
        # 配置真实样本
        real_series = series.float().to("cuda")

        # -----------------
        #  训练生成器
        # -----------------
        optimizer_G.zero_grad() # 新增:生成器梯度清零,避免梯度累加
        # 生成对应batch大小的噪声输入
        z = torch.tensor(np.random.normal(0, 1, (batch_size, 19))).float().to("cuda")
        # 生成批量样本
        gen_series = generator(z).float().to("cuda")
        # 计算生成器损失
        g_loss = adversarial_loss(discriminator(gen_series), valid)
        g_loss.backward()
        optimizer_G.step()

        # ---------------------
        #  训练判别器
        # ---------------------
        optimizer_D.zero_grad()
        # 计算真实样本损失
        real_loss = adversarial_loss(discriminator(real_series), valid)
        # 计算生成样本损失
        fake_loss = adversarial_loss(discriminator(gen_series.detach()), fake)
        d_loss = 0.5 * (real_loss + fake_loss)
        d_loss.backward()
        optimizer_D.step()
        
        # 输出loss
        if i % 200 == 0:
            print(
                "[Epoch %d/%d] [Batch %d/%d] [D loss: %f] [G loss: %f] [real loss/fake loss: %f/%f] \n "
                % (epoch, 2000, i, len(dataloader), d_loss.item(), g_loss.item(), real_loss.item(), fake_loss.item())
            )

其他注意点

  • 原逐样本训练代码中fake标签被错误赋值为1,会导致判别器训练逻辑完全错误,上述修正代码已经统一调整为0
  • 新版PyTorch已经不需要用Variable包裹张量,直接用原生tensor即可,上述代码已经做了兼容简化
  • 生成器输出需要保证维度为(batch_size,15,50),如果生成器输出的是单样本维度,需要检查生成器最后一层的reshape逻辑,确保保留batch维度

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.09.30 04:36:05