Llama模型结合LangChain无法生成SQL查询,求类ChatGPT自然语言输出方案
解决Llama模型SQL问答卡顿及输出格式对齐问题
一、解决SQL查询生成卡顿问题
针对orca-mini-3b模型在SQL生成环节的卡顿,可从以下维度优化:
- 适配硬件调整参数:当前
n_ctx=20000会占用过量内存,远超3B模型的合理上下文范围,建议下调至4096;同时添加n_batch=512参数,提升批量token生成效率,减少硬件等待时间。 - 简化生成策略:关闭
use_query_checker=True,该功能会额外触发一次模型校验,大幅增加耗时;将temperature微调至0.1,在保证回答确定性的同时,降低模型的搜索计算量。 - 关闭流式输出:移除
callback_manager参数,避免逐token打印导致的“感知卡顿”,让模型一次性返回完整结果。
二、对齐ChatGPT风格的自然语言输出
要让Llama生成与ChatGPT格式一致的输出,核心通过提示词约束和结果后处理实现:
- 自定义SQL问答提示词模板:替换SQLDatabaseChain的默认prompt,明确要求模型输出纯自然语言回答,示例如下:
初始化链时传入该prompt:from langchain.prompts import PromptTemplate PROMPT = PromptTemplate( input_variables=["table_info", "question"], template="""给定数据库表结构信息: {table_info} 请用简洁清晰的自然语言回答用户问题:{question} 要求:只输出回答内容,不要包含SQL查询语句或其他无关信息。 """ )db_chain = SQLDatabaseChain.from_llm(llm, db, verbose=True, prompt=PROMPT, use_query_checker=False) - 结果后处理清洗:对模型返回的结果进行字符串过滤,移除可能存在的SQL代码块、调试标记等冗余内容,示例:
def clean_output(raw_output): # 移除SQL代码块 cleaned = raw_output.replace("```sql", "").replace("```", "").strip() # 移除多余的SQL说明文本 if "SQL查询语句:" in cleaned: cleaned = cleaned.split("SQL查询语句:")[0].strip() return cleaned - 使用输出解析器强制格式:借助LangChain的
OutputFixingParser或自定义解析器,定义输出结构,确保模型严格按照要求生成内容。
修改后的完整代码示例
# pip install llama-cpp-python==0.1.78 from langchain.llms import LlamaCpp from langchain.utilities import SQLDatabase from langchain.chains import SQLDatabaseChain from langchain.prompts import PromptTemplate model_path = 'model/orca-mini-3b.ggmlv3.q4_0.bin' def load_model() -> LlamaCpp: llama_model: LlamaCpp = LlamaCpp( model_path=model_path, temperature=0.1, max_tokens=2000, top_p=1, verbose=False, n_ctx=4096, n_batch=512 ) return llama_model # 自定义提示词模板 PROMPT = PromptTemplate( input_variables=["table_info", "question"], template="""给定数据库表结构信息: {table_info} 请用简洁清晰的自然语言回答用户问题:{question} 要求:只输出回答内容,不要包含SQL查询语句或其他无关信息。 """ ) llm = load_model() db = SQLDatabase.from_uri("sqlite:///db.sqlite3") db_chain = SQLDatabaseChain.from_llm(llm, db, verbose=True, prompt=PROMPT, use_query_checker=False) def clean_output(raw_output): cleaned = raw_output.replace("```sql", "").replace("```", "").strip() if "SQL查询语句:" in cleaned: cleaned = cleaned.split("SQL查询语句:")[0].strip() return cleaned try: raw_ans = db_chain.run("数据库中存在哪些表?") ans = clean_output(raw_ans) except Exception as e: print(str(e)) ans = '我当前无法提供服务' print(ans)
内容的提问来源于stack exchange,提问作者Abhishek Sasidharan
相关产品推荐
相关产品推荐

