Llama-2 13B停止判定规则失效:原因排查与修复咨询
问题原因及修复方案
核心原因
- 多token停止词未被正确识别:Llama系列模型的分词器会将部分停止词(比如
###)拆分为多个子token,而你当前使用的StoppingCriteriaSub大概率只检查单个生成token是否匹配停止词,无法识别由多个token组成的完整停止词序列,导致模型生成完整个停止词后仍继续生成。 - 停止词匹配精度问题:微调后的模型生成内容可能存在格式偏差(比如
###前后带有空格),和你设定的停止词不完全一致,导致判定失效。
修复步骤
1. 确认停止词的分词结果
先打印每个停止词对应的token序列,确认是否被拆分为多token:
for sw in stop_words: ids = tokenizer(sw, return_tensors='pt')['input_ids'].squeeze() print(f"停止词 '{sw}' -> token ID: {ids}, 对应分词: {[tokenizer.decode(id) for id in ids]}")
比如###可能被分词为['▁##', '#'](不同版本分词器结果可能有差异),这就需要停止判定逻辑检查连续的多个token是否匹配。
2. 替换为支持多token的停止判定类
自定义一个StoppingCriteria子类,专门检查生成序列的末尾是否完整匹配停止词的token序列:
from transformers import StoppingCriteria, StoppingCriteriaList import torch from typing import List class MultiTokenStoppingCriteria(StoppingCriteria): def __init__(self, stops: List[torch.Tensor], device: str): self.stops = [stop.to(device) for stop in stops] def __call__(self, input_ids: torch.LongTensor, scores: torch.FloatTensor, **kwargs) -> bool: # 遍历所有停止词,检查当前生成序列末尾是否匹配 for stop_seq in self.stops: seq_len = len(stop_seq) if len(input_ids[0]) >= seq_len: if torch.equal(input_ids[0][-seq_len:], stop_seq): return True return False
3. 重新初始化停止判定规则
用新的类替换原来的StoppingCriteriaSub:
stop_words = ["Human:", "Chatbot:", "###"] # 获取每个停止词的token ID序列 stop_words_ids = [tokenizer(stop_word, return_tensors='pt')['input_ids'].squeeze() for stop_word in stop_words] # 初始化多token停止判定 stopping_criteria = StoppingCriteriaList([ MultiTokenStoppingCriteria(stops=stop_words_ids, device="cuda") ]) # 后续生成逻辑不变 generation_config = GenerationConfig(..., stopping_criteria=stopping_criteria)
4. 验证停止词与生成内容的一致性
打印生成的原始结果res,检查模型实际生成的停止词(比如###)前后是否有空格或其他字符。如果存在,需要将停止词调整为完全匹配的格式,比如模型生成的是 ###,就把停止词改成 ###。
额外注意
如果使用Axolotl微调时设置了特定的对话格式,确保停止词和微调数据集中的分隔符完全一致,避免因格式差异导致停止判定失效。
内容的提问来源于stack exchange,提问作者BlackHawk
相关产品推荐
相关产品推荐

