Files
RAG-CUT/backend/rag_cut/retrieval.py
T
2026-07-16 11:12:17 +08:00

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)