BART-LARGE-CNN转ONNX后推理last_hidden_state输出如何解析用于摘要任务
问题解决:ONNX格式BART-LARGE-CNN模型输出解析用于摘要生成
输出属性说明
你观察到的last_hidden_state实际是模型解码器输出后经过lm_head层映射到词汇表维度的logits张量,形状为[批量大小, 当前生成序列长度, 词汇表大小],可直接用于摘要生成的自回归推理逻辑。
生成解析步骤
BART属于自回归序列生成模型,需要多轮迭代推理得到完整摘要,核心流程如下:
- 预处理待摘要原文,得到编码器侧的输入id与注意力掩码
- 初始化解码器侧输入,起始为BART的开始符
<s>对应的id - 每轮推理得到当前步的logits后,在最后一维取概率最高的token id
- 将新生成的token id拼接到解码器输入序列末尾,作为下一轮推理的解码器输入
- 重复迭代直到生成结束符
</s>,或达到预设的最大摘要长度 - 最终将生成的完整id序列解码为自然语言文本即可得到摘要
完整代码示例(Python + ONNX Runtime)
import onnxruntime as ort import numpy as np from transformers import BartTokenizer # 加载对应模型的分词器 tokenizer = BartTokenizer.from_pretrained("facebook/bart-large-cnn") # 加载导出的ONNX模型会话 sess = ort.InferenceSession("你的ONNX模型文件路径", providers=["CPUExecutionProvider"]) # 原文预处理 raw_text = "需要生成摘要的原始长文本" encoder_inputs = tokenizer( raw_text, return_tensors="np", max_length=1024, truncation=True, padding="max_length" ) encoder_input_ids = encoder_inputs["input_ids"] encoder_attention_mask = encoder_inputs["attention_mask"] # 初始化解码器输入,BART起始符id为0 decoder_input_ids = np.array([[tokenizer.bos_token_id]], dtype=np.int64) max_gen_len = 150 for _ in range(max_gen_len): # 执行ONNX推理 outputs = sess.run( None, { "input_ids": encoder_input_ids, "attention_mask": encoder_attention_mask, "decoder_input_ids": decoder_input_ids } ) logits = outputs[0] # 取当前步最后一个位置的预测token next_token_id = logits.argmax(axis=-1)[:, -1:] # 拼接解码器输入 decoder_input_ids = np.concatenate([decoder_input_ids, next_token_id], axis=-1) # 遇到结束符终止生成 if next_token_id[0][0] == tokenizer.eos_token_id: break # 解码得到摘要,跳过特殊符号 summary = tokenizer.decode(decoder_input_ids[0], skip_special_tokens=True) print("生成摘要:", summary)
生成效果与性能优化建议
- 上述示例使用的是贪心搜索策略,如果要提升摘要质量,可替换为集束搜索、top-k采样、top-p核采样等生成策略
- 如果导出ONNX时开启了
use_past参数,可缓存每步生成的past_key_values传入下一轮推理,避免重复计算历史token的注意力权重,大幅降低推理耗时 - 若导出时拆分了单独的编码器、解码器ONNX文件,编码器仅需运行一次即可,不需要每轮生成都重复执行编码器推理
内容的提问来源于stack exchange,提问作者ZWang
相关产品推荐
相关产品推荐

