Initial commit
This commit is contained in:
@@ -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}'")
|
||||
@@ -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}'")
|
||||
@@ -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("未找到匹配的视频")
|
||||
@@ -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()
|
||||
@@ -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()
|
||||
@@ -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}")
|
||||
@@ -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}")
|
||||
|
||||
@@ -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}")
|
||||
|
||||
@@ -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")
|
||||
@@ -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)"}
|
||||
]
|
||||
}
|
||||
]
|
||||
@@ -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": "订单明细表,存储订单中的具体商品信息"
|
||||
}
|
||||
]
|
||||
@@ -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"
|
||||
}
|
||||
]
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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("知识库已清理")
|
||||
@@ -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")
|
||||
Reference in New Issue
Block a user