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

Transformer用于复述生成失效:生成内容异常问题排查求助

问题:Transformer复述生成模型生成结果异常(全BOS/PAD token)

基于PyTorch实现Transformer用于复述生成,训练过程中损失持续下降,但生成结果毫无价值——要么全是BOS token,要么充满无意义符号;仅在训练初期(第2-3epoch)短暂生成过通顺句子,随后迅速退化。以下是具体实现细节及异常表现:

核心代码实现

编码器

class TransformerEncoder(nn.Module):
    def __init__(
        self,
        vocab_size,
        pad_token_id=None,
        embedding_size=256,
        num_heads=8,
        num_layers=3,
        ffnn_size=512,
        dropout=0.1,
    ):
        super(TransformerEncoder, self).__init__()
        self.vocab_size = vocab_size
        self.pad_token_id = pad_token_id

        self.embedding_size = embedding_size
        self.num_heads = num_heads
        self.num_layers = num_layers
        self.ffnn_size = ffnn_size

        self.embed_tokens = TokenEmbedding(vocab_size, embedding_size)
        self.embed_positions = PositionalEmbedding(embedding_size, dropout=dropout)

        encoder_layer = nn.TransformerEncoderLayer(
            embedding_size,
            num_heads,
            ffnn_size,
            dropout,
        )
        encoder_norm = nn.LayerNorm(embedding_size)
        self.encoder = nn.TransformerEncoder(encoder_layer, num_layers, encoder_norm)

    def forward(
        self,
        input_ids,
    ):
        embedded_tokens = self.embed_positions(self.embed_tokens(input_ids))
        # B x T x C -> T x B x C
        embedded_tokens = embedded_tokens.transpose(0, 1)
        memory = self.encoder(embedded_tokens)
        return (memory,)

解码器

class TransformerDecoder(nn.Module):
    def __init__(
        self,
        vocab_size,
        pad_token_id=None,
        embedding_size=256,
        num_heads=8,
        num_layers=3,
        ffnn_size=512,
        dropout=0.1,
    ):
        super(TransformerDecoder, self).__init__()
        self.vocab_size = vocab_size
        self.pad_token_id = pad_token_id

        self.embedding_size = embedding_size
        self.num_heads = num_heads
        self.num_layers = num_layers
        self.ffnn_size = ffnn_size

        self.dropout_module = nn.Dropout(p=dropout)

        self.embed_tokens = TokenEmbedding(vocab_size, embedding_size)
        self.embed_positions = PositionalEmbedding(embedding_size, dropout=dropout)

        decoder_layer = nn.TransformerDecoderLayer(
            embedding_size, num_heads, ffnn_size, dropout
        )
        decoder_norm = nn.LayerNorm(embedding_size)
        self.decoder = nn.TransformerDecoder(decoder_layer, num_layers, decoder_norm)
        self.fc_out = nn.Linear(embedding_size, vocab_size)

    def forward(
        self,
        input_ids,
        encoder_out,
    ):
        seq_len = input_ids.shape[1]
        device = next(self.parameters()).device
        mask = generate_square_subsequent_mask(seq_len).to(device)

        embedded_tokens = self.embed_positions(self.embed_tokens(input_ids))
        # B x T x C -> T x B x C
        embedded_tokens = embedded_tokens.transpose(0, 1)
        output = self.decoder(embedded_tokens, encoder_out[0], tgt_mask=mask)
        # T x B x C -> B x T x C
        output = output.transpose(1, 0)
        return (self.fc_out(output),)

主模型调用逻辑

encoder_outputs = self.encoder(input_ids=input_ids, **kwargs)
decoder_outputs = self.decoder(
    input_ids=decoder_input_ids,
    encoder_out=encoder_outputs,
    **kwargs,
)

标签右移处理

def shift_tokens_right(self, input_ids: torch.Tensor, decoder_start_token_id: int):
   shifted_input_ids = input_ids.new_zeros(input_ids.shape)
   shifted_input_ids[:, 1:] = input_ids[:, :-1].clone()
   shifted_input_ids[:, 0] = decoder_start_token_id
   return shifted_input_ids

损失计算

loss_fct = nn.CrossEntropyLoss(ignore_index=self.pad_token_id)
loss = loss_fct(logits.reshape(-1, logits.shape[-1]), targets.reshape(-1))

异常生成示例

  • 输入:<s> Can I jailbreak iOS 10 ? </s> <pad> <pad> ...(后续为PAD token)
  • 生成结果:<s> <s> <s> ...(全为BOS token)
  • 目标输出:<s> Can you jailbreak iOS 10 ? </s> <pad> <pad> ...

可能的问题排查方向

1. 缺失PAD掩码处理

Transformer的注意力机制会默认关注所有位置,包括PAD token,这会严重干扰有效信息的学习:

  • 编码器需添加src_key_pad_mask,过滤输入中的PAD token:
    # 在encoder forward中添加
    src_key_pad_mask = (input_ids == self.pad_token_id)
    memory = self.encoder(embedded_tokens, src_key_pad_mask=src_key_pad_mask)
    
  • 解码器需同时添加tgt_key_pad_mask(过滤自身输入的PAD)和memory_key_pad_mask(过滤编码器输出中的PAD):
    # 在decoder forward中添加
    tgt_key_pad_mask = (input_ids == self.pad_token_id)
    memory_key_pad_mask = (encoder_input_ids == self.pad_token_id)  # 需传入编码器的原始input_ids
    output = self.decoder(embedded_tokens, encoder_out[0], tgt_mask=mask, 
                          tgt_key_pad_mask=tgt_key_pad_mask, memory_key_pad_mask=memory_key_pad_mask)
    

2. 生成阶段的解码逻辑问题

  • 如果生成时采用贪婪解码,且未设置停止条件(如遇到EOS token就终止生成),模型可能因为早期位置过度拟合BOS,导致循环生成BOS;
  • 可尝试加入温度系数调整概率分布,或改用beam search替代贪婪解码,避免陷入局部最优。

3. 位置编码与嵌入层初始化问题

确认PositionalEmbedding的实现是否正确:比如是否采用正弦余弦编码,或是可学习编码?若位置编码错误,模型无法学习序列的顺序信息,导致生成内容无意义。

4. 训练稳定性问题

  • 尝试使用AdamW优化器替代基础Adam,添加学习率衰减策略(如StepLR或CosineAnnealingLR),避免后期学习率过高导致模型震荡退化;
  • 检查是否存在梯度爆炸/消失:可添加梯度裁剪(torch.nn.utils.clip_grad_norm_),稳定训练过程。

5. 损失计算的对齐问题

确认targets与解码器输出的位置是否完全对齐:解码器输入是右移后的标签,输出的每个位置应对应原标签的下一个位置,若对齐错误,模型会学习到错误的映射关系。


内容的提问来源于stack exchange,提问作者AfonsoSalgadoSousa

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.08 19:05:23