开发准备

创建新的环境

conda create -n test python=3.11 ipykernel

设置国内镜像源

# 设置Pip的全局index-url为阿里云镜像源
pip config set global.index-url https://mirrors.aliyun.com/pypi/simple/

# 将阿里源设置为主机信任源(防止HTTPS证书问题导致报错)
pip config set global.trusted-host mirrors.aliyun.com

pip install jupyterlab

# !pip install --upgrade llama-index
# !pip install llama-index-llms-dashscope
# !pip install llama-index-llms-openai-like
# !pip install llama-index-embeddings-dashscope

 pip install llama-index

pip install llama-cloud-services

设置阿里平台api-key

setx DASHSCOPE_API_KEY "sk-dab4c8deb19xxxx"

1. 大语言模型开发框架的价值是什么?

SDK:Software Development Kit,它是一组软件工具和资源的集合,旨在帮助开发者创建、测试、部署和维护应用程序或软件。

所有开发框架(SDK)的核心价值,都是降低开发、维护成本。

大语言模型开发框架的价值,是让开发者可以更方便地开发基于大语言模型的应用。主要提供两类帮助:

  1. 第三方能力抽象。比如 LLM、向量数据库、搜索接口等
  2. 常用工具、方案封装
  3. 底层实现封装。比如流式接口、超时重连、异步与并行等

好的开发框架,需要具备以下特点:

  1. 可靠性、鲁棒性高
  2. 可维护性高
  3. 可扩展性高
  4. 学习成本低

举些通俗的例子:

  • 与外部功能解依赖
    • 比如可以随意更换 LLM 而不用大量重构代码
    • 更换三方工具也同理
  • 经常变的部分要在外部维护而不是放在代码里
    • 比如 Prompt 模板
  • 各种环境下都适用
    • 比如线程安全
  • 方便调试和测试
    • 至少要能感觉到用了比不用方便吧
    • 合法的输入不会引发框架内部的报错

划重点:选对了框架,事半功倍;反之,事倍功半。

什么是 SDK? https://aws.amazon.com/cn/what-is/sdk/
SDK 和 API 的区别是什么? https://aws.amazon.com/cn/compare/the-difference-between-sdk-and-api/

🌰 举个例子:使用 SDK,4 行代码实现一个简易的 RAG 系统

LlamaIndex 默认的 Embedding 模型是 OpenAIEmbedding(model="text-embedding-ada-002")

import os
from llama_index.core import Settings
from llama_index.llms.openai_like import OpenAILike
from llama_index.llms.dashscope import DashScope, DashScopeGenerationModels
from llama_index.embeddings.dashscope import DashScopeEmbedding, DashScopeTextEmbeddingModels
from llama_index.core import VectorStoreIndex, SimpleDirectoryReader

# LlamaIndex默认使用的大模型被替换为百炼
# Settings.llm = OpenAILike(
#     model="qwen-max",
#     api_base="https://dashscope.aliyuncs.com/compatible-mode/v1",
#     api_key=os.getenv("DASHSCOPE_API_KEY"),
#     is_chat_model=True
# )

#设置llm的模型选择,以及对应的api-key
Settings.llm = DashScope(model_name=DashScopeGenerationModels.QWEN_MAX, api_key=os.getenv("DASHSCOPE_API_KEY"))

# LlamaIndex默认使用的Embedding模型被替换为百炼的Embedding模型
Settings.embed_model = DashScopeEmbedding(
    # model_name="text-embedding-v1"
    model_name=DashScopeTextEmbeddingModels.TEXT_EMBEDDING_V1,
    # api_key=os.getenv("DASHSCOPE_API_KEY")
)

documents = SimpleDirectoryReader("../data").load_data()
index = VectorStoreIndex.from_documents(documents)
query_engine = index.as_query_engine()
response = query_engine.query("deepseek v3有多少参数?")

print(response)

2. LlamaIndex 介绍

官网标题:「 Build AI Knowledge Assistants over your enterprise data 」

LlamaIndex 是一个为开发「知识增强」的大语言模型应用的框架(也就是 SDK)。知识增强,泛指任何在私有或特定领域数据基础上应用大语言模型的情况。例如:

LlamaIndex 有 Python 和 Typescript 两个版本,Python 版的文档相对更完善。

LlamaIndex 是一个开源框架,Github 链接:https://github.com/run-llama

LlamaIndex 的核心模块​​​​​​¶

安装 LlamaIndex

[ ]:

# !pip install llama-index

3.数据加载(Loading)

3.1、加载本地数据

SimpleDirectoryReader 是一个简单的本地文件加载器。它会遍历指定目录,并根据文件扩展名自动加载文件(文本内容)。

支持的文件类型:

  • .csv - comma-separated values
  • .docx - Microsoft Word
  • .epub - EPUB ebook format
  • .hwp - Hangul Word Processor
  • .ipynb - Jupyter Notebook
  • .jpeg.jpg - JPEG image
  • .mbox - MBOX email archive
  • .md - Markdown
  • .mp3.mp4 - audio and video
  • .pdf - Portable Document Format
  • .png - Portable Network Graphics
  • .ppt.pptm.pptx - Microsoft PowerPoint
import json
from pydantic.v1 import BaseModel
from llama_index.core import SimpleDirectoryReader

def show_json(data):
    """用于展示json数据"""
    if isinstance(data, str):
        obj = json.loads(data)
        print(json.dumps(obj, indent=4, ensure_ascii=False))
    elif isinstance(data, dict) or isinstance(data, list):
        print(json.dumps(data, indent=4, ensure_ascii=False))
    elif issubclass(type(data), BaseModel):
        print(json.dumps(data.dict(), indent=4, ensure_ascii=False))

def show_list_obj(data):
    """用于展示一组对象"""
    if isinstance(data, list):
        for item in data:
            show_json(item)
    else:
        raise ValueError("Input is not a list")


reader = SimpleDirectoryReader(
        input_dir="./data", # 目标目录
        recursive=False, # 是否递归遍历子目录
        required_exts=[".pdf"] # (可选)只读取指定后缀的文件
    )
documents = reader.load_data()

print(documents[0].text)
show_json(documents[0].json())

        

注意:对图像、视频、语音类文件,默认不会自动提取其中文字。如需提取,参考下面介绍的 Data Connectors

默认的 PDFReader 效果并不理想,我们可以更换文件加载器

LlamaParse

首先,登录并从 https://cloud.llamaindex.ai ↗ 注册并获取 api-key 。

不确定是否好用,建议使用 https://mineru.net/OpenSourceTools/Extractor

 # 从.env文件加载环境变量
from dotenv import load_dotenv
from llama_cloud_services import LlamaParse
from llama_index.core import SimpleDirectoryReader
import nest_asyncio
import os

# 加载同级目录下的.env文件
load_dotenv()  # 默认加载当前目录下的.env文件

# 如果需要指定.env文件路径,可以使用:
# load_dotenv(".env")  # 指定文件名
# load_dotenv("/path/to/.env")  # 指定完整路径

nest_asyncio.apply()

# 从环境变量获取API密钥(不再在代码中硬编码)
api_key = os.getenv("LLAMA_CLOUD_API_KEY")

# 检查密钥是否加载成功
if not api_key:
    raise ValueError("LLAMA_CLOUD_API_KEY 未在.env文件中找到或为空")

# 设置环境变量(如果需要在其他地方也通过os.environ访问)
# 注意:dotenv已经加载了.env文件到环境变量,所以这步是可选的
os.environ["LLAMA_CLOUD_API_KEY"] = api_key

# 设置解析器
parser = LlamaParse(
    result_type="markdown"
)

file_extractor = {".pdf": parser}

documents = SimpleDirectoryReader(input_dir="./data", required_exts=[".pdf"], file_extractor=file_extractor).load_data()
print(documents[0].text)

3.2、Data Connectors

用于处理更丰富的数据类型,并将其读取为 Document 的形式。

例如:直接读取网页

pip install llama-index-readers-web
from llama_index.readers.web import SimpleWebPageReader

documents = SimpleWebPageReader(html_to_text=True).load_data(
    #["https://edu.guangjuke.com/tx/"]
    ["https://docs.pydantic.dev/2.12/errors/validation_errors/#uuid_type"]
)

print(documents[0].text)

更多 Data Connectors

4. 文本切分与解析(Chunking)

为方便检索,我们通常把 Document 切分为 Node

在 LlamaIndex 中,Node 被定义为一个文本的「chunk」。

4.1、使用 TextSplitters 对文本做切分

例如:TokenTextSplitter 按指定 token 数切分文本


from llama_index.core import Document
from llama_index.core.node_parser import TokenTextSplitter
import json
from pydantic.v1 import BaseModel
from llama_index.core import SimpleDirectoryReader

def show_json(data):
    """用于展示json数据"""
    if isinstance(data, str):
        obj = json.loads(data)
        print(json.dumps(obj, indent=4, ensure_ascii=False))
    elif isinstance(data, dict) or isinstance(data, list):
        print(json.dumps(data, indent=4, ensure_ascii=False))
    elif issubclass(type(data), BaseModel):
        print(json.dumps(data.dict(), indent=4, ensure_ascii=False))

def show_list_obj(data):
    """用于展示一组对象"""
    if isinstance(data, list):
        for item in data:
            show_json(item)
    else:
        raise ValueError("Input is not a list")


reader = SimpleDirectoryReader(
        input_dir="./data", # 目标目录
        recursive=False, # 是否递归遍历子目录
        required_exts=[".pdf"] # (可选)只读取指定后缀的文件
    )
documents = reader.load_data()

#print("dadadada",sep="@@")
print(documents[0].text)
show_json(documents[0].json())


#新增
node_parser = TokenTextSplitter(
    chunk_size=512,  # 每个 chunk 的最大长度
    chunk_overlap=200  # chunk 之间重叠长度
)

nodes = node_parser.get_nodes_from_documents(
    documents, show_progress=False
)

show_json(nodes[1].json())
show_json(nodes[2].json())

LlamaIndex 提供了丰富的 TextSplitter,例如:

  • SentenceSplitter:在切分指定长度的 chunk 同时尽量保证句子边界不被切断;
  • CodeSplitter:根据 AST(编译器的抽象句法树)切分代码,保证代码功能片段完整;
  • SemanticSplitterNodeParser:根据语义相关性对将文本切分为片段。

4.2、使用 NodeParsers 对有结构的文档做解析

例如:HTMLNodeParser解析 HTML 文档

更多的 NodeParser 包括 MarkdownNodeParserJSONNodeParser等等。

from llama_index.core.node_parser import HTMLNodeParser
from llama_index.readers.web import SimpleWebPageReader

documents = SimpleWebPageReader(html_to_text=False).load_data(
    ["https://edu.guangjuke.com/tx/"]
)

# 默认解析 ["p", "h1", "h2", "h3", "h4", "h5", "h6", "li", "b", "i", "u", "section"]
parser = HTMLNodeParser(tags=["span"])  # 可以自定义解析哪些标签
nodes = parser.get_nodes_from_documents(documents)

for node in nodes:
    print(node.text+"\n")

5. 索引(Indexing)与检索(Retrieval)

基础概念:在「检索」相关的上下文中,「索引」即index, 通常是指为了实现快速检索而设计的特定「数据结构」。

索引的具体原理与实现不是本课程的教学重点,感兴趣可以参考:传统索引向量索引

5.1、向量检索

  1. VectorStoreIndex 直接在内存中构建一个 Vector Store 并建索引

from llama_index.core import VectorStoreIndex, SimpleDirectoryReader, Settings
from llama_index.core.node_parser import TokenTextSplitter
from llama_index.embeddings.dashscope import DashScopeEmbedding  # 百炼嵌入模型
from llama_index.llms.dashscope import DashScope  # 百炼LLM(如果后续需要生成回答)
from dotenv import load_dotenv
import os

# 1. 加载环境变量
load_dotenv()  # 从 .env 文件加载

# 2. 设置百炼的API密钥(环境变量名通常为DASHSCOPE_API_KEY)
# 请在您的 .env 文件中添加:DASHSCOPE_API_KEY=您的-sk-密钥
dashscope_api_key = os.getenv("DASHSCOPE_API_KEY")
if not dashscope_api_key:
    raise ValueError("DASHSCOPE_API_KEY 未在.env文件中找到。请前往阿里百炼平台获取。")

# 3. 配置全局设置,使用百炼的嵌入模型
# 关键步骤:替换默认的OpenAI嵌入模型
Settings.embed_model = DashScopeEmbedding(
    model_name="text-embedding-v2",  # 百炼的文本嵌入模型
    api_key=dashscope_api_key,       # 传入密钥
)

# (可选)如果您后续使用查询引擎生成答案,也需要设置LLM
# Settings.llm = DashScope(model="qwen-max", api_key=dashscope_api_key)

# 4. 加载并处理文档(此部分与您的原始代码一致)
documents = SimpleDirectoryReader(
    "./data",
    required_exts=[".pdf"],
).load_data()

node_parser = TokenTextSplitter(chunk_size=512, chunk_overlap=200)
nodes = node_parser.get_nodes_from_documents(documents)

# 5. 构建向量索引
# 此时,VectorStoreIndex 会自动使用上面设置的 DashScopeEmbedding
index = VectorStoreIndex(nodes)

# 6. 创建检索器并进行查询
vector_retriever = index.as_retriever(similarity_top_k=2)
results = vector_retriever.retrieve("deepseek v3数学能力怎么样?")

# 7. 输出结果
if results:
    print(results[0].text)
else:
    print("未检索到相关结果。")

        2.使用自定义的 Vector Store,以 Qdrant 为例:

pip install llama-index-vector-stores-qdrant
#追加到上文代码
from llama_index.core.indices.vector_store.base import VectorStoreIndex
from llama_index.vector_stores.qdrant import QdrantVectorStore
from llama_index.core import StorageContext

from qdrant_client import QdrantClient
from qdrant_client.models import VectorParams, Distance

client = QdrantClient(location=":memory:")
collection_name = "demo"
collection = client.create_collection(
    collection_name=collection_name,
    vectors_config=VectorParams(size=1536, distance=Distance.COSINE)
)

vector_store = QdrantVectorStore(client=client, collection_name=collection_name)
# storage: 指定存储空间
storage_context = StorageContext.from_defaults(vector_store=vector_store)

# 创建 index:通过 Storage Context 关联到自定义的 Vector Store
index = VectorStoreIndex(nodes, storage_context=storage_context)

# 获取 retriever
vector_retriever = index.as_retriever(similarity_top_k=1)

# 检索
results = vector_retriever.retrieve("deepseek v3数学能力怎么样")

print(results[0])

5.2、更多索引与检索方式

LlamaIndex 内置了丰富的检索机制,例如:

5.3、检索后处理

LlamaIndex 的 Node Postprocessors 提供了一系列检索后处理模块。

例如:我们可以用不同模型对检索后的 Nodes 做重排序

完整代码

# 导入所需的库
import os
from dotenv import load_dotenv
from qdrant_client import QdrantClient
from qdrant_client.models import VectorParams, Distance
from llama_index.core import VectorStoreIndex, SimpleDirectoryReader, Settings, StorageContext
from llama_index.core.node_parser import TokenTextSplitter
from llama_index.core.postprocessor import LLMRerank
from llama_index.embeddings.dashscope import DashScopeEmbedding
from llama_index.llms.dashscope import DashScope
from llama_index.vector_stores.qdrant import QdrantVectorStore

# 1. 加载环境变量
load_dotenv()

# 2. 获取阿里百炼API密钥
dashscope_api_key = os.getenv("DASHSCOPE_API_KEY")
if not dashscope_api_key:
    raise ValueError("DASHSCOPE_API_KEY 未在.env文件中找到。请前往阿里百炼平台获取。")

# 3. 配置全局模型设置
# 设置嵌入模型
Settings.embed_model = DashScopeEmbedding(
    model_name="text-embedding-v2",  # 百炼文本嵌入模型
    api_key=dashscope_api_key
)

# 设置LLM(用于重排序)
Settings.llm = DashScope(
    model="qwen-max",  # 可选:qwen-max, qwen-plus, qwen-turbo
    api_key=dashscope_api_key
)

# 4. 加载和分割文档
print("正在加载文档...")
documents = SimpleDirectoryReader(
    "./data",
    required_exts=[".pdf"],
).load_data()

print(f"已加载 {len(documents)} 个文档")
node_parser = TokenTextSplitter(chunk_size=512, chunk_overlap=200)
nodes = node_parser.get_nodes_from_documents(documents)
print(f"文档已分割为 {len(nodes)} 个节点")

# 5. 构建内存中的向量索引
print("\n构建内存中的向量索引...")
memory_index = VectorStoreIndex(nodes)
memory_retriever = memory_index.as_retriever(similarity_top_k=2)
memory_results = memory_retriever.retrieve("deepseek v3数学能力怎么样?")

print("\n=== 内存索引检索结果 ===")
if memory_results:
    for i, result in enumerate(memory_results):
        print(f"\n--- 结果 {i+1} ---")
        print(result.text[:300] + "..." if len(result.text) > 300 else result.text)
else:
    print("未检索到相关结果。")

# 6. 使用Qdrant向量存储
print("\n" + "="*80)
print("使用Qdrant向量存储")
print("="*80)

# 创建Qdrant客户端(内存模式)
client = QdrantClient(location=":memory:")
collection_name = "document_collection"

# 创建集合
collection = client.create_collection(
    collection_name=collection_name,
    vectors_config=VectorParams(size=1536, distance=Distance.COSINE)
)

# 创建Qdrant向量存储
vector_store = QdrantVectorStore(client=client, collection_name=collection_name)
storage_context = StorageContext.from_defaults(vector_store=vector_store)

# 创建索引
qdrant_index = VectorStoreIndex(nodes, storage_context=storage_context)
qdrant_retriever = qdrant_index.as_retriever(similarity_top_k=1)

# 检索
qdrant_results = qdrant_retriever.retrieve("deepseek v3数学能力怎么样")

print(f"\nQdrant存储检索结果(取top-1):")
if qdrant_results:
    print(qdrant_results[0].text[:300] + "..." if len(qdrant_results[0].text) > 300 else qdrant_results[0].text)
else:
    print("未检索到相关结果。")

# 7. 使用Qwen模型进行重排序
print("\n" + "="*80)
print("使用Qwen模型进行重排序")
print("="*80)

# 先进行初步检索(获取更多结果)
preliminary_retriever = qdrant_index.as_retriever(similarity_top_k=5)
preliminary_nodes = preliminary_retriever.retrieve("deepseek v3有多少参数?")

print(f"初步检索到 {len(preliminary_nodes)} 个节点:")
for i, node in enumerate(preliminary_nodes):
    preview = node.text[:150] + "..." if len(node.text) > 150 else node.text
    print(f"[{i}] {preview}\n")

# 使用Qwen模型进行重排序
postprocessor = LLMRerank(top_n=2)
reranked_nodes = postprocessor.postprocess_nodes(
    preliminary_nodes, 
    query_str="deepseek v3有多少参数?"
)

print(f"\n=== 重排序后最相关的 {len(reranked_nodes)} 个段落 ===")
for i, node in enumerate(reranked_nodes):
    print(f"\n--- 段落 {i+1} ---")
    print(node.text[:500] + "..." if len(node.text) > 500 else node.text)

6. 生成回复(QA & Chat)

6.1 单轮问答(Query Engine)

        

流式输出

qa_engine = index.as_query_engine(streaming=True)

response = qa_engine.query("deepseek v3数学能力怎么样?")

response.print_response_stream()

6.2 多轮对话(Chat Engine)

chat_engine = index.as_chat_engine()

response = chat_engine.chat("deepseek v3数学能力怎么样?")

print(response)

response = chat_engine.chat("代码能力呢?")

print(response)

流式输出

chat_engine = index.as_chat_engine()

streaming_response = chat_engine.stream_chat("deepseek v3数学能力怎么样?")

# streaming_response.print_response_stream()

for token in streaming_response.response_gen:

print(token, end="", flush=True)


print("单轮回答")
index = qdrant_index  # 显式指定要使用的索引对象
qa_engine = index.as_query_engine()
response = qa_engine.query("deepseek v3数学能力怎么样?")
print(response)

print("流式输出")
qa_engine = index.as_query_engine(streaming=True)
response = qa_engine.query("deepseek v3数学能力怎么样?")
response.print_response_stream()

print("多轮对话(Chat Engine)¶")
chat_engine = index.as_chat_engine()
response = chat_engine.chat("deepseek v3数学能力怎么样?")
print(response)
response = chat_engine.chat("代码能力呢")
print(response)

print("多轮对话(Chat Engine)流式输出")
chat_engine = index.as_chat_engine()
streaming_response = chat_engine.stream_chat("deepseek v3数学能力怎么样?")
# streaming_response.print_response_stream()
for token in streaming_response.response_gen:
    print(token, end="", flush=True)

7. 底层接口:Prompt、LLM 与 Embedding

7.1 Prompt 模板

PromptTemplate 定义提示词模板
from llama_index.core import PromptTemplate

prompt = PromptTemplate("写一个关于{topic}的笑话")

prompt.format(topic="小明")
ChatPromptTemplate 定义多轮消息模板
from llama_index.core.llms import ChatMessage, MessageRole
from llama_index.core import ChatPromptTemplate

chat_text_qa_msgs = [
    ChatMessage(
        role=MessageRole.SYSTEM,
        content="你叫{name},你必须根据用户提供的上下文回答问题。",
    ),
    ChatMessage(
        role=MessageRole.USER, 
        content=(
            "已知上下文:\n" \
            "{context}\n\n" \
            "问题:{question}"
        )
    ),
]
text_qa_template = ChatPromptTemplate(chat_text_qa_msgs)

print(
    text_qa_template.format(
        name="小明",
        context="这是一个测试",
        question="这是什么"
    )
)

7.2 语言模型

from llama_index.llms.openai import OpenAI

llm = OpenAI(temperature=0, model="gpt-4o")
response = llm.complete(prompt.format(topic="小明"))

print(response.text)
response = llm.complete(
    text_qa_template.format(
        name="小明",
        context="这是一个测试",
        question="你是谁,我们在干嘛"
    )
)

print(response.text)
连接DeepSeek

需要注册 开通deepseek的,key

pip install llama-index-llms-deepseek
import os
from llama_index.llms.deepseek import DeepSeek

llm = DeepSeek(model="deepseek-chat", api_key=os.getenv("DEEPSEEK_API_KEY"), temperature=1.5)

response = llm.complete("写个笑话")
print(response)

设置全局使用的语言模型
from llama_index.core import Settings

Settings.llm = DeepSeek(model="deepseek-chat", api_key=os.getenv("DEEPSEEK_API_KEY"), temperature=1.5)

除 OpenAI 外,LlamaIndex 已集成多个大语言模型,包括云服务 API 和本地部署 API,详见官方文档:Available LLM integrations

7.3 Embedding 模型

from llama_index.embeddings.openai import OpenAIEmbedding
from llama_index.core import Settings

# 全局设定
Settings.embed_model = OpenAIEmbedding(model="text-embedding-3-small", dimensions=512)

LlamaIndex 同样集成了多种 Embedding 模型,包括云服务 API 和开源模型(HuggingFace)等,详见官方文档

8. 基于 LlamaIndex 实现一个功能较完整的 RAG 系统

功能要求:

  • 加载指定目录的文件
  • 支持 RAG-Fusion
  • 使用 Qdrant 向量数据库,并持久化到本地
  • 支持检索后排序
  • 支持多轮对话
import os
from qdrant_client import QdrantClient
from qdrant_client.models import VectorParams, Distance
from llama_index.core import VectorStoreIndex, SimpleDirectoryReader, get_response_synthesizer
from llama_index.vector_stores.qdrant import QdrantVectorStore
from llama_index.core.node_parser import SentenceSplitter
from llama_index.core.response_synthesizers import ResponseMode
from llama_index.core.ingestion import IngestionPipeline
from llama_index.core import Settings
from llama_index.core import StorageContext
from llama_index.core.postprocessor import LLMRerank, SimilarityPostprocessor
from llama_index.core.retrievers import QueryFusionRetriever
from llama_index.core.query_engine import RetrieverQueryEngine
from llama_index.core.chat_engine import CondenseQuestionChatEngine
from llama_index.llms.dashscope import DashScope, DashScopeGenerationModels
from llama_index.embeddings.dashscope import DashScopeEmbedding, DashScopeTextEmbeddingModels


EMBEDDING_DIM = 1536
COLLECTION_NAME = "full_demo"
PATH = "./qdrant_db"

client = QdrantClient(path=PATH)


# 1. 指定全局llm与embedding模型
Settings.llm = DashScope(model_name=DashScopeGenerationModels.QWEN_MAX,api_key=os.getenv("DASHSCOPE_API_KEY"))
Settings.embed_model = DashScopeEmbedding(model_name=DashScopeTextEmbeddingModels.TEXT_EMBEDDING_V1)

# 2. 指定全局文档处理的摄取管道,,SentenceSplitter最大可能保证完整语义
Settings.transformations = [SentenceSplitter(chunk_size=512, chunk_overlap=200)]

# 3. 加载本地文档
documents = SimpleDirectoryReader("./data").load_data()

# 删除旧的 collection
if client.collection_exists(collection_name=COLLECTION_NAME):
    client.delete_collection(collection_name=COLLECTION_NAME)

# 4. 创建 collection
client.create_collection(
    collection_name=COLLECTION_NAME,
    vectors_config=VectorParams(size=EMBEDDING_DIM, distance=Distance.COSINE)
)

# 5. 创建 向量存储
vector_store = QdrantVectorStore(client=client, collection_name=COLLECTION_NAME)

# 6. 指定 向量存储 的 Storage 用于 index
storage_context = StorageContext.from_defaults(vector_store=vector_store)
index = VectorStoreIndex.from_documents(
    documents, storage_context=storage_context
)

# 7. 定义检索后排序模型
reranker = LLMRerank(top_n=2)
# 最终打分低于0.6的文档被过滤掉
sp = SimilarityPostprocessor(similarity_cutoff=0.6)

# 8. 定义 RAG Fusion 检索器
fusion_retriever = QueryFusionRetriever(
    [index.as_retriever()],
    similarity_top_k=5, # 检索召回 top k 结果
    num_queries=3,  # 生成 query 数
    use_async=False,
    # query_gen_prompt="",  # 可以自定义 query 生成的 prompt 模板
)

# 9. 构建单轮 query engine
query_engine = RetrieverQueryEngine.from_args(
    fusion_retriever,
    node_postprocessors=[reranker],
    response_synthesizer=get_response_synthesizer(
        response_mode = ResponseMode.REFINE
    )
)

# 10. 对话引擎
chat_engine = CondenseQuestionChatEngine.from_defaults(
    query_engine=query_engine,
    # condense_question_prompt="" # 可以自定义 chat message prompt 模板
)

# 测试多轮对话
# User: deepseek v3有多少参数
# User: 每次激活多少

while True:
    question=input("User:")
    if question.strip() == "":
        break
    response = chat_engine.chat(question)
    print(f"AI: {response}")

9. Text2SQL / NL2SQL / NL2Chart / ChatBI

9.1 基本介绍

Text2SQL 是一种将自然语言转换为SQL查询语句的技术。

这项技术的意义:让每个人都能像对话一样查询数据库,获取所需信息,而不必学习SQL语法。

9.2 典型应用场景
  • 业务分析师的数据自助服务

  • 智能BI与数据可视化

  • 客服与内部数据库查询

  • 跨部门数据协作与分享

  • 运营数据分析与决策支持

9.3 Text2SQL核心能力与挑战

一个成熟的Text2SQL系统需要具备以下关键能力:

9.4 实现Text2SQL的技术架构

10. 工作流(Workflow)了解

10.1 工作流(Workflow)简介

工作流顾名思义是对一些列工作步骤的抽象。

LlamaIndex 的工作流是事件(event)驱动的:

  • 工作流由 step 组成
  • 每个 step 处理特定的事件
  • step 也会产生新的事件(交由后继的 step 进行处理)
  • 直到产生 StopEvent 整个工作流结束

LlamaIndex Workflows:https://docs.llamaindex.ai/en/stable/module_guides/workflow/

10.2 工作流设计

使用自然语言查询数据库,数据库中包含多张表

工作流设计:

分步说明:

  1. 用户输入自然语言查询
  2. 系统先去检索跟查询相关的表
  3. 根据表的 Schema 让大模型生成 SQL
  4. 用生成的 SQL 查询数据库
  5. 根据查询结果,调用大模型生成自然语言回复

10.3 数据准备

# 下载 WikiTableQuestions
# WikiTableQuestions 是一个为表格问答设计的数据集。其中包含 2,108 个从维基百科提取的 HTML 表格

# !wget "https://github.com/ppasupat/WikiTableQuestions/releases/download/v1.0.2/WikiTableQuestions-1.0.2-compact.zip" -O wiki_data.zip
# !unzip wiki_data.zip

  1. 遍历目录加载表格
import pandas as pd
from pathlib import Path

data_dir = Path("./WikiTableQuestions/csv/200-csv")
csv_files = sorted([f for f in data_dir.glob("*.csv")])
dfs = []
for csv_file in csv_files:
    print(f"processing file: {csv_file}")
    try:
        df = pd.read_csv(csv_file)
        dfs.append(df)
    except Exception as e:
        print(f"Error parsing {csv_file}: {str(e)}")
  1. 为每个表生成一段文字表述(用于检索),保存在 WikiTableQuestions_TableInfo 目录
import pandas as pd
import json
import os
from pathlib import Path
from llama_index.core.prompts import ChatPromptTemplate
from llama_index.core.bridge.pydantic import BaseModel, Field
from llama_index.core.llms import ChatMessage
from llama_index.llms.dashscope import DashScope

#遍历目录加载表格
data_dir = Path("./WikiTableQuestions/csv/200-csv")
csv_files = sorted([f for f in data_dir.glob("*.csv")])
dfs = []
for csv_file in csv_files:
    print(f"processing file: {csv_file}")
    try:
        df = pd.read_csv(csv_file)
        dfs.append(df)
    except Exception as e:
        print(f"Error parsing {csv_file}: {str(e)}")

#为每个表生成一段文字表述(用于检索),保存在 WikiTableQuestions_TableInfo 目录

class TableInfo(BaseModel):
    """Information regarding a structured table."""

    table_name: str = Field(
        ..., description="table name (must be underscores and NO spaces)"
    )
    table_summary: str = Field(
        ..., description="short, concise summary/caption of the table"
    )


prompt_str = """
Give me a summary of the table with the following JSON format.

- The table name must be unique to the table and describe it while being concise. 
- Do NOT output a generic table name (e.g. table, my_table).

Do NOT make the table name one of the following: {exclude_table_name_list}

Table:
{table_str}

Summary: """

prompt_tmpl = ChatPromptTemplate(
    message_templates=[ChatMessage.from_str(prompt_str, role="user")]
)

tableinfo_dir = "WikiTableQuestions_TableInfo"
# !mkdir {tableinfo_dir}


def _get_tableinfo_with_index(idx: int) -> str:
    results_gen = Path(tableinfo_dir).glob(f"{idx}_*")
    results_list = list(results_gen)
    if len(results_list) == 0:
        return None
    elif len(results_list) == 1:
        path = results_list[0]
        with open(path, 'r') as file:
            data = json.load(file)
            return TableInfo.model_validate(data)
    else:
        raise ValueError(
            f"More than one file matching index: {list(results_gen)}"
        )

os.environ["DASHSCOPE_API_KEY"] = "sk-dab4c8deb1xxxx6c13772a"  # 替换为你的 API Key
llm = DashScope(model="qwen-plus", api_key=os.getenv("DASHSCOPE_API_KEY"))  # 初始化 llm


table_names = set()
table_infos = []
for idx, df in enumerate(dfs):
    table_info = _get_tableinfo_with_index(idx)
    if table_info:
        table_infos.append(table_info)
    else:
        while True:
            df_str = df.head(10).to_csv()
            table_info = llm.structured_predict(
                TableInfo,
                prompt_tmpl,
                table_str=df_str,
                exclude_table_name_list=str(list(table_names)),
            )
            table_name = table_info.table_name
            print(f"Processed table: {table_name}")
            if table_name not in table_names:
                table_names.add(table_name)
                break
            else:
                # try again
                print(f"Table name {table_name} already exists, trying again.")
                pass

        out_file = f"{tableinfo_dir}/{idx}_{table_name}.json"
        json.dump(table_info.dict(), open(out_file, "w"))
    table_infos.append(table_info)
  1. 将上述表格存入 SQLite 数据库
# put data into sqlite db
from sqlalchemy import (
    create_engine,
    MetaData,
    Table,
    Column,
    String,
    Integer,
)
import re


# Function to create a sanitized column name
def sanitize_column_name(col_name):
    # Remove special characters and replace spaces with underscores
    return re.sub(r"\W+", "_", col_name)


# Function to create a table from a DataFrame using SQLAlchemy
def create_table_from_dataframe(df: pd.DataFrame, table_name: str, engine, metadata_obj):
    # Sanitize column names
    sanitized_columns = {col: sanitize_column_name(col) for col in df.columns}
    df = df.rename(columns=sanitized_columns)

    # Dynamically create columns based on DataFrame columns and data types
    columns = [
        Column(col, String if dtype == "object" else Integer)
        for col, dtype in zip(df.columns, df.dtypes)
    ]

    # Create a table with the defined columns
    table = Table(table_name, metadata_obj, *columns)

    # Create the table in the database
    metadata_obj.create_all(engine)

    # Insert data from DataFrame into the table
    with engine.connect() as conn:
        for _, row in df.iterrows():
            insert_stmt = table.insert().values(**row.to_dict())
            conn.execute(insert_stmt)
        conn.commit()


# engine = create_engine("sqlite:///:memory:")
engine = create_engine("sqlite:///wiki_table_questions.db")
metadata_obj = MetaData()
for idx, df in enumerate(dfs):
    tableinfo = _get_tableinfo_with_index(idx)
    print(f"Creating table: {tableinfo.table_name}")
    create_table_from_dataframe(df, tableinfo.table_name, engine, metadata_obj)

链接SQLite数据库

打开navicat,选择链接SQLite

然后选择对应db文件的目录

10.4 构建基础工具

  1. 创建基于表的描述的向量索引
import os
from llama_index.core import Settings
from llama_index.llms.dashscope import DashScope, DashScopeGenerationModels
from llama_index.embeddings.dashscope import DashScopeEmbedding, DashScopeTextEmbeddingModels
from llama_index.core.objects import (
    SQLTableNodeMapping,
    ObjectIndex,
    SQLTableSchema,
)
from llama_index.core import SQLDatabase, VectorStoreIndex

# 设置全局模型
Settings.llm = DashScope(model_name=DashScopeGenerationModels.QWEN_MAX, api_key=os.getenv("DASHSCOPE_API_KEY"))
Settings.embed_model = DashScopeEmbedding(model_name=DashScopeTextEmbeddingModels.TEXT_EMBEDDING_V1)

sql_database = SQLDatabase(engine)

table_node_mapping = SQLTableNodeMapping(sql_database)
table_schema_objs = [
    SQLTableSchema(table_name=t.table_name, context_str=t.table_summary)
    for t in table_infos
]  # add a SQLTableSchema for each table

obj_index = ObjectIndex.from_objects(
    table_schema_objs,
    table_node_mapping,
    VectorStoreIndex,
)
obj_retriever = obj_index.as_retriever(similarity_top_k=3)
  1. 创建 SQL 查询器
from llama_index.core.retrievers import SQLRetriever
from typing import List

sql_retriever = SQLRetriever(sql_database)


def get_table_context_str(table_schema_objs: List[SQLTableSchema]):
    """Get table context string."""
    context_strs = []
    for table_schema_obj in table_schema_objs:
        table_info = sql_database.get_single_table_info(
            table_schema_obj.table_name
        )
        if table_schema_obj.context_str:
            table_opt_context = " The table description is: "
            table_opt_context += table_schema_obj.context_str
            table_info += table_opt_context

        context_strs.append(table_info)
    return "\n\n".join(context_strs)
  1. 创建 Text2SQL 的提示词(系统默认模板),和输出结果解析器(从生成的文本中抽取SQL)
    from llama_index.core.prompts.default_prompts import DEFAULT_TEXT_TO_SQL_PROMPT
    from llama_index.core import PromptTemplate
    from llama_index.core.llms import ChatResponse
    
    def parse_response_to_sql(chat_response: ChatResponse) -> str:
        """Parse response to SQL."""
        response = chat_response.message.content
        sql_query_start = response.find("SQLQuery:")
        if sql_query_start != -1:
            response = response[sql_query_start:]
            # TODO: move to removeprefix after Python 3.9+
            if response.startswith("SQLQuery:"):
                response = response[len("SQLQuery:") :]
        sql_result_start = response.find("SQLResult:")
        if sql_result_start != -1:
            response = response[:sql_result_start]
        return response.strip().strip("```").strip()
    
    
    text2sql_prompt = DEFAULT_TEXT_TO_SQL_PROMPT.partial_format(
        dialect=engine.dialect.name
    )
    print(text2sql_prompt.template)

  1. 创建自然语言回复生成模板
response_synthesis_prompt_str = (
    "Given an input question, synthesize a response from the query results.\n"
    "Query: {query_str}\n"
    "SQL: {sql_query}\n"
    "SQL Response: {context_str}\n"
    "Response: "
)
response_synthesis_prompt = PromptTemplate(
    response_synthesis_prompt_str,
)

10.5 定义工作流

from llama_index.core.workflow import (
    Workflow,
    StartEvent,
    StopEvent,
    step,
    Context,
    Event,
)

# 事件:找到数据库中相关的表
class TableRetrieveEvent(Event):
    """Result of running table retrieval."""

    table_context_str: str
    query: str

# 事件:文本转 SQL
class TextToSQLEvent(Event):
    """Text-to-SQL event."""

    sql: str
    query: str


class TextToSQLWorkflow1(Workflow):
    """Text-to-SQL Workflow that does query-time table retrieval."""

    def __init__(
        self,
        obj_retriever,
        text2sql_prompt,
        sql_retriever,
        response_synthesis_prompt,
        llm,
        *args,
        **kwargs
    ) -> None:
        """Init params."""
        super().__init__(*args, **kwargs)
        self.obj_retriever = obj_retriever
        self.text2sql_prompt = text2sql_prompt
        self.sql_retriever = sql_retriever
        self.response_synthesis_prompt = response_synthesis_prompt
        self.llm = llm

    @step
    def retrieve_tables(
        self, ctx: Context, ev: StartEvent
    ) -> TableRetrieveEvent:
        """Retrieve tables."""
        table_schema_objs = self.obj_retriever.retrieve(ev.query)
        table_context_str = get_table_context_str(table_schema_objs)
        print("====\n"+table_context_str+"\n====")
        return TableRetrieveEvent(
            table_context_str=table_context_str, query=ev.query
        )

    @step
    def generate_sql(
        self, ctx: Context, ev: TableRetrieveEvent
    ) -> TextToSQLEvent:
        """Generate SQL statement."""
        fmt_messages = self.text2sql_prompt.format_messages(
            query_str=ev.query, schema=ev.table_context_str
        )
        chat_response = self.llm.chat(fmt_messages)
        sql = parse_response_to_sql(chat_response)
        print("====\n"+sql+"\n====")
        return TextToSQLEvent(sql=sql, query=ev.query)

    @step
    def generate_response(self, ctx: Context, ev: TextToSQLEvent) -> StopEvent:
        """Run SQL retrieval and generate response."""
        retrieved_rows = self.sql_retriever.retrieve(ev.sql)
        print("====\n"+str(retrieved_rows)+"\n====")
        fmt_messages = self.response_synthesis_prompt.format_messages(
            sql_query=ev.sql,
            context_str=str(retrieved_rows),
            query_str=ev.query,
        )
        chat_response = llm.chat(fmt_messages)
        return StopEvent(result=chat_response)
workflow = TextToSQLWorkflow1(
    obj_retriever,
    text2sql_prompt,
    sql_retriever,
    response_synthesis_prompt,
    llm,
    verbose=True,
)
response = await workflow.run(
    query="What was the year that The Notorious B.I.G was signed to Bad Boy?"
)
print(str(response))

10.6 可视化工作流

pip install llama-index-utils-workflow

from llama_index.utils.workflow import draw_all_possible_flows

draw_all_possible_flows(
    TextToSQLWorkflow1, filename="text_to_sql_table_retrieval.html"
)

10.7 工作流管理框架意义是什么

思考以下情况:

  • step 的执行顺序有逻辑分支
  • step 的执行有循环
  • step 的执行可以并行
  • 一个 step 的触发条件依赖前面若干 step 的结果,且它们之间可能有循环或者并行

所以,工作流管理框架的意思是便于将单个事件的处理逻辑和事件之间的执行顺序独立开

关于 LlamaIndex 工作流的更详细文档:https://docs.llamaindex.ai/en/stable/examples/workflow/workflows_cookbook/

11. LlamaIndex 的更多功能

以上内容涉及较多背景知识,暂时不在本课展开,相关知识会在后面课程中逐一详细讲解。

此外,LlamaIndex 针对生产级的 RAG 系统中遇到的各个方面的细节问题,总结了很多高端技巧(Advanced Topics),对实战很有参考价值,非常推荐有能力的同学阅读。

完整正确代码

"""
基于LlamaIndex和阿里百炼的Text-to-SQL工作流实现
本代码实现了一个完整的表格问答系统,支持从CSV文件加载数据、生成表格描述、
存储到SQLite数据库,并通过自然语言查询生成SQL语句并执行
"""

# ========================================
# 1. 标准库导入
# ========================================
import os
import json
import re
import asyncio
from pathlib import Path

# ========================================
# 2. 第三方库导入
# ========================================
import pandas as pd
from sqlalchemy import create_engine, MetaData, Table, Column, String, Integer

# ========================================
# 3. LlamaIndex相关导入
# ========================================
from llama_index.core import Settings, SQLDatabase, VectorStoreIndex
from llama_index.core.objects import (
    SQLTableNodeMapping,
    ObjectIndex,
    SQLTableSchema,
)
from llama_index.core.prompts import ChatPromptTemplate
from llama_index.core.bridge.pydantic import BaseModel, Field
from llama_index.core.llms import ChatMessage
from llama_index.core.retrievers import SQLRetriever
from llama_index.core.prompts.default_prompts import DEFAULT_TEXT_TO_SQL_PROMPT
from llama_index.core import PromptTemplate
from llama_index.core.llms import ChatResponse
from llama_index.core.workflow import (
    Workflow,
    StartEvent,
    StopEvent,
    step,
    Context,
    Event,
)
from llama_index.utils.workflow import draw_all_possible_flows

# ========================================
# 4. 阿里百炼模型导入
# ========================================
from llama_index.llms.dashscope import DashScope, DashScopeGenerationModels
from llama_index.embeddings.dashscope import DashScopeEmbedding, DashScopeTextEmbeddingModels

# ========================================
# 5. 数据加载模块:遍历目录加载CSV表格
# ========================================
# 定义数据目录路径
data_dir = Path("WikiTableQuestions/csv/200-csv")

# 获取所有CSV文件并按文件名排序
csv_files = sorted([f for f in data_dir.glob("*.csv")])

# 加载所有CSV文件到DataFrame列表
dfs = []
for csv_file in csv_files:
    print(f"处理文件中: {csv_file}")
    try:
        df = pd.read_csv(csv_file)
        dfs.append(df)
    except Exception as e:
        print(f"解析文件 {csv_file} 时出错: {str(e)}")


# ========================================
# 6. 表格信息描述生成模块
# ========================================
# 定义表格信息数据模型
class TableInfo(BaseModel):
    """表格信息的结构化定义"""
    table_name: str = Field(
        ..., description="表格名称(必须使用下划线,不能有空格)"
    )
    table_summary: str = Field(
        ..., description="表格的简短、简洁摘要/描述"
    )


# 定义生成表格描述的提示词模板
prompt_str = """
请为以下表格提供一个摘要,使用以下JSON格式输出。

- 表格名称必须唯一且描述准确,同时要简洁。
- 不要输出通用的表格名称(例如table、my_table)。

请不要使用以下表格名称: {exclude_table_name_list}

表格:
{table_str}

摘要: """

# 创建聊天提示词模板
prompt_tmpl = ChatPromptTemplate(
    message_templates=[ChatMessage.from_str(prompt_str, role="user")]
)

# 表格信息保存目录
tableinfo_dir = "WikiTableQuestions_TableInfo"


# ========================================
# 7. 表格信息缓存管理函数
# ========================================
def _get_tableinfo_with_index(idx: int) -> TableInfo:
    """
    根据索引从缓存目录获取已保存的表格信息

    参数:
        idx: 表格索引

    返回:
        TableInfo对象,如果不存在则返回None
    """
    results_gen = Path(tableinfo_dir).glob(f"{idx}_*")
    results_list = list(results_gen)

    if len(results_list) == 0:
        return None
    elif len(results_list) == 1:
        path = results_list[0]
        with open(path, 'r', encoding='utf-8') as file:
            data = json.load(file)
            return TableInfo.model_validate(data)
    else:
        raise ValueError(
            f"找到多个匹配索引 {idx} 的文件: {list(results_gen)}"
        )


# ========================================
# 8. 阿里百炼模型初始化
# ========================================
# 设置阿里百炼API密钥
os.environ["DASHSCOPE_API_KEY"] = "sk-dab4c8deb19xxxxc13772a"  # 替换为你的API Key

# 初始化LLM模型
llm = DashScope(model="qwen-plus", api_key=os.getenv("DASHSCOPE_API_KEY"))

# ========================================
# 9. 表格信息生成主流程
# ========================================
table_names = set()  # 用于确保表格名称唯一
table_infos = []  # 存储所有表格信息

for idx, df in enumerate(dfs):
    # 尝试从缓存加载表格信息
    table_info = _get_tableinfo_with_index(idx)

    if table_info:
        # 如果缓存存在,直接使用
        table_infos.append(table_info)
    else:
        # 如果缓存不存在,使用LLM生成
        while True:
            # 提取前10行数据作为表格样本
            df_str = df.head(10).to_csv()

            # 使用LLM生成表格信息
            table_info = llm.structured_predict(
                TableInfo,
                prompt_tmpl,
                table_str=df_str,
                exclude_table_name_list=str(list(table_names)),
            )

            table_name = table_info.table_name
            print(f"已处理表格: {table_name}")

            # 确保表格名称唯一
            if table_name not in table_names:
                table_names.add(table_name)
                break
            else:
                # 如果名称重复,重新生成
                print(f"表格名称 {table_name} 已存在,重新生成...")

        # 保存生成的表格信息到文件
        out_file = f"{tableinfo_dir}/{idx}_{table_name}.json"
        json.dump(table_info.dict(), open(out_file, "w", encoding='utf-8'))

    table_infos.append(table_info)


# ========================================
# 10. 表格数据存储到SQLite数据库
# ========================================
def sanitize_column_name(col_name: str) -> str:
    """
    清理列名,移除特殊字符并用下划线替换空格

    参数:
        col_name: 原始列名

    返回:
        清理后的列名
    """
    return re.sub(r"\W+", "_", col_name)


def create_table_from_dataframe(df: pd.DataFrame, table_name: str, engine, metadata_obj):
    """
    将DataFrame数据存储到SQLAlchemy表中

    参数:
        df: 数据框
        table_name: 表名
        engine: SQLAlchemy引擎
        metadata_obj: 元数据对象
    """
    # 清理列名
    sanitized_columns = {col: sanitize_column_name(col) for col in df.columns}
    df = df.rename(columns=sanitized_columns)

    # 动态创建列,根据数据类型选择String或Integer
    columns = [
        Column(col, String if dtype == "object" else Integer)
        for col, dtype in zip(df.columns, df.dtypes)
    ]

    # 创建表
    table = Table(table_name, metadata_obj, *columns)

    # 在数据库中创建表
    metadata_obj.create_all(engine)

    # 将数据插入表中
    with engine.connect() as conn:
        for _, row in df.iterrows():
            insert_stmt = table.insert().values(**row.to_dict())
            conn.execute(insert_stmt)
        conn.commit()


# 创建SQLite数据库引擎
engine = create_engine("sqlite:///wiki_table_questions.db")
metadata_obj = MetaData()

# 将所有DataFrame存储到数据库
for idx, df in enumerate(dfs):
    tableinfo = _get_tableinfo_with_index(idx)
    if tableinfo:
        print(f"创建表格: {tableinfo.table_name}")
        create_table_from_dataframe(df, tableinfo.table_name, engine, metadata_obj)

# ========================================
# 11. 创建基于表格描述的向量索引
# ========================================
# 设置全局LLM和嵌入模型
Settings.llm = DashScope(model_name=DashScopeGenerationModels.QWEN_MAX,
                         api_key=os.getenv("DASHSCOPE_API_KEY"))
Settings.embed_model = DashScopeEmbedding(
    model_name=DashScopeTextEmbeddingModels.TEXT_EMBEDDING_V1
)

# 创建SQL数据库对象
sql_database = SQLDatabase(engine)

# 创建表格节点映射和表格模式对象
table_node_mapping = SQLTableNodeMapping(sql_database)
table_schema_objs = [
    SQLTableSchema(table_name=t.table_name, context_str=t.table_summary)
    for t in table_infos
]

# 创建对象索引
obj_index = ObjectIndex.from_objects(
    table_schema_objs,
    table_node_mapping,
    VectorStoreIndex,
)
obj_retriever = obj_index.as_retriever(similarity_top_k=3)

# ========================================
# 12. SQL检索器创建
# ========================================
sql_retriever = SQLRetriever(sql_database)


def get_table_context_str(table_schema_objs: list) -> str:
    """
    获取表格上下文字符串

    参数:
        table_schema_objs: SQLTableSchema对象列表

    返回:
        拼接后的表格上下文字符串
    """
    context_strs = []
    for table_schema_obj in table_schema_objs:
        table_info = sql_database.get_single_table_info(
            table_schema_obj.table_name
        )
        if table_schema_obj.context_str:
            table_opt_context = " 表格描述为: "
            table_opt_context += table_schema_obj.context_str
            table_info += table_opt_context
        context_strs.append(table_info)
    return "\n\n".join(context_strs)


# ========================================
# 13. Text-to-SQL相关配置
# ========================================
def parse_response_to_sql(chat_response: ChatResponse) -> str:
    """
    从聊天响应中解析出SQL语句

    参数:
        chat_response: LLM的聊天响应

    返回:
        解析出的SQL语句
    """
    response = chat_response.message.content
    sql_query_start = response.find("SQLQuery:")

    if sql_query_start != -1:
        response = response[sql_query_start:]
        if response.startswith("SQLQuery:"):
            response = response[len("SQLQuery:"):]

    sql_result_start = response.find("SQLResult:")
    if sql_result_start != -1:
        response = response[:sql_result_start]

    return response.strip().strip("```").strip()


# 创建Text-to-SQL提示词模板
text2sql_prompt = DEFAULT_TEXT_TO_SQL_PROMPT.partial_format(
    dialect=engine.dialect.name
)

# ========================================
# 14. 自然语言回复生成模板
# ========================================
response_synthesis_prompt_str = (
    "根据输入的问题,从查询结果中合成一个回答。\n"
    "查询: {query_str}\n"
    "SQL: {sql_query}\n"
    "SQL响应: {context_str}\n"
    "回答: "
)
response_synthesis_prompt = PromptTemplate(
    response_synthesis_prompt_str,
)

# 重新初始化LLM(使用QWEN_MAX模型)
llm = DashScope(model_name=DashScopeGenerationModels.QWEN_MAX,
                api_key=os.getenv("DASHSCOPE_API_KEY"))


# ========================================
# 15. 工作流定义
# ========================================
# 事件定义:表格检索事件
class TableRetrieveEvent(Event):
    """表格检索结果事件"""
    table_context_str: str
    query: str


# 事件定义:Text-to-SQL事件
class TextToSQLEvent(Event):
    """文本转SQL事件"""
    sql: str
    query: str


class TextToSQLWorkflow1(Workflow):
    """Text-to-SQL工作流,支持查询时的表格检索"""

    def __init__(
            self,
            obj_retriever,
            text2sql_prompt,
            sql_retriever,
            response_synthesis_prompt,
            llm,
            *args,
            **kwargs
    ) -> None:
        """初始化参数"""
        super().__init__(*args, **kwargs)
        self.obj_retriever = obj_retriever
        self.text2sql_prompt = text2sql_prompt
        self.sql_retriever = sql_retriever
        self.response_synthesis_prompt = response_synthesis_prompt
        self.llm = llm

    @step
    def retrieve_tables(self, ctx: Context, ev: StartEvent) -> TableRetrieveEvent:
        """检索相关表格"""
        table_schema_objs = self.obj_retriever.retrieve(ev.query)
        table_context_str = get_table_context_str(table_schema_objs)
        print("====\n" + table_context_str + "\n====")
        return TableRetrieveEvent(
            table_context_str=table_context_str, query=ev.query
        )

    @step
    def generate_sql(self, ctx: Context, ev: TableRetrieveEvent) -> TextToSQLEvent:
        """生成SQL语句"""
        fmt_messages = self.text2sql_prompt.format_messages(
            query_str=ev.query, schema=ev.table_context_str
        )
        chat_response = self.llm.chat(fmt_messages)
        sql = parse_response_to_sql(chat_response)
        print("====\n" + sql + "\n====")
        return TextToSQLEvent(sql=sql, query=ev.query)

    @step
    def generate_response(self, ctx: Context, ev: TextToSQLEvent) -> StopEvent:
        """执行SQL检索并生成自然语言回答"""
        retrieved_rows = self.sql_retriever.retrieve(ev.sql)
        print("====\n" + str(retrieved_rows) + "\n====")
        fmt_messages = self.response_synthesis_prompt.format_messages(
            sql_query=ev.sql,
            context_str=str(retrieved_rows),
            query_str=ev.query,
        )
        chat_response = llm.chat(fmt_messages)
        return StopEvent(result=chat_response)


# ========================================
# 16. 异步主函数
# ========================================
async def main():
    """主异步函数,执行完整的工作流"""
    workflow = TextToSQLWorkflow1(
        obj_retriever,
        text2sql_prompt,
        sql_retriever,
        response_synthesis_prompt,
        llm,
        verbose=True,
    )

    # ========================================
    # 17. 工作流可视化
    # ========================================
    # 修正:draw_all_possible_flows 需要传入工作流实例,而不是类
    draw_all_possible_flows(
        workflow,  # 传入实例而不是类
        filename="text_to_sql_table_retrieval.html"
    )

    # 执行示例查询
    response = await workflow.run(
        query="What was the year that The Notorious B.I.G was signed to Bad Boy?"
    )

    print(str(response))
    return response  # 返回结果以便检查


# ========================================
# 18. 程序入口点
# ========================================
if __name__ == "__main__":
    result = asyncio.run(main())

Logo

欢迎加入DeepSeek 技术社区。在这里,你可以找到志同道合的朋友,共同探索AI技术的奥秘。

更多推荐