Source code for evidence_fetcher.refinement

"""Deterministic query refinement from explicit relevance feedback."""

from __future__ import annotations

from collections import Counter
from collections.abc import Sequence
from typing import Protocol

from evidence_fetcher.models import Evidence, FeedbackLabel, RelevanceFeedback
from evidence_fetcher.ranking import (
    _document_text,
    _document_vector,
    _inverse_document_frequency,
    _tokenize,
)

MAX_EXPANSION_TERMS = 4

_STOP_WORDS = frozenset(
    {
        "a",
        "an",
        "and",
        "are",
        "as",
        "at",
        "be",
        "by",
        "for",
        "from",
        "how",
        "in",
        "is",
        "it",
        "of",
        "on",
        "or",
        "that",
        "the",
        "this",
        "to",
        "with",
    }
)


[docs] class QueryRefiner(Protocol): """Produce a provider search query from evidence and relevance feedback.""" def refine( self, query: str, evidence: Sequence[Evidence], feedback: Sequence[RelevanceFeedback], ) -> str: """Return a refined provider query."""
[docs] class TfidfQueryRefiner: """Small lexical query refiner based on feedback-marked evidence."""
[docs] def __init__(self, *, max_terms: int = MAX_EXPANSION_TERMS) -> None: if max_terms < 1: raise ValueError("max_terms must be at least 1") self._max_terms = max_terms
def refine( self, query: str, evidence: Sequence[Evidence], feedback: Sequence[RelevanceFeedback], ) -> str: """Append deterministic expansion terms to the original query.""" original_terms = _useful_terms(query) material_feedback = [ item for item in feedback if item.label in {FeedbackLabel.MORE, FeedbackLabel.RELEVANT} ] if not material_feedback: raise ValueError("relevant feedback is needed to refine the query") evidence_by_url = {item.url: item for item in evidence} documents = [_tokenize(_document_text(item)) for item in evidence] idf = _inverse_document_frequency(documents) positive_weights: Counter[str] = Counter() negative_weights: Counter[str] = Counter() for item in feedback: evidence_item = evidence_by_url.get(item.url) if evidence_item is None: raise ValueError("feedback URL is not present in the evidence set") weights = _term_weights(evidence_item, idf) match item.label: case FeedbackLabel.MORE | FeedbackLabel.RELEVANT: positive_weights.update(weights) case FeedbackLabel.LESS: negative_weights.update( {term: weight * 0.25 for term, weight in weights.items()} ) case FeedbackLabel.IRRELEVANT: negative_weights.update(weights) case FeedbackLabel.UNSURE: continue candidates = [] for term, weight in positive_weights.items(): if term in original_terms: continue negative_weight = negative_weights.get(term, 0.0) if negative_weight >= weight: continue candidates.append((term, weight - negative_weight)) candidates.sort(key=lambda item: (-item[1], item[0])) expansion_terms = [term for term, _ in candidates[: self._max_terms]] if not expansion_terms: raise ValueError("relevant feedback did not produce expansion terms") return " ".join([query.strip(), *expansion_terms])
def _term_weights(evidence: Evidence, idf: dict[str, float]) -> dict[str, float]: tokens = [ term for term in _tokenize(_document_text(evidence)) if term not in _STOP_WORDS and len(term) > 2 ] return _document_vector(Counter(tokens), idf) def _useful_terms(text: str) -> set[str]: return { term for term in _tokenize(text) if term not in _STOP_WORDS and len(term) > 2 }