pppp24
rust docs RAG pipeline, eval suite, and data ingestion
005e9fd
Raw
History Blame Contribute Delete
7.23 kB
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__)
@functools.cache
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))
@functools.cache
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,
)
@functools.cache
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}")