当前位置:首页 > 文章列表 > 科技周边 > 人工智能 > Transformers generate 如何用 stopping criteria 停在自定义标记

Transformers generate 如何用 stopping criteria 停在自定义标记

来源:17golang原创 2026-09-14 20:35:38 0浏览 收藏

如果你想让 Hugging Face Transformers 的 generate() 在模型生成自定义标记(例如 )时停下,关键不是比较解码后的字符串,而是用同一个 tokenizer 把标记转成 token 序列,再在 StoppingCriteria 中比较每行 input_ids 的尾部。标记可能被拆成多个 token,逐字符判断很容易漏掉边界。

官方地址:https://huggingface.co/docs/transformers/main_classes/text_generation

要点速览
  • 自定义规则返回每个 batch 行一个布尔值,True 表示该行可以停止。
  • 停止标记必须使用当前模型的 tokenizer 编码,不能手写固定 token id。
  • max_new_tokens 仍要保留,防止模型永远生成不到标记。

把自定义标记变成可比较的 token 序列

Transformers 的停止规则接收的是 token 化后的 input_ids,不是尚未解码的字符串。以 为例,模型词表可能把它编码成一个 token,也可能编码成多个 token,所以初始化规则时要保存完整序列。生成每增加一个 token,就比较当前序列最后几位是否与它一致。

Transformers 自定义 END 标记经过 tokenizer 变成 stop_ids 并与 input_ids 尾部比较的结构示意图
图1:自定义标记经过 tokenizer 后,与 input_ids 尾部形成静态比较关系的操作示意图。

这也解释了为什么不要在循环里每次 decode 全量文本再查字符串:全量解码更慢,而且空格、特殊 token 和 token 边界可能让字符串判断与模型实际输出不同。

用 StoppingCriteria 判断最后生成的 token

下面的规则只依赖 input_ids,因此不需要打开分数输出。返回值按 batch 行计算,批量生成时某一行先遇到标记,不会强迫其他行立即结束。

import torch
from transformers import StoppingCriteria

class StopOnMarker(StoppingCriteria):
    def __init__(self, tokenizer, marker=""):
        # 用目标模型自己的 tokenizer,保留标记可能被拆分出的全部 token。
        encoded = tokenizer(marker, add_special_tokens=False, return_tensors="pt")
        self.stop_ids = encoded.input_ids[0]

    def __call__(self, input_ids, scores, **kwargs):
        # 生成长度还不够时,每个 batch 行都继续生成。
        size = self.stop_ids.numel()
        if input_ids.shape[1] 

这里的 scores 参数仍然要保留,因为它属于停止规则的统一调用签名;本文没有读取它。如果你的规则要按概率、置信度或 logits 停止,官方文档要求在 generate() 中同时设置 return_dict_in_generate=Trueoutput_scores=True

把规则接入 generate,并保留长度兜底

把实例放入 StoppingCriteriaList 后传给 generate()。同时设置 max_new_tokens,因为停止标记可能没有被模型生成,或者模型的输出格式发生变化。

generate 接收 StoppingCriteriaList、StopOnMarker 并与 EOS 和 max_new_tokens 组成停止边界的结构示意图
图2:generate 接收自定义停止规则并与 EOS、max_new_tokens 共同构成停止边界的结果示意图。
from transformers import StoppingCriteriaList

marker_rule = StopOnMarker(tokenizer, marker="")
criteria = StoppingCriteriaList([marker_rule])

outputs = model.generate(
    **inputs,
    stopping_criteria=criteria,
    max_new_tokens=128,  # 标记缺失时仍限制本次生成的最大 token 数。
    do_sample=False,
)

text = tokenizer.batch_decode(outputs, skip_special_tokens=True)[0]
text = text.split("", 1)[0].rstrip()  # 展示层移除控制标记。
print(text)

在 decoder-only 模型中,input_ids 通常包含提示词和新生成内容。上面的规则只看序列尾部,因此提示词中间出现 不会触发;但不要让提示词本身以这个标记结尾,否则第一次检查就可能命中。对不同长度的 batch 提示词,还应让 tokenizer 正确生成 attention mask,并在业务层验证每行最终是否真的包含控制标记。

多条停止规则的边界检查

现象原因处理方式
生成到标记仍不停标记编码与当前 tokenizer 不一致在同一 tokenizer 上重新编码,不要复制别的模型 token id
刚开始就停止提示词最后已经是完整标记修改提示词结尾,或为规则保存生成起点后只检查新增长度
输出带着 停止发生在标记 token 已写入序列之后解码后在展示层裁剪,不要改写停止判定
模型一直生成模型没有输出标记或规则不适合当前批处理保留 max_new_tokens,并记录每行是否命中标记

如果还要同时支持 EOS、时间限制或其他业务条件,可以把多个规则放入同一个 StoppingCriteriaList。每条规则都应返回与 batch 行对应的布尔张量;不要返回单个 Python 布尔值,否则批量推理时无法表达“只停止其中一行”。

常见问题

自定义标记必须注册成特殊 token 吗?

不必须。只要它能被当前 tokenizer 编码,尾部 token 序列就可以匹配;如果业务还要求跳过特殊 token 解码,则再按模型的 special tokens 配置决定是否注册。

为什么不用 generate 的 stop_strings 参数?

如果当前 Transformers 版本和模型配置支持 stop_strings,它更省代码;需要按 batch 行组合多个条件、读取额外状态或兼容旧版本时,自定义 StoppingCriteria 更可控。

规则依赖 scores 时要改什么?

保留同样的调用签名,并在 generate() 中开启 return_dict_in_generate=Trueoutput_scores=True,否则规则可能拿不到需要的分数。

版本声明
本文转载于:17golang原创 如有侵犯,请联系study_golang@163.com删除
Go base64.NewEncoder 关闭前不调用 Close 会少多少数据Go base64.NewEncoder 关闭前不调用 Close 会少多少数据
上一篇
Go base64.NewEncoder 关闭前不调用 Close 会少多少数据
Go json.RawMessage 延迟解析时如何避免共享底层字节
下一篇
Go json.RawMessage 延迟解析时如何避免共享底层字节
查看更多
最新文章
查看更多
课程推荐
  • 前端进阶之JavaScript设计模式
    前端进阶之JavaScript设计模式
    设计模式是开发人员在软件开发过程中面临一般问题时的解决方案,代表了最佳的实践。本课程的主打内容包括JS常见设计模式以及具体应用场景,打造一站式知识长龙服务,适合有JS基础的同学学习。
    543次学习
  • GO语言核心编程课程
    GO语言核心编程课程
    本课程采用真实案例,全面具体可落地,从理论到实践,一步一步将GO核心编程技术、编程思想、底层实现融会贯通,使学习者贴近时代脉搏,做IT互联网时代的弄潮儿。
    516次学习
  • 简单聊聊mysql8与网络通信
    简单聊聊mysql8与网络通信
    如有问题加微信:Le-studyg;在课程中,我们将首先介绍MySQL8的新特性,包括性能优化、安全增强、新数据类型等,帮助学生快速熟悉MySQL8的最新功能。接着,我们将深入解析MySQL的网络通信机制,包括协议、连接管理、数据传输等,让
    500次学习
  • JavaScript正则表达式基础与实战
    JavaScript正则表达式基础与实战
    在任何一门编程语言中,正则表达式,都是一项重要的知识,它提供了高效的字符串匹配与捕获机制,可以极大的简化程序设计。
    487次学习
  • 从零制作响应式网站—Grid布局
    从零制作响应式网站—Grid布局
    本系列教程将展示从零制作一个假想的网络科技公司官网,分为导航,轮播,关于我们,成功案例,服务流程,团队介绍,数据部分,公司动态,底部信息等内容区块。网站整体采用CSSGrid布局,支持响应式,有流畅过渡和展现动画。
    485次学习
查看更多
AI推荐
  • PubMedQA数据集详解:生物医学问答基准、功能与应用指南
    PubMedQA
    深入了解PubMedQA生物医学问答数据集,涵盖其核心功能、使用方法及在临床决策、药物研发等场景的应用,助力提升NLP模型性能。
    26次使用
  • H2O EvalGPT:开源LLM大模型评估与排行榜工具
    H2O EvalGPT
    H2O EvalGPT是H2O.ai推出的开源LLM评估平台,提供详细的大模型性能排行榜、行业特定基准测试及A/B测试功能,助您快速选择最适合项目的高性能大语言模型。
    130次使用
  • LMArena是什么?伯克利AI模型评估平台使用指南与功能解析
    LMArena
    LMArena是加州大学伯克利分校推出的AI模型匿名评测平台。通过盲测投票机制,用户可对比不同大模型回答并生成实时排行榜,助力开发者优化模型及用户选择最佳AI工具。
    57次使用
  • 斯坦福HELM:大语言模型Holistic Evaluation整体评估框架详解
    HELM
    深入了解斯坦福推出的HELM(Holistic Evaluation of Language Models)大模型评测体系。本文解析其核心功能、安装配置步骤及应用场景,涵盖准确性、公平性、鲁棒性等多维度指标,助力开发者全面优化语言模型性能。
    22次使用
  • OpenCompass大模型评测体系详解:功能、使用指南与应用场景
    OpenCompass
    OpenCompass是上海AI实验室推出的开源大模型评测平台,提供CompassKit、CompassHub和CompassRank三大核心组件,支持LLM及多模态模型的一站式标准化评估与排行榜查询。
    80次使用
微信登录更方便
  • 密码登录
  • 注册账号
登录即同意 用户协议隐私政策
返回登录
  • 重置密码