LangChain中HuggingFace模型的enforce_stop_tokens机制及用法疑问
问题背景
在LangChain调用HuggingFaceHub模型时,会遇到生成文本无法按预期停止的问题。查看LangChain源码可知,HuggingFacePipeline类的_call方法通过enforce_stop_tokens函数处理停止词:
class HuggingFacePipeline(LLM): ... def _call( ... if stop is not None: # This is a bit hacky, but I can't figure out a better way to enforce # stop tokens when making calls to huggingface_hub. text = enforce_stop_tokens(text, stop) return text
而enforce_stop_tokens的核心逻辑是用正则分割文本后取第一段(来自langchain/llms/utils.py):
re.split("|".join(stop), text)[0]
以GPT2模型测试为例,生成的文本如下:
from transformers import pipeline from transformers import GPT2LMHeadModel, AutoTokenizer tokenizer = AutoTokenizer.from_pretrained('gpt2') model = GPT2LMHeadModel.from_pretrained('gpt2') generator = pipeline('text-generation', model=model, tokenizer=tokenizer) output = generator("Hey Pizza! ") # 输出结果 [{'generated_text': 'Hey Pizza! 」\n\n「Hurry up, leave the place! 」\n\n「Oi! 」\n\nWhile eating pizza and then, Yuigahama came in contact with Ruriko in the middle of the'}]
使用普通单词(如["up", "then"])作为停止词时,会在文本中间随机截断,无法在合理的生成结束位置停止,不符合需求。
解决方案:使用模型专属标记或明确分隔符
要让enforce_stop_tokens实现有效截断,必须使用模型训练时的结束标记或语义明确的专属分隔符,具体方式如下:
1. 用模型自带的特殊令牌
不同模型有自身定义的结束令牌,比如GPT2的<|endoftext|>,这是模型训练时的默认结束标记。直接将其作为停止词传入即可:
from langchain.llms import HuggingFacePipeline from transformers import pipeline, GPT2LMHeadModel, AutoTokenizer tokenizer = AutoTokenizer.from_pretrained('gpt2') tokenizer.pad_token = tokenizer.eos_token # 补全pad token配置 model = GPT2LMHeadModel.from_pretrained('gpt2') generator = pipeline( 'text-generation', model=model, tokenizer=tokenizer, eos_token_id=tokenizer.eos_token_id, max_new_tokens=100 ) llm = HuggingFacePipeline(pipeline=generator) result = llm.predict("Hey Pizza! ", stop=["<|endoftext|>"]) print(result)
如果模型生成时不主动输出该标记,需在生成配置中指定eos_token_id,确保模型触发结束逻辑。
2. 自定义语义明确的分隔符
如果模型没有合适的自带标记,可以在prompt模板末尾添加自定义分隔符(如### END ###),并将其作为停止词:
prompt = "Hey Pizza! ### END ###" stop_tokens = ["### END ###"] result = llm.predict(prompt, stop=stop_tokens) # 截断后会保留分隔符之前的有效内容
3. 利用文本中已有的结束标记
观察模型生成的文本,选取其中语义明确的结束符号作为停止词。比如GPT2测试输出中的」是对话结束标记,可直接使用:
stop_tokens = ["」"] result = llm.predict("Hey Pizza! ", stop=stop_tokens)
这样就能在对话段落结束的位置精准截断。
关键注意点
enforce_stop_tokens是事后截断逻辑:先让模型生成完整文本,再用正则匹配停止词分割取第一段。因此停止词必须是模型生成文本中明确出现、且语义上代表生成结束的字符串,普通单词因频繁出现在文本中,会导致截断位置混乱。
内容的提问来源于stack exchange,提问作者alvas

