File size: 7,555 Bytes
6ba7d76
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
"""IOL-AI 2026 β€” v4b: the OFFICIAL-BASELINE simple prompt + our count-fix + robustness.

A/B partner to script.py (v7). Hypothesis from real leaderboard data: the ONLY thing
that ever cracked exact_match is the official baseline's dead-simple prompt (EM 0.073,
score 0.123). Every elaborate script β€” ours, v5, the multi-agent β€” sits near EM 0. So
v4b strips the reasoning/marker prompt back to baseline style and keeps only the parts
the data supports:
  * the COUNT FIX: answers come from the model (and context for matching), never
    truncated to a query regex β€” the bug that zeroed matching/fill_blanks in v1-v3;
  * single-sequence greedy decode (no OOM), incremental writes, per-row try/except,
    and a time-budget guard, so it can never time out into an empty submission.

Same model as v7 (Qwen2.5-14B-AWQ, loaded from "."). Only the PROMPT + parse differ.
"""

import os
os.environ.setdefault("HF_HUB_OFFLINE", "1")
os.environ.setdefault("TRANSFORMERS_OFFLINE", "1")

import re
import csv
import json

MODEL_DIR = os.environ.get("IOL_MODEL_DIR", ".")
TEST_CSV = os.environ.get("IOL_TEST_CSV", "/tmp/data/test.csv")
OUT_CSV = os.environ.get("IOL_OUT_CSV", "submission.csv")
MAX_NEW_TOKENS = int(os.environ.get("IOL_MAX_NEW_TOKENS", "512"))   # no CoT -> short + fast
QUANT = os.environ.get("IOL_QUANT", "4bit")
TIME_BUDGET_S = float(os.environ.get("IOL_TIME_BUDGET_S", "1620"))  # 27 min safety valve

# Deliberately close to the official baseline's system prompt (the only known EM winner),
# plus one clause about the multi-item count (the fix) and per-task answer form.
SYSTEM_PROMPT = (
    "You solve International Linguistics Olympiad problems. Everything you need is in the "
    "data given; use ONLY it, not outside knowledge of any language. Answer EVERY "
    "sub-question β€” a problem may have many (e.g. one answer per numbered item in the "
    "data, a dozen or more). Put each answer on its own line, in order, with NO numbering "
    "and NO extra text. Answer in the language the query asks; for matching give only the "
    "option letter; for numbers give digits or the written-out number as asked. Copy the "
    "exact characters and diacritics from the data. Never leave an item blank."
)

TASK_HINT = {
    "translation": "Translation: each line is only the translated word/phrase.",
    "text_to_num": "Numbers: each line is only digits (e.g. 42).",
    "num_to_text": "Numbers: each line is only the number in the target language's words.",
    "match_letters": "Matching: each line is only the option letter (A, B, C, ...); one per item in the data.",
    "matching": "Matching: each line is only the option letter; one per item in the data.",
    "fill_blank": "Fill in the blank: each line is only the missing form.",
    "fill_blanks": "Fill in the blank: each line is only the missing form.",
}


def build_messages(row):
    context = (row.get("context") or "").strip()
    query = (row.get("query") or "").strip()
    ttype = (row.get("task_type") or "").strip().lower()
    system = SYSTEM_PROMPT
    hint = TASK_HINT.get(ttype)
    if hint:
        system = system + "\n\n" + hint
    return [
        {"role": "system", "content": system},
        {"role": "user", "content": context + "\n\n" + query},
    ]


def detect_count(context, query):
    """Best-effort sub-question count β€” a HINT and minimum pad, never a truncation.
    Matching queries number nothing, so fall back to the numbered items in the context."""
    q = re.findall(r"(?m)^\s*(\d+)[\.\)]", query)
    if q:
        return len(q)
    par = re.findall(r"\((\d+)\)", query)
    if par:
        return len(set(par))
    c = re.findall(r"(?m)^\s*(\d+)[\.\)]", context)
    if c:
        return len(c)
    return 1


def parse_answers(text, min_count=1):
    """Baseline-style: every non-empty line is an answer (strip numbering/bullets/labels).
    No marker, no comma-splitting β€” the simple prompt tells the model one answer per line.
    Skip obvious header lines. Pad up to min_count; never truncate; never blank."""
    answers = []
    for ln in text.splitlines():
        c = ln.strip()
        if not c:
            continue
        if re.fullmatch(r"(?i)(final\s+)?answers?\s*:?", c):   # header like "Answers:"
            continue
        c = re.sub(r"^\s*(?:\d+[\.\):]|[-*β€’])\s*", "", c).strip()
        c = re.sub(r"^(?:answer|ans|translation|result)s?\s*[:\-]\s*", "", c, flags=re.I).strip()
        c = c.strip("\"'β€œβ€β€˜β€™` ").strip()
        if c:
            answers.append(c)
    if len(answers) < min_count:
        answers += ["?"] * (min_count - len(answers))
    return answers if answers else ["?"]


def _already_quantized(model_dir):
    cfg = os.path.join(model_dir, "config.json")
    try:
        with open(cfg, encoding="utf-8") as f:
            return "quantization_config" in json.load(f)
    except Exception:
        return False


def load_model():
    import torch
    from transformers import AutoTokenizer, AutoModelForCausalLM

    tok = AutoTokenizer.from_pretrained(MODEL_DIR)
    if tok.pad_token_id is None:
        tok.pad_token = tok.eos_token
    if not torch.cuda.is_available():
        return tok, AutoModelForCausalLM.from_pretrained(
            MODEL_DIR, torch_dtype=torch.float32).eval()   # CPU dev fallback

    kwargs = dict(torch_dtype=torch.float16, device_map="auto")  # T4 has no bf16
    if _already_quantized(MODEL_DIR):
        pass  # AWQ auto-detected from config.json
    elif QUANT == "4bit":
        from transformers import BitsAndBytesConfig
        kwargs["quantization_config"] = BitsAndBytesConfig(
            load_in_4bit=True, bnb_4bit_compute_dtype=torch.float16,
            bnb_4bit_quant_type="nf4", bnb_4bit_use_double_quant=True)
    return tok, AutoModelForCausalLM.from_pretrained(MODEL_DIR, **kwargs).eval()


def generate_one(tok, model, messages):
    import torch
    dev = model.device if hasattr(model, "device") else "cpu"
    ids = tok.apply_chat_template(
        messages, add_generation_prompt=True, return_tensors="pt").to(dev)
    with torch.no_grad():
        gen = model.generate(ids, max_new_tokens=MAX_NEW_TOKENS, do_sample=False,
                             pad_token_id=tok.pad_token_id)
    return tok.decode(gen[0][ids.shape[-1]:], skip_special_tokens=True).strip()


def main():
    import time
    tok, model = load_model()

    with open(TEST_CSV, newline="", encoding="utf-8") as f:
        rows = list(csv.DictReader(f))

    fout = open(OUT_CSV, "w", newline="", encoding="utf-8")
    writer = csv.DictWriter(fout, fieldnames=["id", "pred"])
    writer.writeheader()
    fout.flush()

    start_t = time.time()
    for k, r in enumerate(rows):
        context = (r.get("context") or "").strip()
        query = (r.get("query") or "").strip()
        min_count = detect_count(context, query)
        try:
            if (time.time() - start_t) > TIME_BUDGET_S:
                raise TimeoutError("time budget exceeded")   # -> placeholder, keep moving
            text = generate_one(tok, model, build_messages(r))
            answers = parse_answers(text, min_count)
        except Exception as e:
            print("row %s fallback: %r" % (r.get("id"), e), flush=True)
            answers = ["?"] * min_count
        writer.writerow({"id": r["id"], "pred": json.dumps(answers, ensure_ascii=False)})
        fout.flush()
        print("%d/%d done" % (k + 1, len(rows)), flush=True)

    fout.close()
    print("wrote %s (%d rows)" % (OUT_CSV, len(rows)), flush=True)


if __name__ == "__main__":
    main()