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
6869189099
commit
2a274e93d8
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
|
||||
SPARSE_PREFETCH_K = 200
|
||||
RETRIEVE_K = 150
|
||||
RERANK_LIMIT = 17
|
||||
RERANK_LIMIT = 25
|
||||
KEYWORD_BOOST_EXTRA = 10
|
||||
|
||||
|
||||
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:
|
||||
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:
|
||||
parts = [question.text.strip()]
|
||||
base = question.search_text.strip() if question.search_text else question.text.strip()
|
||||
parts = [base]
|
||||
if question.keywords:
|
||||
parts.extend(question.keywords)
|
||||
if question.search_text:
|
||||
parts = [question.search_text]
|
||||
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(
|
||||
client: AsyncQdrantClient,
|
||||
dense_vectors: list[list[float]],
|
||||
|
|
@ -315,25 +362,24 @@ async def rerank_points(
|
|||
query: str,
|
||||
points: list[Any],
|
||||
) -> list[Any]:
|
||||
rerank_candidates = points[:RERANK_LIMIT]
|
||||
rerank_targets = [point.payload.get("page_content") for point in rerank_candidates]
|
||||
scores = await get_rerank_scores(client, query, rerank_targets)
|
||||
if not points:
|
||||
return []
|
||||
targets = [point.payload.get("page_content") for point in points]
|
||||
scores = await get_rerank_scores(client, query, targets)
|
||||
|
||||
if not scores:
|
||||
logger.warning("Reranker unavailable, returning RRF order")
|
||||
return rerank_candidates
|
||||
if not scores or len(scores) != len(points):
|
||||
logger.warning("Reranker unavailable or score mismatch, returning RRF order")
|
||||
return points
|
||||
|
||||
reranked_candidates = [
|
||||
return [
|
||||
point
|
||||
for _, point in sorted(
|
||||
zip(scores, rerank_candidates, strict=True),
|
||||
zip(scores, points),
|
||||
key=lambda item: item[0],
|
||||
reverse=True,
|
||||
)
|
||||
]
|
||||
|
||||
return reranked_candidates
|
||||
|
||||
|
||||
@app.get("/health")
|
||||
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_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:
|
||||
hyde_texts = question.hyde[:2]
|
||||
hyde_vectors = await embed_dense_batch(client, hyde_texts)
|
||||
dense_vectors.extend(hyde_vectors)
|
||||
extra_vecs = await embed_dense_batch(client, extra_texts)
|
||||
dense_vectors.extend(extra_vecs)
|
||||
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)
|
||||
|
||||
|
|
@ -373,20 +428,20 @@ async def search(payload: SearchAPIRequest) -> SearchAPIResponse:
|
|||
|
||||
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}
|
||||
remaining = [p for p in all_points if id(p) not in reranked_ids]
|
||||
final_points = reranked + remaining
|
||||
|
||||
seen: set[str] = set()
|
||||
message_ids: list[str] = []
|
||||
for point in final_points:
|
||||
msg_score: dict[str, float] = {}
|
||||
for rank, point in enumerate(reranked):
|
||||
score = 1.0 / (rank + 1)
|
||||
for mid in extract_message_ids(point):
|
||||
if mid not in seen:
|
||||
seen.add(mid)
|
||||
message_ids.append(mid)
|
||||
message_ids = message_ids[:50]
|
||||
msg_score[mid] = msg_score.get(mid, 0.0) + score
|
||||
for rank, point in enumerate(rerank_tail):
|
||||
score = 1.0 / (60 + rank + 1)
|
||||
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)])
|
||||
|
||||
|
|
|
|||
Loading…
Reference in a new issue