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:
parent
3561a064d5
commit
6f122eabd0
1 changed files with 87 additions and 32 deletions
119
search/main.py
119
search/main.py
|
|
@ -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)])
|
||||||
|
|
||||||
|
|
|
||||||
Loading…
Reference in a new issue