使用torch.no_grad()时TransformerEncoder评估阶段序列长度异常
问题:TransformerEncoder在eval模式下输出序列长度异常变化
问题描述
构建TransformerEncoder模型后,训练模式下输入形状为(32,64)的张量,输出符合预期的(32,64,256);但切换到eval模式并使用with torch.no_grad()上下文时,输出序列长度随机变化(如(32,31,256)、(32,25,256)),无法得到固定的(32,64,256)输出。
可能原因及解决方案
1. 模型forward方法存在语法错误
从提供的代码来看,forward函数未正确缩进为TransEnc类的成员方法,且缺少返回语句,这会导致模型使用默认的nn.Module.forward,行为不可控。
修正代码:
class TransEnc(nn.Module): def __init__(self, ntoken: int, encoder_embedding_dim: int, max_item_count: int, encoder_num_heads: int, encoder_hidden_dim: int, encoder_num_layers: int, padding_idx: int, dropout: float = 0.2): super().__init__() self.encoder_embedding = nn.Embedding(ntoken, encoder_embedding_dim, padding_idx=padding_idx) self.pos_encoder = PositionalEncoding(encoder_embedding_dim, max_item_count, dropout) encoder_layers = nn.TransformerEncoderLayer(encoder_embedding_dim, encoder_num_heads, encoder_hidden_dim, dropout, batch_first=True) self.transformer_encoder = nn.TransformerEncoder(encoder_layers, encoder_num_layers) self.encoder_embedding_dim = encoder_embedding_dim # 修正缩进,确保是类的成员方法 def forward(self, src: torch.Tensor, src_key_padding_mask: torch.Tensor = None) -> torch.Tensor: src = self.encoder_embedding(src.long()) * math.sqrt(self.encoder_embedding_dim) src = self.pos_encoder(src) src = self.transformer_encoder(src, src_key_padding_mask=src_key_padding_mask) return src # 添加返回语句
2. PositionalEncoding实现异常
如果自定义的PositionalEncoding类在eval模式下对输入序列进行了动态截断(比如仅保留非padding部分),会导致输出长度变化。
确保PositionalEncoding不修改序列长度:
使用PyTorch官方标准的位置编码实现,仅添加位置信息,不改变输入形状:
import math import torch import torch.nn as nn class PositionalEncoding(nn.Module): def __init__(self, d_model: int, max_len: int, dropout: float = 0.1): super().__init__() self.dropout = nn.Dropout(p=dropout) position = torch.arange(max_len).unsqueeze(1) div_term = torch.exp(torch.arange(0, d_model, 2) * (-math.log(10000.0) / d_model)) pe = torch.zeros(max_len, 1, d_model) pe[:, 0, 0::2] = torch.sin(position * div_term) pe[:, 0, 1::2] = torch.cos(position * div_term) self.register_buffer('pe', pe) def forward(self, x: torch.Tensor) -> torch.Tensor: # 适配batch_first=True的输入格式:(batch_size, seq_len, d_model) x = x + self.pe[:x.size(1)] return self.dropout(x)
3. src_key_padding_mask格式错误
PyTorch的TransformerEncoderLayer(batch_first=True时)要求src_key_padding_mask形状为(batch_size, seq_len),且为布尔类型(True表示对应位置是padding token)。若mask形状或类型错误,可能导致模型行为异常。
验证并修正mask:
# 确保mask的形状和类型正确 src_key_padding_mask = (src == tokenizer.pad_token_id).bool() print(f"mask shape: {src_key_padding_mask.shape}") # 应输出torch.Size([32, 64])
4. 定位形状变化的具体步骤
在forward方法中添加打印语句,确认每一步的张量形状,找到长度变化的环节:
def forward(self, src: torch.Tensor, src_key_padding_mask: torch.Tensor = None) -> torch.Tensor: print(f"Input src shape: {src.shape}") src = self.encoder_embedding(src.long()) * math.sqrt(self.encoder_embedding_dim) print(f"After embedding shape: {src.shape}") src = self.pos_encoder(src) print(f"After positional encoding shape: {src.shape}") src = self.transformer_encoder(src, src_key_padding_mask=src_key_padding_mask) print(f"Transformer output shape: {src.shape}") return src
内容的提问来源于stack exchange,提问作者Leon
相关产品推荐
相关产品推荐

