登录
推荐 文章 Go 技术 课程 下载 专题 AI
首页 >  科技周边 >  人工智能

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

来源:17golang原创

时间:2026-09-14 20:35:38 120浏览 收藏

如果你想让 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删除
相关阅读
更多>
最新阅读
更多>
课程推荐
更多>