如何使用停用词列表实现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
相关产品推荐
相关产品推荐

