79 lines
3.0 KiB
Python
79 lines
3.0 KiB
Python
"""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)
|