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

如何使用停用词列表实现transformers自回归模型的生成提前终止?

解决方案

你的代码存在3处核心问题,修正后即可实现生成到}立即停止的功能:

  • 停止判断逻辑错误:input_ids是形状为(batch_size, seq_len)的张量,存储全批次已生成的完整token序列,原判断input_ids in self.keywords不成立,需要逐序列检查末尾是否匹配任意停止词的id序列
  • 参数传递错误:generate方法接收的停止条件参数名为小写开头的stopping_criteria,且必须传入StoppingCriteriaList包装后的实例,不能直接传自定义停止条件对象
  • 停止词id处理问题:tokenizer.encode返回id列表,不同停止词长度可能不同,需要做后缀匹配,不能直接做包含判断

可运行完整代码

import torch
from transformers import StoppingCriteria, StoppingCriteriaList, AutoModelForCausalLM, AutoTokenizer

# 可替换为你使用的GPT-Neo模型名
model_name = 'gpt2'
dtype = torch.float16
tokenizer = AutoTokenizer.from_pretrained(model_name)
model = AutoModelForCausalLM.from_pretrained(
    model_name, 
    pad_token_id=tokenizer.eos_token_id, 
    torch_dtype=dtype
).eval()

class KeywordsStoppingCriteria(StoppingCriteria):
    def __init__(self, keywords_ids: list):
        # 停止词按长度倒序,优先匹配更长的停止词避免误判
        self.keywords = sorted(keywords_ids, key=lambda x: len(x), reverse=True)

    def __call__(self, input_ids: torch.LongTensor, scores: torch.FloatTensor, **kwargs) -> bool:
        for seq in input_ids:
            seq_len = len(seq)
            for keyword in self.keywords:
                kw_len = len(keyword)
                if seq_len < kw_len:
                    continue
                # 匹配序列末尾和停止词id
                if torch.all(seq[-kw_len:] == torch.tensor(keyword, device=seq.device)):
                    return True
        return False

stop_words = ['}', ' }', '\n']
# encode时关闭特殊字符添加避免额外id干扰
stop_ids = [tokenizer.encode(w, add_special_tokens=False) for w in stop_words]
stop_ids.append([tokenizer.eos_token_id])
stop_criteria = KeywordsStoppingCriteria(stop_ids)

# 先对输入做tokenize,适配新版本transformers规范
inputs = tokenizer('some text:{', return_tensors='pt')
output = model.generate(
    **inputs,
    max_new_tokens=100, # 按需调整最大生成长度
    stopping_criteria=StoppingCriteriaList([stop_criteria])
)
# 解码输出跳过特殊token
print(tokenizer.decode(output[0], skip_special_tokens=True))

如果需要适配多批次生成场景,仅需调整__call__方法的判断逻辑,改为所有批次序列都匹配到停止词时再返回True即可。

内容的提问来源于stack exchange,提问作者Lei Zhang

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.10.01 18:15:03