Spaces:
Sleeping
Sleeping
| import functools | |
| import json | |
| import logging | |
| import os | |
| import chromadb | |
| from llama_index.core import QueryBundle, VectorStoreIndex | |
| from llama_index.core.retrievers import BaseRetriever, VectorIndexRetriever | |
| from llama_index.core.schema import NodeWithScore, TextNode | |
| from llama_index.core.vector_stores import FilterOperator, MetadataFilter, MetadataFilters | |
| from llama_index.embeddings.huggingface import HuggingFaceEmbedding | |
| from llama_index.postprocessor.sbert_rerank import SentenceTransformerRerank | |
| from llama_index.retrievers.bm25 import BM25Retriever | |
| from llama_index.vector_stores.chroma import ChromaVectorStore | |
| from rag.config import ( | |
| BOOKS, | |
| CANDIDATE_TOP_K, | |
| CHROMA_COLLECTION, | |
| CHROMA_DIR, | |
| CHROMA_DISTANCE, | |
| CHUNKS_PATH, | |
| EMBED_MODEL, | |
| HEADING_SEPARATOR, | |
| HF_INDEX_REPO_ENV, | |
| QUERY_INSTRUCTION, | |
| RERANK_MODEL, | |
| ) | |
| from rag.types import ChunkMetadata | |
| RETRIEVAL_MODES = ("vector", "bm25", "hybrid", "rerank") | |
| logger = logging.getLogger(__name__) | |
| def make_embed_model() -> HuggingFaceEmbedding: | |
| return HuggingFaceEmbedding( | |
| model_name=EMBED_MODEL, | |
| query_instruction=QUERY_INSTRUCTION, | |
| ) | |
| def flatten_metadata(metadata: ChunkMetadata) -> dict[str, str | bool]: | |
| code_tags = metadata["code_tags"] | |
| return { | |
| "book": metadata["book"], | |
| "book_title": metadata["book_title"], | |
| "chapter": metadata["chapter"], | |
| "part": metadata["part"], | |
| "heading_path": HEADING_SEPARATOR.join(metadata["heading_path"]), | |
| "path": metadata["path"], | |
| "url": metadata["url"], | |
| "has_code": metadata["has_code"], | |
| "code_tags": " | ".join(code_tags), | |
| } | |
| def download_index_if_missing() -> None: | |
| if CHROMA_DIR.exists() and CHUNKS_PATH.exists(): | |
| return | |
| repo_id = os.environ.get(HF_INDEX_REPO_ENV, "").strip() | |
| if not repo_id: | |
| raise RuntimeError( | |
| f"No index at {CHROMA_DIR} and {HF_INDEX_REPO_ENV} is unset. " | |
| "Build one with `uv run python -m ingest.build_index`, or set " | |
| f"{HF_INDEX_REPO_ENV}." | |
| ) | |
| from huggingface_hub import snapshot_download | |
| # We download the whole data directory, not just the collection, because BM25 | |
| # builds its term index from `chunks.jsonl` rather than from the vector | |
| # database. | |
| logger.info("downloading index from %s", repo_id) | |
| CHROMA_DIR.parent.mkdir(parents=True, exist_ok=True) | |
| snapshot_download(repo_id=repo_id, repo_type="dataset", local_dir=str(CHROMA_DIR.parent)) | |
| def open_collection(fresh: bool = False): | |
| client = chromadb.PersistentClient(path=str(CHROMA_DIR)) | |
| if fresh: | |
| try: | |
| client.delete_collection(CHROMA_COLLECTION) | |
| except Exception: | |
| pass | |
| return client.get_or_create_collection( | |
| CHROMA_COLLECTION, metadata={"hnsw:space": CHROMA_DISTANCE} | |
| ) | |
| def make_vector_store(fresh: bool = False) -> ChromaVectorStore: | |
| return ChromaVectorStore(chroma_collection=open_collection(fresh=fresh)) | |
| def load_nodes() -> tuple[TextNode, ...]: | |
| with CHUNKS_PATH.open(encoding="utf-8") as handle: | |
| chunks = [json.loads(line) for line in handle] | |
| return tuple( | |
| TextNode(id_=chunk["id"], text=chunk["text"], metadata=flatten_metadata(chunk["metadata"])) | |
| for chunk in chunks | |
| ) | |
| def book_filter(book: str | None) -> MetadataFilters | None: | |
| if not book or book not in BOOKS: | |
| return None | |
| return MetadataFilters( | |
| filters=[MetadataFilter(key="book", value=book, operator=FilterOperator.EQ)] | |
| ) | |
| class RerankingRetriever(BaseRetriever): | |
| def __init__(self, retriever: BaseRetriever, reranker: SentenceTransformerRerank): | |
| self._retriever = retriever | |
| self._reranker = reranker | |
| super().__init__() | |
| def _retrieve(self, query_bundle: QueryBundle) -> list[NodeWithScore]: | |
| candidates = self._retriever.retrieve(query_bundle) | |
| return self._reranker.postprocess_nodes(candidates, query_bundle=query_bundle) | |
| # The constant from the original reciprocal-rank-fusion paper, and the same value | |
| # LlamaIndex uses. | |
| RRF_K = 60.0 | |
| class HybridRetriever(BaseRetriever): | |
| """Vector and BM25 candidates merged by reciprocal rank fusion. | |
| We write the fusion out ourselves rather than using `QueryFusionRetriever`. | |
| Its fusion is the same arithmetic, but it also generates query variations with | |
| an LLM, which would spend the user's key inside retrieval and may duplicate the | |
| query the agent already wrote. | |
| """ | |
| def __init__(self, retrievers: list[BaseRetriever], top_k: int): | |
| self._retrievers = retrievers | |
| self._top_k = top_k | |
| super().__init__() | |
| def _retrieve(self, query_bundle: QueryBundle) -> list[NodeWithScore]: | |
| scores: dict[str, float] = {} | |
| found: dict[str, NodeWithScore] = {} | |
| for retriever in self._retrievers: | |
| ranked = sorted( | |
| retriever.retrieve(query_bundle), key=lambda n: n.score or 0.0, reverse=True | |
| ) | |
| for rank, node in enumerate(ranked): | |
| key = node.node.hash | |
| found[key] = node | |
| scores[key] = scores.get(key, 0.0) + 1.0 / (rank + RRF_K) | |
| best = sorted(scores, key=lambda key: scores[key], reverse=True)[: self._top_k] | |
| return [NodeWithScore(node=found[key].node, score=scores[key]) for key in best] | |
| def _fused(final_top_k: int, book: str | None) -> HybridRetriever: | |
| embed_model = make_embed_model() | |
| filters = book_filter(book) | |
| index = VectorStoreIndex.from_vector_store( | |
| vector_store=make_vector_store(), embed_model=embed_model | |
| ) | |
| return HybridRetriever( | |
| retrievers=[ | |
| VectorIndexRetriever( | |
| index=index, | |
| similarity_top_k=CANDIDATE_TOP_K, | |
| embed_model=embed_model, | |
| filters=filters, | |
| ), | |
| BM25Retriever.from_defaults( | |
| nodes=list(load_nodes()), | |
| similarity_top_k=CANDIDATE_TOP_K, | |
| filters=filters, | |
| ), | |
| ], | |
| top_k=final_top_k, | |
| ) | |
| def load_retriever( | |
| mode: str, top_k: int, book: str | None = None, rerank_model: str | None = None | |
| ) -> BaseRetriever: | |
| if mode == "vector": | |
| embed_model = make_embed_model() | |
| index = VectorStoreIndex.from_vector_store( | |
| vector_store=make_vector_store(), embed_model=embed_model | |
| ) | |
| return VectorIndexRetriever( | |
| index=index, | |
| similarity_top_k=top_k, | |
| embed_model=embed_model, | |
| filters=book_filter(book), | |
| ) | |
| if mode == "bm25": | |
| return BM25Retriever.from_defaults( | |
| nodes=list(load_nodes()), similarity_top_k=top_k, filters=book_filter(book) | |
| ) | |
| if mode == "hybrid": | |
| return _fused(top_k, book) | |
| if mode == "rerank": | |
| return RerankingRetriever( | |
| retriever=_fused(CANDIDATE_TOP_K, book), | |
| reranker=SentenceTransformerRerank(model=rerank_model or RERANK_MODEL, top_n=top_k), | |
| ) | |
| raise ValueError(f"Unknown retrieval mode {mode!r}. Expected one of: {RETRIEVAL_MODES}") | |