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
相关产品推荐
相关产品推荐

