277 lines
10 KiB
Python
277 lines
10 KiB
Python
|
|
"""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
|