File size: 10,077 Bytes
8129b09
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
#!/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