如何解决CodeLlama生成内容冗余、重复输入的问题?
解决CodeLlama-7b-Instruct-hf生成SQL时重复冗余内容的方案
核心问题分析
导致模型重复输入内容的主要原因包括:
- Tokenizer与Instruct模型不匹配,无法正确识别指令边界
- Prompt格式不符合CodeLlama-Instruct的官方对话规范
- 生成参数配置不当,未有效触发停止机制
有效解决方案
1. 使用匹配的Tokenizer
模型为codellama/CodeLlama-7b-Instruct-hf,需使用对应版本的Tokenizer,而非基础版的codellama/CodeLlama-7b-hf,确保特殊标记(如对话分隔符)被正确识别。
tokenizer = AutoTokenizer.from_pretrained(base_model) # base_model为"codellama/CodeLlama-7b-Instruct-hf" # 补充设置pad_token(CodeLlama默认无pad_token) tokenizer.pad_token = tokenizer.eos_token
2. 遵循官方对话Prompt格式
CodeLlama-Instruct要求用<s>[INST] ... [/INST]包裹指令,明确区分输入与输出边界,避免模型混淆上下文。修改后的Prompt如下:
eval_prompt = """<s>[INST] You are a powerful text-to-SQL model. Your job is to answer questions about a database. You are given a question and context regarding one or more tables. You must output the SQL query that answers the question. Only return the SQL query. ### Input: Which Class has a Frequency MHz larger than 91.5, and a City of license of hyannis, nebraska? ### Context: CREATE TABLE table_name_12 (class VARCHAR, frequency_mhz VARCHAR, city_of_license VARCHAR) [/INST]"""
3. 优化生成参数配置
通过确定性生成参数+明确停止机制,避免冗余内容:
model.eval() with torch.no_grad(): outputs = model.generate( **model_input, max_new_tokens=100, eos_token_id=tokenizer.eos_token_id, pad_token_id=tokenizer.pad_token_id, do_sample=False, # 关闭采样,确保确定性输出 temperature=0.0, # 温度设为0,消除随机性 stop_strings=["</s>"] # 指定特殊结束标记作为停止条件 ) # 提取纯响应内容,过滤输入Prompt部分 response = tokenizer.decode(outputs[0], skip_special_tokens=True).split("[/INST]")[-1].strip() print(response)
4. 自定义StoppingCriteria(可选)
若内置停止机制无效,可实现自定义停止逻辑,检测到特定Token时终止生成:
from transformers import StoppingCriteria, StoppingCriteriaList class SQLStopCriteria(StoppingCriteria): def __init__(self, stop_ids): self.stop_ids = stop_ids def __call__(self, input_ids, scores, **kwargs): # 检查最后生成的Token是否属于停止标记集合 return input_ids[0][-1] in self.stop_ids # 配置停止Token集合 stop_ids = [tokenizer.eos_token_id] stopping_criteria = StoppingCriteriaList([SQLStopCriteria(stop_ids)]) # 在generate中使用自定义停止规则 outputs = model.generate( **model_input, max_new_tokens=100, stopping_criteria=stopping_criteria, do_sample=False, temperature=0.0 )
内容的提问来源于stack exchange,提问作者MJava
相关产品推荐
相关产品推荐

