215 lines
7.3 KiB
Python
215 lines
7.3 KiB
Python
"""
|
|||
|
|
Few-shot 经验样本的 Chroma 向量库存储与检索。
|
||
|
|
|
||
|
|
与 SchemaIndexer 一致:使用项目统一 ``get_embedder()``,持久化目录默认
|
||
|
|
``./data/embeddings/chroma_fewshot``,集合名 ``fewshot_samples``。
|
||
|
|
"""
|
||
|
|
|
||
|
|
from __future__ import annotations
|
||
|
|
|
||
|
|
import json
|
||
|
|
import logging
|
||
|
|
import os
|
||
|
|
from pathlib import Path
|
||
|
|
from typing import Any, Dict, List, Optional, Tuple
|
||
|
|
|
||
|
|
import chromadb
|
||
|
|
from chromadb.config import Settings as ChromaSettings
|
||
|
|
|
||
|
|
logger = logging.getLogger(__name__)
|
||
|
|
|
||
|
|
DEFAULT_FEWSHOT_CHROMA_DIR = "./data/embeddings/chroma_fewshot"
|
||
|
|
COLLECTION_NAME = "fewshot_samples"
|
||
|
|
# Chroma metadata 单值不宜过大,SQL/说明超长时截断
|
||
|
|
_MAX_SQL_META = 16000
|
||
|
|
_MAX_EXPLAIN_META = 6000
|
||
|
|
_MAX_QEN_META = 2000
|
||
|
|
|
||
|
|
|
||
|
|
class FewShotChromaStore:
|
||
|
|
"""Few-shot JSONL → Chroma 写入与按向量检索。"""
|
||
|
|
|
||
|
|
def __init__(
|
||
|
|
self,
|
||
|
|
embedder: Any,
|
||
|
|
persist_dir: Optional[str] = None,
|
||
|
|
collection_name: Optional[str] = None,
|
||
|
|
):
|
||
|
|
self.embedder = embedder
|
||
|
|
# Chroma 的 collection 名;显式参数优先,否则读 FEWSHOT_CHROMA_COLLECTION,再回退默认
|
||
|
|
explicit = (collection_name or "").strip()
|
||
|
|
from_env = (os.getenv("FEWSHOT_CHROMA_COLLECTION") or "").strip()
|
||
|
|
self.collection_name = explicit or from_env or COLLECTION_NAME
|
||
|
|
self.persist_dir = Path(
|
||
|
|
(persist_dir or os.getenv("FEWSHOT_CHROMA_PATH") or DEFAULT_FEWSHOT_CHROMA_DIR).strip()
|
||
|
|
)
|
||
|
|
self.persist_dir.mkdir(parents=True, exist_ok=True)
|
||
|
|
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(
|
||
|
|
"[OK] FewShotChromaStore: path=%s collection=%s count=%s",
|
||
|
|
self.persist_dir,
|
||
|
|
self.collection_name,
|
||
|
|
self.collection.count(),
|
||
|
|
)
|
||
|
|
|
||
|
|
def count(self) -> int:
|
||
|
|
return int(self.collection.count())
|
||
|
|
|
||
|
|
def clear(self) -> None:
|
||
|
|
"""删除集合内全部文档(用于 force rebuild)。"""
|
||
|
|
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("[OK] Few-shot Chroma 集合已清空并重建")
|
||
|
|
|
||
|
|
def build_from_samples(
|
||
|
|
self,
|
||
|
|
samples: List["ExperienceSample"],
|
||
|
|
*,
|
||
|
|
force_rebuild: bool = False,
|
||
|
|
batch_size: int = 32,
|
||
|
|
) -> int:
|
||
|
|
"""
|
||
|
|
将已解析的样本列表写入 Chroma(向量由 question_zh 编码)。
|
||
|
|
|
||
|
|
Returns:
|
||
|
|
写入条数
|
||
|
|
"""
|
||
|
|
if not samples:
|
||
|
|
logger.warning("Few-shot Chroma: 无样本,跳过构建")
|
||
|
|
return 0
|
||
|
|
|
||
|
|
if force_rebuild and self.count() > 0:
|
||
|
|
self.clear()
|
||
|
|
|
||
|
|
if not force_rebuild and self.count() > 0:
|
||
|
|
logger.info("Few-shot Chroma 已有 %s 条,跳过构建(加 --force 可重建)", self.count())
|
||
|
|
return self.count()
|
||
|
|
|
||
|
|
texts = [(s.question_zh or "").strip() or " " for s in samples]
|
||
|
|
# Chroma 要求 ids 为唯一字符串
|
||
|
|
ids = [str(s.qid) for s in samples]
|
||
|
|
metadatas: List[Dict[str, Any]] = []
|
||
|
|
for s in samples:
|
||
|
|
metadatas.append(
|
||
|
|
{
|
||
|
|
"qid": str(s.qid),
|
||
|
|
"question_en": ((s.question_en or "")[:_MAX_QEN_META]),
|
||
|
|
"sql": (s.sql or "")[:_MAX_SQL_META],
|
||
|
|
"explanation": (s.explanation or "")[:_MAX_EXPLAIN_META],
|
||
|
|
"rating": int(s.rating) if s.rating is not None else -1,
|
||
|
|
"difficulty": (s.difficulty or "medium")[:32],
|
||
|
|
"tags_json": json.dumps(s.tags or [], ensure_ascii=False)[:4000],
|
||
|
|
}
|
||
|
|
)
|
||
|
|
|
||
|
|
logger.info("Few-shot Chroma: 计算 %s 条 embedding...", len(texts))
|
||
|
|
embeddings = self.embedder.encode(
|
||
|
|
texts,
|
||
|
|
batch_size=min(batch_size, len(texts)),
|
||
|
|
normalize=True,
|
||
|
|
show_progress=True,
|
||
|
|
)
|
||
|
|
|
||
|
|
self.collection.add(
|
||
|
|
embeddings=embeddings.tolist(),
|
||
|
|
documents=texts,
|
||
|
|
metadatas=metadatas,
|
||
|
|
ids=ids,
|
||
|
|
)
|
||
|
|
logger.info("[OK] Few-shot Chroma 索引完成: %s 条", len(ids))
|
||
|
|
return len(ids)
|
||
|
|
|
||
|
|
def search_raw(
|
||
|
|
self,
|
||
|
|
query: str,
|
||
|
|
top_k: int = 20,
|
||
|
|
) -> List[Tuple[float, Dict[str, Any], str]]:
|
||
|
|
"""
|
||
|
|
Returns:
|
||
|
|
(score, metadata_dict, document) 列表,score 同 SchemaIndexer 为 1-distance
|
||
|
|
"""
|
||
|
|
n = self.count()
|
||
|
|
if n == 0:
|
||
|
|
return []
|
||
|
|
q_emb = self.embedder.encode(
|
||
|
|
[(query or "").strip() or " "],
|
||
|
|
batch_size=1,
|
||
|
|
normalize=True,
|
||
|
|
show_progress=False,
|
||
|
|
)
|
||
|
|
k = min(max(top_k, 1), n)
|
||
|
|
results = self.collection.query(
|
||
|
|
query_embeddings=q_emb.tolist(),
|
||
|
|
n_results=k,
|
||
|
|
)
|
||
|
|
out: List[Tuple[float, Dict[str, Any], str]] = []
|
||
|
|
if not results["ids"] or not results["ids"][0]:
|
||
|
|
return out
|
||
|
|
for tid, dist, meta, doc in zip(
|
||
|
|
results["ids"][0],
|
||
|
|
results["distances"][0],
|
||
|
|
results["metadatas"][0],
|
||
|
|
results["documents"][0],
|
||
|
|
):
|
||
|
|
score = 1.0 - float(dist)
|
||
|
|
m = dict(meta or {})
|
||
|
|
m["qid"] = tid
|
||
|
|
out.append((score, m, doc or ""))
|
||
|
|
return out
|
||
|
|
|
||
|
|
def get_all_as_samples(self) -> List[Any]:
|
||
|
|
"""导出集合中全部样本(用于统计/按标签列举;条数大时慎用)。"""
|
||
|
|
if self.count() == 0:
|
||
|
|
return []
|
||
|
|
data = self.collection.get(include=["metadatas", "documents"])
|
||
|
|
ids = data.get("ids") or []
|
||
|
|
metas = data.get("metadatas") or []
|
||
|
|
docs = data.get("documents") or []
|
||
|
|
out: List[Any] = []
|
||
|
|
for i, qid in enumerate(ids):
|
||
|
|
meta = dict(metas[i] or {}) if i < len(metas) else {}
|
||
|
|
meta["qid"] = qid
|
||
|
|
doc = docs[i] if i < len(docs) else ""
|
||
|
|
out.append(sample_from_chroma_metadata(meta, doc))
|
||
|
|
return out
|
||
|
|
|
||
|
|
|
||
|
|
def sample_from_chroma_metadata(meta: Dict[str, Any], document: str) -> "ExperienceSample":
|
||
|
|
from utils.fewshot_selector import ExperienceSample
|
||
|
|
|
||
|
|
tags_raw = meta.get("tags_json") or "[]"
|
||
|
|
try:
|
||
|
|
tags = json.loads(tags_raw) if isinstance(tags_raw, str) else []
|
||
|
|
except json.JSONDecodeError:
|
||
|
|
tags = []
|
||
|
|
if not isinstance(tags, list):
|
||
|
|
tags = []
|
||
|
|
r = meta.get("rating")
|
||
|
|
try:
|
||
|
|
ri = int(r) if r is not None else -1
|
||
|
|
except (TypeError, ValueError):
|
||
|
|
ri = -1
|
||
|
|
rating = ri if ri >= 0 else None
|
||
|
|
qen = meta.get("question_en")
|
||
|
|
return ExperienceSample(
|
||
|
|
qid=str(meta.get("qid", "")),
|
||
|
|
question_zh=document or str(meta.get("question_zh", "")),
|
||
|
|
question_en=qen if qen else None,
|
||
|
|
sql=str(meta.get("sql", "")),
|
||
|
|
explanation=str(meta.get("explanation", "")),
|
||
|
|
rating=rating,
|
||
|
|
tags=tags,
|
||
|
|
difficulty=str(meta.get("difficulty", "medium") or "medium"),
|
||
|
|
)
|