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
相关产品推荐
相关产品推荐

