如何修改create_sql_chain代码以适配GPT-4及其他GPT模型?
适配GPT-4修复LangChain SQL查询链问题
你的代码切换到GPT-4后失效,核心原因是GPT-4属于聊天模型,而你误用了针对文本补全模型的OpenAI类,同时Prompt变量不匹配、格式不适配聊天模型也是关键问题。以下是针对性修改方案:
关键修改点
- 替换LLM实例:用
ChatOpenAI替代OpenAI类,适配GPT-4这类聊天模型 - 对齐Prompt变量:将Prompt中的
{input}改为{question},与调用时传入的参数名一致 - 强化输出约束:在Prompt中明确要求仅返回SQL代码,避免GPT-4生成额外解释
- 修正参数匹配:确保Prompt的输入变量与链式调用时传入的参数完全对应
修改后的完整代码
from langchain.chains import create_sql_query_chain from langchain.chat_models import ChatOpenAI # 替换为聊天模型专用类 from langchain.sql_database import SQLDatabase from langchain.prompts import PromptTemplate connection_string = "" db = SQLDatabase.from_uri(connection_string) # 用ChatOpenAI实例化GPT-4 llm = ChatOpenAI(temperature=0, verbose=True, model='gpt-4') # 优化Prompt,强化仅返回SQL的要求 seed_prompt = """ 你是专业SQL生成工具,仅返回符合要求的SQL代码,不添加任何额外解释、说明或标记。 示例: Question: "查询用户表中所有活跃用户的姓名" SQLQuery: "SELECT name FROM users WHERE is_active = 1" """ restrictions = """ 必须遵守以下规则: 1. 禁止使用LIMIT语句,改用TOP语句 2. 数值结果需格式化为###,###,###,### 3. 仅返回与问题直接相关的列 4. 若表或列不存在,返回"table or column could not be found" Question: {question} SQLQuery: """ prompt = seed_prompt + restrictions PROMPT = PromptTemplate( input_variables=["question"], # 与调用时的参数名对齐 template=prompt ) database_chain = create_sql_query_chain(llm, db, prompt=PROMPT) # 调用时传入的参数名与Prompt变量保持一致 sql_query = database_chain.invoke({"question": x}) # 清理可能的多余格式(如首尾引号) sql_query = sql_query.strip().strip('"') print(sql_query)
额外提示
如果仍存在格式问题,可尝试用ChatPromptTemplate构建对话式Prompt,明确区分系统指令和用户问题,进一步适配GPT-4的对话逻辑;同时确保你的LangChain版本为最新,避免版本兼容问题。
内容的提问来源于stack exchange,提问作者YOLOLJJ
相关产品推荐
相关产品推荐

