PyTorch图像字幕Decoder LSTM输入尺寸疑问咨询
解答:图像字幕Decoder中LSTM输入维度不匹配的问题
嘿,我刚入门PyTorch做图像字幕的时候也碰到过一模一样的困惑!咱们来拆解一下问题,然后看看常见的解决思路:
你观察得很准确——直接把字幕嵌入(维度embed_size)和Encoder输出的上下文特征(通常是更大的维度,比如2048维,来自CNN的最后一层)拼接后,维度肯定会超过LSTM定义的输入尺寸,这时候确实不能直接喂进去。你大概率是没注意到示例代码里藏着的维度映射或者状态初始化的逻辑,这是图像字幕Decoder的关键细节之一。
常见的两种解决方案如下:
方案1:用上下文特征初始化LSTM的隐藏状态
这是最常用的一种方式:我们不把上下文特征和字幕嵌入拼接,而是把它转换成LSTM的初始隐藏状态(h0)和细胞状态(c0),这样LSTM的输入就只有字幕嵌入,维度完全匹配。
举个代码例子:
import torch import torch.nn as nn class DecoderRNN(nn.Module): def __init__(self, embed_size, hidden_size, vocab_size, encoder_dim): super().__init__() self.embedding = nn.Embedding(vocab_size, embed_size) # LSTM输入维度就是embed_size,和字幕嵌入一致 self.lstm = nn.LSTM(embed_size, hidden_size, batch_first=True) # 两个线性层,把Encoder的高维特征映射到LSTM的隐藏状态维度 self.init_hidden = nn.Linear(encoder_dim, hidden_size) self.init_cell = nn.Linear(encoder_dim, hidden_size) self.fc = nn.Linear(hidden_size, vocab_size) def forward(self, features, captions): # features: 来自Encoder的上下文特征,shape=(batch_size, encoder_dim) # captions: 输入字幕序列,shape=(batch_size, seq_len) embeddings = self.embedding(captions) # shape=(batch_size, seq_len, embed_size) # 初始化LSTM的h0和c0,注意要扩展维度适配LSTM的输入格式 h0 = self.init_hidden(features).unsqueeze(0) # shape=(1, batch_size, hidden_size) c0 = self.init_cell(features).unsqueeze(0) # 把嵌入序列和初始状态传入LSTM outputs, _ = self.lstm(embeddings, (h0, c0)) outputs = self.fc(outputs) # 映射到词汇表维度,shape=(batch_size, seq_len, vocab_size) return outputs
方案2:拼接后用线性层映射到LSTM输入维度
如果你确实需要把上下文特征和每个时间步的字幕嵌入结合起来,可以先把上下文特征复制到和字幕序列一样的长度,拼接后用一个线性层把维度降到embed_size,再喂给LSTM。
代码示例:
import torch import torch.nn as nn class DecoderRNN(nn.Module): def __init__(self, embed_size, hidden_size, vocab_size, encoder_dim): super().__init__() self.embedding = nn.Embedding(vocab_size, embed_size) # 线性层:把拼接后的高维特征映射到LSTM需要的embed_size self.feature_proj = nn.Linear(encoder_dim + embed_size, embed_size) self.lstm = nn.LSTM(embed_size, hidden_size, batch_first=True) self.fc = nn.Linear(hidden_size, vocab_size) def forward(self, features, captions): embeddings = self.embedding(captions) # shape=(batch_size, seq_len, embed_size) # 把上下文特征扩展到序列长度维度,方便拼接 features_expanded = features.unsqueeze(1).repeat(1, embeddings.size(1), 1) # shape=(batch_size, seq_len, encoder_dim) # 拼接嵌入和上下文特征 concat_features = torch.cat([embeddings, features_expanded], dim=-1) # shape=(batch_size, seq_len, embed_size+encoder_dim) # 映射到LSTM输入维度 lstm_input = self.feature_proj(concat_features) # shape=(batch_size, seq_len, embed_size) outputs, _ = self.lstm(lstm_input) outputs = self.fc(outputs) return outputs
你可以回头看看示例代码,大概率是用了第一种方案——把Encoder特征用来初始化LSTM状态,而不是直接拼接。很多时候这些线性层的定义可能比较简洁,容易被忽略~
内容的提问来源于stack exchange,提问作者Arijit
相关产品推荐
相关产品推荐

