Initial commit

This commit is contained in:
2026-05-12 09:41:56 +08:00
commit 572283e101
936 changed files with 133949 additions and 0 deletions
+209
View File
@@ -0,0 +1,209 @@
import json
import os
import numpy as np
from pymilvus import connections, MilvusClient, FieldSchema, CollectionSchema, DataType, Collection, AnnSearchRequest, RRFRanker
from pymilvus.model.hybrid import BGEM3EmbeddingFunction
# 1. 初始化设置
COLLECTION_NAME = "dragon_hybrid_demo"
MILVUS_URI = "http://localhost:19530" # 服务器模式
DATA_PATH = "../../data/C4/metadata/dragon.json" # 相对路径
BATCH_SIZE = 50
# 2. 连接 Milvus 并初始化嵌入模型
print(f"--> 正在连接到 Milvus: {MILVUS_URI}")
connections.connect(uri=MILVUS_URI)
print("--> 正在初始化 BGE-M3 嵌入模型...")
ef = BGEM3EmbeddingFunction(use_fp16=False, device="cpu")
print(f"--> 嵌入模型初始化完成。密集向量维度: {ef.dim['dense']}")
# 3. 创建 Collection
milvus_client = MilvusClient(uri=MILVUS_URI)
if milvus_client.has_collection(COLLECTION_NAME):
print(f"--> 正在删除已存在的 Collection '{COLLECTION_NAME}'...")
milvus_client.drop_collection(COLLECTION_NAME)
fields = [
FieldSchema(name="pk", dtype=DataType.VARCHAR, is_primary=True, auto_id=True, max_length=100),
FieldSchema(name="img_id", dtype=DataType.VARCHAR, max_length=100),
FieldSchema(name="path", dtype=DataType.VARCHAR, max_length=256),
FieldSchema(name="title", dtype=DataType.VARCHAR, max_length=256),
FieldSchema(name="description", dtype=DataType.VARCHAR, max_length=4096),
FieldSchema(name="category", dtype=DataType.VARCHAR, max_length=64),
FieldSchema(name="location", dtype=DataType.VARCHAR, max_length=128),
FieldSchema(name="environment", dtype=DataType.VARCHAR, max_length=64),
FieldSchema(name="sparse_vector", dtype=DataType.SPARSE_FLOAT_VECTOR),
FieldSchema(name="dense_vector", dtype=DataType.FLOAT_VECTOR, dim=ef.dim["dense"])
]
# 如果集合不存在,则创建它及索引
if not milvus_client.has_collection(COLLECTION_NAME):
print(f"--> 正在创建 Collection '{COLLECTION_NAME}'...")
schema = CollectionSchema(fields, description="关于龙的混合检索示例")
# 创建集合
collection = Collection(name=COLLECTION_NAME, schema=schema, consistency_level="Strong")
print("--> Collection 创建成功。")
# 4. 创建索引
print("--> 正在为新集合创建索引...")
sparse_index = {"index_type": "SPARSE_INVERTED_INDEX", "metric_type": "IP"}
collection.create_index("sparse_vector", sparse_index)
print("稀疏向量索引创建成功。")
dense_index = {"index_type": "AUTOINDEX", "metric_type": "IP"}
collection.create_index("dense_vector", dense_index)
print("密集向量索引创建成功。")
collection = Collection(COLLECTION_NAME)
# 5. 加载数据并插入
collection.load()
print(f"--> Collection '{COLLECTION_NAME}' 已加载到内存。")
if collection.is_empty:
print(f"--> Collection 为空,开始插入数据...")
if not os.path.exists(DATA_PATH):
raise FileNotFoundError(f"数据文件未找到: {DATA_PATH}")
with open(DATA_PATH, 'r', encoding='utf-8') as f:
dataset = json.load(f)
docs, metadata = [], []
for item in dataset:
parts = [
item.get('title', ''),
item.get('description', ''),
item.get('location', ''),
item.get('environment', ''),
# *item.get('combat_details', {}).get('combat_style', []),
# *item.get('combat_details', {}).get('abilities_used', []),
# item.get('scene_info', {}).get('time_of_day', '')
]
docs.append(' '.join(filter(None, parts)))
metadata.append(item)
print(f"--> 数据加载完成,共 {len(docs)} 条。")
print("--> 正在生成向量嵌入...")
embeddings = ef(docs)
print("--> 向量生成完成。")
print("--> 正在分批插入数据...")
# 为每个字段准备批量数据
img_ids = [doc["img_id"] for doc in metadata]
paths = [doc["path"] for doc in metadata]
titles = [doc["title"] for doc in metadata]
descriptions = [doc["description"] for doc in metadata]
categories = [doc["category"] for doc in metadata]
locations = [doc["location"] for doc in metadata]
environments = [doc["environment"] for doc in metadata]
# 获取向量
sparse_vectors = embeddings["sparse"]
dense_vectors = embeddings["dense"]
# 插入数据
collection.insert([
img_ids,
paths,
titles,
descriptions,
categories,
locations,
environments,
sparse_vectors,
dense_vectors
])
collection.flush()
print(f"--> 数据插入完成,总数: {collection.num_entities}")
else:
print(f"--> Collection 中已有 {collection.num_entities} 条数据,跳过插入。")
# 6. 执行搜索
search_query = "悬崖上的巨龙"
search_filter = 'category in ["western_dragon", "chinese_dragon", "movie_character"]'
top_k = 5
print(f"\n{'='*20} 开始混合搜索 {'='*20}")
print(f"查询: '{search_query}'")
print(f"过滤器: '{search_filter}'")
query_embeddings = ef([search_query])
dense_vec = query_embeddings["dense"][0]
sparse_vec = query_embeddings["sparse"]._getrow(0)
# 打印向量信息
print("\n=== 向量信息 ===")
print(f"密集向量维度: {len(dense_vec)}")
print(f"密集向量前5个元素: {dense_vec[:5]}")
print(f"密集向量范数: {np.linalg.norm(dense_vec):.4f}")
print(f"\n稀疏向量维度: {sparse_vec.shape[1]}")
print(f"稀疏向量非零元素数量: {sparse_vec.nnz}")
print("稀疏向量前5个非零元素:")
for i in range(min(5, sparse_vec.nnz)):
print(f" - 索引: {sparse_vec.indices[i]}, 值: {sparse_vec.data[i]:.4f}")
density = (sparse_vec.nnz / sparse_vec.shape[1] * 100)
print(f"\n稀疏向量密度: {density:.8f}%")
# 定义搜索参数
search_params = {"metric_type": "IP", "params": {}}
# 先执行单独的搜索
print("\n--- [单独] 密集向量搜索结果 ---")
dense_results = collection.search(
[dense_vec],
anns_field="dense_vector",
param=search_params,
limit=top_k,
expr=search_filter,
output_fields=["title", "path", "description", "category", "location", "environment"]
)[0]
for i, hit in enumerate(dense_results):
print(f"{i+1}. {hit.entity.get('title')} (Score: {hit.distance:.4f})")
print(f" 路径: {hit.entity.get('path')}")
print(f" 描述: {hit.entity.get('description')[:100]}...")
print("\n--- [单独] 稀疏向量搜索结果 ---")
sparse_results = collection.search(
[sparse_vec],
anns_field="sparse_vector",
param=search_params,
limit=top_k,
expr=search_filter,
output_fields=["title", "path", "description", "category", "location", "environment"]
)[0]
for i, hit in enumerate(sparse_results):
print(f"{i+1}. {hit.entity.get('title')} (Score: {hit.distance:.4f})")
print(f" 路径: {hit.entity.get('path')}")
print(f" 描述: {hit.entity.get('description')[:100]}...")
print("\n--- [混合] 稀疏+密集向量搜索结果 ---")
# 创建 RRF 融合器
rerank = RRFRanker(k=60)
# 创建搜索请求
dense_req = AnnSearchRequest([dense_vec], "dense_vector", search_params, limit=top_k)
sparse_req = AnnSearchRequest([sparse_vec], "sparse_vector", search_params, limit=top_k)
# 执行混合搜索
results = collection.hybrid_search(
[sparse_req, dense_req],
rerank=rerank,
limit=top_k,
output_fields=["title", "path", "description", "category", "location", "environment"]
)[0]
# 打印最终结果
for i, hit in enumerate(results):
print(f"{i+1}. {hit.entity.get('title')} (Score: {hit.distance:.4f})")
print(f" 路径: {hit.entity.get('path')}")
print(f" 描述: {hit.entity.get('description')[:100]}...")
# 7. 清理资源
milvus_client.release_collection(collection_name=COLLECTION_NAME)
print(f"已从内存中释放 Collection: '{COLLECTION_NAME}'")
milvus_client.drop_collection(COLLECTION_NAME)
print(f"已删除 Collection: '{COLLECTION_NAME}'")
+329
View File
@@ -0,0 +1,329 @@
import json
import os
import numpy as np
import torch
from transformers import AutoModel, AutoProcessor
from sklearn.feature_extraction.text import TfidfVectorizer
from scipy.sparse import csr_matrix
from pymilvus import connections, MilvusClient, FieldSchema, CollectionSchema, DataType, Collection, AnnSearchRequest, RRFRanker
# 1. 初始化设置
COLLECTION_NAME = "dragon_siglip_demo"
MILVUS_URI = "http://localhost:19530" # 服务器模式
DATA_PATH = "../../data/C4/metadata/dragon.json" # 相对路径
BATCH_SIZE = 50
# 2. 自定义SigLIP嵌入函数类
class SigLIPEmbeddingFunction:
def __init__(self, model_name="google/siglip-base-patch16-256-multilingual", device="cpu"):
"""
初始化SigLIP嵌入函数
Args:
model_name: SigLIP模型名称
device: 设备类型 ("cpu""cuda")
"""
self.model_name = model_name
self.device = device
print(f"--> 正在加载 SigLIP 模型: {model_name}")
self.model = AutoModel.from_pretrained(model_name)
self.processor = AutoProcessor.from_pretrained(model_name)
self.model.to(device)
self.model.eval()
# 初始化TF-IDF作为稀疏向量生成器
self.tfidf_vectorizer = TfidfVectorizer(
max_features=10000, # 限制词汇表大小以节省空间
stop_words='english',
ngram_range=(1, 2)
)
self.tfidf_fitted = False
# 获取文本编码器的输出维度
with torch.no_grad():
dummy_text = ["test"]
inputs = self.processor(text=dummy_text, padding="max_length", return_tensors="pt")
outputs = self.model.text_model(**{k: v.to(device) for k, v in inputs.items() if k != 'pixel_values'})
self.dense_dim = outputs.pooler_output.shape[-1]
print(f"--> SigLIP 模型加载完成。密集向量维度: {self.dense_dim}")
@property
def dim(self):
"""返回维度信息,兼容原BGE-M3接口"""
return {
"dense": self.dense_dim,
"sparse": self.tfidf_vectorizer.max_features if self.tfidf_fitted else 10000
}
def fit_sparse(self, docs):
"""拟合稀疏向量模型(TF-IDF"""
print("--> 正在拟合 TF-IDF 模型...")
self.tfidf_vectorizer.fit(docs)
self.tfidf_fitted = True
print(f"--> TF-IDF 模型拟合完成。词汇表大小: {len(self.tfidf_vectorizer.vocabulary_)}")
def encode_text_dense(self, texts):
"""使用SigLIP编码文本为密集向量"""
if isinstance(texts, str):
texts = [texts]
dense_vectors = []
batch_size = 8 # 减小批次大小以节省内存
with torch.no_grad():
for i in range(0, len(texts), batch_size):
batch_texts = texts[i:i + batch_size]
inputs = self.processor(text=batch_texts, padding="max_length", truncation=True, return_tensors="pt")
inputs = {k: v.to(self.device) for k, v in inputs.items() if k != 'pixel_values'}
outputs = self.model.text_model(**inputs)
embeddings = outputs.pooler_output
# 归一化向量
embeddings = torch.nn.functional.normalize(embeddings, p=2, dim=1)
dense_vectors.extend(embeddings.cpu().numpy())
return np.array(dense_vectors)
def encode_text_sparse(self, texts):
"""使用TF-IDF编码文本为稀疏向量"""
if not self.tfidf_fitted:
raise ValueError("请先调用 fit_sparse() 方法拟合TF-IDF模型")
if isinstance(texts, str):
texts = [texts]
sparse_matrix = self.tfidf_vectorizer.transform(texts)
return sparse_matrix
def __call__(self, texts):
"""主调用方法,返回密集和稀疏向量"""
if isinstance(texts, str):
texts = [texts]
# 如果还没有拟合稀疏模型,先拟合
if not self.tfidf_fitted:
self.fit_sparse(texts)
dense_vectors = self.encode_text_dense(texts)
sparse_vectors = self.encode_text_sparse(texts)
return {
"dense": dense_vectors,
"sparse": sparse_vectors
}
# 3. 连接 Milvus 并初始化嵌入模型
print(f"--> 正在连接到 Milvus: {MILVUS_URI}")
connections.connect(uri=MILVUS_URI)
print("--> 正在初始化 SigLIP 嵌入模型...")
ef = SigLIPEmbeddingFunction(device="cpu") # 如果有GPU可以改为"cuda"
# 4. 创建 Collection
milvus_client = MilvusClient(uri=MILVUS_URI)
if milvus_client.has_collection(COLLECTION_NAME):
print(f"--> 正在删除已存在的 Collection '{COLLECTION_NAME}'...")
milvus_client.drop_collection(COLLECTION_NAME)
fields = [
FieldSchema(name="pk", dtype=DataType.VARCHAR, is_primary=True, auto_id=True, max_length=100),
FieldSchema(name="img_id", dtype=DataType.VARCHAR, max_length=100),
FieldSchema(name="path", dtype=DataType.VARCHAR, max_length=256),
FieldSchema(name="title", dtype=DataType.VARCHAR, max_length=256),
FieldSchema(name="description", dtype=DataType.VARCHAR, max_length=4096),
FieldSchema(name="category", dtype=DataType.VARCHAR, max_length=64),
FieldSchema(name="location", dtype=DataType.VARCHAR, max_length=128),
FieldSchema(name="environment", dtype=DataType.VARCHAR, max_length=64),
FieldSchema(name="sparse_vector", dtype=DataType.SPARSE_FLOAT_VECTOR),
FieldSchema(name="dense_vector", dtype=DataType.FLOAT_VECTOR, dim=ef.dim["dense"])
]
# 如果集合不存在,则创建它及索引
if not milvus_client.has_collection(COLLECTION_NAME):
print(f"--> 正在创建 Collection '{COLLECTION_NAME}'...")
schema = CollectionSchema(fields, description="使用SigLIP的龙混合检索示例")
# 创建集合
collection = Collection(name=COLLECTION_NAME, schema=schema, consistency_level="Strong")
print("--> Collection 创建成功。")
# 5. 创建索引
print("--> 正在为新集合创建索引...")
sparse_index = {"index_type": "SPARSE_INVERTED_INDEX", "metric_type": "IP"}
collection.create_index("sparse_vector", sparse_index)
print("稀疏向量索引创建成功。")
dense_index = {"index_type": "AUTOINDEX", "metric_type": "IP"}
collection.create_index("dense_vector", dense_index)
print("密集向量索引创建成功。")
collection = Collection(COLLECTION_NAME)
# 6. 加载数据并插入
collection.load()
print(f"--> Collection '{COLLECTION_NAME}' 已加载到内存。")
if collection.is_empty:
print(f"--> Collection 为空,开始插入数据...")
if not os.path.exists(DATA_PATH):
raise FileNotFoundError(f"数据文件未找到: {DATA_PATH}")
with open(DATA_PATH, 'r', encoding='utf-8') as f:
dataset = json.load(f)
docs, metadata = [], []
for item in dataset:
parts = [
item.get('title', ''),
item.get('description', ''),
item.get('location', ''),
item.get('environment', ''),
# *item.get('combat_details', {}).get('combat_style', []),
# *item.get('combat_details', {}).get('abilities_used', []),
# item.get('scene_info', {}).get('time_of_day', '')
]
docs.append(' '.join(filter(None, parts)))
metadata.append(item)
print(f"--> 数据加载完成,共 {len(docs)} 条。")
print("--> 正在生成向量嵌入...")
embeddings = ef(docs)
print("--> 向量生成完成。")
print("--> 正在分批插入数据...")
# 为每个字段准备批量数据
img_ids = [doc["img_id"] for doc in metadata]
paths = [doc["path"] for doc in metadata]
titles = [doc["title"] for doc in metadata]
descriptions = [doc["description"] for doc in metadata]
categories = [doc["category"] for doc in metadata]
locations = [doc["location"] for doc in metadata]
environments = [doc["environment"] for doc in metadata]
# 获取向量 - 注意SigLIP返回的格式与BGE-M3不同
sparse_vectors = []
dense_vectors = embeddings["dense"].tolist()
# 将稀疏矩阵转换为Milvus可接受的格式
sparse_matrix = embeddings["sparse"]
for i in range(sparse_matrix.shape[0]):
row = sparse_matrix.getrow(i)
# 创建稀疏向量字典格式
sparse_dict = {}
for j in range(row.nnz):
sparse_dict[row.indices[j]] = float(row.data[j])
sparse_vectors.append(sparse_dict)
# 插入数据
collection.insert([
img_ids,
paths,
titles,
descriptions,
categories,
locations,
environments,
sparse_vectors,
dense_vectors
])
collection.flush()
print(f"--> 数据插入完成,总数: {collection.num_entities}")
else:
print(f"--> Collection 中已有 {collection.num_entities} 条数据,跳过插入。")
# 7. 执行搜索
search_query = "悬崖上的巨龙"
search_filter = 'category in ["western_dragon", "chinese_dragon", "movie_character"]'
top_k = 5
print(f"\n{'='*20} 开始混合搜索 {'='*20}")
print(f"查询: '{search_query}'")
print(f"过滤器: '{search_filter}'")
# 生成查询向量
query_embeddings = ef([search_query])
dense_vec = query_embeddings["dense"][0].tolist()
# 处理稀疏向量
sparse_matrix = query_embeddings["sparse"]
sparse_row = sparse_matrix.getrow(0)
sparse_dict = {}
for j in range(sparse_row.nnz):
sparse_dict[sparse_row.indices[j]] = float(sparse_row.data[j])
# 打印向量信息
print("\n=== 向量信息 ===")
print(f"密集向量维度: {len(dense_vec)}")
print(f"密集向量前5个元素: {dense_vec[:5]}")
print(f"密集向量范数: {np.linalg.norm(dense_vec):.4f}")
print(f"\n稀疏向量维度: {sparse_matrix.shape[1]}")
print(f"稀疏向量非零元素数量: {sparse_row.nnz}")
print("稀疏向量前5个非零元素:")
for i, (idx, val) in enumerate(list(sparse_dict.items())[:5]):
print(f" - 索引: {idx}, 值: {val:.4f}")
density = (sparse_row.nnz / sparse_matrix.shape[1] * 100)
print(f"\n稀疏向量密度: {density:.8f}%")
# 定义搜索参数
search_params = {"metric_type": "IP", "params": {}}
# 先执行单独的搜索
print("\n--- [单独] 密集向量搜索结果 ---")
dense_results = collection.search(
[dense_vec],
anns_field="dense_vector",
param=search_params,
limit=top_k,
expr=search_filter,
output_fields=["title", "path", "description", "category", "location", "environment"]
)[0]
for i, hit in enumerate(dense_results):
print(f"{i+1}. {hit.entity.get('title')} (Score: {hit.distance:.4f})")
print(f" 路径: {hit.entity.get('path')}")
print(f" 描述: {hit.entity.get('description')[:100]}...")
print("\n--- [单独] 稀疏向量搜索结果 ---")
sparse_results = collection.search(
[sparse_dict],
anns_field="sparse_vector",
param=search_params,
limit=top_k,
expr=search_filter,
output_fields=["title", "path", "description", "category", "location", "environment"]
)[0]
for i, hit in enumerate(sparse_results):
print(f"{i+1}. {hit.entity.get('title')} (Score: {hit.distance:.4f})")
print(f" 路径: {hit.entity.get('path')}")
print(f" 描述: {hit.entity.get('description')[:100]}...")
print("\n--- [混合] 稀疏+密集向量搜索结果 ---")
# 创建 RRF 融合器
rerank = RRFRanker(k=60)
# 创建搜索请求
dense_req = AnnSearchRequest([dense_vec], "dense_vector", search_params, limit=top_k)
sparse_req = AnnSearchRequest([sparse_dict], "sparse_vector", search_params, limit=top_k)
# 执行混合搜索
results = collection.hybrid_search(
[sparse_req, dense_req],
rerank=rerank,
limit=top_k,
output_fields=["title", "path", "description", "category", "location", "environment"]
)[0]
# 打印最终结果
for i, hit in enumerate(results):
print(f"{i+1}. {hit.entity.get('title')} (Score: {hit.distance:.4f})")
print(f" 路径: {hit.entity.get('path')}")
print(f" 描述: {hit.entity.get('description')[:100]}...")
# 8. 清理资源
milvus_client.release_collection(collection_name=COLLECTION_NAME)
print(f"已从内存中释放 Collection: '{COLLECTION_NAME}'")
milvus_client.drop_collection(COLLECTION_NAME)
print(f"已删除 Collection: '{COLLECTION_NAME}'")
+111
View File
@@ -0,0 +1,111 @@
import os
from langchain_deepseek import ChatDeepSeek
from langchain_community.document_loaders import BiliBiliLoader
from langchain.chains.query_constructor.base import AttributeInfo
from langchain.retrievers.self_query.base import SelfQueryRetriever
from langchain_community.vectorstores import Chroma
from langchain_huggingface import HuggingFaceEmbeddings
import logging
logging.basicConfig(level=logging.INFO)
# 1. 初始化视频数据
video_urls = [
"https://www.bilibili.com/video/BV1Bo4y1A7FU",
"https://www.bilibili.com/video/BV1ug4y157xA",
"https://www.bilibili.com/video/BV1yh411V7ge",
]
bili = []
try:
loader = BiliBiliLoader(video_urls=video_urls)
docs = loader.load()
for doc in docs:
original = doc.metadata
# 提取基本元数据字段
metadata = {
'title': original.get('title', '未知标题'),
'author': original.get('owner', {}).get('name', '未知作者'),
'source': original.get('bvid', '未知ID'),
'view_count': original.get('stat', {}).get('view', 0),
'length': original.get('duration', 0),
}
doc.metadata = metadata
bili.append(doc)
except Exception as e:
print(f"加载BiliBili视频失败: {str(e)}")
if not bili:
print("没有成功加载任何视频,程序退出")
exit()
# 2. 创建向量存储
embed_model = HuggingFaceEmbeddings(model_name="BAAI/bge-small-zh-v1.5")
vectorstore = Chroma.from_documents(bili, embed_model)
# 3. 配置元数据字段信息
metadata_field_info = [
AttributeInfo(
name="title",
description="视频标题(字符串)",
type="string",
),
AttributeInfo(
name="author",
description="视频作者(字符串)",
type="string",
),
AttributeInfo(
name="view_count",
description="视频观看次数(整数)",
type="integer",
),
AttributeInfo(
name="length",
description="视频长度(整数)",
type="integer"
)
]
# 4. 创建自查询检索器
llm = ChatDeepSeek(
model="deepseek-chat",
temperature=0,
api_key=os.getenv("DEEPSEEK_API_KEY")
)
retriever = SelfQueryRetriever.from_llm(
llm=llm,
vectorstore=vectorstore,
document_contents="记录视频标题、作者、观看次数等信息的视频元数据",
metadata_field_info=metadata_field_info,
enable_limit=True,
verbose=True
)
# 5. 执行查询示例
queries = [
"时间最短的视频",
"时长大于600秒的视频"
]
for query in queries:
print(f"\n--- 查询: '{query}' ---")
results = retriever.invoke(query)
if results:
for doc in results:
title = doc.metadata.get('title', '未知标题')
author = doc.metadata.get('author', '未知作者')
view_count = doc.metadata.get('view_count', '未知')
length = doc.metadata.get('length', '未知')
print(f"标题: {title}")
print(f"作者: {author}")
print(f"观看次数: {view_count}")
print(f"时长: {length}")
print("="*50)
else:
print("未找到匹配的视频")
+220
View File
@@ -0,0 +1,220 @@
import os
import sys
import sqlite3
# 添加text2sql模块路径
sys.path.append(os.path.join(os.path.dirname(__file__), 'text2sql'))
from text2sql.text2sql_agent import SimpleText2SQLAgent
def setup_demo():
"""设置演示环境"""
print("=== Text2SQL框架演示 ===\n")
# 检查API密钥
api_key = os.getenv("DEEPSEEK_API_KEY")
if not api_key:
print("先设置DEEPSEEK_API_KEY环境变量")
return None
# 创建演示数据库
print("创建演示数据库...")
db_path = create_demo_database()
# 初始化Text2SQL代理
print("初始化Text2SQL代理...")
agent = SimpleText2SQLAgent(api_key=api_key)
# 连接数据库
print("连接数据库...")
if not agent.connect_database(db_path):
print("数据库连接失败!")
return None
# 加载知识库
print("加载知识库...")
try:
agent.load_knowledge_base()
print("知识库加载成功!")
except Exception as e:
print(f"知识库加载失败: {str(e)}")
return None
return agent, db_path
def create_demo_database():
"""创建演示数据库"""
db_path = "text2sql_demo.db"
if os.path.exists(db_path):
os.remove(db_path)
conn = sqlite3.connect(db_path)
cursor = conn.cursor()
# 创建用户表
cursor.execute("""
CREATE TABLE users (
id INTEGER PRIMARY KEY,
name TEXT NOT NULL,
email TEXT UNIQUE,
age INTEGER,
city TEXT
)
""")
# 创建产品表
cursor.execute("""
CREATE TABLE products (
id INTEGER PRIMARY KEY,
name TEXT NOT NULL,
category TEXT,
price REAL,
stock INTEGER
)
""")
# 创建订单表
cursor.execute("""
CREATE TABLE orders (
id INTEGER PRIMARY KEY,
user_id INTEGER,
product_id INTEGER,
quantity INTEGER,
order_date TEXT,
total_price REAL,
FOREIGN KEY (user_id) REFERENCES users(id),
FOREIGN KEY (product_id) REFERENCES products(id)
)
""")
# 插入示例数据
users_data = [
(1, '张三', 'zhangsan@email.com', 25, '北京'),
(2, '李四', 'lisi@email.com', 32, '上海'),
(3, '王五', 'wangwu@email.com', 28, '广州'),
(4, '赵六', 'zhaoliu@email.com', 35, '深圳'),
(5, '陈七', 'chenqi@email.com', 29, '杭州'),
]
products_data = [
(1, 'iPhone 15', '电子产品', 7999.0, 50),
(2, 'MacBook Pro', '电子产品', 12999.0, 20),
(3, 'Nike运动鞋', '服装', 599.0, 100),
(4, '办公椅', '家具', 899.0, 30),
(5, '台灯', '家具', 199.0, 80),
(6, 'iPad', '电子产品', 3999.0, 40),
(7, 'Adidas外套', '服装', 399.0, 60),
]
orders_data = [
(1, 1, 1, 1, '2024-01-15', 7999.0),
(2, 2, 3, 2, '2024-01-16', 1198.0),
(3, 3, 5, 1, '2024-01-17', 199.0),
(4, 1, 2, 1, '2024-01-18', 12999.0),
(5, 4, 4, 1, '2024-01-19', 899.0),
(6, 5, 6, 1, '2024-01-20', 3999.0),
(7, 2, 7, 1, '2024-01-21', 399.0),
]
cursor.executemany("INSERT INTO users VALUES (?, ?, ?, ?, ?)", users_data)
cursor.executemany("INSERT INTO products VALUES (?, ?, ?, ?, ?)", products_data)
cursor.executemany("INSERT INTO orders VALUES (?, ?, ?, ?, ?, ?)", orders_data)
conn.commit()
conn.close()
print(f"演示数据库已创建: {db_path}")
return db_path
def run_demo_queries(agent):
"""运行演示查询"""
demo_questions = [
"查询所有用户的姓名和邮箱",
"年龄大于30的用户有哪些",
"哪些产品的库存少于50",
"查询来自北京的用户的所有订单",
"统计每个城市的用户数量",
"查询价格在500-8000之间的产品"
]
print("\n开始运行演示查询...\n")
success_count = 0
for i, question in enumerate(demo_questions, 1):
print(f"问题 {i}: {question}")
print("-" * 60)
try:
result = agent.query(question)
if result["success"]:
print(f"成功! SQL: {result['sql']}")
if isinstance(result["results"], dict) and "rows" in result["results"]:
count = result["results"]["count"]
print(f"返回 {count} 行数据")
# 显示前2行数据
if count > 0:
for j, row in enumerate(result["results"]["rows"][:2]):
row_str = " | ".join(f"{k}: {v}" for k, v in row.items())
print(f" {j+1}. {row_str}")
if count > 2:
print(f" ... 还有 {count - 2}")
else:
print(f"结果: {result['results']}")
success_count += 1
else:
print(f"失败: {result['error']}")
print(f"SQL: {result['sql']}")
except Exception as e:
print(f"执行错误: {str(e)}")
print()
# 输出统计
total_count = len(demo_questions)
def cleanup(agent, db_path):
"""清理资源"""
print("\n清理资源...")
if agent:
agent.cleanup()
if os.path.exists(db_path):
os.remove(db_path)
print(f"已删除演示数据库: {db_path}")
def main():
"""主函数"""
# 设置演示环境
setup_result = setup_demo()
if setup_result is None:
return
agent, db_path = setup_result
try:
# 运行演示查询
run_demo_queries(agent)
finally:
# 清理资源
cleanup(agent, db_path)
if __name__ == "__main__":
main()
+377
View File
@@ -0,0 +1,377 @@
import os
import json
import sqlite3
import numpy as np
from typing import List, Dict, Any
from sentence_transformers import SentenceTransformer
from sklearn.feature_extraction.text import TfidfVectorizer
from sklearn.metrics.pairwise import cosine_similarity
from pymilvus import connections, MilvusClient, FieldSchema, CollectionSchema, DataType, Collection
class BGESmallEmbeddingFunction:
"""BGE-Small中文嵌入函数,用于Text2SQL知识库向量化"""
def __init__(self, model_name="BAAI/bge-small-zh-v1.5", device="cpu"):
self.model_name = model_name
self.device = device
self.model = SentenceTransformer(model_name, device=device)
self.dense_dim = self.model.get_sentence_embedding_dimension()
def encode_text(self, texts):
"""编码文本为密集向量"""
if isinstance(texts, str):
texts = [texts]
embeddings = self.model.encode(
texts,
normalize_embeddings=True,
batch_size=16,
convert_to_numpy=True
)
return embeddings
@property
def dim(self):
"""返回向量维度"""
return self.dense_dim
class SimpleKnowledgeBase:
"""简化的知识库,使用BGE-Small进行向量检索"""
def __init__(self, milvus_uri: str = "http://localhost:19530"):
self.milvus_uri = milvus_uri
self.collection_name = "text2sql_knowledge_base"
self.milvus_client = None
self.collection = None
self.embedding_function = BGESmallEmbeddingFunction(
model_name="BAAI/bge-small-zh-v1.5",
device="cpu"
)
self.sql_examples = []
self.table_schemas = []
self.data_loaded = False
def connect_milvus(self):
"""连接Milvus数据库"""
connections.connect(uri=self.milvus_uri)
self.milvus_client = MilvusClient(uri=self.milvus_uri)
return True
def create_collection(self):
"""创建Milvus集合"""
if not self.milvus_client:
self.connect_milvus()
if self.milvus_client.has_collection(self.collection_name):
self.milvus_client.drop_collection(self.collection_name)
fields = [
FieldSchema(name="pk", dtype=DataType.VARCHAR, is_primary=True, auto_id=True, max_length=100),
FieldSchema(name="content_type", dtype=DataType.VARCHAR, max_length=50),
FieldSchema(name="question", dtype=DataType.VARCHAR, max_length=1000),
FieldSchema(name="sql", dtype=DataType.VARCHAR, max_length=2000),
FieldSchema(name="description", dtype=DataType.VARCHAR, max_length=1000),
FieldSchema(name="table_name", dtype=DataType.VARCHAR, max_length=100),
FieldSchema(name="embedding", dtype=DataType.FLOAT_VECTOR, dim=self.embedding_function.dim)
]
schema = CollectionSchema(fields, description="Text2SQL知识库")
self.collection = Collection(name=self.collection_name, schema=schema, consistency_level="Strong")
index_params = {"index_type": "AUTOINDEX", "metric_type": "IP", "params": {}}
self.collection.create_index("embedding", index_params)
return True
def load_data(self):
"""加载知识库数据"""
data_dir = os.path.join(os.path.dirname(__file__), "data")
self.load_sql_examples(data_dir)
self.load_table_schemas(data_dir)
self.vectorize_and_store()
self.data_loaded = True
def load_sql_examples(self, data_dir: str):
"""加载SQL示例"""
sql_examples_path = os.path.join(data_dir, "qsql_examples.json")
default_examples = [
{"question": "查询所有用户信息", "sql": "SELECT * FROM users", "description": "获取用户记录", "database": "sqlite"},
{"question": "年龄大于30的用户", "sql": "SELECT * FROM users WHERE age > 30", "description": "年龄筛选", "database": "sqlite"},
{"question": "统计用户总数", "sql": "SELECT COUNT(*) as user_count FROM users", "description": "用户计数", "database": "sqlite"},
{"question": "查询库存不足的产品", "sql": "SELECT * FROM products WHERE stock < 50", "description": "库存筛选", "database": "sqlite"},
{"question": "查询用户订单信息", "sql": "SELECT u.name, p.name, o.quantity FROM orders o JOIN users u ON o.user_id = u.id JOIN products p ON o.product_id = p.id", "description": "订单详情", "database": "sqlite"},
{"question": "按城市统计用户", "sql": "SELECT city, COUNT(*) as count FROM users GROUP BY city", "description": "城市分组", "database": "sqlite"}
]
if os.path.exists(sql_examples_path):
with open(sql_examples_path, 'r', encoding='utf-8') as f:
self.sql_examples = json.load(f)
else:
self.sql_examples = default_examples
os.makedirs(data_dir, exist_ok=True)
with open(sql_examples_path, 'w', encoding='utf-8') as f:
json.dump(self.sql_examples, f, ensure_ascii=False, indent=2)
def load_table_schemas(self, data_dir: str):
"""加载表结构信息"""
schema_path = os.path.join(data_dir, "table_schemas.json")
default_schemas = [
{
"table_name": "users",
"description": "用户信息表",
"columns": [
{"name": "id", "type": "INTEGER", "description": "用户ID"},
{"name": "name", "type": "VARCHAR", "description": "用户姓名"},
{"name": "age", "type": "INTEGER", "description": "用户年龄"},
{"name": "email", "type": "VARCHAR", "description": "邮箱地址"},
{"name": "city", "type": "VARCHAR", "description": "所在城市"},
{"name": "created_at", "type": "DATETIME", "description": "创建时间"}
]
},
{
"table_name": "products",
"description": "产品信息表",
"columns": [
{"name": "id", "type": "INTEGER", "description": "产品ID"},
{"name": "product_name", "type": "VARCHAR", "description": "产品名称"},
{"name": "category", "type": "VARCHAR", "description": "产品类别"},
{"name": "price", "type": "DECIMAL", "description": "产品价格"},
{"name": "stock", "type": "INTEGER", "description": "库存数量"},
{"name": "description", "type": "TEXT", "description": "产品描述"}
]
},
{
"table_name": "orders",
"description": "订单信息表",
"columns": [
{"name": "id", "type": "INTEGER", "description": "订单ID"},
{"name": "user_id", "type": "INTEGER", "description": "用户ID"},
{"name": "product_id", "type": "INTEGER", "description": "产品ID"},
{"name": "quantity", "type": "INTEGER", "description": "购买数量"},
{"name": "total_price", "type": "DECIMAL", "description": "总价格"},
{"name": "order_date", "type": "DATETIME", "description": "订单日期"}
]
}
]
if os.path.exists(schema_path):
with open(schema_path, 'r', encoding='utf-8') as f:
self.table_schemas = json.load(f)
else:
self.table_schemas = default_schemas
os.makedirs(data_dir, exist_ok=True)
with open(schema_path, 'w', encoding='utf-8') as f:
json.dump(self.table_schemas, f, ensure_ascii=False, indent=2)
def vectorize_and_store(self):
"""向量化数据并存储到Milvus"""
self.create_collection()
all_texts = []
all_metadata = []
for example in self.sql_examples:
text = f"问题: {example['question']} SQL: {example['sql']} 描述: {example.get('description', '')}"
all_texts.append(text)
all_metadata.append({
"content_type": "sql_example",
"question": example['question'],
"sql": example['sql'],
"description": example.get('description', ''),
"table_name": ""
})
for schema in self.table_schemas:
columns_desc = ", ".join([f"{col['name']} ({col['type']}): {col.get('description', '')}"
for col in schema['columns']])
text = f"{schema['table_name']}: {schema['description']} 字段: {columns_desc}"
all_texts.append(text)
all_metadata.append({
"content_type": "table_schema",
"question": "",
"sql": "",
"description": schema['description'],
"table_name": schema['table_name']
})
embeddings = self.embedding_function.encode_text(all_texts)
insert_data = []
for i, (embedding, metadata) in enumerate(zip(embeddings, all_metadata)):
insert_data.append([
metadata["content_type"],
metadata["question"],
metadata["sql"],
metadata["description"],
metadata["table_name"],
embedding.tolist()
])
self.collection.insert(insert_data)
self.collection.flush()
self.collection.load()
def search(self, query: str, top_k: int = 5) -> List[Dict[str, Any]]:
"""搜索相关的知识库信息"""
if not self.data_loaded:
self.load_data()
query_embedding = self.embedding_function.encode_text([query])[0]
search_params = {"metric_type": "IP", "params": {}}
results = self.collection.search(
[query_embedding.tolist()],
anns_field="embedding",
param=search_params,
limit=top_k,
output_fields=["content_type", "question", "sql", "description", "table_name"]
)[0]
formatted_results = []
for hit in results:
result = {
"score": float(hit.distance),
"content_type": hit.entity.get("content_type"),
"question": hit.entity.get("question"),
"sql": hit.entity.get("sql"),
"description": hit.entity.get("description"),
"table_name": hit.entity.get("table_name")
}
formatted_results.append(result)
return formatted_results
def _fallback_search(self, query: str, top_k: int) -> List[Dict[str, Any]]:
"""降级搜索方法(简单文本匹配)"""
results = []
query_lower = query.lower()
for example in self.sql_examples:
question_lower = example['question'].lower()
sql_lower = example['sql'].lower()
score = 0
for word in query_lower.split():
if word in question_lower:
score += 2
if word in sql_lower:
score += 1
if score > 0:
results.append({
"score": score,
"content_type": "sql_example",
"question": example['question'],
"sql": example['sql'],
"description": example.get('description', ''),
"table_name": ""
})
results.sort(key=lambda x: x['score'], reverse=True)
return results[:top_k]
def add_sql_example(self, question: str, sql: str, description: str = ""):
"""添加新的SQL示例"""
new_example = {
"question": question,
"sql": sql,
"description": description,
"database": "sqlite"
}
self.sql_examples.append(new_example)
data_dir = os.path.join(os.path.dirname(__file__), "data")
sql_examples_path = os.path.join(data_dir, "qsql_examples.json")
with open(sql_examples_path, 'w', encoding='utf-8') as f:
json.dump(self.sql_examples, f, ensure_ascii=False, indent=2)
if self.collection and self.data_loaded:
text = f"问题: {question} SQL: {sql} 描述: {description}"
embedding = self.embedding_function.encode_text([text])[0]
insert_data = [[
"sql_example",
question,
sql,
description,
"",
embedding.tolist()
]]
self.collection.insert(insert_data)
self.collection.flush()
def cleanup(self):
"""清理资源"""
if self.collection:
self.collection.release()
if self.milvus_client and self.milvus_client.has_collection(self.collection_name):
self.milvus_client.drop_collection(self.collection_name)
def demo():
"""简单演示"""
# 模型测试
embedding_function = BGESmallEmbeddingFunction()
test_texts = ["查询用户", "统计数据"]
embeddings = embedding_function.encode_text(test_texts)
print(f"向量维度: {embeddings.shape}")
# 数据库查询演示
db_path = "demo.db"
if os.path.exists(db_path):
os.remove(db_path)
conn = sqlite3.connect(db_path)
cursor = conn.cursor()
cursor.execute("CREATE TABLE users (id INTEGER PRIMARY KEY, name TEXT, age INTEGER, city TEXT)")
users_data = [(1, '张三', 25, '北京'), (2, '李四', 32, '上海'), (3, '王五', 35, '深圳')]
cursor.executemany("INSERT INTO users VALUES (?, ?, ?, ?)", users_data)
conn.commit()
# 执行查询
test_sqls = [
("查询所有用户", "SELECT * FROM users"),
("年龄大于30的用户", "SELECT * FROM users WHERE age > 30"),
("统计用户总数", "SELECT COUNT(*) FROM users")
]
for i, (question, sql) in enumerate(test_sqls, 1):
print(f"\n问题 {i}: {question}")
print("-" * 40)
print(f"SQL: {sql}")
cursor.execute(sql)
rows = cursor.fetchall()
if rows:
print(f"返回 {len(rows)} 行数据")
for j, row in enumerate(rows[:2], 1):
print(f" {j}. {row}")
if len(rows) > 2:
print(f" ... 还有 {len(rows) - 2}")
else:
print("无数据返回")
conn.close()
os.remove(db_path)
if __name__ == "__main__":
demo()
+149
View File
@@ -0,0 +1,149 @@
import os
from langchain_deepseek import ChatDeepSeek
from langchain_community.document_loaders import BiliBiliLoader
from langchain.chains.query_constructor.base import AttributeInfo
from openai import OpenAI
from langchain_community.vectorstores import Chroma
from langchain_huggingface import HuggingFaceEmbeddings
import logging
logging.basicConfig(level=logging.INFO)
# 1. 初始化视频数据
video_urls = [
"https://www.bilibili.com/video/BV1Bo4y1A7FU",
"https://www.bilibili.com/video/BV1ug4y157xA",
"https://www.bilibili.com/video/BV1yh411V7ge",
]
bili = []
try:
loader = BiliBiliLoader(video_urls=video_urls)
docs = loader.load()
for doc in docs:
original = doc.metadata
# 提取基本元数据字段
metadata = {
'title': original.get('title', '未知标题'),
'author': original.get('owner', {}).get('name', '未知作者'),
'source': original.get('bvid', '未知ID'),
'view_count': original.get('stat', {}).get('view', 0),
'length': original.get('duration', 0),
}
doc.metadata = metadata
bili.append(doc)
except Exception as e:
print(f"加载BiliBili视频失败: {str(e)}")
if not bili:
print("没有成功加载任何视频,程序退出")
exit()
# 2. 创建向量存储
embed_model = HuggingFaceEmbeddings(model_name="BAAI/bge-small-zh-v1.5")
vectorstore = Chroma.from_documents(bili, embed_model)
# 3. 配置元数据字段信息
metadata_field_info = [
AttributeInfo(
name="title",
description="视频标题(字符串)",
type="string",
),
AttributeInfo(
name="author",
description="视频作者(字符串)",
type="string",
),
AttributeInfo(
name="view_count",
description="视频观看次数(整数)",
type="integer",
),
AttributeInfo(
name="length",
description="视频长度(整数)",
type="integer"
)
]
# 4. 初始化LLM客户端
client = OpenAI(
base_url="https://api.deepseek.com",
api_key=os.getenv("DEEPSEEK_API_KEY")
)
# 5. 获取所有文档用于排序
all_documents = vectorstore.similarity_search("", k=len(bili))
# 6. 执行查询示例
queries = [
"时间最短的视频",
"播放量最高的视频"
]
for query in queries:
print(f"\n--- 原始查询: '{query}' ---")
# 使用大模型将自然语言转换为排序指令
prompt = f"""你是一个智能助手,请将用户的问题转换成一个用于排序视频的JSON指令。
你需要识别用户想要排序的字段和排序方向。
- 排序字段必须是 'view_count' (观看次数) 或 'length' (时长) 之一。
- 排序方向必须是 'asc' (升序) 或 'desc' (降序) 之一。
例如:
- '时间最短的视频''哪个视频时间最短' 应转换为 {{"sort_by": "length", "order": "asc"}}
- '播放量最高的视频''哪个视频最火' 应转换为 {{"sort_by": "view_count", "order": "desc"}}
请根据以下问题生成JSON指令:
原始问题: "{query}"
JSON指令:"""
response = client.chat.completions.create(
model="deepseek-chat",
messages=[
{"role": "user", "content": prompt}
],
temperature=0,
response_format={"type": "json_object"}
)
try:
import json
instruction_str = response.choices[0].message.content
instruction = json.loads(instruction_str)
print(f"--- 生成的排序指令: {instruction} ---")
sort_by = instruction.get('sort_by')
order = instruction.get('order')
if sort_by in ['length', 'view_count'] and order in ['asc', 'desc']:
# 在代码中执行排序
reverse_order = (order == 'desc')
sorted_docs = sorted(all_documents, key=lambda doc: doc.metadata.get(sort_by, 0), reverse=reverse_order)
# 获取排序后的第一个结果
if sorted_docs:
doc = sorted_docs[0]
title = doc.metadata.get('title', '未知标题')
author = doc.metadata.get('author', '未知作者')
view_count = doc.metadata.get('view_count', '未知')
length = doc.metadata.get('length', '未知')
print(f"标题: {title}")
print(f"作者: {author}")
print(f"观看次数: {view_count}")
print(f"时长: {length}")
print("="*50)
else:
print("没有找到任何视频")
else:
print("生成的指令无效,无法执行排序")
except (json.JSONDecodeError, KeyError) as e:
print(f"解析或执行指令失败: {e}")
+72
View File
@@ -0,0 +1,72 @@
import os
from langchain_core.prompts import ChatPromptTemplate
from langchain_core.output_parsers import StrOutputParser
from langchain_deepseek import ChatDeepSeek
from langchain_core.runnables import RunnableBranch
llm = ChatDeepSeek(
model="deepseek-chat",
temperature=0,
api_key=os.getenv("DEEPSEEK_API_KEY")
)
# 1. 设置不同菜系的处理链
sichuan_prompt = ChatPromptTemplate.from_template(
"你是一位川菜大厨。请用正宗的川菜做法,回答关于「{question}」的问题。"
)
sichuan_chain = sichuan_prompt | llm | StrOutputParser()
cantonese_prompt = ChatPromptTemplate.from_template(
"你是一位粤菜大厨。请用经典的粤菜做法,回答关于「{question}」的问题。"
)
cantonese_chain = cantonese_prompt | llm | StrOutputParser()
# 定义备用通用链
general_prompt = ChatPromptTemplate.from_template(
"你是一个美食助手。请回答关于「{question}」的问题。"
)
general_chain = general_prompt | llm | StrOutputParser()
# 2. 创建路由链
classifier_prompt = ChatPromptTemplate.from_template(
"""根据用户问题中提到的菜品,将其分类为:['川菜', '粤菜', 或 '其他']。
不要解释你的理由,只返回一个单词的分类结果。
问题: {question}"""
)
classifier_chain = classifier_prompt | llm | StrOutputParser()
# 定义路由分支
router_branch = RunnableBranch(
(lambda x: "川菜" in x["topic"], sichuan_chain),
(lambda x: "粤菜" in x["topic"], cantonese_chain),
general_chain # 默认选项
)
# 组合成完整路由链
full_router_chain = {"topic": classifier_chain, "question": lambda x: x["question"]} | router_branch
print("完整的路由链创建成功。\n")
# 3. 运行演示查询
demo_questions = [
{"question": "麻婆豆腐怎么做?"}, # 应该路由到川菜
{"question": "白切鸡的正宗做法是什么?"}, # 应该路由到粤菜
{"question": "番茄炒蛋需要放糖吗?"} # 应该路由到其他
]
for i, item in enumerate(demo_questions, 1):
question = item["question"]
print(f"\n--- 问题 {i}: {question} ---")
try:
# 获取路由决策
topic = classifier_chain.invoke({"question": question})
print(f"路由决策: {topic}")
# 执行完整链
result = full_router_chain.invoke(item)
print(f"回答: {result}")
except Exception as e:
print(f"执行错误: {e}")
+83
View File
@@ -0,0 +1,83 @@
import os
from langchain_core.prompts import PromptTemplate
from langchain_core.output_parsers import StrOutputParser
from langchain_deepseek import ChatDeepSeek
from langchain_huggingface import HuggingFaceEmbeddings
from langchain_core.runnables import RunnableLambda, RunnablePassthrough, RunnablePassthrough
from langchain_community.utils.math import cosine_similarity
import numpy as np
# 1. 定义路由描述
sichuan_route_prompt = "你是一位处理川菜的专家。用户的问题是关于麻辣、辛香、重口味的菜肴,例如水煮鱼、麻婆豆腐、鱼香肉丝、宫保鸡丁、花椒、海椒等。"
cantonese_route_prompt = "你是一位处理粤菜的专家。用户的问题是关于清淡、鲜美、原汁原味的菜肴,例如白切鸡、老火靓汤、虾饺、云吞面等。"
route_prompts = [sichuan_route_prompt, cantonese_route_prompt]
route_names = ["川菜", "粤菜"]
# 初始化嵌入模型,并对路由描述进行向量化
embeddings = HuggingFaceEmbeddings(model_name="BAAI/bge-small-zh-v1.5")
route_prompt_embeddings = embeddings.embed_documents(route_prompts)
print(f"已定义 {len(route_names)} 个路由: {', '.join(route_names)}")
# 2. 定义不同路由的目标链
llm = ChatDeepSeek(
model="deepseek-chat",
temperature=0,
api_key=os.getenv("DEEPSEEK_API_KEY")
)
# 定义川菜和粤菜处理链
sichuan_chain = (
PromptTemplate.from_template("你是一位川菜大厨。请用正宗的川菜做法,回答关于「{query}」的问题。")
| llm
| StrOutputParser()
)
cantonese_chain = (
PromptTemplate.from_template("你是一位粤菜大厨。请用经典的粤菜做法,回答关于「{query}」的问题。")
| llm
| StrOutputParser()
)
route_map = { "川菜": sichuan_chain, "粤菜": cantonese_chain }
print("川菜和粤菜的处理链创建成功。\n")
# 3. 创建路由函数
def route(info):
# 对用户查询进行嵌入
query_embedding = embeddings.embed_query(info["query"])
# 计算与各路由提示的余弦相似度
similarity_scores = cosine_similarity([query_embedding], route_prompt_embeddings)[0]
# 找到最相似的路由
chosen_route_index = np.argmax(similarity_scores)
chosen_route_name = route_names[chosen_route_index]
print(f"路由决策: 检测到问题与“{chosen_route_name}”最相似。")
# 获取对应的处理链
chosen_chain = route_map[chosen_route_name]
# 直接调用选中的链并返回结果
return chosen_chain.invoke(info)
# 创建完整的路由链
full_chain = RunnableLambda(route)
# 4. 运行演示查询
demo_queries = [
"水煮鱼怎么做才嫩?", # 应该路由到川菜
"如何做一碗清淡的云吞面?", # 应该路由到粤菜
"麻婆豆腐的核心调料是什么?", # 应该路由到川菜
]
for i, query in enumerate(demo_queries, 1):
print(f"\n--- 问题 {i}: {query} ---")
try:
# 传入字典,full_chain 会直接返回最终答案
result = full_chain.invoke({"query": query})
print(f"回答: {result}")
except Exception as e:
print(f"执行错误: {e}")
+186
View File
@@ -0,0 +1,186 @@
import os
from langchain_community.vectorstores import FAISS
from langchain.retrievers import ContextualCompressionRetriever
from langchain.retrievers.document_compressors import LLMChainExtractor
from langchain_community.embeddings import HuggingFaceBgeEmbeddings
from langchain.text_splitter import RecursiveCharacterTextSplitter
from langchain_community.document_loaders import TextLoader
from langchain_deepseek import ChatDeepSeek
# 导入ColBERT重排器需要的模块
from langchain.retrievers.document_compressors.base import BaseDocumentCompressor
from langchain.retrievers.document_compressors import DocumentCompressorPipeline
from langchain_core.documents import Document
from typing import Sequence
import torch
from transformers import AutoTokenizer, AutoModel
import torch.nn.functional as F
class ColBERTReranker(BaseDocumentCompressor):
"""ColBERT重排器"""
def __init__(self, **kwargs):
super().__init__(**kwargs)
model_name = "bert-base-uncased"
# 加载模型和分词器
object.__setattr__(self, 'tokenizer', AutoTokenizer.from_pretrained(model_name))
object.__setattr__(self, 'model', AutoModel.from_pretrained(model_name))
self.model.eval()
print(f"ColBERT模型加载完成")
def encode_text(self, texts):
"""ColBERT文本编码"""
inputs = self.tokenizer(
texts,
return_tensors="pt",
padding=True,
truncation=True,
max_length=128
)
with torch.no_grad():
outputs = self.model(**inputs)
embeddings = outputs.last_hidden_state
embeddings = F.normalize(embeddings, p=2, dim=-1)
return embeddings
def calculate_colbert_similarity(self, query_emb, doc_embs, query_mask, doc_masks):
"""ColBERT相似度计算(MaxSim操作)"""
scores = []
for i, doc_emb in enumerate(doc_embs):
doc_mask = doc_masks[i:i+1]
# 计算相似度矩阵
similarity_matrix = torch.matmul(query_emb, doc_emb.unsqueeze(0).transpose(-2, -1))
# 应用文档mask
doc_mask_expanded = doc_mask.unsqueeze(1)
similarity_matrix = similarity_matrix.masked_fill(~doc_mask_expanded.bool(), -1e9)
# MaxSim操作
max_sim_per_query_token = similarity_matrix.max(dim=-1)[0]
# 应用查询mask
query_mask_expanded = query_mask.unsqueeze(0)
max_sim_per_query_token = max_sim_per_query_token.masked_fill(~query_mask_expanded.bool(), 0)
# 求和得到最终分数
colbert_score = max_sim_per_query_token.sum(dim=-1).item()
scores.append(colbert_score)
return scores
def compress_documents(
self,
documents: Sequence[Document],
query: str,
callbacks=None,
) -> Sequence[Document]:
"""对文档进行ColBERT重排序"""
if len(documents) == 0:
return documents
# 编码查询
query_inputs = self.tokenizer(
[query],
return_tensors="pt",
padding=True,
truncation=True,
max_length=128
)
with torch.no_grad():
query_outputs = self.model(**query_inputs)
query_embeddings = F.normalize(query_outputs.last_hidden_state, p=2, dim=-1)
# 编码文档
doc_texts = [doc.page_content for doc in documents]
doc_inputs = self.tokenizer(
doc_texts,
return_tensors="pt",
padding=True,
truncation=True,
max_length=128
)
with torch.no_grad():
doc_outputs = self.model(**doc_inputs)
doc_embeddings = F.normalize(doc_outputs.last_hidden_state, p=2, dim=-1)
# 计算ColBERT相似度
scores = self.calculate_colbert_similarity(
query_embeddings,
doc_embeddings,
query_inputs['attention_mask'],
doc_inputs['attention_mask']
)
# 排序并返回前5个
scored_docs = list(zip(documents, scores))
scored_docs.sort(key=lambda x: x[1], reverse=True)
reranked_docs = [doc for doc, _ in scored_docs[:5]]
return reranked_docs
# 初始化配置
hf_bge_embeddings = HuggingFaceBgeEmbeddings(
model_name="BAAI/bge-large-zh-v1.5"
)
llm = ChatDeepSeek(
model="deepseek-chat",
temperature=0.1,
api_key=os.getenv("DEEPSEEK_API_KEY")
)
# 1. 加载和处理文档
loader = TextLoader("../../data/C4/txt/ai.txt", encoding="utf-8")
documents = loader.load()
text_splitter = RecursiveCharacterTextSplitter(chunk_size=500, chunk_overlap=100)
docs = text_splitter.split_documents(documents)
# 2. 创建向量存储和基础检索器
vectorstore = FAISS.from_documents(docs, hf_bge_embeddings)
base_retriever = vectorstore.as_retriever(search_kwargs={"k": 20})
# 3. 设置ColBERT重排序器
reranker = ColBERTReranker()
# 4. 设置LLM压缩器
compressor = LLMChainExtractor.from_llm(llm)
# 5. 使用DocumentCompressorPipeline组装压缩管道
# 流程: ColBERT重排 -> LLM压缩
pipeline_compressor = DocumentCompressorPipeline(
transformers=[reranker, compressor]
)
# 6. 创建最终的压缩检索器
final_retriever = ContextualCompressionRetriever(
base_compressor=pipeline_compressor,
base_retriever=base_retriever
)
# 7. 执行查询并展示结果
query = "AI还有哪些缺陷需要克服?"
print(f"\n{'='*20} 开始执行查询 {'='*20}")
print(f"查询: {query}\n")
# 7.1 基础检索结果
print(f"--- (1) 基础检索结果 (Top 20) ---")
base_results = base_retriever.get_relevant_documents(query)
for i, doc in enumerate(base_results):
print(f" [{i+1}] {doc.page_content[:100]}...\n")
# 7.2 使用管道压缩器的最终结果
print(f"\n--- (2) 管道压缩后结果 (ColBERT重排 + LLM压缩) ---")
final_results = final_retriever.get_relevant_documents(query)
for i, doc in enumerate(final_results):
print(f" [{i+1}] {doc.page_content}\n")
+17
View File
@@ -0,0 +1,17 @@
"""
简化的Text2SQL框架
基于RAGFlow方案实现的Text2SQL框架
"""
__version__ = "1.0.0"
__author__ = "RAG Team"
from .knowledge_base import SimpleKnowledgeBase
from .sql_generator import SimpleSQLGenerator
from .text2sql_agent import SimpleText2SQLAgent
__all__ = [
"SimpleKnowledgeBase",
"SimpleSQLGenerator",
"SimpleText2SQLAgent"
]
@@ -0,0 +1,57 @@
[
{
"table_name": "users",
"table_description": "用户信息表,存储注册用户的基本信息",
"columns": [
{"name": "id", "description": "用户唯一标识符,主键", "type": "INT"},
{"name": "name", "description": "用户姓名,不能为空", "type": "VARCHAR(100)"},
{"name": "email", "description": "用户邮箱地址,必须唯一", "type": "VARCHAR(150)"},
{"name": "age", "description": "用户年龄", "type": "INT"},
{"name": "created_at", "description": "用户注册时间", "type": "TIMESTAMP"}
]
},
{
"table_name": "orders",
"table_description": "订单表,记录用户的购买订单信息",
"columns": [
{"name": "id", "description": "订单唯一标识符,主键", "type": "INT"},
{"name": "user_id", "description": "下单用户的ID,外键关联users表", "type": "INT"},
{"name": "product_name", "description": "购买的产品名称", "type": "VARCHAR(200)"},
{"name": "quantity", "description": "购买数量", "type": "INT"},
{"name": "price", "description": "订单总价格", "type": "DECIMAL(10,2)"},
{"name": "order_date", "description": "下单时间", "type": "TIMESTAMP"}
]
},
{
"table_name": "products",
"table_description": "产品表,存储商城中所有产品的信息",
"columns": [
{"name": "id", "description": "产品唯一标识符,主键", "type": "INT"},
{"name": "name", "description": "产品名称", "type": "VARCHAR(200)"},
{"name": "category", "description": "产品分类", "type": "VARCHAR(100)"},
{"name": "price", "description": "产品单价", "type": "DECIMAL(10,2)"},
{"name": "stock", "description": "库存数量", "type": "INT"},
{"name": "description", "description": "产品详细描述", "type": "TEXT"}
]
},
{
"table_name": "categories",
"table_description": "产品分类表,定义产品的分类信息",
"columns": [
{"name": "id", "description": "分类唯一标识符,主键", "type": "INT"},
{"name": "name", "description": "分类名称,必须唯一", "type": "VARCHAR(100)"},
{"name": "description", "description": "分类描述", "type": "TEXT"}
]
},
{
"table_name": "order_items",
"table_description": "订单明细表,存储订单中包含的具体商品信息",
"columns": [
{"name": "id", "description": "订单明细唯一标识符,主键", "type": "INT"},
{"name": "order_id", "description": "关联的订单ID,外键", "type": "INT"},
{"name": "product_id", "description": "关联的产品ID,外键", "type": "INT"},
{"name": "quantity", "description": "该商品在订单中的数量", "type": "INT"},
{"name": "unit_price", "description": "该商品的单价(下单时的价格)", "type": "DECIMAL(10,2)"}
]
}
]
+27
View File
@@ -0,0 +1,27 @@
[
{
"table_name": "users",
"ddl_statement": "CREATE TABLE users (id INT PRIMARY KEY AUTO_INCREMENT, name VARCHAR(100) NOT NULL, email VARCHAR(150) UNIQUE, age INT, created_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP)",
"description": "用户信息表,存储用户基本信息"
},
{
"table_name": "orders",
"ddl_statement": "CREATE TABLE orders (id INT PRIMARY KEY AUTO_INCREMENT, user_id INT, product_name VARCHAR(200), quantity INT, price DECIMAL(10,2), order_date TIMESTAMP DEFAULT CURRENT_TIMESTAMP, FOREIGN KEY (user_id) REFERENCES users(id))",
"description": "订单表,存储用户订单信息"
},
{
"table_name": "products",
"ddl_statement": "CREATE TABLE products (id INT PRIMARY KEY AUTO_INCREMENT, name VARCHAR(200) NOT NULL, category VARCHAR(100), price DECIMAL(10,2), stock INT DEFAULT 0, description TEXT)",
"description": "产品表,存储产品基本信息"
},
{
"table_name": "categories",
"ddl_statement": "CREATE TABLE categories (id INT PRIMARY KEY AUTO_INCREMENT, name VARCHAR(100) NOT NULL UNIQUE, description TEXT)",
"description": "产品分类表"
},
{
"table_name": "order_items",
"ddl_statement": "CREATE TABLE order_items (id INT PRIMARY KEY AUTO_INCREMENT, order_id INT, product_id INT, quantity INT, unit_price DECIMAL(10,2), FOREIGN KEY (order_id) REFERENCES orders(id), FOREIGN KEY (product_id) REFERENCES products(id))",
"description": "订单明细表,存储订单中的具体商品信息"
}
]
+62
View File
@@ -0,0 +1,62 @@
[
{
"question": "查询所有用户的姓名和邮箱",
"sql": "SELECT name, email FROM users",
"database": "ecommerce"
},
{
"question": "查找年龄大于25岁的用户",
"sql": "SELECT * FROM users WHERE age > 25",
"database": "ecommerce"
},
{
"question": "查询每个用户的订单数量",
"sql": "SELECT u.name, COUNT(o.id) as order_count FROM users u LEFT JOIN orders o ON u.id = o.user_id GROUP BY u.id, u.name",
"database": "ecommerce"
},
{
"question": "查找最近7天的订单",
"sql": "SELECT * FROM orders WHERE order_date >= DATE_SUB(NOW(), INTERVAL 7 DAY)",
"database": "ecommerce"
},
{
"question": "查询总销售额最高的前5个产品",
"sql": "SELECT p.name, SUM(oi.quantity * oi.unit_price) as total_sales FROM products p JOIN order_items oi ON p.id = oi.product_id GROUP BY p.id, p.name ORDER BY total_sales DESC LIMIT 5",
"database": "ecommerce"
},
{
"question": "查询某个用户的所有订单",
"sql": "SELECT o.*, u.name as user_name FROM orders o JOIN users u ON o.user_id = u.id WHERE u.name = '张三'",
"database": "ecommerce"
},
{
"question": "查询价格在100到500之间的产品",
"sql": "SELECT * FROM products WHERE price BETWEEN 100 AND 500",
"database": "ecommerce"
},
{
"question": "查询库存少于10的产品",
"sql": "SELECT * FROM products WHERE stock < 10",
"database": "ecommerce"
},
{
"question": "查询每个分类的产品数量",
"sql": "SELECT category, COUNT(*) as product_count FROM products GROUP BY category",
"database": "ecommerce"
},
{
"question": "查询订单总金额大于1000的订单",
"sql": "SELECT o.*, SUM(oi.quantity * oi.unit_price) as total_amount FROM orders o JOIN order_items oi ON o.id = oi.order_id GROUP BY o.id HAVING total_amount > 1000",
"database": "ecommerce"
},
{
"question": "查询没有下过订单的用户",
"sql": "SELECT u.* FROM users u LEFT JOIN orders o ON u.id = o.user_id WHERE o.id IS NULL",
"database": "ecommerce"
},
{
"question": "查询平均订单金额",
"sql": "SELECT AVG(total_amount) as avg_order_amount FROM (SELECT o.id, SUM(oi.quantity * oi.unit_price) as total_amount FROM orders o JOIN order_items oi ON o.id = oi.order_id GROUP BY o.id) as order_totals",
"database": "ecommerce"
}
]
+184
View File
@@ -0,0 +1,184 @@
import json
import os
from typing import List, Dict, Any
from pymilvus import MilvusClient, FieldSchema, CollectionSchema, DataType
from pymilvus.model.hybrid import BGEM3EmbeddingFunction
class SimpleKnowledgeBase:
"""知识库"""
def __init__(self, milvus_uri: str = "http://localhost:19530"):
self.milvus_uri = milvus_uri
self.client = MilvusClient(uri=milvus_uri)
self.embedding_function = BGEM3EmbeddingFunction(use_fp16=False, device="cpu")
self.collection_name = "text2sql_kb"
self._setup_collection()
def _setup_collection(self):
"""设置集合"""
if self.client.has_collection(self.collection_name):
self.client.drop_collection(self.collection_name)
# 定义字段
fields = [
FieldSchema(name="pk", dtype=DataType.VARCHAR, is_primary=True, auto_id=True, max_length=100),
FieldSchema(name="content", dtype=DataType.VARCHAR, max_length=4096),
FieldSchema(name="type", dtype=DataType.VARCHAR, max_length=32), # ddl, qsql, description
FieldSchema(name="dense_vector", dtype=DataType.FLOAT_VECTOR, dim=self.embedding_function.dim["dense"])
]
schema = CollectionSchema(fields, description="Text2SQL知识库")
# 创建集合
self.client.create_collection(
collection_name=self.collection_name,
schema=schema,
consistency_level="Strong"
)
# 创建索引
index_params = self.client.prepare_index_params()
index_params.add_index(
field_name="dense_vector",
index_type="AUTOINDEX",
metric_type="IP"
)
self.client.create_index(
collection_name=self.collection_name,
index_params=index_params
)
def load_data(self):
"""加载所有知识库数据"""
data_dir = os.path.join(os.path.dirname(__file__), "data")
# 加载DDL数据
ddl_path = os.path.join(data_dir, "ddl_examples.json")
if os.path.exists(ddl_path):
with open(ddl_path, 'r', encoding='utf-8') as f:
ddl_data = json.load(f)
self._add_ddl_data(ddl_data)
# 加载Q->SQL数据
qsql_path = os.path.join(data_dir, "qsql_examples.json")
if os.path.exists(qsql_path):
with open(qsql_path, 'r', encoding='utf-8') as f:
qsql_data = json.load(f)
self._add_qsql_data(qsql_data)
# 加载描述数据
desc_path = os.path.join(data_dir, "db_descriptions.json")
if os.path.exists(desc_path):
with open(desc_path, 'r', encoding='utf-8') as f:
desc_data = json.load(f)
self._add_description_data(desc_data)
# 加载集合到内存
self.client.load_collection(collection_name=self.collection_name)
print("知识库数据加载完成")
def _add_ddl_data(self, data: List[Dict]):
"""添加DDL数据"""
contents = []
types = []
for item in data:
content = f"表名: {item.get('table_name', '')}\n"
content += f"DDL: {item.get('ddl_statement', '')}\n"
content += f"描述: {item.get('description', '')}"
contents.append(content)
types.append("ddl")
self._insert_data(contents, types)
def _add_qsql_data(self, data: List[Dict]):
"""添加Q->SQL数据"""
contents = []
types = []
for item in data:
content = f"问题: {item.get('question', '')}\n"
content += f"SQL: {item.get('sql', '')}"
contents.append(content)
types.append("qsql")
self._insert_data(contents, types)
def _add_description_data(self, data: List[Dict]):
"""添加描述数据"""
contents = []
types = []
for item in data:
content = f"表名: {item.get('table_name', '')}\n"
content += f"表描述: {item.get('table_description', '')}\n"
columns = item.get('columns', [])
if columns:
content += "字段信息:\n"
for col in columns:
content += f" - {col.get('name', '')}: {col.get('description', '')} ({col.get('type', '')})\n"
contents.append(content)
types.append("description")
self._insert_data(contents, types)
def _insert_data(self, contents: List[str], types: List[str]):
"""插入数据"""
if not contents:
return
# 生成嵌入
embeddings = self.embedding_function(contents)
# 构建插入数据,每一行是一个字典
data_to_insert = []
for i in range(len(contents)):
data_to_insert.append({
"content": contents[i],
"type": types[i],
"dense_vector": embeddings["dense"][i]
})
# 插入数据
result = self.client.insert(
collection_name=self.collection_name,
data=data_to_insert
)
def search(self, query: str, top_k: int = 5) -> List[Dict[str, Any]]:
"""搜索相关内容"""
self.client.load_collection(collection_name=self.collection_name)
query_embeddings = self.embedding_function([query])
search_results = self.client.search(
collection_name=self.collection_name,
data=query_embeddings["dense"],
anns_field="dense_vector",
search_params={"metric_type": "IP"},
limit=top_k,
output_fields=["content", "type"]
)
results = []
for hit in search_results[0]:
results.append({
"content": hit["entity"]["content"],
"type": hit["entity"]["type"],
"score": hit["distance"]
})
return results
def cleanup(self):
"""清理资源"""
try:
self.client.drop_collection(self.collection_name)
except:
pass
+113
View File
@@ -0,0 +1,113 @@
import os
from typing import List, Dict, Any
from langchain_deepseek import ChatDeepSeek
from langchain.schema import HumanMessage, SystemMessage
class SimpleSQLGenerator:
"""简化的SQL生成器"""
def __init__(self, api_key: str = None):
self.llm = ChatDeepSeek(
model="deepseek-chat",
temperature=0,
api_key=api_key or os.getenv("DEEPSEEK_API_KEY")
)
def generate_sql(self, user_query: str, knowledge_results: List[Dict[str, Any]]) -> str:
"""生成SQL语句"""
# 构建上下文
context = self._build_context(knowledge_results)
# 构建提示
prompt = f"""你是一个SQL专家。请根据以下信息将用户问题转换为SQL查询语句。
数据库信息:
{context}
用户问题:{user_query}
要求:
1. 只返回SQL语句,不要包含任何解释
2. 确保SQL语法正确
3. 使用上下文中提供的表名和字段名
4. 如果需要JOIN,请根据表结构进行合理关联
SQL语句:"""
messages = [HumanMessage(content=prompt)]
response = self.llm.invoke(messages)
# 清理SQL语句
sql = response.content.strip()
if sql.startswith("```sql"):
sql = sql[6:]
if sql.startswith("```"):
sql = sql[3:]
if sql.endswith("```"):
sql = sql[:-3]
return sql.strip()
def fix_sql(self, original_sql: str, error_message: str, knowledge_results: List[Dict[str, Any]]) -> str:
"""修复SQL语句"""
context = self._build_context(knowledge_results)
prompt = f"""请修复以下SQL语句的错误。
数据库信息:
{context}
原始SQL
{original_sql}
错误信息:
{error_message}
请返回修复后的SQL语句(只返回SQL,不要解释):"""
messages = [HumanMessage(content=prompt)]
response = self.llm.invoke(messages)
# 清理SQL语句
fixed_sql = response.content.strip()
if fixed_sql.startswith("```sql"):
fixed_sql = fixed_sql[6:]
if fixed_sql.startswith("```"):
fixed_sql = fixed_sql[3:]
if fixed_sql.endswith("```"):
fixed_sql = fixed_sql[:-3]
return fixed_sql.strip()
def _build_context(self, knowledge_results: List[Dict[str, Any]]) -> str:
"""构建上下文信息"""
context = ""
# 按类型分组
ddl_info = []
qsql_examples = []
descriptions = []
for result in knowledge_results:
if result["type"] == "ddl":
ddl_info.append(result["content"])
elif result["type"] == "qsql":
qsql_examples.append(result["content"])
elif result["type"] == "description":
descriptions.append(result["content"])
# 构建上下文
if ddl_info:
context += "=== 表结构信息 ===\n"
context += "\n".join(ddl_info) + "\n\n"
if descriptions:
context += "=== 表和字段描述 ===\n"
context += "\n".join(descriptions) + "\n\n"
if qsql_examples:
context += "=== 查询示例 ===\n"
context += "\n".join(qsql_examples) + "\n\n"
return context
+213
View File
@@ -0,0 +1,213 @@
import sqlite3
import os
from typing import Dict, Any, List, Tuple
from .knowledge_base import SimpleKnowledgeBase
from .sql_generator import SimpleSQLGenerator
class SimpleText2SQLAgent:
"""Text2SQL代理"""
def __init__(self, milvus_uri: str = "http://localhost:19530", api_key: str = None):
"""初始化代理"""
self.knowledge_base = SimpleKnowledgeBase(milvus_uri)
self.sql_generator = SimpleSQLGenerator(api_key)
self.db_path = None
self.connection = None
# 配置参数
self.max_retry_count = 3
self.top_k_retrieval = 5
self.max_result_rows = 100
def connect_database(self, db_path: str) -> bool:
"""连接SQLite数据库"""
try:
self.db_path = db_path
self.connection = sqlite3.connect(db_path)
print(f"成功连接到数据库: {db_path}")
return True
except Exception as e:
print(f"数据库连接失败: {str(e)}")
return False
def load_knowledge_base(self):
"""加载知识库"""
self.knowledge_base.load_data()
def query(self, user_question: str) -> Dict[str, Any]:
"""执行Text2SQL查询"""
if not self.connection:
return {
"success": False,
"error": "数据库未连接",
"sql": None,
"results": None
}
print(f"\n=== 处理查询: {user_question} ===")
# 1. 从知识库检索
print("检索知识库...")
knowledge_results = self.knowledge_base.search(user_question, self.top_k_retrieval)
print(f"检索到 {len(knowledge_results)} 条相关信息")
# 2. 生成SQL
print("生成SQL...")
sql = self.sql_generator.generate_sql(user_question, knowledge_results)
print(f"生成的SQL: {sql}")
# 3. 执行SQL(带重试)
retry_count = 0
while retry_count < self.max_retry_count:
print(f"执行SQL (尝试 {retry_count + 1}/{self.max_retry_count})...")
success, result = self._execute_sql(sql)
if success:
print("SQL执行成功!")
return {
"success": True,
"error": None,
"sql": sql,
"results": result,
"retry_count": retry_count
}
else:
print(f"SQL执行失败: {result}")
if retry_count < self.max_retry_count - 1:
print("尝试修复SQL...")
sql = self.sql_generator.fix_sql(sql, result, knowledge_results)
print(f"修复后的SQL: {sql}")
retry_count += 1
return {
"success": False,
"error": f"超过最大重试次数 ({self.max_retry_count})",
"sql": sql,
"results": None,
"retry_count": retry_count
}
def _execute_sql(self, sql: str) -> Tuple[bool, Any]:
"""执行SQL语句"""
try:
cursor = self.connection.cursor()
# 添加LIMIT限制
if sql.strip().upper().startswith('SELECT') and 'LIMIT' not in sql.upper():
sql = f"{sql.rstrip(';')} LIMIT {self.max_result_rows}"
cursor.execute(sql)
if sql.strip().upper().startswith('SELECT'):
# 查询语句
columns = [desc[0] for desc in cursor.description]
rows = cursor.fetchall()
results = []
for row in rows:
result_row = {}
for i, value in enumerate(row):
result_row[columns[i]] = value
results.append(result_row)
cursor.close()
return True, {
"columns": columns,
"rows": results,
"count": len(results)
}
else:
# 非查询语句
self.connection.commit()
cursor.close()
return True, "SQL执行成功"
except Exception as e:
return False, str(e)
def add_example(self, question: str, sql: str):
"""添加新的Q->SQL示例"""
# 简化版本:直接保存到文件
data_dir = os.path.join(os.path.dirname(__file__), "data")
qsql_path = os.path.join(data_dir, "qsql_examples.json")
try:
import json
# 读取现有数据
if os.path.exists(qsql_path):
with open(qsql_path, 'r', encoding='utf-8') as f:
data = json.load(f)
else:
data = []
# 添加新示例
data.append({
"question": question,
"sql": sql,
"database": "sqlite"
})
# 保存
with open(qsql_path, 'w', encoding='utf-8') as f:
json.dump(data, f, ensure_ascii=False, indent=2)
print(f"已添加新示例: {question}")
except Exception as e:
print(f"添加示例失败: {str(e)}")
def get_table_info(self) -> List[Dict[str, Any]]:
"""获取数据库表信息"""
if not self.connection:
return []
try:
cursor = self.connection.cursor()
# 获取所有表名
cursor.execute("SELECT name FROM sqlite_master WHERE type='table'")
tables = cursor.fetchall()
table_info = []
for table in tables:
table_name = table[0]
# 获取表结构
cursor.execute(f"PRAGMA table_info({table_name})")
columns = cursor.fetchall()
table_info.append({
"table_name": table_name,
"columns": [
{
"name": col[1],
"type": col[2],
"nullable": not col[3],
"default": col[4],
"primary_key": bool(col[5])
}
for col in columns
]
})
cursor.close()
return table_info
except Exception as e:
print(f"获取表信息失败: {str(e)}")
return []
def cleanup(self):
"""清理资源"""
if self.connection:
self.connection.close()
self.connection = None
print("数据库连接已关闭")
self.knowledge_base.cleanup()
print("知识库已清理")
+193
View File
@@ -0,0 +1,193 @@
import os
from langchain_community.vectorstores import FAISS
from langchain.retrievers import ContextualCompressionRetriever
from langchain.retrievers.document_compressors import LLMChainExtractor
from langchain_community.embeddings import HuggingFaceBgeEmbeddings
from langchain.text_splitter import RecursiveCharacterTextSplitter
from langchain_community.document_loaders import TextLoader
from langchain_deepseek import ChatDeepSeek
# 导入ColBERT重排器需要的模块
from langchain.retrievers.document_compressors.base import BaseDocumentCompressor
from langchain.retrievers.document_compressors import DocumentCompressorPipeline
from langchain_core.documents import Document
from typing import Sequence
import torch
from transformers import AutoTokenizer, AutoModel
import torch.nn.functional as F
class ColBERTReranker(BaseDocumentCompressor):
"""ColBERT重排器"""
def __init__(self, **kwargs):
super().__init__(**kwargs)
model_name = "bert-base-uncased"
# 加载模型和分词器
object.__setattr__(self, 'tokenizer', AutoTokenizer.from_pretrained(model_name))
object.__setattr__(self, 'model', AutoModel.from_pretrained(model_name))
self.model.eval()
print(f"ColBERT模型加载完成")
def encode_text(self, texts):
"""ColBERT文本编码"""
inputs = self.tokenizer(
texts,
return_tensors="pt",
padding=True,
truncation=True,
max_length=128
)
with torch.no_grad():
outputs = self.model(**inputs)
embeddings = outputs.last_hidden_state
embeddings = F.normalize(embeddings, p=2, dim=-1)
return embeddings
def calculate_colbert_similarity(self, query_emb, doc_embs, query_mask, doc_masks):
"""ColBERT相似度计算(MaxSim操作)"""
scores = []
for i, doc_emb in enumerate(doc_embs):
doc_mask = doc_masks[i:i+1]
# 计算相似度矩阵
similarity_matrix = torch.matmul(query_emb, doc_emb.unsqueeze(0).transpose(-2, -1))
# 应用文档mask
doc_mask_expanded = doc_mask.unsqueeze(1)
similarity_matrix = similarity_matrix.masked_fill(~doc_mask_expanded.bool(), -1e9)
# MaxSim操作
max_sim_per_query_token = similarity_matrix.max(dim=-1)[0]
# 应用查询mask
query_mask_expanded = query_mask.unsqueeze(0)
max_sim_per_query_token = max_sim_per_query_token.masked_fill(~query_mask_expanded.bool(), 0)
# 求和得到最终分数
colbert_score = max_sim_per_query_token.sum(dim=-1).item()
scores.append(colbert_score)
return scores
def compress_documents(
self,
documents: Sequence[Document],
query: str,
callbacks=None,
) -> Sequence[Document]:
"""对文档进行ColBERT重排序"""
if len(documents) == 0:
return documents
# 编码查询
query_inputs = self.tokenizer(
[query],
return_tensors="pt",
padding=True,
truncation=True,
max_length=128
)
with torch.no_grad():
query_outputs = self.model(**query_inputs)
query_embeddings = F.normalize(query_outputs.last_hidden_state, p=2, dim=-1)
# 编码文档
doc_texts = [doc.page_content for doc in documents]
doc_inputs = self.tokenizer(
doc_texts,
return_tensors="pt",
padding=True,
truncation=True,
max_length=128
)
with torch.no_grad():
doc_outputs = self.model(**doc_inputs)
doc_embeddings = F.normalize(doc_outputs.last_hidden_state, p=2, dim=-1)
# 计算ColBERT相似度
scores = self.calculate_colbert_similarity(
query_embeddings,
doc_embeddings,
query_inputs['attention_mask'],
doc_inputs['attention_mask']
)
# 排序并返回前5个
scored_docs = list(zip(documents, scores))
scored_docs.sort(key=lambda x: x[1], reverse=True)
reranked_docs = [doc for doc, _ in scored_docs[:5]]
return reranked_docs
# 初始化配置
hf_bge_embeddings = HuggingFaceBgeEmbeddings(
model_name="BAAI/bge-large-zh-v1.5"
)
llm = ChatDeepSeek(
model="deepseek-chat",
temperature=0.1,
api_key=os.getenv("DEEPSEEK_API_KEY")
)
# 1. 加载和处理文档
loader = TextLoader("../../data/C4/txt/ai.txt", encoding="utf-8")
documents = loader.load()
# 优化分块策略:减少重叠,使用中文友好的分隔符
text_splitter = RecursiveCharacterTextSplitter(
separators=["\n\n", "\n", "", "", "", "", "", " ", ""],
chunk_size=300,
chunk_overlap=20
)
docs = text_splitter.split_documents(documents)
# 2. 创建向量存储和基础检索器
vectorstore = FAISS.from_documents(docs, hf_bge_embeddings)
base_retriever = vectorstore.as_retriever(search_kwargs={"k": 20})
# 3. 设置ColBERT重排序器
reranker = ColBERTReranker()
# 4. 设置LLM压缩器
compressor = LLMChainExtractor.from_llm(llm)
# 5. 使用DocumentCompressorPipeline组装压缩管道
# 流程: ColBERT重排 -> LLM压缩
pipeline_compressor = DocumentCompressorPipeline(
transformers=[reranker, compressor]
)
# 6. 创建最终的压缩检索器
final_retriever = ContextualCompressionRetriever(
base_compressor=pipeline_compressor,
base_retriever=base_retriever
)
# 7. 执行查询并展示结果
query = "AI还有哪些缺陷需要克服?"
print(f"\n{'='*20} 开始执行查询 {'='*20}")
print(f"查询: {query}\n")
# 7.1 基础检索结果
print(f"--- (1) 基础检索结果 (Top 20) ---")
base_results = base_retriever.get_relevant_documents(query)
for i, doc in enumerate(base_results):
print(f" [{i+1}] {doc.page_content[:100]}...\n")
# 7.2 使用管道压缩器的最终结果
print(f"\n--- (2) 管道压缩后结果 (ColBERT重排 + LLM压缩) ---")
final_results = final_retriever.get_relevant_documents(query)
for i, doc in enumerate(final_results):
print(f" [{i+1}] {doc.page_content}\n")