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

如何在Transformers的BART模型中输出注意力张量?

解决BART模型生成时无法获取注意力张量的问题

问题原因

model.generate() 默认仅返回生成的token id序列,即便初始化模型时设置了 output_attentions=True,也不会自动返回生成过程中的注意力张量。必须在调用generate()时显式指定相关参数。

修改后的代码

from transformers import AutoTokenizer, BartForConditionalGeneration

model = BartForConditionalGeneration.from_pretrained("facebook/bart-large-cnn", output_attentions=True)
tokenizer = AutoTokenizer.from_pretrained("facebook/bart-large-cnn")

ARTICLE_TO_SUMMARIZE = (
    "PG&E stated it scheduled the blackouts in response to forecasts for high winds "
    "amid dry conditions. The aim is to reduce the risk of wildfires. Nearly 800 thousand customers were "
    "scheduled to be affected by the shutoffs which were expected to last through at least midday tomorrow."
)
inputs = tokenizer([ARTICLE_TO_SUMMARIZE], max_length=1024, return_tensors="pt")

# 生成时指定返回完整结果字典和注意力参数
generation_output = model.generate(
    inputs["input_ids"],
    num_beams=2,
    min_length=0,
    max_length=20,
    return_dict_in_generate=True,  # 关键:返回包含所有生成信息的字典
    output_attentions=True         # 关键:触发注意力张量的计算与返回
)

# 提取生成结果和注意力张量
summary_ids = generation_output.sequences
encoder_attentions = generation_output.encoder_attentions
decoder_attentions = generation_output.decoder_attentions
cross_attentions = generation_output.cross_attentions

# 解码并打印摘要
summary = tokenizer.batch_decode(summary_ids, skip_special_tokens=True, clean_up_tokenization_spaces=False)[0]
print("生成的摘要:", summary)
print("编码器注意力层数:", len(encoder_attentions))
print("解码器注意力层数:", len(decoder_attentions))
print("交叉注意力层数:", len(cross_attentions))

参数说明

  • return_dict_in_generate=True:让generate返回完整的生成结果字典,而非仅token序列。
  • output_attentions=True:在生成流程中启用注意力张量的计算和返回。
  • 返回的注意力张量分为三类:
    • encoder_attentions:编码器各层的自注意力张量
    • decoder_attentions:解码器各层的自注意力张量
    • cross_attentions:解码器各层对编码器输出的交叉注意力张量

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.02 09:42:13