如何反向还原tokenizer.apply_chat_template()生成的prompt字符串至原始对话数组,以及了解pipeline的底层解析逻辑?
其实吧,这个问题得拆成两部分来看——先解决怎么把生成的prompt字符串转回到原始对话数组,再聊聊pipeline到底是怎么在底层处理解析的。
一、反向还原prompt字符串到原始对话数组
首先得明确:没有通用的一键还原方法,因为不同模型的对话模板格式千差万别(比如Llama 2、GPT、Qwen的模板规则都不一样)。但核心思路很简单:先搞清楚你用的tokenizer的chat_template是什么样的,再跟着模板规则写对应的解析逻辑就行。
举个实际例子,比如你用的是Llama 2的tokenizer,它的默认聊天模板大概是这样的:
{% set loop_messages = messages %} {% for message in loop_messages %} {% if message['role'] == 'user' %} {{ bos_token + '[INST] ' + message['content'].strip() + ' [/INST] ' }} {% elif message['role'] == 'assistant' %} {{ message['content'].strip() + ' ' + eos_token }} {% endif %} {% endfor %}
如果用这个模板生成的prompt字符串是类似<s>[INST] Random prompt. [/INST]这样的,我们就可以用正则表达式或者字符串拆分来解析:
import re # 假设这是tokenizer生成后的prompt字符串 generated_prompt = "<s>[INST] Random prompt. [/INST]" # 针对Llama 2模板的解析逻辑 def parse_llama_prompt(prompt_str): # 先去掉首尾的特殊token cleaned = prompt_str.replace("<s>", "").replace("</s>", "").strip() # 用模板里的[INST]和[/INST]作为分隔符提取用户内容 user_contents = re.findall(r'\[INST\] (.*?) \[/INST\]', cleaned) # 组装回原始对话数组格式 conversation = [] for content in user_contents: conversation.append({"role": "user", "content": content.strip()}) # 如果是多轮对话,还得识别助手回复的部分,比如模板里[/INST]后面的内容 # 这里以单轮为例,多轮的话可以继续拆分eos token前后的内容 return conversation # 测试解析效果 original_convo = parse_llama_prompt(generated_prompt) print(original_convo) # 输出: [{'role': 'user', 'content': 'Random prompt.'}]
关键步骤就是先查看你的tokenizer的聊天模板——直接打印print(tokenizer.chat_template)就能看到具体规则,然后针对性地写拆分逻辑就行。比如有的模型用USER:、ASSISTANT:作为角色前缀,那解析时就按这些前缀来拆分每一轮对话。
二、聊聊pipeline的底层解析逻辑
你说不用pipeline的话,模型返回的就是一串字符串,得自己手动解析。其实pipeline本质上就是帮你把这些“重复又繁琐的脏活”给封装了,它的底层逻辑核心还是依赖tokenizer的聊天模板和特殊token:
- 对话状态跟踪:pipeline内部会一直维护着你最开始用的那种对话数组格式,每次你输入新的用户消息,它会先用
apply_chat_template把对话数组转成模型能理解的prompt字符串,再传给模型生成内容。 - 生成结果的拆分:模型输出字符串后,pipeline会对照聊天模板里的角色分隔符(比如Llama 2的
[/INST]、GPT的ASSISTANT:),精准提取出属于助手的回复部分,还会自动去掉多余的特殊token(比如结束符eos token)。 - 特殊token辅助判断:它会用tokenizer的
eos_token来判断模型是否完成了当前轮次的回复,一旦检测到这个token,就立刻停止生成,只保留token之前的有效内容。
我给你模拟个简化版的pipeline内部逻辑,你一看就懂:
def mock_pipeline_logic(conversation, tokenizer, model): # 第一步:把对话数组转成prompt字符串 prompt = tokenizer.apply_chat_template(conversation, tokenize=False, add_generation_prompt=True) # 第二步:模型生成内容 inputs = tokenizer(prompt, return_tensors="pt").to(model.device) outputs = model.generate(**inputs, max_new_tokens=100) # 第三步:解析生成的结果 generated_str = tokenizer.decode(outputs[0], skip_special_tokens=False) # 按Llama 2模板拆分助手回复:找[/INST]之后到eos token之前的内容 start_pos = generated_str.find("[/INST]") + len("[/INST]") end_pos = generated_str.find(tokenizer.eos_token, start_pos) assistant_reply = generated_str[start_pos:end_pos].strip() # 更新对话数组 conversation.append({"role": "assistant", "content": assistant_reply}) return conversation, assistant_reply
其实你自己手动解析的时候,完全可以照着这个思路来——先定位模板里的角色分隔符,再提取对应区间的内容就行,和pipeline干的活没差多少。
备注:内容来源于stack exchange,提问作者A.A

