Pipeline
Composable baseline and advanced retrieval-augmented generation pipelines.
"""Composable baseline and advanced retrieval-augmented generation pipelines."""
from dataclasses import replace
from typing import List, Optional, Sequence
from .context import ContextPacker
from .generation import ExtractiveGenerator, Generator
from .models import Answer, Document, SearchResult, TraceEvent
from .rerankers import CrossFeatureReranker
from .retrievers import (
AdaptiveRetriever,
BM25Retriever,
GraphExpandedRetriever,
HashingSemanticRetriever,
HybridRetriever,
MultiQueryRetriever,
Retriever,
simple_query_variants,
)
from .text import chunk_documents
class RAGPipeline:
"""Retrieve, optionally rerank, pack evidence, and generate with a trace."""
def __init__(
self,
retriever: Retriever,
generator: Generator,
context_packer: ContextPacker,
reranker: Optional[CrossFeatureReranker] = None,
retrieval_k: int = 12,
rerank_k: int = 8,
) -> None:
self.retriever = retriever
self.generator = generator
self.context_packer = context_packer
self.reranker = reranker
self.retrieval_k = retrieval_k
self.rerank_k = rerank_k
def search(self, query: str, k: Optional[int] = None) -> List[SearchResult]:
"""Run retrieval and reranking without context packing or generation."""
requested = k or self.rerank_k
candidates = self.retriever.search(query, max(self.retrieval_k, requested))
if self.reranker is not None:
return self.reranker.rerank(query, candidates, requested)
return candidates[:requested]
def ask(self, query: str) -> Answer:
route = "fixed"
route_method = getattr(self.retriever, "route_for", None)
if callable(route_method):
route = str(route_method(query))
raw = self.retriever.search(query, self.retrieval_k)
trace = [
TraceEvent(
stage="route",
detail=f"selected {route} retrieval",
values={"route": route},
),
TraceEvent(
stage="retrieve",
detail=f"retrieved {len(raw)} candidates",
values={"count": len(raw), "k": self.retrieval_k},
),
]
ranked = raw
if self.reranker is not None:
ranked = self.reranker.rerank(query, raw, self.rerank_k)
trace.append(
TraceEvent(
stage="rerank",
detail=f"retained {len(ranked)} reranked candidates",
values={"count": len(ranked), "k": self.rerank_k},
)
)
packed = self.context_packer.pack(ranked)
trace.append(
TraceEvent(
stage="pack",
detail=f"packed {len(packed)} non-redundant chunks",
values={
"count": len(packed),
"max_tokens": self.context_packer.max_tokens,
},
)
)
answer = self.generator.generate(query, packed)
trace.append(
TraceEvent(
stage="generate",
detail="abstained" if answer.abstained else "returned grounded evidence",
values={
"citations": len(answer.citations),
"confidence": round(answer.confidence, 4),
"abstained": answer.abstained,
},
)
)
return replace(answer, trace=tuple(trace))
def build_baseline_pipeline(documents: Sequence[Document]) -> RAGPipeline:
"""A 2020-style retrieve-then-generate baseline with sparse retrieval."""
chunks = chunk_documents(documents, chunk_size=120, overlap=20)
return RAGPipeline(
retriever=BM25Retriever(chunks),
generator=ExtractiveGenerator(max_sentences=2),
context_packer=ContextPacker(max_tokens=450, max_chunks=5, relevance_weight=0.9),
reranker=None,
retrieval_k=5,
rerank_k=5,
)
def build_advanced_pipeline(documents: Sequence[Document]) -> RAGPipeline:
"""A modern modular stack: hybrid, multi-query, graph routing, reranking, MMR."""
chunks = chunk_documents(documents, chunk_size=90, overlap=18)
sparse = BM25Retriever(chunks)
semantic = HashingSemanticRetriever(chunks)
hybrid = HybridRetriever(
(("sparse", sparse, 1.0), ("semantic", semantic, 1.0)),
rrf_constant=30,
)
multi_query = MultiQueryRetriever(hybrid, simple_query_variants, rrf_constant=30)
graph = GraphExpandedRetriever(multi_query, chunks, expansion_weight=0.30)
adaptive = AdaptiveRetriever(sparse=sparse, hybrid=multi_query, graph=graph)
return RAGPipeline(
retriever=adaptive,
generator=ExtractiveGenerator(max_sentences=2),
context_packer=ContextPacker(
max_tokens=560,
max_chunks=6,
relevance_weight=0.78,
duplicate_threshold=0.86,
),
reranker=CrossFeatureReranker(),
retrieval_k=14,
rerank_k=8,
)