IOL-AI-v3 / iol_common.py
ShayanShamsi's picture
Add files using upload-large-folder tool
8129b09 verified
Raw
History Blame Contribute Delete
10.1 kB
#!/usr/bin/env python3
"""Shared prompt-building / parsing / voting utilities for IOL-AI experiments."""
import json
import re
import unicodedata
from collections import Counter
# ---------------------------------------------------------------------------
# Query splitting
# ---------------------------------------------------------------------------
_ITEM_RE = re.compile(r"(?m)^\s*([0-9]+[.)]|\([0-9]+\)|[0-9]+:)\s*")
def split_query(query: str):
"""Return (header, [item texts]) from a numbered query."""
query = (query or "").strip()
matches = list(_ITEM_RE.finditer(query))
if not matches:
return query, [query] if query else ["?"]
header = query[: matches[0].start()].strip()
items = []
for i, m in enumerate(matches):
end = matches[i + 1].start() if i + 1 < len(matches) else len(query)
items.append(query[m.end():end].strip())
return header, items
_PAREN_RE = re.compile(r"\((\d+)\)")
_RANGE_RE = re.compile(r"\(?\b(\d+)\s*[-–—]\s*(\d+)\)?")
_QUESTION_HINT = re.compile(
r"(?i)\b(what|which|why|how|who|other|explain|mean)\b|\?")
def infer_labels(context: str, query: str):
"""Return the list of item labels (as strings, in sheet order) for a row."""
query = (query or "").strip()
lines = [ln.strip() for ln in query.splitlines()]
first = lines[0] if lines else ""
# 1. explicit labels in the query body: ^-anchored item numbers and/or
# parenthesised (n) placeholders, in order of first appearance
explicit = []
for x in re.finditer(r"(?m)^\s*(\d+)[.):]\s|\((\d+)\)", query):
lab = x.group(1) or x.group(2)
if lab not in explicit:
explicit.append(lab)
# 2. a range in the instruction line ("Fill the blanks (1-10)",
# "Write the numbers (13–16)...") wins only when it covers and extends
# the explicit labels — otherwise it's a cross-reference, not the items
m = _RANGE_RE.search(first)
if m:
a, b = int(m.group(1)), int(m.group(2))
if 0 < a <= b and b - a < 60:
rng = [str(i) for i in range(a, b + 1)]
if not explicit or (set(explicit) <= set(rng) and len(rng) > len(explicit)):
return rng
if explicit:
return explicit
# 4. unnumbered one-item-per-line list after an instruction line
body = [ln for ln in lines[1:] if ln]
if first.rstrip().endswith(":") and body:
return [str(i) for i in range(1, len(body) + 1)]
# 5. single free-form question (no list anywhere in the query)
if len(body) == 0 and _QUESTION_HINT.search(first):
return ["1"]
# 6. numbered/lettered items in the context (match_letters)
m = list(_ITEM_RE.finditer(context or ""))
if m:
return [re.sub(r"\D", "", x.group(1)) for x in m]
# 7. fall back: any body lines, else one item
if body:
return [str(i) for i in range(1, len(body) + 1)]
return ["1"]
# ---------------------------------------------------------------------------
# Prompts
# ---------------------------------------------------------------------------
_INDUCTION = {
"translation": (
"Below is a problem sheet from a linguistics exam. Your task is to determine as "
"much information about the language as possible, purely from the information "
"provided. Systematically determine the vocabulary meaning of each word, the "
"syntactic structure (such as word order), the morphology (including any verb "
"conjugations), and the meaning of any affixes or subwords. Test every piece of "
"information you determine against every example provided."
),
"fill_blanks": (
"Below is a problem sheet from a linguistics exam. Your task is to determine as "
"much information about the language as possible, purely from the information "
"provided. Systematically determine the morphological and phonological patterns "
"of the language (such as noun declension), including the meaning of any subwords "
"or affixes. Test every piece of information you determine against every example "
"provided."
),
"number": (
"Below is a problem sheet from a linguistics exam. Your task is to determine as "
"much information about the language and its number system as possible, purely "
"from the information provided. Determine the vocabulary meaning of each number "
"word, the base of the number system, the word order, and any other patterns in "
"how numbers are composed. Test every piece of information you determine against "
"every example provided."
),
"match_letters": (
"Below is a problem sheet from a linguistics exam. Work out the correspondences "
"between the items and their meanings, and the vocabulary, morphology and "
"structure of the language. Test every correspondence you determine against "
"every example provided."
),
}
def induction_text(task_type: str) -> str:
if task_type in ("text_to_num", "num_to_text"):
return _INDUCTION["number"]
return _INDUCTION.get(task_type, _INDUCTION["translation"])
def answer_format_note(task_type: str, labels) -> str:
note = ""
if task_type == "match_letters":
note = "For each numbered item give ONLY the letter of its correct match. "
elif task_type == "text_to_num":
note = "For each numbered item give the value in digits only. "
elif task_type == "fill_blanks":
note = "For each numbered blank give ONLY the missing form. "
keys = ", ".join(f'"{l}": ""' for l in labels)
return (
f"{note}Answer every item. Give your final answer STRICTLY as a single JSON "
f"object with one key per item number, plus an \"explanation\" key briefly "
f"stating the rules you discovered (2-3 sentences, no reasoning trace):\n"
f"{{{keys}, \"explanation\": \"\"}}"
)
def direct_prompt(context: str, query: str, task_type: str, labels) -> str:
return (
f"{induction_text(task_type)}\n\n{context.strip()}\n\n{query.strip()}\n\n"
f"{answer_format_note(task_type, labels)}"
)
def induction_prompt(context: str, task_type: str) -> str:
return f"{induction_text(task_type)}\n\n{context.strip()}"
def application_prompt(context, task_type, rules, query, labels):
return (
f"{induction_prompt(context, task_type)}\n\n"
f"Here is an analysis of the language:\n{rules.strip()}\n\n"
f"Based on this analysis, solve the following puzzle:\n{query.strip()}\n\n"
f"{answer_format_note(task_type, labels)}"
)
# ---------------------------------------------------------------------------
# Output parsing
# ---------------------------------------------------------------------------
def strip_think(text: str) -> str:
text = re.sub(r"<think>.*?</think>", "", text, flags=re.DOTALL)
# unclosed think block: keep tail after last </think>, else keep the body
if "</think>" in text:
text = text.split("</think>")[-1]
return text.replace("<think>", "")
def extract_json(text: str):
cands = re.findall(r"```(?:json)?\s*(\{.*?\})\s*```", text, re.DOTALL)
cands += re.findall(r"\{(?:[^{}]|\{[^{}]*\})*\}", text, re.DOTALL)
for c in reversed(cands):
try:
o = json.loads(c)
if isinstance(o, dict):
return o
except Exception:
continue
return {}
def parse_items(text: str, labels):
"""labels: list of item labels (or an int for 1..n). -> list[str] same length."""
if isinstance(labels, int):
labels = [str(i) for i in range(1, labels + 1)]
n_items = len(labels)
text = strip_think(text)
obj = extract_json(text)
answers = [""] * n_items
if obj:
for i, lab in enumerate(labels):
for key in (lab, str(lab), int(lab) if str(lab).isdigit() else lab,
str(i + 1), i + 1):
if key in obj and str(obj[key]).strip():
answers[i] = str(obj[key]).strip()
break
if not any(answers):
lines = [ln.strip() for ln in text.splitlines() if ln.strip()]
lines = [ln for ln in lines if not ln.lower().startswith(("here", "based on", "```"))]
for i in range(min(n_items, len(lines))):
answers[i] = re.sub(r"^\s*[0-9]+[.):]\s*", "", lines[i])
return answers
# ---------------------------------------------------------------------------
# Self-consistency vote
# ---------------------------------------------------------------------------
def _norm_vote(answer: str) -> str:
answer = unicodedata.normalize("NFC", answer)
return answer.strip().strip("\"'").rstrip(".").strip().lower()
def _chrf3(a: str, b: str, n: int = 3) -> float:
a, b = a.lower(), b.lower()
if not a and not b:
return 1.0
if not a or not b:
return 0.0
total = 0.0
for k in range(1, n + 1):
ag = Counter(a[i:i + k] for i in range(len(a) - k + 1))
bg = Counter(b[i:i + k] for i in range(len(b) - k + 1))
if not ag or not bg:
continue
inter = sum((ag & bg).values())
p = inter / max(sum(ag.values()), 1)
r = inter / max(sum(bg.values()), 1)
total += 0.0 if (p + r) == 0 else 2 * p * r / (p + r)
return total / n
def majority_vote(candidates):
valid = [c for c in candidates if isinstance(c, str) and c.strip()]
if not valid:
return ""
groups = {}
for c in valid:
groups.setdefault(_norm_vote(c), []).append(c)
counts = {k: len(v) for k, v in groups.items()}
top = max(counts.values())
winners = [k for k, n in counts.items() if n == top]
if len(winners) == 1:
return Counter(groups[winners[0]]).most_common(1)[0][0]
best, best_score = valid[0], -1.0
for c in valid:
s = sum(_chrf3(c, o) for o in valid if o is not c)
if s > best_score:
best, best_score = c, s
return best