Bert2Bert文本摘要模型预测输出重复token问题求助
基于Transformer实现了Bert2Bert(Encoder-Decoder架构)用于文本摘要任务,训练阶段损失值持续下降,但在预测阶段模型输出重复的token序列(如[2,2,2,……])。以下是模型、训练及预测代码:
模型代码
class BERT2BERT(nn.Module): def __init__(self, encoder_model_name='HooshvareLab/bert-base-parsbert-uncased', decoder_model_name='HooshvareLab/bert-base-parsbert-uncased',tokenizer=None,vocab_size=0): super(BERT2BERT, self).__init__() self.vocab_size=vocab_size # Encoder: ParsBERT transformer layers self.encoder = BertModel.from_pretrained(encoder_model_name) # Decoder: ParsBERT transformer layers with modifications decoder_config = BertConfig.from_pretrained(decoder_model_name) decoder_config.is_decoder = True # Enable cross-attention decoder_config.add_cross_attention = True decoder_config.decoder_start_token_id=tokenizer.cls_token_id self.decoder = BertModel(config=decoder_config) self.lm_head = nn.Linear(decoder_config.hidden_size, self.vocab_size) # Adjust weights #self.initialize_decoder_weights() # def initialize_decoder_weights(self): # for name, param in self.decoder.named_parameters(): # if "cross_attention" in name: # # Random initialization for cross-attention layers # if param.requires_grad: # nn.init.xavier_uniform_(param.data) # else: # # Use pre-trained weights from ParsBERT # param.data.copy_(self.encoder.state_dict().get(name, param.data)) def forward(self, input_ids, attention_mask, decoder_input_ids, decoder_attention_mask): # Encoder: Encode the input sequence encoder_outputs = self.encoder(input_ids=input_ids, attention_mask=attention_mask) encoded_sequence = encoder_outputs.last_hidden_state # Decoder: Decode the sequence conditioned on the encoder outputs decoder_outputs = self.decoder( input_ids=decoder_input_ids, attention_mask=decoder_attention_mask, encoder_hidden_states=encoded_sequence, encoder_attention_mask=attention_mask, ) logits = self.lm_head(decoder_outputs.last_hidden_state) return logits
训练阶段代码
def train_model(model,tokenizer, input_texts, target_summaries): optimizer = AdamW(model.parameters(), lr=LEARNING_RATE) padding_idx = tokenizer.pad_token_id loss_fn = nn.CrossEntropyLoss(ignore_index=padding_idx) # Create PyTorch DataLoader train_data = list(zip(input_texts, target_summaries)) train_loader = torch.utils.data.DataLoader(train_data, batch_size=BATCH_SIZE, shuffle=True) for epoch in range(EPOCH): model.train() total=len(train_loader) index=0 for batch in train_loader: optimizer.zero_grad() # Tokenize and convert to tensors input_texts, target_texts = batch inputs = encode_batch(input_texts, tokenizer) targets = encode_batch(target_texts, tokenizer) input_ids = inputs['input_ids'] attention_mask = inputs['attention_mask'] decoder_input_ids = targets['input_ids'] decoder_attention_mask = targets['attention_mask'] # Forward pass outputs = model( input_ids=input_ids, attention_mask=attention_mask, decoder_input_ids=decoder_input_ids, decoder_attention_mask=decoder_attention_mask ) logits = outputs labels = decoder_input_ids logits=logits.view(-1, logits.size(-1)) loss = loss_fn(logits, labels.view(-1)) loss.backward() torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0) optimizer.step() index+=1 print(f"Epoch {epoch + 1}, Total Batches:{total} ,Batch:{index}, Loss: {loss.item()}") torch.save(model.state_dict(), "model.model")
预测阶段代码
def evaluate_model(model, tokenizer, test_texts, reference_summaries): #model.load_state_dict(torch.load(MODEL_PATH)) model.eval() rou = Rouge() predictions = [] for text in test_texts: # Tokenize the input text inputs = tokenizer( text, return_tensors="pt", max_length=512, truncation=True, padding="max_length", ) # Generate decoder inputs (start with <sos> token) decoder_input_ids = torch.tensor([[tokenizer.cls_token_id]]) decoder_attention_mask = torch.ones_like(decoder_input_ids) output_tokens = [] with torch.no_grad(): for _ in range(50): # Generate up to 50 tokens logits = model( input_ids=inputs["input_ids"], attention_mask=inputs["attention_mask"], decoder_input_ids=decoder_input_ids, decoder_attention_mask=decoder_attention_mask, ) # Get the token with the highest probability next_token = torch.argmax(logits[:, -1, :], dim=-1) if ( next_token.item() == tokenizer.sep_token_id ): # Stop if end-of-sequence token is generated break output_tokens.append(next_token.item()) # Update decoder inputs decoder_input_ids = torch.cat( [decoder_input_ids, next_token.unsqueeze(0)], dim=1 ) decoder_attention_mask = torch.ones_like(decoder_input_ids) # Decode the generated tokens into text prediction = tokenizer.decode(output_tokens, skip_special_tokens=True) predictions.append(prediction) # Compute ROUGE scores results = rou.get_scores(predictions, reference_summaries, avg=True) return results, predictions
修正训练时的decoder输入与标签逻辑
当前训练直接把target的完整input_ids作为decoder输入,这不符合seq2seq的训练逻辑——模型应该基于前序token预测下一个token。需改为移位标签:# 替换训练代码中decoder输入和标签的赋值 decoder_input_ids = targets['input_ids'][:, :-1] # 去掉最后一个token作为decoder输入 decoder_attention_mask = targets['attention_mask'][:, :-1] labels = targets['input_ids'][:, 1:] # 去掉第一个token作为预测标签损失计算保持原维度调整逻辑即可,这样模型训练时学习的是“根据前序token生成下一个token”的能力,避免退化到输出重复序列。
替换贪心解码为更鲁棒的生成策略
用torch.argmax做贪心解码容易陷入局部最优,导致重复输出高概率token。建议采用束搜索或Top-K/Top-P采样,也可以直接使用Hugging Face内置的generate方法:# 替换预测代码中的手动生成循环 outputs = model.generate( input_ids=inputs["input_ids"], attention_mask=inputs["attention_mask"], max_length=50, num_beams=5, # 束搜索避免重复 early_stopping=True, eos_token_id=tokenizer.sep_token_id, decoder_start_token_id=tokenizer.cls_token_id ) prediction = tokenizer.decode(outputs[0], skip_special_tokens=True)若坚持手动实现,可改用Top-K采样:
# 替换原next_token计算逻辑 logits = logits[:, -1, :] temperature = 0.7 # 温度系数调整概率分布平滑度 logits = logits / temperature top_k = 50 top_k_logits = torch.topk(logits, top_k)[0] indices_to_remove = logits < torch.min(top_k_logits) logits[indices_to_remove] = -float('Inf') probabilities = torch.softmax(logits, dim=-1) next_token = torch.multinomial(probabilities, 1)启用decoder权重初始化逻辑
当前注释掉了initialize_decoder_weights方法,导致decoder的cross-attention层未正确初始化,其他层也未加载预训练权重,训练稳定性不足。建议启用该方法,或者直接用Hugging Face的EncoderDecoderModel构建模型,它会自动处理权重对齐:from transformers import EncoderDecoderModel model = EncoderDecoderModel.from_encoder_decoder_pretrained( encoder_model_name, decoder_model_name ) model.config.decoder_start_token_id = tokenizer.cls_token_id model.config.eos_token_id = tokenizer.sep_token_id model.config.pad_token_id = tokenizer.pad_token_id确保decoder自注意力的因果掩码生效
当用BertModel作为decoder时,设置is_decoder=True后模型会自动添加因果掩码(屏蔽未来token),但需确保传入的decoder_attention_mask是(batch_size, seq_len)的形状,避免干扰因果掩码的生成。
内容的提问来源于stack exchange,提问作者rasoul mohammadi

