初始化项目,由ModelHub XC社区提供模型
Model: pei39/iol-qwen2.5-14b-sft-awq Source: Original Platform
This commit is contained in:
276
rag_resources/retriever.py
Normal file
276
rag_resources/retriever.py
Normal file
@@ -0,0 +1,276 @@
|
||||
"""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
|
||||
Reference in New Issue
Block a user