Files
ai-g3sb-backman2.0/backend/utils/fewshot_chroma_store.py
T

244 lines
8.5 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.
"""
Few-shot 经验样本的 Chroma 向量库存储与检索。
与 SchemaIndexer 共用 ``get_embedder()``。Few-shot 默认使用 ``PersistentClient``
(``FEWSHOT_CHROMA_PATH`` / ``./data/embeddings/chroma_fewshot``),与灌库脚本写入目录一致;
仅当 ``FEWSHOT_CHROMA_EPHEMERAL=true`` 或显式 ``persist_to_disk=False`` 时使用内存 Chroma。
集合名默认 ``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,
*,
persist_to_disk: Optional[bool] = None,
):
"""
Args:
embedder: 编码器实例。
persist_dir: Chroma 持久化根目录;内存模式下仍解析为配置引用路径。
collection_name: 集合名;可空,空则读环境变量或默认。
persist_to_disk: 为 True/False 时强制磁盘或内存;为 None 时默认磁盘,除非环境变量
``FEWSHOT_CHROMA_EPHEMERAL`` 为真。
"""
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()
)
env_ephemeral = os.getenv("FEWSHOT_CHROMA_EPHEMERAL", "").lower() in (
"1",
"true",
"yes",
)
if persist_to_disk is None:
self.persist_to_disk = not env_ephemeral
else:
self.persist_to_disk = bool(persist_to_disk)
if self.persist_to_disk:
self.persist_dir.mkdir(parents=True, exist_ok=True)
self.client = chromadb.PersistentClient(
path=str(self.persist_dir),
settings=ChromaSettings(anonymized_telemetry=False),
)
else:
self.client = chromadb.EphemeralClient(
settings=ChromaSettings(anonymized_telemetry=False),
)
self.collection = self.client.get_or_create_collection(
name=self.collection_name,
metadata={"hnsw:space": "cosine"},
)
backend = (
f"磁盘 {self.persist_dir}"
if self.persist_to_disk
else f"内存(配置路径={self.persist_dir})"
)
logger.info(
f"[OK] FewShotChromaStore: {backend} collection={self.collection_name} "
f"count={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"),
)