Spaces:
Sleeping
Sleeping
Upload folder using huggingface_hub
Browse files- rag_system/cache.py +7 -2
- rag_system/models.py +22 -0
- rag_system/query_engine.py +54 -6
- rag_system/retriever.py +44 -8
rag_system/cache.py
CHANGED
|
@@ -101,6 +101,10 @@ def _exact_key(query: str, collection_key: str, params_key: str) -> str:
|
|
| 101 |
return EXACT_PREFIX + hashlib.sha256(payload.encode()).hexdigest()[:32]
|
| 102 |
|
| 103 |
|
|
|
|
|
|
|
|
|
|
|
|
|
| 104 |
def _vector_bytes(vec: list[float]) -> bytes:
|
| 105 |
return np.array(vec, dtype=np.float32).tobytes()
|
| 106 |
|
|
@@ -146,9 +150,10 @@ def get_semantic(query_vec: list[float], collection_key: str, params_key: str) -
|
|
| 146 |
|
| 147 |
vec = _vector_bytes(query_vec)
|
| 148 |
k = 4
|
|
|
|
| 149 |
query = (
|
| 150 |
f"@{COLLECTION_FIELD}:{{{collection_key}}} "
|
| 151 |
-
f"@{PARAMS_FIELD}:{{{
|
| 152 |
f"=>[KNN {k} @{VECTOR_FIELD} $vec AS score]"
|
| 153 |
)
|
| 154 |
|
|
@@ -223,7 +228,7 @@ def set_semantic(
|
|
| 223 |
mapping={
|
| 224 |
"query": query,
|
| 225 |
COLLECTION_FIELD: collection_key,
|
| 226 |
-
PARAMS_FIELD: params_key,
|
| 227 |
VECTOR_FIELD: _vector_bytes(query_vec),
|
| 228 |
PAYLOAD_FIELD: payload,
|
| 229 |
},
|
|
|
|
| 101 |
return EXACT_PREFIX + hashlib.sha256(payload.encode()).hexdigest()[:32]
|
| 102 |
|
| 103 |
|
| 104 |
+
def _tag_hash(value: str) -> str:
|
| 105 |
+
return hashlib.sha1(value.encode()).hexdigest()[:16]
|
| 106 |
+
|
| 107 |
+
|
| 108 |
def _vector_bytes(vec: list[float]) -> bytes:
|
| 109 |
return np.array(vec, dtype=np.float32).tobytes()
|
| 110 |
|
|
|
|
| 150 |
|
| 151 |
vec = _vector_bytes(query_vec)
|
| 152 |
k = 4
|
| 153 |
+
params_tag = _tag_hash(params_key)
|
| 154 |
query = (
|
| 155 |
f"@{COLLECTION_FIELD}:{{{collection_key}}} "
|
| 156 |
+
f"@{PARAMS_FIELD}:{{{params_tag}}} "
|
| 157 |
f"=>[KNN {k} @{VECTOR_FIELD} $vec AS score]"
|
| 158 |
)
|
| 159 |
|
|
|
|
| 228 |
mapping={
|
| 229 |
"query": query,
|
| 230 |
COLLECTION_FIELD: collection_key,
|
| 231 |
+
PARAMS_FIELD: _tag_hash(params_key),
|
| 232 |
VECTOR_FIELD: _vector_bytes(query_vec),
|
| 233 |
PAYLOAD_FIELD: payload,
|
| 234 |
},
|
rag_system/models.py
CHANGED
|
@@ -47,6 +47,10 @@ class QueryRequest(BaseModel):
|
|
| 47 |
retrieval_mode: RetrievalMode = RetrievalMode.HYBRID
|
| 48 |
embedding_mode: Optional[EmbeddingMode] = None
|
| 49 |
top_k: Optional[int] = None
|
|
|
|
|
|
|
|
|
|
|
|
|
| 50 |
doc_collections: Optional[list[str]] = None # per-doc sub-collections; None = legacy single-collection mode
|
| 51 |
history: list[ChatMessage] = Field(default_factory=list)
|
| 52 |
stream: bool = False
|
|
@@ -56,6 +60,24 @@ class QueryRequest(BaseModel):
|
|
| 56 |
def sanitize_query(cls,v):
|
| 57 |
return v.strip()
|
| 58 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 59 |
class SourceDocument(BaseModel):
|
| 60 |
doc_id: str
|
| 61 |
content: str
|
|
|
|
| 47 |
retrieval_mode: RetrievalMode = RetrievalMode.HYBRID
|
| 48 |
embedding_mode: Optional[EmbeddingMode] = None
|
| 49 |
top_k: Optional[int] = None
|
| 50 |
+
top_k_retrieval: Optional[int] = None
|
| 51 |
+
mmr_lambda: Optional[float] = None
|
| 52 |
+
bm25_weight: Optional[float] = None
|
| 53 |
+
vector_weight: Optional[float] = None
|
| 54 |
doc_collections: Optional[list[str]] = None # per-doc sub-collections; None = legacy single-collection mode
|
| 55 |
history: list[ChatMessage] = Field(default_factory=list)
|
| 56 |
stream: bool = False
|
|
|
|
| 60 |
def sanitize_query(cls,v):
|
| 61 |
return v.strip()
|
| 62 |
|
| 63 |
+
@field_validator("top_k", "top_k_retrieval")
|
| 64 |
+
@classmethod
|
| 65 |
+
def validate_top_k(cls, v):
|
| 66 |
+
if v is None:
|
| 67 |
+
return v
|
| 68 |
+
if v < 1:
|
| 69 |
+
raise ValueError("top_k values must be >= 1")
|
| 70 |
+
return v
|
| 71 |
+
|
| 72 |
+
@field_validator("mmr_lambda", "bm25_weight", "vector_weight")
|
| 73 |
+
@classmethod
|
| 74 |
+
def validate_weights(cls, v):
|
| 75 |
+
if v is None:
|
| 76 |
+
return v
|
| 77 |
+
if v < 0 or v > 1:
|
| 78 |
+
raise ValueError("weights must be between 0 and 1")
|
| 79 |
+
return v
|
| 80 |
+
|
| 81 |
class SourceDocument(BaseModel):
|
| 82 |
doc_id: str
|
| 83 |
content: str
|
rag_system/query_engine.py
CHANGED
|
@@ -71,9 +71,26 @@ def _cache_collection_key(collections: list[str]) -> str:
|
|
| 71 |
return hashlib.sha1(raw.encode()).hexdigest()[:16]
|
| 72 |
|
| 73 |
|
| 74 |
-
def
|
| 75 |
-
|
| 76 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 77 |
|
| 78 |
# Query rewriting
|
| 79 |
async def rewrite_query(query: str) -> str:
|
|
@@ -195,7 +212,14 @@ async def query(
|
|
| 195 |
mode_val = request.retrieval_mode.value if hasattr(request.retrieval_mode, "value") else str(request.retrieval_mode)
|
| 196 |
cache_allowed = settings.cache_enabled and _is_try_docs_scope(collections)
|
| 197 |
cache_collection_key = _cache_collection_key(collections)
|
| 198 |
-
cache_params_key =
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 199 |
cache_query_vec = None
|
| 200 |
|
| 201 |
# 1. Input guardrail
|
|
@@ -250,6 +274,10 @@ async def query(
|
|
| 250 |
collections=scoped,
|
| 251 |
mode=request.retrieval_mode.value if hasattr(request.retrieval_mode, "value") else str(request.retrieval_mode),
|
| 252 |
k_per_collection=k_per,
|
|
|
|
|
|
|
|
|
|
|
|
|
| 253 |
use_reranker=True,
|
| 254 |
expand_context=True,
|
| 255 |
)
|
|
@@ -260,6 +288,10 @@ async def query(
|
|
| 260 |
collection=collections[0],
|
| 261 |
mode=request.retrieval_mode,
|
| 262 |
top_k=request.top_k,
|
|
|
|
|
|
|
|
|
|
|
|
|
| 263 |
use_reranker=True,
|
| 264 |
expand_context=True,
|
| 265 |
)
|
|
@@ -368,7 +400,14 @@ async def pipeline_stream_query(request: QueryRequest) -> AsyncIterator[str]:
|
|
| 368 |
embedding_mode = resolve_embedding_mode_for_collections(collections, request.embedding_mode)
|
| 369 |
cache_allowed = settings.cache_enabled and _is_try_docs_scope(collections)
|
| 370 |
cache_collection_key = _cache_collection_key(collections)
|
| 371 |
-
cache_params_key =
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 372 |
|
| 373 |
yield emit("pipeline_start", "in_progress", {
|
| 374 |
"query": request.query,
|
|
@@ -451,7 +490,8 @@ async def pipeline_stream_query(request: QueryRequest) -> AsyncIterator[str]:
|
|
| 451 |
# --- Retrieval ---
|
| 452 |
yield emit("retrieval_start", "in_progress", {
|
| 453 |
"mode": mode_val,
|
| 454 |
-
"top_k": request.top_k or settings.
|
|
|
|
| 455 |
"collections": len(scoped),
|
| 456 |
})
|
| 457 |
|
|
@@ -462,6 +502,10 @@ async def pipeline_stream_query(request: QueryRequest) -> AsyncIterator[str]:
|
|
| 462 |
collections=scoped,
|
| 463 |
mode=mode_val,
|
| 464 |
k_per_collection=k_per,
|
|
|
|
|
|
|
|
|
|
|
|
|
| 465 |
use_reranker=True,
|
| 466 |
expand_context=True,
|
| 467 |
)
|
|
@@ -471,6 +515,10 @@ async def pipeline_stream_query(request: QueryRequest) -> AsyncIterator[str]:
|
|
| 471 |
collection=scoped[0],
|
| 472 |
mode=request.retrieval_mode,
|
| 473 |
top_k=request.top_k,
|
|
|
|
|
|
|
|
|
|
|
|
|
| 474 |
use_reranker=True,
|
| 475 |
expand_context=True,
|
| 476 |
)
|
|
|
|
| 71 |
return hashlib.sha1(raw.encode()).hexdigest()[:16]
|
| 72 |
|
| 73 |
|
| 74 |
+
def _fmt_param(value: Optional[float]) -> str:
|
| 75 |
+
if value is None:
|
| 76 |
+
return "-"
|
| 77 |
+
return f"{value:.3f}"
|
| 78 |
+
|
| 79 |
+
|
| 80 |
+
def _cache_params_key_v2(
|
| 81 |
+
mode: str,
|
| 82 |
+
top_k: Optional[int],
|
| 83 |
+
top_k_retrieval: Optional[int],
|
| 84 |
+
mmr_lambda: Optional[float],
|
| 85 |
+
bm25_weight: Optional[float],
|
| 86 |
+
vector_weight: Optional[float],
|
| 87 |
+
) -> str:
|
| 88 |
+
k_final = top_k if top_k is not None else settings.top_k_rerank
|
| 89 |
+
k_retrieve = top_k_retrieval if top_k_retrieval is not None else settings.top_k_retrieval
|
| 90 |
+
return (
|
| 91 |
+
f"{mode}:{k_final}:{k_retrieve}:"
|
| 92 |
+
f"{_fmt_param(mmr_lambda)}:{_fmt_param(bm25_weight)}:{_fmt_param(vector_weight)}"
|
| 93 |
+
)
|
| 94 |
|
| 95 |
# Query rewriting
|
| 96 |
async def rewrite_query(query: str) -> str:
|
|
|
|
| 212 |
mode_val = request.retrieval_mode.value if hasattr(request.retrieval_mode, "value") else str(request.retrieval_mode)
|
| 213 |
cache_allowed = settings.cache_enabled and _is_try_docs_scope(collections)
|
| 214 |
cache_collection_key = _cache_collection_key(collections)
|
| 215 |
+
cache_params_key = _cache_params_key_v2(
|
| 216 |
+
mode_val,
|
| 217 |
+
request.top_k,
|
| 218 |
+
request.top_k_retrieval,
|
| 219 |
+
request.mmr_lambda,
|
| 220 |
+
request.bm25_weight,
|
| 221 |
+
request.vector_weight,
|
| 222 |
+
)
|
| 223 |
cache_query_vec = None
|
| 224 |
|
| 225 |
# 1. Input guardrail
|
|
|
|
| 274 |
collections=scoped,
|
| 275 |
mode=request.retrieval_mode.value if hasattr(request.retrieval_mode, "value") else str(request.retrieval_mode),
|
| 276 |
k_per_collection=k_per,
|
| 277 |
+
top_k_retrieval=request.top_k_retrieval,
|
| 278 |
+
mmr_lambda=request.mmr_lambda,
|
| 279 |
+
bm25_weight=request.bm25_weight,
|
| 280 |
+
vector_weight=request.vector_weight,
|
| 281 |
use_reranker=True,
|
| 282 |
expand_context=True,
|
| 283 |
)
|
|
|
|
| 288 |
collection=collections[0],
|
| 289 |
mode=request.retrieval_mode,
|
| 290 |
top_k=request.top_k,
|
| 291 |
+
top_k_retrieval=request.top_k_retrieval,
|
| 292 |
+
mmr_lambda=request.mmr_lambda,
|
| 293 |
+
bm25_weight=request.bm25_weight,
|
| 294 |
+
vector_weight=request.vector_weight,
|
| 295 |
use_reranker=True,
|
| 296 |
expand_context=True,
|
| 297 |
)
|
|
|
|
| 400 |
embedding_mode = resolve_embedding_mode_for_collections(collections, request.embedding_mode)
|
| 401 |
cache_allowed = settings.cache_enabled and _is_try_docs_scope(collections)
|
| 402 |
cache_collection_key = _cache_collection_key(collections)
|
| 403 |
+
cache_params_key = _cache_params_key_v2(
|
| 404 |
+
mode_val,
|
| 405 |
+
request.top_k,
|
| 406 |
+
request.top_k_retrieval,
|
| 407 |
+
request.mmr_lambda,
|
| 408 |
+
request.bm25_weight,
|
| 409 |
+
request.vector_weight,
|
| 410 |
+
)
|
| 411 |
|
| 412 |
yield emit("pipeline_start", "in_progress", {
|
| 413 |
"query": request.query,
|
|
|
|
| 490 |
# --- Retrieval ---
|
| 491 |
yield emit("retrieval_start", "in_progress", {
|
| 492 |
"mode": mode_val,
|
| 493 |
+
"top_k": request.top_k or settings.top_k_rerank,
|
| 494 |
+
"top_k_retrieval": request.top_k_retrieval or settings.top_k_retrieval,
|
| 495 |
"collections": len(scoped),
|
| 496 |
})
|
| 497 |
|
|
|
|
| 502 |
collections=scoped,
|
| 503 |
mode=mode_val,
|
| 504 |
k_per_collection=k_per,
|
| 505 |
+
top_k_retrieval=request.top_k_retrieval,
|
| 506 |
+
mmr_lambda=request.mmr_lambda,
|
| 507 |
+
bm25_weight=request.bm25_weight,
|
| 508 |
+
vector_weight=request.vector_weight,
|
| 509 |
use_reranker=True,
|
| 510 |
expand_context=True,
|
| 511 |
)
|
|
|
|
| 515 |
collection=scoped[0],
|
| 516 |
mode=request.retrieval_mode,
|
| 517 |
top_k=request.top_k,
|
| 518 |
+
top_k_retrieval=request.top_k_retrieval,
|
| 519 |
+
mmr_lambda=request.mmr_lambda,
|
| 520 |
+
bm25_weight=request.bm25_weight,
|
| 521 |
+
vector_weight=request.vector_weight,
|
| 522 |
use_reranker=True,
|
| 523 |
expand_context=True,
|
| 524 |
)
|
rag_system/retriever.py
CHANGED
|
@@ -49,13 +49,30 @@ def bm25_retrieve(query: str, collection: str,k: int) -> list[tuple[Document,flo
|
|
| 49 |
def _rrf_score(rank: int, k: int = 60) -> float:
|
| 50 |
return 1.0 / (k + rank + 1) #here 1 is added to handle rank 1 which here comes as 0
|
| 51 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 52 |
#Hybrid Retrieval
|
| 53 |
def hybrid_retrieve(
|
| 54 |
query: str,
|
| 55 |
collection: str,
|
| 56 |
-
k: int
|
|
|
|
|
|
|
| 57 |
) -> list[tuple[Document,float]]:
|
| 58 |
"""Reciprocal Rank fusion of BM25 and Vector results"""
|
|
|
|
| 59 |
pool_size = k*3 #casting a wide net before fusing
|
| 60 |
vec_results = similarity_search_with_scores(query,collection,k=pool_size)
|
| 61 |
bm25_results = bm25_retrieve(query,collection,k=pool_size)
|
|
@@ -65,12 +82,12 @@ def hybrid_retrieve(
|
|
| 65 |
|
| 66 |
for rank, (doc, _) in enumerate(vec_results):
|
| 67 |
did = doc.metadata.get("doc_id",id(doc))
|
| 68 |
-
rrf_scores[did] = rrf_scores.get(did,0) +
|
| 69 |
doc_map[did] = doc
|
| 70 |
|
| 71 |
for rank, (doc, _) in enumerate(bm25_results):
|
| 72 |
did = doc.metadata.get("doc_id",id(doc))
|
| 73 |
-
rrf_scores[did] = rrf_scores.get(did,0) +
|
| 74 |
doc_map[did] = doc
|
| 75 |
|
| 76 |
sorted_ids = sorted(rrf_scores, key=lambda x: rrf_scores[x], reverse=True)[:k]
|
|
@@ -81,14 +98,15 @@ async def mmr_retrieve(
|
|
| 81 |
query: str,
|
| 82 |
collection: str,
|
| 83 |
k: int,
|
| 84 |
-
lambda_mult: float = None,
|
|
|
|
| 85 |
) -> list[tuple[Document, float]]:
|
| 86 |
"""
|
| 87 |
Maximal Marginal Relevance: balance relevance vs diversity.
|
| 88 |
"""
|
| 89 |
# 1. Setup parameters
|
| 90 |
lam = lambda_mult or settings.mmr_lambda
|
| 91 |
-
fetch_k = settings.top_k_retrieval
|
| 92 |
store = get_store(collection)
|
| 93 |
|
| 94 |
# 2. Get original scores to map them back later
|
|
@@ -195,10 +213,14 @@ async def retrieve(
|
|
| 195 |
collection: str = "default",
|
| 196 |
mode: str = "hybrid",
|
| 197 |
top_k: Optional[int] = None,
|
|
|
|
|
|
|
|
|
|
|
|
|
| 198 |
use_reranker: bool = True,
|
| 199 |
expand_context: bool = True,
|
| 200 |
) -> list[tuple[Document,float]]:
|
| 201 |
-
k_retrieve = settings.top_k_retrieval
|
| 202 |
k_final = top_k or settings.top_k_rerank
|
| 203 |
|
| 204 |
if mode == "vector":
|
|
@@ -206,10 +228,16 @@ async def retrieve(
|
|
| 206 |
elif mode == "bm25":
|
| 207 |
results = bm25_retrieve(query,collection,k=k_retrieve)
|
| 208 |
elif mode == "mmr":
|
| 209 |
-
results = await mmr_retrieve(query,collection,k=k_final)
|
| 210 |
return results # MMR alrady handles diversity , skip reranker
|
| 211 |
else: #go with hybrid
|
| 212 |
-
results = hybrid_retrieve(
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 213 |
|
| 214 |
#Rerank
|
| 215 |
if use_reranker:
|
|
@@ -259,6 +287,10 @@ async def multi_collection_retrieve(
|
|
| 259 |
collections: list[str],
|
| 260 |
mode: str = "hybrid",
|
| 261 |
k_per_collection: int = 5,
|
|
|
|
|
|
|
|
|
|
|
|
|
| 262 |
use_reranker: bool = True,
|
| 263 |
expand_context: bool = True,
|
| 264 |
) -> list[tuple[Document, float]]:
|
|
@@ -275,6 +307,10 @@ async def multi_collection_retrieve(
|
|
| 275 |
collection=coll,
|
| 276 |
mode=mode,
|
| 277 |
top_k=k_per_collection,
|
|
|
|
|
|
|
|
|
|
|
|
|
| 278 |
use_reranker=False, # defer reranking until after merge
|
| 279 |
expand_context=expand_context,
|
| 280 |
)
|
|
|
|
| 49 |
def _rrf_score(rank: int, k: int = 60) -> float:
|
| 50 |
return 1.0 / (k + rank + 1) #here 1 is added to handle rank 1 which here comes as 0
|
| 51 |
|
| 52 |
+
|
| 53 |
+
def _resolve_hybrid_weights(
|
| 54 |
+
bm25_weight: Optional[float],
|
| 55 |
+
vector_weight: Optional[float],
|
| 56 |
+
) -> tuple[float, float]:
|
| 57 |
+
bw = settings.bm25_weight if bm25_weight is None else bm25_weight
|
| 58 |
+
vw = settings.vector_weight if vector_weight is None else vector_weight
|
| 59 |
+
total = bw + vw
|
| 60 |
+
if total <= 0:
|
| 61 |
+
bw = settings.bm25_weight
|
| 62 |
+
vw = settings.vector_weight
|
| 63 |
+
total = bw + vw
|
| 64 |
+
return bw / total, vw / total
|
| 65 |
+
|
| 66 |
#Hybrid Retrieval
|
| 67 |
def hybrid_retrieve(
|
| 68 |
query: str,
|
| 69 |
collection: str,
|
| 70 |
+
k: int,
|
| 71 |
+
bm25_weight: Optional[float] = None,
|
| 72 |
+
vector_weight: Optional[float] = None,
|
| 73 |
) -> list[tuple[Document,float]]:
|
| 74 |
"""Reciprocal Rank fusion of BM25 and Vector results"""
|
| 75 |
+
bw, vw = _resolve_hybrid_weights(bm25_weight, vector_weight)
|
| 76 |
pool_size = k*3 #casting a wide net before fusing
|
| 77 |
vec_results = similarity_search_with_scores(query,collection,k=pool_size)
|
| 78 |
bm25_results = bm25_retrieve(query,collection,k=pool_size)
|
|
|
|
| 82 |
|
| 83 |
for rank, (doc, _) in enumerate(vec_results):
|
| 84 |
did = doc.metadata.get("doc_id",id(doc))
|
| 85 |
+
rrf_scores[did] = rrf_scores.get(did,0) + vw * _rrf_score(rank) #check the existing score first and then add the fresh score
|
| 86 |
doc_map[did] = doc
|
| 87 |
|
| 88 |
for rank, (doc, _) in enumerate(bm25_results):
|
| 89 |
did = doc.metadata.get("doc_id",id(doc))
|
| 90 |
+
rrf_scores[did] = rrf_scores.get(did,0) + bw * _rrf_score(rank)
|
| 91 |
doc_map[did] = doc
|
| 92 |
|
| 93 |
sorted_ids = sorted(rrf_scores, key=lambda x: rrf_scores[x], reverse=True)[:k]
|
|
|
|
| 98 |
query: str,
|
| 99 |
collection: str,
|
| 100 |
k: int,
|
| 101 |
+
lambda_mult: Optional[float] = None,
|
| 102 |
+
fetch_k: Optional[int] = None,
|
| 103 |
) -> list[tuple[Document, float]]:
|
| 104 |
"""
|
| 105 |
Maximal Marginal Relevance: balance relevance vs diversity.
|
| 106 |
"""
|
| 107 |
# 1. Setup parameters
|
| 108 |
lam = lambda_mult or settings.mmr_lambda
|
| 109 |
+
fetch_k = fetch_k or settings.top_k_retrieval
|
| 110 |
store = get_store(collection)
|
| 111 |
|
| 112 |
# 2. Get original scores to map them back later
|
|
|
|
| 213 |
collection: str = "default",
|
| 214 |
mode: str = "hybrid",
|
| 215 |
top_k: Optional[int] = None,
|
| 216 |
+
top_k_retrieval: Optional[int] = None,
|
| 217 |
+
mmr_lambda: Optional[float] = None,
|
| 218 |
+
bm25_weight: Optional[float] = None,
|
| 219 |
+
vector_weight: Optional[float] = None,
|
| 220 |
use_reranker: bool = True,
|
| 221 |
expand_context: bool = True,
|
| 222 |
) -> list[tuple[Document,float]]:
|
| 223 |
+
k_retrieve = top_k_retrieval or settings.top_k_retrieval
|
| 224 |
k_final = top_k or settings.top_k_rerank
|
| 225 |
|
| 226 |
if mode == "vector":
|
|
|
|
| 228 |
elif mode == "bm25":
|
| 229 |
results = bm25_retrieve(query,collection,k=k_retrieve)
|
| 230 |
elif mode == "mmr":
|
| 231 |
+
results = await mmr_retrieve(query,collection,k=k_final,lambda_mult=mmr_lambda,fetch_k=k_retrieve)
|
| 232 |
return results # MMR alrady handles diversity , skip reranker
|
| 233 |
else: #go with hybrid
|
| 234 |
+
results = hybrid_retrieve(
|
| 235 |
+
query,
|
| 236 |
+
collection,
|
| 237 |
+
k=k_retrieve,
|
| 238 |
+
bm25_weight=bm25_weight,
|
| 239 |
+
vector_weight=vector_weight,
|
| 240 |
+
)
|
| 241 |
|
| 242 |
#Rerank
|
| 243 |
if use_reranker:
|
|
|
|
| 287 |
collections: list[str],
|
| 288 |
mode: str = "hybrid",
|
| 289 |
k_per_collection: int = 5,
|
| 290 |
+
top_k_retrieval: Optional[int] = None,
|
| 291 |
+
mmr_lambda: Optional[float] = None,
|
| 292 |
+
bm25_weight: Optional[float] = None,
|
| 293 |
+
vector_weight: Optional[float] = None,
|
| 294 |
use_reranker: bool = True,
|
| 295 |
expand_context: bool = True,
|
| 296 |
) -> list[tuple[Document, float]]:
|
|
|
|
| 307 |
collection=coll,
|
| 308 |
mode=mode,
|
| 309 |
top_k=k_per_collection,
|
| 310 |
+
top_k_retrieval=top_k_retrieval,
|
| 311 |
+
mmr_lambda=mmr_lambda,
|
| 312 |
+
bm25_weight=bm25_weight,
|
| 313 |
+
vector_weight=vector_weight,
|
| 314 |
use_reranker=False, # defer reranking until after merge
|
| 315 |
expand_context=expand_context,
|
| 316 |
)
|