从零实现文本嵌入条件GAN遇通道不匹配错误,求架构修正
错误修复与架构优化方案
核心错误:通道维度完全不匹配
你遇到的RuntimeError本质是生成器第一个转置卷积层的输入通道数与实际输入的通道数完全不匹配:
- 转置卷积层期望输入有456通道,但你的输入通道数是1049088——这个异常大的数值说明你错误地把拼接后的噪声+文本嵌入向量展平成了通道维度,而非保持合理的通道数。
直接修复步骤
1. 修正生成器输入的维度转换
假设噪声z形状为[batch_size, z_dim],文本嵌入text_emb形状为[batch_size, 768],拼接后要转换成卷积层可处理的4D张量(通道维度在前):
def forward(self, z, text_emb): # 拼接噪声与文本嵌入:[batch_size, z_dim + 768] x = torch.cat([z, text_emb], dim=1) # 转换为4D张量:[batch_size, 通道数, 1, 1],通道数=z_dim+768 x = x.view(x.size(0), x.size(1), 1, 1) # 后续转置卷积层处理 x = self.conv1(x) ...
2. 修正生成器第一个转置卷积层的输入通道数
生成器第一个转置卷积的in_channels必须等于z_dim + 768(比如z_dim设为256时,in_channels=256+768=1024),示例生成器适配512×512图像输出:
import torch.nn as nn import torch class Generator(nn.Module): def __init__(self, z_dim=256, text_dim=768, img_channels=3): super().__init__() # 从1×1特征图开始,逐步提升分辨率到512×512 self.conv1 = nn.ConvTranspose2d(z_dim + text_dim, 512, kernel_size=4, stride=2, padding=1) self.bn1 = nn.BatchNorm2d(512) self.conv2 = nn.ConvTranspose2d(512, 256, 4, 2, 1) self.bn2 = nn.BatchNorm2d(256) self.conv3 = nn.ConvTranspose2d(256, 128, 4, 2, 1) self.bn3 = nn.BatchNorm2d(128) self.conv4 = nn.ConvTranspose2d(128, 64, 4, 2, 1) self.bn4 = nn.BatchNorm2d(64) self.conv5 = nn.ConvTranspose2d(64, 32, 4, 2, 1) self.bn5 = nn.BatchNorm2d(32) self.conv6 = nn.ConvTranspose2d(32, 16, 4, 2, 1) self.bn6 = nn.BatchNorm2d(16) self.conv7 = nn.ConvTranspose2d(16, img_channels, 4, 2, 1) self.relu = nn.ReLU(inplace=True) self.tanh = nn.Tanh() def forward(self, z, text_emb): x = torch.cat([z, text_emb], dim=1) x = x.view(x.size(0), x.size(1), 1, 1) x = self.relu(self.bn1(self.conv1(x))) # [batch,512,2,2] x = self.relu(self.bn2(self.conv2(x))) # [batch,256,4,4] x = self.relu(self.bn3(self.conv3(x))) # [batch,128,8,8] x = self.relu(self.bn4(self.conv4(x))) # [batch,64,16,16] x = self.relu(self.bn5(self.conv5(x))) # [batch,32,32,32] x = self.relu(self.bn6(self.conv6(x))) # [batch,16,64,64] x = self.tanh(self.conv7(x)) # [batch,3,512,512],输出归一化到[-1,1] return x
架构其他潜在问题修正
1. 判别器的文本嵌入注入方式
判别器需要同时处理图像和文本条件,推荐采用「文本嵌入投影后与图像特征拼接」的方式,确保维度匹配:
class Discriminator(nn.Module): def __init__(self, text_dim=768, img_channels=3): super().__init__() # 图像特征提取:逐步降低分辨率、提升通道数 self.conv1 = nn.Conv2d(img_channels, 16, 4, 2, 1) self.conv2 = nn.Conv2d(16, 32, 4, 2, 1) self.bn2 = nn.BatchNorm2d(32) self.conv3 = nn.Conv2d(32, 64, 4, 2, 1) self.bn3 = nn.BatchNorm2d(64) self.conv4 = nn.Conv2d(64, 128, 4, 2, 1) self.bn4 = nn.BatchNorm2d(128) self.conv5 = nn.Conv2d(128, 256, 4, 2, 1) self.bn5 = nn.BatchNorm2d(256) # 文本嵌入投影:将768维嵌入投影到与图像特征匹配的通道数 self.text_proj = nn.Linear(text_dim, 256) # 最终判别层:拼接图像+文本特征后输出判别结果 self.final_conv = nn.Conv2d(256 + 256, 1, 4, 2, 1) self.leaky_relu = nn.LeakyReLU(0.2, inplace=True) def forward(self, img, text_emb): # 处理图像:[batch,3,512,512] → [batch,256,16,16] x = self.leaky_relu(self.conv1(img)) x = self.leaky_relu(self.bn2(self.conv2(x))) x = self.leaky_relu(self.bn3(self.conv3(x))) x = self.leaky_relu(self.bn4(self.conv4(x))) x = self.leaky_relu(self.bn5(self.conv5(x))) # 处理文本嵌入:广播到与图像特征相同的空间尺寸 text_feat = self.text_proj(text_emb) text_feat = text_feat.view(text_feat.size(0), text_feat.size(1), 1, 1) text_feat = text_feat.repeat(1, 1, x.size(2), x.size(3)) # 拼接特征并判别 x = torch.cat([x, text_feat], dim=1) x = self.final_conv(x) return x.view(x.size(0), -1)
2. 训练流程的关键细节
- 确保噪声
z的形状为[batch_size, z_dim],文本嵌入text_emb为[batch_size,768],避免维度不匹配 - 图像数据需归一化到
[-1,1],匹配生成器最后一层的Tanh输出 - 采用条件GAN损失函数,训练时需将文本嵌入同时传入生成器和判别器:
def train_step(gen, disc, real_imgs, text_embs, z, opt_gen, opt_disc, criterion): batch_size = real_imgs.size(0) # 训练判别器 opt_disc.zero_grad() real_pred = disc(real_imgs, text_embs) real_loss = criterion(real_pred, torch.ones_like(real_pred)) fake_imgs = gen(z, text_embs) fake_pred = disc(fake_imgs.detach(), text_embs) fake_loss = criterion(fake_pred, torch.zeros_like(fake_pred)) disc_loss = (real_loss + fake_loss) / 2 disc_loss.backward() opt_disc.step() # 训练生成器 opt_gen.zero_grad() fake_pred = disc(fake_imgs, text_embs) gen_loss = criterion(fake_pred, torch.ones_like(fake_pred)) gen_loss.backward() opt_gen.step() return disc_loss.item(), gen_loss.item()
内容的提问来源于stack exchange,提问作者Ploutos Connecte
相关产品推荐
相关产品推荐

