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

使用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.

错误原因分析

  1. 张量形状不匹配:Embedding输出为(8,64)(序列长度8,嵌入维度64),噪声张量为(56,64),在dim=1拼接时,第一个维度(8 vs 56)无法对齐,触发报错。
  2. 卷积层输入格式错误: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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.06 11:02:10