跳到正文
LlamaIndex:产品、工程与评测·· 2023-08-17精选AI 评分62

LlamaIndex 教程:如何微调 Llama 2 用于 Text-to-SQL 应用

Easily Finetune Llama 2 for Your Text-to-SQL Applications

AI 导读

LlamaIndex 发布教程,展示如何用 text-to-SQL 数据集微调 Llama 2 并接入 LlamaIndex 对任意 SQL 数据库做结构化分析。

推荐理由

原文给出了微调 Llama 2 做 text-to-SQL 的完整技术栈和前后对比示例,可帮助读者评估微调对结构化分析场景的实际收益。

正文 · AI 翻译

Llama 2 是开源 LLM 发展中的一个巨大里程碑。最大的模型及其微调变体位居 Hugging Face Open LLM Leaderboard 榜首。多项基准测试表明,它在性能上正接近 GPT-3.5(在某些情况下甚至超越)。所有这些都意味着,开源 LLM 正日益成为复杂 LLM 应用(从 RAG 系统到智能体)中可行且可靠的选择。

立即探索我们的免费和付费计划。

背景:Llama-2–7B 不擅长文本转 SQL

然而,最小的 Llama 2 模型(70 亿参数)的一个缺点是它不太擅长生成 SQL,这使得它在结构化分析用例中不切实际。例如,我们尝试提示 Llama 2 根据以下提示模板生成正确的 SQL 语句:

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.

### Input:
{input}

### Context:
{context}

### Response:

这里我们插入了来自 sql-create-context 数据集 的一个样本条目。

input: In 1981 which team picked overall 148?
context: CREATE TABLE table_name_8 (team VARCHAR, year VARCHAR, overall_pick VARCHAR)

同时,这里是生成的输出与正确输出的对比:

Generated output: SELECT * FROM `table_name_8` WHERE '1980' = YEAR AND TEAM = "Boston Celtics" ORDER BY OVERALL_PICK DESC LIMIT 1;

Correct output: SELECT team FROM table_name_8 WHERE year = 1981 AND overall_pick = "148"

这显然不理想。与 ChatGPT 和 GPT-4 不同,Llama 2 无法可靠地生成格式良好且正确的 SQL 输出。

这正是微调发挥作用的地方——给定一个合适的文本转 SQL 数据语料库,我们可以教 Llama 2 更好地从自然语言生成 SQL 输出。在高层次上,微调涉及以某种方式修改模型的权重。微调模型有不同的方法,从更新网络的所有参数,到更新参数子集,再到仅微调额外参数(例如 LoRA 的工作原理)。

一旦模型微调完成,它仍然可以接入下游 LLM 应用。这正是本教程旨在展示的内容。它比我们现有的主要关注“上下文学习”和“检索增强”用例的教程更深入一步——冻结模型本身,但专注于将数据编排到输入提示中。微调可能有很高的学习曲线,并且需要大量计算。本教程尽可能简化入门过程。

教程概述

在本教程中,我们向您展示如何在文本转 SQL 数据集上微调 Llama 2,然后利用 LlamaIndex 的功能,将其用于针对任何 SQL 数据库的结构化分析。

以下是我们使用的技术栈:

特别提及来自 Anyscale 的精彩 Llama 2 教程,它启发了这个项目。

我们所有的材料都可以在我们的 Github 仓库中找到:https://github.com/run-llama/modal_finetune_sql(再次强调这是改编自 doppel-bot)。此外,完整教程可以在我们的 Jupyter notebook 指南 中找到。请务必查看!

如上所述,执行微调确实需要不少步骤。我们的目标是使其尽可能简单直接,以便遵循和开箱即用。我们没有涵盖 Modal、PEFT、微调过程本身等的所有细节,但我们给出了一个大致的概述。

当然,我们也可以使用更高级的 API(例如 OpenAI、Lamini)来完成这项任务。后续教程还有很大的空间来涵盖这些主题!

第 1 步:加载用于微调 LLaMa 的训练数据

这里的第一步是打开 Jupyter notebook。该 notebook 由一系列可运行的脚本组成,每个脚本执行加载数据所需的步骤。

我们的代码在编排的每一步都使用 Modal,而 Modal 最好直接用在 Python 脚本之上。这就是为什么这些单元格中有很多并不包含自己的 Python 代码块。

首先,我们使用 Modal 加载 b-mc2/sql-create-context 数据集。这是一个简单的任务,只需加载数据集并将其格式化为 .jsonl 文件。

modal run src.load_data_sql --data-dir "data_sql"

如我们所见,底层任务其实相当简单:


@stub.function(
    retries=Retries(
        max_retries=3,
        initial_delay=5.0,
        backoff_coefficient=2.0,
    ),
    timeout=60 * 60 * 2,
    network_file_systems={VOL_MOUNT_PATH.as_posix(): output_vol},
    cloud="gcp",
)
def load_data_sql(data_dir: str = "data_sql"):
    from datasets import load_dataset

    dataset = load_dataset("b-mc2/sql-create-context")

    dataset_splits = {"train": dataset["train"]}
    out_path = get_data_path(data_dir)

    out_path.parent.mkdir(parents=True, exist_ok=True)

    for key, ds in dataset_splits.items():
        with open(out_path, "w") as f:
            for item in ds:
                newitem = {
                    "input": item["question"],
                    "context": item["context"],
                    "output": item["answer"],
                }
                f.write(json.dumps(newitem) + "\n")

第 2 步:运行微调脚本

下一步是在解析后的数据集上运行我们的微调脚本。

modal run src.finetune_sql --data-dir "data_sql" --model-dir "model_sql"

微调脚本执行以下步骤。

将数据集划分为训练集和验证集

train_val = data["train"].train_test_split(test_size=val_set_size, shuffle=True, seed=42)
train_data = train_val["train"].shuffle().map(generate_and_tokenize_prompt)
val_data = train_val["test"].shuffle().map(generate_and_tokenize_prompt)

将每个划分格式化为(输入提示,标签)元组:输入查询和上下文被格式化为相同的输入提示。然后对输入提示进行分词,并将标签设置为与输入提示完全相同——这允许模型基于下一 token 预测进行训练。

def generate_and_tokenize_prompt(data_point):
  full_prompt = generate_prompt_sql(
      data_point["input"],
      data_point["context"],
      data_point["output"],
  )
  tokenized_full_prompt = tokenize(full_prompt)
  if not train_on_inputs:
      raise NotImplementedError("not implemented yet")
  return tokenized_full_prompt

输入提示与本博客顶部给出的完全相同。

运行微调脚本时,模型会保存在由 model_dir 指定的远程云目录中(如果未指定,则设置为默认值)。

第 3 步:评估

模型已完成微调,可以从云端提供服务。我们可以使用 sql-create-context 中的样本数据运行一些基本评估,以比较微调模型与基线 Llama 2 模型的性能。

modal run src.eval_sql::main

结果表明微调模型有巨大提升:

Input 1: {'input': 'Which region (year) has Abigail at number 7, Sophia at number 1 and Aaliyah at number 5?', 'context': 'CREATE TABLE table_name_12 (region__year_ VARCHAR, no_5 VARCHAR, no_7 VARCHAR, no_1 VARCHAR)', 'output': 'SELECT region__year_ FROM table_name_12 WHERE no_7 = "abigail" AND no_1 = "sophia" AND
no_5 = "aaliyah"'}
Output 1 (finetuned model): SELECT region__year_ FROM table_name_12 WHERE no_7 = "abigail" AND no_1 = "aaliyah" AND no_5 = "sophia"
Output 1 (base model): SELECT * FROM table_name_12 WHERE region__year = '2018' AND no_5 = 'Abigail' AND no_7 = 'Sophia' AND no_1 = 'Aaliyah';


Input 2: {'input': 'Name the result/games for 54741', 'context': 'CREATE TABLE table_21436373_11 (result_games VARCHAR, attendance VARCHAR)', 'output': 'SELECT result_games FROM table_21436373_11 WHERE attendance = 54741'}
Output 2 (finetuned model): SELECT result_games FROM table_21436373_11 WHERE attendance = "54741"
Output 2 (base model): SELECT * FROM table_21436373_11 WHERE result_games = 'name' AND attendance > 0;

基础模型会产生格式错误的输出或错误的 SQL 语句,

而微调模型能够产生与预期输出接近得多的结果。

第 4 步:将微调模型与 LlamaIndex 集成

现在,我们可以在 LlamaIndex 中使用该模型,对任何数据库执行 text-to-SQL。

我们首先定义一个测试 SQL 数据库,然后可以用它来测试模型的推理能力。

我们创建一个玩具 city_stats 表,其中包含城市名称、人口和国家信息,并填入几个示例城市。

db_file = "cities.db"
engine = create_engine(f"sqlite:///{db_file}")
metadata_obj = MetaData()
# create city SQL table
table_name = "city_stats"
city_stats_table = Table(
    table_name,
    metadata_obj,
    Column("city_name", String(16), primary_key=True),
    Column("population", Integer),
    Column("country", String(16), nullable=False),
)
metadata_obj.create_all(engine)

它存储在一个 cities.db 文件中。

然后,我们可以使用 Modal 将微调后的模型和该数据库文件加载到 LlamaIndex 中的 NLSQLTableQueryEngine——这个查询引擎允许用户轻松开始对给定数据库执行 text-to-SQL。

modal run src.inference_sql_llamaindex::main --query "Which city has the highest population?" --sqlite-file-path "nbs/cities.db" --model-dir "model_sql" --use-finetuned-model True

我们会得到类似如下的响应:

SQL Query: SELECT MAX(population) FROM city_stats WHERE country = "United States"
Response: [(2679000,)]

结论

基本上就是这样!本教程提供了一种非常高层级的方式,帮助你开始微调 Llama 2 模型以生成 SQL 语句,并端到端展示了如何将其接入你使用 LlamaIndex 的 text-to-SQL 工作流。

资源

为完整起见,我们在此再次链接所有资源。

教程仓库:https://github.com/run-llama/modal_finetune_sql(改编自 doppel-bot)。

Jupyter notebook 指南。

技术栈:

特别提及:来自 Anyscale 的 Llama 2 教程。

来源:LlamaIndex:产品、工程与评测 · llamaindex.ai