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

构建带标签输入的MNIST GAN时遇矩阵维度不匹配错误求助

解决GAN生成器输入维度不匹配问题

错误原因

报错RuntimeError: mat1 and mat2 shapes cannot be multiplied (1x110 and 100x256)的核心问题是输入维度不匹配:你将10维的one-hot标签向量与100维的噪声向量拼接后,得到了1×110的输入张量,但生成器的第一个全连接层nn.Linear(100, 256)仅接受100维输入,两者维度冲突导致矩阵乘法无法执行。

解决方案

修改生成器第一个全连接层的输入维度为拼接后的总维度(10+100=110),后续层的维度保持不变即可。

修改后的完整代码

import torch
import torch.nn as nn
import torch.nn.functional as F

class Generator(nn.Module):
    def __init__(self):
        super(Generator, self).__init__()
        # 调整第一个全连接层输入维度为110(标签10维+噪声100维)
        self.fc1 = nn.Linear(110, 256)
        self.fc2 = nn.Linear(256, 512)
        self.fc3 = nn.Linear(512, 1024)
        self.fc4 = nn.Linear(1024, 784)

    def forward(self, x):
        x = F.relu(self.fc1(x))
        x = F.relu(self.fc2(x))
        x = F.relu(self.fc3(x))
        x = torch.tanh(self.fc4(x))
        return x

# 初始化生成器实例
generator = Generator()

# 设置目标标签
label = 3

# 生成one-hot编码标签
one_hot = torch.zeros(1, 10)
one_hot[0][label] = 1

# 生成噪声向量
noise = torch.randn(1, 100)

# 拼接标签与噪声
noise_with_label = torch.cat([one_hot, noise], dim=1)

# 生成图像
generated_image = generator(noise_with_label)
# 验证输出形状(MNIST图像展平后为784维)
print(generated_image.shape)  # 输出应为 torch.Size([1, 784])

额外提示

  • 若需要批量生成图像,只需同步调整one-hot标签和噪声向量的批量维度(例如torch.zeros(32,10)和torch.randn(32,100)),生成器会自动兼容批量输入
  • 训练阶段需保证生成器的输入始终是标签向量+噪声向量的拼接结果,维度维持110维一致

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.01 16:16:05