使用CNN与Transposed CNN实现Text2Image遇张量拼接错误,求解决
Text2Image网络实现的张量拼接错误解决
问题概述
需实现基于CNN、转置CNN和Embedding层的Text2Image网络,当前代码运行时报错,推测源于Embedding结果与噪声张量的拼接操作。
导入依赖库
import torch from torch import nn
输入文本预处理
text = "A cat wearing glasses and playing the guitar " # 简单文本预处理 word_to_ix = {"A": 0, "cat": 1, "wearing": 2, "glasses": 3, "and": 4, "playing": 5, "the": 6, "guitar":7} lookup_tensor = torch.tensor(list(word_to_ix.values()), dtype = torch.long) # 用整数表示单词的张量 vocab_size = len(lookup_tensor)
网络架构实现
class TextToImage(nn.Module): def __init__(self, vocab_size): super(TextToImage, self).__init__() self.vocab_size = vocab_size self.noise = torch.rand((56,64)) # 定义层 # Embedding层 self.embed = nn.Embedding(num_embeddings=self.vocab_size, embedding_dim = 64) # CNN层 self.conv2d_1 = nn.Conv2d(in_channels=64, out_channels=3, kernel_size=(3, 3), stride=(2, 2), padding='valid') self.conv2d_2 = nn.Conv2d(in_channels=3, out_channels=16, kernel_size=(3, 3), stride=(2, 2), padding='valid') # 转置CNN层 self.conv2dTran_1 = nn.ConvTranspose2d(in_channels=16, out_channels=16, kernel_size=(3, 3), stride=(1, 1), padding=1) self.conv2dTran_2 = nn.ConvTranspose2d(in_channels=16, out_channels=3, kernel_size=(3, 3), stride=(2, 2), padding=0) self.conv2dTran_3 = nn.ConvTranspose2d(in_channels=6, out_channels=3, kernel_size=(4, 4), stride=(2, 2), padding=0) self.relu = torch.nn.ReLU(inplace=False) self.dropout = torch.nn.Dropout(0.4) def forward(self, text_tensor): # 将输入文本张量传入Embedding层 emb = self.embed(text_tensor) # 将Embedding结果与噪声张量拼接,转为3维 combine1 = torch.cat((emb, self.noise), dim=1, out=None) # 将含噪Embedding传入CNN和转置CNN层 conv2d_1 = self.conv2d_1(combine1) conv2d_2 = self.conv2d_2(conv2d_1) dropout = self.dropout(conv2d_2) conv2dTran_1 = self.conv2dTran_1(dropout) conv2dTran_2 = self.conv2dTran_2(conv2dTran_1) # 按架构图的跳连接拼接输出 combine2 = torch.cat((conv2d_1, conv2dTran_2), dim=1, out=None) conv2dTran_3 = self.conv2dTran_3(combine2) # 将拼接结果传入最终层,输出命名为image image = self.relu(conv2dTran_3) return image
期望输出
torch.Size( [3, 64, 64] )
测试代码
texttoimage = TextToImage(vocab_size=vocab_size) output = texttoimage(lookup_tensor) output.size()
报错信息
RuntimeError: Sizes of tensors must match except in dimension 1. Expected size 8 but got size 56 for tensor number 1 in the list.
错误原因分析
- 张量形状不匹配:Embedding输出为
(8,64)(序列长度8,嵌入维度64),噪声张量为(56,64),在dim=1拼接时,第一个维度(8 vs 56)无法对齐,触发报错。 - 卷积层输入格式错误:PyTorch的Conv2d要求输入为4D张量
(batch_size, channels, height, width),当前代码中直接传入2D张量,后续即使拼接成功也会报错。
解决方案
1. 调整张量维度与拼接逻辑
修改forward函数,将Embedding结果转换为卷积层要求的4D格式,并生成匹配形状的噪声张量,在空间维度拼接:
def forward(self, text_tensor): # Embedding输出形状:(seq_len, embed_dim) = (8,64) emb = self.embed(text_tensor) # 转换为4D卷积输入格式:(batch_size, channels, height, width) # 将嵌入维度作为通道,序列长度作为高度,宽度设为64匹配最终输出 emb = emb.T.unsqueeze(0).unsqueeze(-1).repeat(1,1,1,64) # shape: (1,64,8,64) # 生成匹配形状的噪声,避免在__init__中固定形状导致批量不兼容 noise = torch.randn_like(emb)[:, :, :56, :] # shape: (1,64,56,64) # 在高度维度拼接,得到(1,64,64,64) combine1 = torch.cat((emb, noise), dim=2) # 卷积层添加激活函数,增强非线性表达 conv2d_1 = self.relu(self.conv2d_1(combine1)) conv2d_2 = self.relu(self.conv2d_2(conv2d_1)) dropout = self.dropout(conv2d_2) conv2dTran_1 = self.relu(self.conv2dTran_1(dropout)) conv2dTran_2 = self.relu(self.conv2dTran_2(conv2dTran_1)) # 跳连接需保证张量形状匹配,若不匹配可调整卷积参数或用Upsample对齐 combine2 = torch.cat((conv2d_1, conv2dTran_2), dim=1) conv2dTran_3 = self.conv2dTran_3(combine2) # 移除batch维度,得到期望输出形状 image = self.relu(conv2dTran_3).squeeze(0) return image
2. 调整卷积层参数以保证形状传递
修改卷积和转置卷积的参数,确保中间张量形状匹配,支持跳连接拼接:
# 修改CNN层,用same padding保持宽高比例 self.conv2d_1 = nn.Conv2d(in_channels=64, out_channels=3, kernel_size=(3, 3), stride=(2, 2), padding='same') self.conv2d_2 = nn.Conv2d(in_channels=3, out_channels=16, kernel_size=(3, 3), stride=(2, 2), padding='same') # 修改转置CNN层,调整padding和output_padding恢复形状,匹配跳连接要求 self.conv2dTran_1 = nn.ConvTranspose2d(in_channels=16, out_channels=16, kernel_size=(3, 3), stride=(1, 1), padding=1) self.conv2dTran_2 = nn.ConvTranspose2d(in_channels=16, out_channels=3, kernel_size=(3, 3), stride=(2, 2), padding=1, output_padding=1) self.conv2dTran_3 = nn.ConvTranspose2d(in_channels=6, out_channels=3, kernel_size=(4, 4), stride=(2, 2), padding=1, output_padding=1)
3. 其他优化建议
- 噪声张量建议在
forward中生成或定义为可学习参数(nn.Parameter(torch.randn(1,64,56,64))),适配不同批量大小。 - 所有卷积/转置卷积层后添加激活函数,提升网络表达能力。
测试验证
运行修改后的代码,输出将符合预期:
texttoimage = TextToImage(vocab_size=vocab_size) output = texttoimage(lookup_tensor) print(output.size()) # 输出: torch.Size([3, 64, 64])
内容的提问来源于stack exchange,提问作者Mohammed
相关产品推荐
相关产品推荐

