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

HuggingFace摘要生成如何逐词获取logits实现自定义选词?

DistilBART自定义逐词生成相关问题解答

关于model(**input).logits[:, -1, :]的返回值含义

这个切片拿到的logits,对应当前传入序列最后一个位置的下一个待生成token的预测分布,既不是固定对应第一个待生成token,也不是序列里已有最后一个token本身的输出:

  • 如果你传入的解码器序列只有模型规定的解码器起始符,此时拿到的就是摘要第一个待生成token的候选logits
  • 如果你已经生成了N个token,解码器输入是「起始符 + 已生成的N个token」,此时拿到的就是第N+1个待生成token的候选logits

逐步骤自定义选词的实现方法

封装好的generate()方法做不了逐步骤自定义选词,你可以手动实现自回归生成循环,逻辑很简单:

  • 第一步先单独编码原始文本,把编码器输出缓存住,不用每步重复编码浪费算力
  • 初始化解码器输入为解码器起始token
  • 循环执行:拿当前解码器输入+固定的编码器输出跑模型前向,取最后位置的logits,按你自己的规则选下一个token,把选中的token拼到解码器输入尾部,碰到EOS token或者达到预设最大长度就退出循环

参考实现代码如下:

import torch
from transformers import AutoTokenizer, AutoModelForSeq2SeqLM

# 加载模型和分词器
model_name = "sshleifer/distilbart-cnn-12-6"
tokenizer = AutoTokenizer.from_pretrained(model_name)
model = AutoModelForSeq2SeqLM.from_pretrained(model_name)

# 处理待摘要的原始文本,提前算好编码器输出
raw_text = "替换成你要做摘要的原始文本内容"
inputs = tokenizer(raw_text, return_tensors="pt", max_length=1024, truncation=True)
encoder_outputs = model.get_encoder()(**inputs)

# 初始化生成参数
decoder_input_ids = torch.tensor([[model.config.decoder_start_token_id]])
eos_token_id = model.config.eos_token_id
max_generate_length = 128  # 设置最大长度上限,避免循环跑飞
generated_token_ids = []

# 逐词生成循环
for _ in range(max_generate_length):
    model_outputs = model(
        encoder_outputs=encoder_outputs,
        decoder_input_ids=decoder_input_ids
    )
    # 取下一个token的logits分布
    next_token_logits = model_outputs.logits[:, -1, :]

    # 这里替换成你自己的选词逻辑:可以过滤低概率候选、加自定义权重、按业务规则排序选词都可以
    select_token_id = your_custom_token_selection(next_token_logits)

    generated_token_ids.append(select_token_id.item())
    # 碰到结束符直接终止生成
    if select_token_id.item() == eos_token_id:
        break
    # 把选中的token拼到解码器输入,供下一轮生成使用
    decoder_input_ids = torch.cat(
        [decoder_input_ids, select_token_id.unsqueeze(0)],
        dim=-1
    )

# 解码得到最终摘要
final_summary = tokenizer.decode(generated_token_ids, skip_special_tokens=True)

优化提示:如果对生成速度有要求,可以在每轮前向时传入past_key_values参数缓存解码器已经计算过的注意力键值对,不用每轮重新计算整个解码器序列的前向传播,生成速度能提升数倍。

内容的提问来源于stack exchange,提问作者Diego Miguel

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.09.03 08:27:23