best: score 0.5496 (recall 0.5698, ndcg 0.4690)

Search improvements on top of v1.0 index:
- RERANK_LIMIT 17 → 25
- prefilter with keyword-boosted stragglers (KEYWORD_BOOST_EXTRA=10)
- dense/sparse queries prefer search_text, keywords always in sparse
- variants+hyde as extra dense queries (up to 3)
- message_id score aggregation (rerank head full RRF, tail with k=60)

Co-Authored-By: Claude Sonnet 4.6 <noreply@anthropic.com>
This commit is contained in:
q 2026-04-19 12:42:29 +03:00
parent 92cde65e42
commit c0f2d52f70

View file

@ -165,7 +165,8 @@ app = FastAPI(title="Search Service", version="0.1.0", lifespan=lifespan)
DENSE_PREFETCH_K = 80 DENSE_PREFETCH_K = 80
SPARSE_PREFETCH_K = 200 SPARSE_PREFETCH_K = 200
RETRIEVE_K = 150 RETRIEVE_K = 150
RERANK_LIMIT = 17 RERANK_LIMIT = 25
KEYWORD_BOOST_EXTRA = 10
async def embed_dense(client: httpx.AsyncClient, text: str) -> list[float]: async def embed_dense(client: httpx.AsyncClient, text: str) -> list[float]:
@ -211,18 +212,64 @@ def embed_sparse_sync(text: str) -> SparseVector:
def build_dense_query(question: Question) -> str: def build_dense_query(question: Question) -> str:
return question.text.strip() q = question.search_text.strip() if question.search_text else question.text.strip()
return q
def build_sparse_query(question: Question) -> str: def build_sparse_query(question: Question) -> str:
parts = [question.text.strip()] base = question.search_text.strip() if question.search_text else question.text.strip()
parts = [base]
if question.keywords: if question.keywords:
parts.extend(question.keywords) parts.extend(question.keywords)
if question.search_text:
parts = [question.search_text]
return " ".join(parts) return " ".join(parts)
def _build_keyword_set(question: Question) -> list[str]:
tokens: list[str] = []
if question.keywords:
tokens.extend(kw.lower() for kw in question.keywords if kw)
if question.entities:
for field in (
question.entities.people,
question.entities.emails,
question.entities.documents,
question.entities.names,
question.entities.links,
):
tokens.extend(e.lower() for e in (field or []) if e)
return tokens
def prefilter_for_rerank(
points: list[Any],
question: Question,
) -> tuple[list[Any], list[Any]]:
"""Select candidates for reranking: top by RRF + keyword-boosted stragglers."""
if not points:
return [], []
head = points[:RERANK_LIMIT]
tail = points[RERANK_LIMIT:]
keywords = _build_keyword_set(question)
if not keywords or not tail:
return head, tail
extra: list[Any] = []
remaining_tail: list[Any] = []
for p in tail:
if len(extra) >= KEYWORD_BOOST_EXTRA:
remaining_tail.append(p)
continue
content = ((p.payload or {}).get("page_content") or "").lower()
if any(kw in content for kw in keywords):
extra.append(p)
else:
remaining_tail.append(p)
return head + extra, remaining_tail
async def qdrant_search( async def qdrant_search(
client: AsyncQdrantClient, client: AsyncQdrantClient,
dense_vectors: list[list[float]], dense_vectors: list[list[float]],
@ -315,25 +362,24 @@ async def rerank_points(
query: str, query: str,
points: list[Any], points: list[Any],
) -> list[Any]: ) -> list[Any]:
rerank_candidates = points[:RERANK_LIMIT] if not points:
rerank_targets = [point.payload.get("page_content") for point in rerank_candidates] return []
scores = await get_rerank_scores(client, query, rerank_targets) targets = [point.payload.get("page_content") for point in points]
scores = await get_rerank_scores(client, query, targets)
if not scores: if not scores or len(scores) != len(points):
logger.warning("Reranker unavailable, returning RRF order") logger.warning("Reranker unavailable or score mismatch, returning RRF order")
return rerank_candidates return points
reranked_candidates = [ return [
point point
for _, point in sorted( for _, point in sorted(
zip(scores, rerank_candidates, strict=True), zip(scores, points),
key=lambda item: item[0], key=lambda item: item[0],
reverse=True, reverse=True,
) )
] ]
return reranked_candidates
@app.get("/health") @app.get("/health")
async def health() -> dict[str, str]: async def health() -> dict[str, str]:
@ -358,13 +404,22 @@ async def search(payload: SearchAPIRequest) -> SearchAPIResponse:
dense_vector, sparse_vector = await asyncio.gather(dense_task, sparse_task) dense_vector, sparse_vector = await asyncio.gather(dense_task, sparse_task)
dense_vectors = [dense_vector] dense_vectors = [dense_vector]
if question.hyde and len(question.hyde) > 0: extra_texts: list[str] = []
for v in (question.variants or []):
q_v = v.strip()
if q_v and q_v not in extra_texts:
extra_texts.append(q_v)
for h in (question.hyde or []):
q_h = h.strip()
if q_h and q_h not in extra_texts:
extra_texts.append(q_h)
extra_texts = extra_texts[:3]
if extra_texts:
try: try:
hyde_texts = question.hyde[:2] extra_vecs = await embed_dense_batch(client, extra_texts)
hyde_vectors = await embed_dense_batch(client, hyde_texts) dense_vectors.extend(extra_vecs)
dense_vectors.extend(hyde_vectors)
except Exception as e: except Exception as e:
logger.warning(f"HyDE embedding failed: {e}") logger.warning(f"Extra dense embedding failed: {e}")
all_points = await qdrant_search(qdrant, dense_vectors, sparse_vector) all_points = await qdrant_search(qdrant, dense_vectors, sparse_vector)
@ -373,20 +428,20 @@ async def search(payload: SearchAPIRequest) -> SearchAPIResponse:
all_points = list(all_points) all_points = list(all_points)
reranked = await rerank_points(client, query, all_points) rerank_pool, rerank_tail = prefilter_for_rerank(all_points, question)
reranked = await rerank_points(client, query, rerank_pool)
final_points = reranked + rerank_tail
reranked_ids = {id(p) for p in reranked} msg_score: dict[str, float] = {}
remaining = [p for p in all_points if id(p) not in reranked_ids] for rank, point in enumerate(reranked):
final_points = reranked + remaining score = 1.0 / (rank + 1)
seen: set[str] = set()
message_ids: list[str] = []
for point in final_points:
for mid in extract_message_ids(point): for mid in extract_message_ids(point):
if mid not in seen: msg_score[mid] = msg_score.get(mid, 0.0) + score
seen.add(mid) for rank, point in enumerate(rerank_tail):
message_ids.append(mid) score = 1.0 / (60 + rank + 1)
message_ids = message_ids[:50] for mid in extract_message_ids(point):
msg_score[mid] = msg_score.get(mid, 0.0) + score
message_ids = sorted(msg_score, key=lambda m: msg_score[m], reverse=True)[:50]
return SearchAPIResponse(results=[SearchAPIItem(message_ids=message_ids)]) return SearchAPIResponse(results=[SearchAPIItem(message_ids=message_ids)])