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

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.15 21:55:54