Files
ai-g3sb-backman2.0/backend/schema/indexer.py
T
2026-04-15 13:07:39 +08:00

257 lines
8.3 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
"""
Schema向量索引构建器 - 基于ChromaDB
"""
import chromadb
from chromadb.config import Settings as ChromaSettings
from typing import Any, List, Dict, Optional
import logging
from pathlib import Path
from schema.manager import SchemaManager
logger = logging.getLogger(__name__)
class SchemaIndexer:
"""
Schema向量索引器
功能:
1. 为表结构构建向量索引
2. 基于问题的表检索
3. Chroma 使用持久化(PersistentClient),向量与元数据写入 ``persist_dir``,进程重启后可复用
"""
def __init__(
self,
embedder: Any,
persist_dir: str = "./data/embeddings",
collection_name: str = "schema_tables",
):
"""
初始化索引器
Args:
embedder: Embedding模型实例
persist_dir: Chroma 持久化根目录(磁盘路径)
collection_name: 集合名称
"""
self.embedder = embedder
self.persist_dir = Path(persist_dir)
self.persist_dir.mkdir(parents=True, exist_ok=True)
self.collection_name = collection_name
# 持久化 Chroma:数据落盘至 persist_dir
self.client = chromadb.PersistentClient(
path=str(self.persist_dir),
settings=ChromaSettings(anonymized_telemetry=False),
)
# 获取或创建集合
self.collection = self.client.get_or_create_collection(
name=self.collection_name,
metadata={"hnsw:space": "cosine"}, # 使用余弦相似度
)
logger.info(
f"[OK] 初始化SchemaIndexer(Chroma持久化): collection={self.collection_name}, "
f"path={self.persist_dir}"
)
def build_index(
self,
schema_manager: SchemaManager,
batch_size: int = 32,
force_rebuild: bool = False,
) -> bool:
"""
为Schema构建向量索引
Args:
schema_manager: SchemaManager实例
batch_size: 批处理大小
force_rebuild: 是否强制重建(默认为False,增量更新)
Returns:
True=成功,False=已存在且未强制重建
"""
existing_ids = set(self.collection.get()["ids"]) if self.collection.count() > 0 else set()
if not force_rebuild and existing_ids:
logger.info(f"索引已存在({len(existing_ids)}条记录),跳过构建")
return False
if force_rebuild and existing_ids:
logger.info(f"强制重建索引,删除{len(existing_ids)}条旧记录")
# Chroma 新版本要求 delete 必须带 ids/where;整库清空用删集合再建
self.client.delete_collection(self.collection_name)
self.collection = self.client.get_or_create_collection(
name=self.collection_name,
metadata={"hnsw:space": "cosine"},
)
# 准备数据
table_texts = []
table_ids = []
table_metadatas = []
for table_info in schema_manager.get_table_for_embedding():
table_texts.append(table_info["text"])
table_ids.append(table_info["name"])
table_metadatas.append({
"table_name": table_info["name"],
"column_count": len(schema_manager.get_table(table_info["name"]).columns),
})
# 批量计算embedding
logger.info(f"计算{len(table_texts)}张表的embedding...")
embeddings = self.embedder.encode(
table_texts, batch_size=batch_size, normalize=True
)
# 存入ChromaDB
self.collection.add(
embeddings=embeddings.tolist(),
documents=table_texts,
metadatas=table_metadatas,
ids=table_ids,
)
logger.info(f"[OK] 索引构建完成,共{len(table_ids)}张表")
return True
def search(
self,
query: str,
top_k: int = 20,
score_threshold: float = 0.3,
) -> List[Dict]:
"""
检索相关表
Args:
query: 查询文本(用户问题)
top_k: 返回前K个结果
score_threshold: 相似度阈值(0-1),低于此值的结果会被过滤
Returns:
检索结果列表,按相似度降序排列
[
{
"table_name": "表名",
"score": 0.95,
"document": "表描述文本",
"metadata": {...}
},
...
]
"""
# 编码查询文本(与建库时一致:L2 归一化 + 余弦空间)
query_embedding = self.embedder.encode([query], normalize=True)
# 执行检索
results = self.collection.query(
query_embeddings=query_embedding.tolist(),
n_results=min(top_k, self.collection.count()),
)
# 格式化结果
formatted = []
if results["ids"] and len(results["ids"][0]) > 0:
for idx, (table_id, distance, metadata, document) in enumerate(zip(
results["ids"][0],
results["distances"][0],
results["metadatas"][0],
results["documents"][0],
)):
score = 1 - distance # 余弦相似度(已归一化,转换为0-1,越大越相似)
if score >= score_threshold:
formatted.append({
"table_name": table_id,
"score": float(score),
"document": document,
"metadata": metadata,
"rank": idx + 1,
})
logger.debug(f"检索 '{query[:50]}...' -> 找到{len(formatted)}个相关表(阈值={score_threshold})")
return formatted
def search_by_table_names(self, table_names: List[str]) -> List[Dict]:
"""
直接通过表名获取表信息
Args:
table_names: 表名列表
Returns:
表信息列表
"""
results = []
for name in table_names:
try:
result = self.collection.get(ids=[name], include=["metadatas", "documents"])
if result["ids"]:
results.append({
"table_name": name,
"score": 1.0,
"document": result["documents"][0],
"metadata": result["metadatas"][0],
})
except Exception as e:
logger.warning(f"获取表信息失败 {name}: {e}")
return results
def get_all_tables(self) -> List[str]:
"""获取索引中的所有表名"""
result = self.collection.get()
return result["ids"] if result["ids"] else []
def delete_table(self, table_name: str):
"""从索引中删除表"""
self.collection.delete(ids=[table_name])
logger.info(f"已删除表索引: {table_name}")
def clear(self):
"""清空索引"""
if self.collection.count() > 0:
self.client.delete_collection(self.collection_name)
self.collection = self.client.get_or_create_collection(
name=self.collection_name,
metadata={"hnsw:space": "cosine"},
)
logger.info("索引已清空")
def count(self) -> int:
"""获取索引中的表数量"""
return self.collection.count()
# ========== 统计与调试 ==========
def get_statistics(self) -> Dict:
"""获取索引统计信息"""
count = self.collection.count()
result = self.collection.get(include=["metadatas"])
metas = result.get("metadatas") or []
def _col_count(m: Optional[Dict]) -> int:
if not m:
return 0
v = m.get("column_count", 0)
try:
return int(v)
except (TypeError, ValueError):
return 0
total_columns = sum(_col_count(m) for m in metas)
return {
"indexed_tables": count,
"total_columns": total_columns,
"avg_columns": total_columns / count if count > 0 else 0,
"persist_dir": str(self.persist_dir),
"chroma_mode": "persistent",
}