@@ -0,0 +1,78 @@
|
||||
"""Small dependency-free lexical retriever for the recall demo."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import math
|
||||
import re
|
||||
from collections import Counter
|
||||
|
||||
from rag_cut.models import Chunk
|
||||
|
||||
LATIN_TOKEN_RE = re.compile(r"[a-z0-9_]+", re.I)
|
||||
CJK_RE = re.compile(r"[\u3400-\u9fff]")
|
||||
|
||||
|
||||
def _tokens(text: str) -> list[str]:
|
||||
normalized = text.lower()
|
||||
tokens = LATIN_TOKEN_RE.findall(normalized)
|
||||
cjk = CJK_RE.findall(normalized)
|
||||
tokens.extend(cjk)
|
||||
tokens.extend("".join(cjk[i : i + 2]) for i in range(len(cjk) - 1))
|
||||
return tokens
|
||||
|
||||
|
||||
def _search_text(chunk: Chunk) -> str:
|
||||
meta = chunk.meta
|
||||
fields = [
|
||||
meta.get("heading") or "",
|
||||
meta.get("nearest_heading") or "",
|
||||
meta.get("embedding_text") or "",
|
||||
" ".join(meta.get("keywords") or []),
|
||||
chunk.content,
|
||||
]
|
||||
return "\n".join(str(value) for value in fields if value)
|
||||
|
||||
|
||||
def recall_chunks(query: str, chunks: list[Chunk], top_k: int = 5) -> tuple[list[dict], int]:
|
||||
"""Rank retrievable chunks with a compact BM25-style lexical score."""
|
||||
candidates = [chunk for chunk in chunks if chunk.meta.get("retrieval", True)]
|
||||
query_tokens = _tokens(query)
|
||||
if not candidates or not query_tokens:
|
||||
return [], len(candidates)
|
||||
|
||||
documents = [_tokens(_search_text(chunk)) for chunk in candidates]
|
||||
document_frequency = Counter(token for tokens in documents for token in set(tokens))
|
||||
average_length = sum(len(tokens) for tokens in documents) / len(documents)
|
||||
query_frequency = Counter(query_tokens)
|
||||
scored: list[tuple[float, Chunk]] = []
|
||||
|
||||
for chunk, tokens in zip(candidates, documents):
|
||||
frequency = Counter(tokens)
|
||||
length_normalizer = 1.2 * (0.25 + 0.75 * len(tokens) / max(average_length, 1))
|
||||
score = 0.0
|
||||
for token, query_count in query_frequency.items():
|
||||
term_count = frequency[token]
|
||||
if not term_count:
|
||||
continue
|
||||
inverse_frequency = math.log(1 + (len(documents) - document_frequency[token] + 0.5) / (document_frequency[token] + 0.5))
|
||||
score += inverse_frequency * ((term_count * 2.2) / (term_count + length_normalizer)) * min(query_count, 2)
|
||||
if score > 0:
|
||||
scored.append((score, chunk))
|
||||
|
||||
scored.sort(key=lambda item: (-item[0], item[1].index))
|
||||
results = []
|
||||
for rank, (score, chunk) in enumerate(scored[:top_k], start=1):
|
||||
results.append(
|
||||
{
|
||||
"rank": rank,
|
||||
"chunk_index": chunk.index,
|
||||
"score": round(score, 6),
|
||||
"content": chunk.content,
|
||||
"heading": chunk.meta.get("heading") or chunk.meta.get("nearest_heading"),
|
||||
"pages": chunk.meta.get("pages") or ([chunk.meta["page"]] if chunk.meta.get("page") is not None else []),
|
||||
"block_types": chunk.block_types,
|
||||
"parent_chunk_id": chunk.meta.get("parent_chunk_id"),
|
||||
"is_sub_chunk": bool(chunk.meta.get("is_sub_chunk")),
|
||||
}
|
||||
)
|
||||
return results, len(candidates)
|
||||
Reference in New Issue
Block a user