"""Offline hybrid BM25/character-TF-IDF retrieval over the public book corpus.""" from __future__ import annotations import json import math import re import unicodedata from collections import Counter from pathlib import Path from typing import Any, Iterable from sklearn.feature_extraction.text import TfidfVectorizer RESOURCE_DIR = Path(__file__).resolve().parent FAMILY_KEYWORDS = { "writing_system": ( "alphabet", "braille", "character", "decipher", "glyph", "letter", "orthography", "script", "symbol", "writing system", ), "phonetics": ( "accent", "consonant", "metre", "phonetic", "pronounce", "rhyme", "stress", "syllable", "tone", "vowel", ), "phonology": ( "alternation", "correspondence", "sound change", "sound rule", "underlying form", ), "noun_morphology": ( "case", "gender", "noun", "plural", "possession", "possessive", "singular", ), "verb_morphology": ( "agreement", "aspect", "conjug", "object", "person", "subject", "tense", "verb", ), "syntax": ("clause", "focus", "sentence", "subject", "object", "word order"), "number_system": ("arithmetic", "digit", "number", "numeral", "numerals"), "kinship_orientation": ( "brother", "daughter", "direction", "east", "family", "father", "kinship", "map", "mother", "north", "sister", "son", "south", "west", ), "matching": ("correspondence", "match", "matching", "random order"), "translation": ("translate", "translation"), "fill_blanks": ("blank", "fill in", "gap", "missing"), } TOPIC_FAMILIES = { "writing systems and script decipherment": {"writing_system"}, "phonetics, stress, tone, and versification": {"phonetics"}, "phonological rules and sound correspondences": {"phonology", "phonetics"}, "noun morphology and noun phrases": {"noun_morphology"}, "verb morphology and argument structure": {"verb_morphology"}, "syntax, word order, focus, and alignment": {"syntax"}, "semantics and graph-based matching": {"matching"}, "number systems": {"number_system"}, "orientation, kinship, and other structural problems": {"kinship_orientation"}, "general problem-solving methodology": set(FAMILY_KEYWORDS), } def normalize(text: Any) -> str: return " ".join(unicodedata.normalize("NFKC", str(text)).casefold().split()) def tokenize(text: Any) -> list[str]: return re.findall(r"[^\W_]+", normalize(text), flags=re.UNICODE) def read_jsonl(path: Path) -> list[dict[str, Any]]: with path.open("r", encoding="utf-8") as handle: return [json.loads(line) for line in handle if line.strip()] def infer_task_families(context: str, query: str) -> list[str]: """Infer likely families from the problem itself, never from a gold answer.""" text = normalize(context + "\n" + query) scored: list[tuple[int, str]] = [] for family, keywords in FAMILY_KEYWORDS.items(): score = sum(1 for keyword in keywords if keyword in text) if score: scored.append((score, family)) return [family for _, family in sorted(scored, key=lambda item: (-item[0], item[1]))] class BM25: def __init__(self, documents: Iterable[str], k1: float = 1.5, b: float = 0.75): self.k1 = k1 self.b = b self.tokens = [tokenize(document) for document in documents] self.lengths = [len(tokens) for tokens in self.tokens] self.average_length = sum(self.lengths) / max(len(self.lengths), 1) self.term_frequencies = [Counter(tokens) for tokens in self.tokens] document_frequency: Counter[str] = Counter() for tokens in self.tokens: document_frequency.update(set(tokens)) document_count = len(self.tokens) self.idf = { term: math.log( 1.0 + (document_count - frequency + 0.5) / (frequency + 0.5) ) for term, frequency in document_frequency.items() } def scores(self, query: str) -> list[float]: query_terms = Counter(tokenize(query)) scores: list[float] = [] for frequencies, length in zip(self.term_frequencies, self.lengths): score = 0.0 normalization = self.k1 * ( 1.0 - self.b + self.b * length / max(self.average_length, 1.0) ) for term, query_frequency in query_terms.items(): frequency = frequencies.get(term, 0) if not frequency: continue score += ( self.idf.get(term, 0.0) * frequency * (self.k1 + 1.0) / (frequency + normalization) * min(query_frequency, 3) ) scores.append(score) return scores class CharNgramTFIDF: """Cosine similarity over normalized character 3-5 grams.""" def __init__(self, documents: Iterable[str]): self.vectorizer = TfidfVectorizer( analyzer="char_wb", ngram_range=(3, 5), preprocessor=normalize, lowercase=False, sublinear_tf=True, norm="l2", ) self.matrix = self.vectorizer.fit_transform(documents) def scores(self, query: str) -> list[float]: query_vector = self.vectorizer.transform([query]) return (self.matrix @ query_vector.T).toarray().ravel().tolist() def _clip(text: Any, limit: int) -> str: value = str(text or "").strip() if len(value) <= limit: return value clipped = value[:limit].rsplit("\n", 1)[0].rstrip() if len(clipped) < limit // 2: clipped = value[:limit].rsplit(" ", 1)[0].rstrip() return clipped + "\n[excerpt truncated]" class BookRetriever: """Retrieve methods and analogous examples from langsci/420 only.""" def __init__(self, resource_dir: Path = RESOURCE_DIR): self.methods = read_jsonl(resource_dir / "book_methods.jsonl") self.examples = [ row for row in read_jsonl(resource_dir / "book_examples.jsonl") if row.get("string_only_usable", False) ] self.method_bm25 = BM25(row["retrieval_text"] for row in self.methods) self.example_bm25 = BM25(row["retrieval_text"] for row in self.examples) self.method_char_tfidf = CharNgramTFIDF( row["retrieval_text"] for row in self.methods ) self.example_char_tfidf = CharNgramTFIDF( row["retrieval_text"] for row in self.examples ) @staticmethod def _topic_bonus(row: dict[str, Any], families: set[str]) -> float: return 1.25 * len( families & TOPIC_FAMILIES.get(str(row.get("topic", "")), set()) ) def retrieve( self, row: dict[str, Any], top_methods: int = 2, top_examples: int = 2, char_tfidf_weight: float = 3.0, ) -> dict[str, Any]: context = str(row.get("context", "")) query = str(row.get("query", "")) test_text = context + "\n" + query families = infer_task_families(context, query) family_set = set(families) expanded_query = test_text + "\n" + " ".join(families) ranked_examples: list[tuple[float, dict[str, Any]]] = [] example_bm25_scores = self.example_bm25.scores(expanded_query) example_char_scores = self.example_char_tfidf.scores(expanded_query) for bm25_score, char_score, example in zip( example_bm25_scores, example_char_scores, self.examples ): adjusted = ( bm25_score + char_tfidf_weight * char_score + self._topic_bonus(example, family_set) ) if example.get("solution_tier") == "detailed_worked": adjusted += 0.35 ranked_examples.append((adjusted, example)) ranked_examples.sort(key=lambda item: (-item[0], item[1]["id"])) selected_examples = ranked_examples[:top_examples] linked_method_ids = { method_id for _, example in ranked_examples[:max(top_examples * 3, 5)] for method_id in example.get("method_ids", []) } ranked_methods: list[tuple[float, dict[str, Any]]] = [] method_bm25_scores = self.method_bm25.scores(expanded_query) method_char_scores = self.method_char_tfidf.scores(expanded_query) for bm25_score, char_score, method in zip( method_bm25_scores, method_char_scores, self.methods ): adjusted = ( bm25_score + char_tfidf_weight * char_score + self._topic_bonus(method, family_set) ) if method["id"] in linked_method_ids: adjusted += 0.75 ranked_methods.append((adjusted, method)) ranked_methods.sort(key=lambda item: (-item[0], item[1]["id"])) return { "inferred_task_families": families, "char_tfidf_weight": char_tfidf_weight, "methods": [method for _, method in ranked_methods[:top_methods]], "examples": [example for _, example in selected_examples], } def format_for_prompt(self, result: dict[str, Any], max_chars: int = 16000) -> str: blocks = [ "BOOK REFERENCE MATERIAL\n" "Use it for transferable methods and analogies. Do not copy a conclusion " "unless it is supported by the current problem." ] families = result.get("inferred_task_families", []) if families: blocks.append("Likely families inferred from the current text: " + ", ".join(families)) for method in result["methods"]: blocks.append( f"METHOD — {method['section_title']}\n" + _clip(method.get("text"), 2800) ) for example in result["examples"]: blocks.append( "\n".join( [ f"ANALOGOUS WORKED EXAMPLE — {example.get('language', 'unknown language')}", "Problem context:", _clip(example.get("context"), 2400), "Problem query:", _clip(example.get("query"), 1200), "Worked solution:", _clip(example.get("reasoning_trace"), 3600), ] ) ) output = "\n\n".join(blocks) if len(output) > max_chars: output = output[:max_chars].rsplit("\n", 1)[0].rstrip() output += "\n[retrieved material truncated]" return output