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

微调Pegasus模型时,如何正确设置Decoder提示并仅计算摘要损失?

解决Pegasus微调时仅计算摘要部分损失的方案

要实现Decoder接收实体链提示、仅优化摘要生成部分的目标,核心是通过标签掩码屏蔽实体链区域的损失计算——Hugging Face模型默认会忽略labels中值为-100的位置的损失。以下是具体实现步骤:

1. 数据预处理:构建带掩码的标签

在预处理训练数据时,需要将实体链与摘要拼接成完整目标序列,再把实体链对应的标签位置设为-100,只保留摘要部分的真实标签。

示例代码:

from transformers import PegasusTokenizer
import torch

tokenizer = PegasusTokenizer.from_pretrained("google/pegasus-xsum")

def preprocess_data(examples):
    # 假设输入数据包含document(原文)、entity_chain(实体链提示)、summary(目标摘要)三个字段
    entity_chains = examples["entity_chain"]
    summaries = examples["summary"]
    documents = examples["document"]
    
    # 1. 构建完整目标序列:实体链 + 摘要
    full_targets = [f"{ec} {summary}" for ec, summary in zip(entity_chains, summaries)]
    
    # 2. 编码目标序列得到原始labels
    raw_labels = tokenizer(
        full_targets,
        padding="max_length",
        truncation=True,
        max_length=512,
        return_tensors="pt"
    ).input_ids
    
    # 3. 生成带掩码的labels:实体链部分设为-100
    masked_labels = []
    for ec, label in zip(entity_chains, raw_labels):
        # 计算实体链对应的token长度
        ec_token_len = len(tokenizer.encode(ec, add_special_tokens=False))
        # 前ec_token_len个位置设为-100,后续保留原标签
        masked_label = torch.where(
            torch.arange(len(label)) < ec_token_len,
            torch.tensor(-100, dtype=torch.long),
            label
        )
        # 处理padding长度不足的情况
        if len(masked_label) < 512:
            masked_label = torch.cat([masked_label, torch.tensor([-100]*(512-len(masked_label)), dtype=torch.long)])
        masked_labels.append(masked_label)
    
    # 4. 编码Encoder输入(原文)
    encoder_inputs = tokenizer(
        documents,
        padding="max_length",
        truncation=True,
        max_length=1024,
        return_tensors="pt"
    )
    
    # 5. 编码Decoder输入(实体链+摘要的完整序列,适配Pegasus输入要求)
    decoder_inputs = tokenizer(
        full_targets,
        padding="max_length",
        truncation=True,
        max_length=512,
        return_tensors="pt"
    )
    
    return {
        "input_ids": encoder_inputs["input_ids"],
        "attention_mask": encoder_inputs["attention_mask"],
        "decoder_input_ids": decoder_inputs["input_ids"],
        "labels": torch.stack(masked_labels)
    }

2. 训练阶段:自动忽略掩码区域损失

不管用Trainer还是原生PyTorch循环,模型都会自动跳过labels中-100的位置计算损失:

用Trainer训练

from transformers import PegasusForConditionalGeneration, Trainer, TrainingArguments

model = PegasusForConditionalGeneration.from_pretrained("google/pegasus-xsum")

training_args = TrainingArguments(
    output_dir="./pegasus_entity_chain_finetune",
    per_device_train_batch_size=4,
    learning_rate=5e-5,
    num_train_epochs=3,
    logging_dir="./logs",
)

# 假设train_dataset是预处理后的数据集
trainer = Trainer(
    model=model,
    args=training_args,
    train_dataset=train_dataset,
)

trainer.train()

原生PyTorch循环训练

import torch
from torch.utils.data import DataLoader

device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
model = model.to(device)
optimizer = torch.optim.AdamW(model.parameters(), lr=5e-5)

train_dataloader = DataLoader(train_dataset, batch_size=4, shuffle=True)

for epoch in range(3):
    model.train()
    total_loss = 0
    for batch in train_dataloader:
        optimizer.zero_grad()
        # 把batch数据移到设备上
        batch = {k: v.to(device) for k, v in batch.items()}
        # 前向传播
        outputs = model(**batch)
        loss = outputs.loss
        total_loss += loss.item()
        # 反向传播+优化
        loss.backward()
        optimizer.step()
    print(f"Epoch {epoch+1}, Average Loss: {total_loss/len(train_dataloader)}")

3. 推理阶段:基于实体链生成摘要

推理时直接把实体链作为Decoder的初始输入,模型会自动生成后续摘要:

model.eval()
entity_chain = "<s> [ENTITYCHAIN] Frozen | Disney [SUMMARY]"
# 编码实体链作为Decoder输入
decoder_input_ids = tokenizer.encode(entity_chain, return_tensors="pt").to(device)
# 编码原文作为Encoder输入
input_ids = tokenizer.encode("原文内容...", return_tensors="pt").to(device)

# 生成摘要
generated_ids = model.generate(
    input_ids=input_ids,
    decoder_input_ids=decoder_input_ids,
    max_length=150,
    num_beams=4,
    early_stopping=True
)
generated_summary = tokenizer.decode(generated_ids[0], skip_special_tokens=True)
print(generated_summary)

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.21 05:37:21