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

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.13 09:35:10