用 SQL 与 Jina Reranker v2 构建 RAG
本 Notebook 展示如何创建一个简单的检索增强生成(RAG)系统:信息来源是 SQL 数据库,而不是文档存储。
工作原理
• 从 SQL 数据库提取并保存表定义,即 SQL 导出文件中的 CREATE 语句。本教程已完成这一步,并把表定义作为列表保存在内存。扩大规模时可能需要更完善的存储方式。
• 用户用自然语言提出查询。
• Jina AI 的 SQL 感知重排模型 Jina Reranker v2 根据表定义与用户查询的相关性排序。
• 把用户查询及排名前三的表定义放入提示词,交给 Mistral 7B Instruct v0.1,要求它生成适合任务的 SQL 查询。
• Mistral Instruct 生成 SQL,随后对数据库执行查询并取得结果。
• 将 SQL 查询结果转换成 JSON,连同用户原始问题、SQL 语句,以及“用自然语言回答用户”的要求,组成新提示词交给 Mistral Instruct。
• 向用户返回 Mistral Instruct 生成的自然语言回答。
数据库
教程使用 GitHub 上公开的小型电子游戏销量数据库。这里选择 SQLite 版本,因为它小巧、跨平台,而且 Python 内置了接口支持。
软件与硬件要求
Jina Reranker v2 在本地运行。若使用 Google Colab,请选择可访问 GPU 的运行时;若在本机运行,需要 Python 3(原作者使用 Python 3.11 编写教程)。有支持 CUDA 的 GPU 时,运行速度会明显更快。
教程大量使用开源 LlamaIndex RAG 框架,并通过 Hugging Face Inference API 访问 Mistral 7B Instruct v0.1。需要 Hugging Face 账号和至少有 READ 权限的访问令牌。
> Google Colab 已安装 SQLite。本机若未安装,请按 SQLite 网站 的说明安装。Python 已内置 SQLite 接口,无须另外安装 Python 模块。
准备环境
安装依赖
首先安装所需的 Python 模块:
!pip install -qU transformers einops llama-index llama-index-postprocessor-jinaai-rerank llama-index-llms-huggingface "huggingface_hub[inference]"
下载数据库
将 GitHub 上的 SQLite 数据库 videogames.db 下载到本地。系统没有 wget 时,可从数据库链接 下载,并放在运行 Notebook 的同一目录:
!wget https://github.com/bbrumm/databasestar/raw/main/sample_databases/sample_db_videogames/sqlite/videogames.db
下载并运行 Jina Reranker v2
以下代码下载 jina-reranker-v2-base-multilingual,并在本地运行:
from transformers import AutoModelForSequenceClassification
reranker_model = AutoModelForSequenceClassification.from_pretrained(
"jinaai/jina-reranker-v2-base-multilingual",
torch_dtype="auto",
trust_remote_code=True,
)
reranker_model.to("cuda") # or 'cpu' if no GPU is available
reranker_model.eval()
设置 Mistral Instruct 接口
使用 LlamaIndex 创建一个连接对象,连接 Hugging Face 推理 API 以及在其上运行的 mistralai/Mixtral-8x7B-Instruct-v0.1。
先在 Hugging Face 账户设置 中取得访问令牌。
按下方提示输入:
import getpass
print("Paste your Hugging Face access token here: ")
hf_token = getpass.getpass()
接着初始化 LlamaIndex 的 HuggingFaceInferenceAPI,保存为 mistral_llm:
from llama_index.llms.huggingface import HuggingFaceInferenceAPI
mistral_llm = HuggingFaceInferenceAPI(model_name="mistralai/Mixtral-8x7B-Instruct-v0.1", token=hf_token)
使用具备 SQL 感知能力的 Jina Reranker v2
原作者从 GitHub 的数据库导入文件中提取了八个表定义。运行以下代码,将它们放进名为 table_declarations 的 Python 列表:
table_declarations = [
"CREATE TABLE platform (\n\tid INTEGER PRIMARY KEY,\n\tplatform_name TEXT DEFAULT NULL\n);",
"CREATE TABLE genre (\n\tid INTEGER PRIMARY KEY,\n\tgenre_name TEXT DEFAULT NULL\n);",
"CREATE TABLE publisher (\n\tid INTEGER PRIMARY KEY,\n\tpublisher_name TEXT DEFAULT NULL\n);",
"CREATE TABLE region (\n\tid INTEGER PRIMARY KEY,\n\tregion_name TEXT DEFAULT NULL\n);",
"CREATE TABLE game (\n\tid INTEGER PRIMARY KEY,\n\tgenre_id INTEGER,\n\tgame_name TEXT DEFAULT NULL,\n\tCONSTRAINT fk_gm_gen FOREIGN KEY (genre_id) REFERENCES genre(id)\n);",
"CREATE TABLE game_publisher (\n\tid INTEGER PRIMARY KEY,\n\tgame_id INTEGER DEFAULT NULL,\n\tpublisher_id INTEGER DEFAULT NULL,\n\tCONSTRAINT fk_gpu_gam FOREIGN KEY (game_id) REFERENCES game(id),\n\tCONSTRAINT fk_gpu_pub FOREIGN KEY (publisher_id) REFERENCES publisher(id)\n);",
"CREATE TABLE game_platform (\n\tid INTEGER PRIMARY KEY,\n\tgame_publisher_id INTEGER DEFAULT NULL,\n\tplatform_id INTEGER DEFAULT NULL,\n\trelease_year INTEGER DEFAULT NULL,\n\tCONSTRAINT fk_gpl_gp FOREIGN KEY (game_publisher_id) REFERENCES game_publisher(id),\n\tCONSTRAINT fk_gpl_pla FOREIGN KEY (platform_id) REFERENCES platform(id)\n);",
"CREATE TABLE region_sales (\n\tregion_id INTEGER DEFAULT NULL,\n\tgame_platform_id INTEGER DEFAULT NULL,\n\tnum_sales REAL,\n CONSTRAINT fk_rs_gp FOREIGN KEY (game_platform_id) REFERENCES game_platform(id),\n\tCONSTRAINT fk_rs_reg FOREIGN KEY (region_id) REFERENCES region(id)\n);",
]
下面定义一个函数,接收自然语言查询和表定义列表,用 Jina Reranker v2 对各表评分,再按分数从高到低返回:
from typing import List, Tuple
def rank_tables(query: str, table_specs: List[str], top_n: int = 0) -> List[Tuple[float, str]]:
"""
Get sorted pairs of scores and table specifications, then return the top N,
or all if top_n is 0 or default.
"""
pairs = [[query, table_spec] for table_spec in table_specs]
scores = reranker_model.compute_score(pairs)
scored_tables = [(score, table_spec) for score, table_spec in zip(scores, table_specs)]
scored_tables.sort(key=lambda x: x[0], reverse=True)
if top_n and top_n < len(scored_tables):
return scored_tables[0:top_n]
return scored_tables
Jina Reranker v2 会对每个表定义评分。默认情况下,函数返回全部表及其分数。可选参数 top_n 将结果限制为用户指定的数量,优先返回最高分项。
先定义一个查询来试用:
user_query = "Identify the top 10 platforms by total sales."
运行 rank_tables 获取表定义列表。将 top_n 设为 3,把结果赋给 ranked_tables,然后查看:
ranked_tables = rank_tables(user_query, table_declarations, top_n=3)
ranked_tables
输出应包括 region_sales、platform、game_platform 三个表;它们看起来都是查找答案的合理数据来源。
使用 Mistral Instruct 生成 SQL
接下来让 Mistral Instruct v0.1 根据重排器排名前三的表定义,编写回答用户问题的 SQL。
先使用 LlamaIndex 的 PromptTemplate 创建提示词:
from llama_index.core import PromptTemplate
make_sql_prompt_tmpl_text = """
Generate a SQL query to answer the following question from the user:
\"{query_str}\"
The SQL query should use only tables with the following SQL definitions:
Table 1:
{table_1}
Table 2:
{table_2}
Table 3:
{table_3}
Make sure you ONLY output an SQL query and no explanation.
"""
make_sql_prompt_tmpl = PromptTemplate(make_sql_prompt_tmpl_text)
调用 format,把用户查询和 Jina Reranker v2 选出的三个表定义填入模板:
make_sql_prompt = make_sql_prompt_tmpl.format(
query_str=user_query, table_1=ranked_tables[0][1], table_2=ranked_tables[1][1], table_3=ranked_tables[2][1]
)
可以查看实际将发送给 Mistral Instruct 的文本:
print(make_sql_prompt)
现在发送提示词并取得模型响应:
response = mistral_llm.complete(make_sql_prompt)
sql_query = str(response)
print(sql_query)
执行 SQL 查询
使用 Python 内置 SQLite 接口,对 videogames.db 执行以上查询:
import sqlite3
con = sqlite3.connect("videogames.db")
cur = con.cursor()
sql_response = cur.execute(sql_query).fetchall()
SQLite 接口详情参见 Python 3 文档。
查看结果:
sql_response
可以自行编写 SQL 查询,验证结果是否正确。数据库以浮点数保存销量;原作者推测单位是千份或百万份。
得到自然语言回答
现在使用新的提示词模板,把用户查询、SQL 语句和查询结果再次传给 Mistral Instruct。
与前面一样,先通过 LlamaIndex 创建新模板:
rag_prompt_tmpl_str = """
Use the information in the JSON table to answer the following user query.
Do not explain anything, just answer concisely. Use natural language in your
answer, not computer formatting.
USER QUERY: {query_str}
JSON table:
{json_table}
This table was generated by the following SQL query:
{sql_query}
Answer ONLY using the information in the table and the SQL query, and if the
table does not provide the information to answer the question, answer
"No Information".
"""
rag_prompt_tmpl = PromptTemplate(rag_prompt_tmpl_str)
把 SQL 输出转换为 Mistral Instruct v0.1 能理解的 JSON 格式。
填入模板字段:
import json
rag_prompt = rag_prompt_tmpl.format(
query_str="Identify the top 10 platforms by total sales", json_table=json.dumps(sql_response), sql_query=sql_query
)
请求 Mistral Instruct 生成自然语言回答:
rag_response = mistral_llm.complete(rag_prompt)
print(str(rag_response))
自行尝试
将上述流程整理为一个包含异常捕获的函数:
def answer_sql(user_query: str) -> str:
try:
ranked_tables = rank_tables(user_query, table_declarations, top_n=3)
except Exception as e:
print(f"Ranking failed.\nUser query:\n{user_query}\n\n")
raise (e)
make_sql_prompt = make_sql_prompt_tmpl.format(
query_str=user_query, table_1=ranked_tables[0][1], table_2=ranked_tables[1][1], table_3=ranked_tables[2][1]
)
try:
response = mistral_llm.complete(make_sql_prompt)
except Exception as e:
print(f"SQL query generation failed\nPrompt:\n{make_sql_prompt}\n\n")
raise (e)
# Backslash removal is a necessary hack because sometimes Mistral puts them
# in its generated code.
sql_query = str(response).replace("\\", "")
try:
sql_response = sqlite3.connect("videogames.db").cursor().execute(sql_query).fetchall()
except Exception as e:
print(f"SQL querying failed. Query:\n{sql_query}\n\n")
raise (e)
rag_prompt = rag_prompt_tmpl.format(query_str=user_query, json_table=json.dumps(sql_response), sql_query=sql_query)
try:
rag_response = mistral_llm.complete(rag_prompt)
return str(rag_response)
except Exception as e:
print(f"Answer generation failed. Prompt:\n{rag_prompt}\n\n")
raise (e)
试运行:
print(answer_sql("Identify the top 10 platforms by total sales."))
尝试其他问题:
print(answer_sql("Summarize sales by region."))
print(answer_sql("List the publisher with the largest number of published games."))
print(answer_sql("Display the year with most games released."))
print(answer_sql("What is the most popular game genre on the Wii platform?"))
print(answer_sql("What is the most popular game genre of 2012?"))
尝试自己的问题:
print(answer_sql("<INSERT QUESTION OR INSTRUCTION HERE>"))
回顾与结论
本教程展示了一个基础自然语言问答 RAG 系统,以 SQL 数据库作为信息来源。在该实现中,同一个大语言模型(Mistral Instruct v0.1)既负责生成 SQL 查询,也负责组织自然语言回答。
这里的数据库只是很小的示例。扩大规模时,仅对表定义列表排序可能不够,可以采用两阶段流程:先用嵌入模型和向量存储召回更多结果,再用重排模型把结果缩减到生成模型提示词可容纳的数量。
Notebook 假设任何请求最多只涉及三个表,实际场景显然不总是如此。Mistral 7B Instruct v0.1 也不能保证输出正确、甚至可执行的 SQL;生产环境需要更深入的错误处理。
更完善的错误处理、更长的输入上下文,以及面向 SQL 任务专门训练的生成模型,都可能显著影响实际应用效果。
这个示例说明,RAG 的思路可以延伸到结构化数据库,从而大幅拓宽应用范围。
来源:RAG backed by SQL and Jina Reranker v2。作者 Scott Martens(Jina AI)。Hugging Face Cookbook 采用 Apache 2.0。模型和数据库分别遵循其各自许可。











暂无评论内容