File size: 7,305 Bytes
5460e4e
 
 
 
 
 
 
 
 
 
 
 
 
2a1a805
5460e4e
 
 
 
 
 
 
2a1a805
 
 
5460e4e
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
2a1a805
 
 
 
 
5460e4e
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
37ccd6e
 
 
5460e4e
37ccd6e
5460e4e
 
37ccd6e
5460e4e
 
37ccd6e
5460e4e
 
 
 
 
 
 
 
 
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
import os
os.environ["HF_HUB_OFFLINE"] = "1"
os.environ["TRANSFORMERS_OFFLINE"] = "1"
os.environ["TOKENIZERS_PARALLELISM"] = "false"

# --- runtime bootstrap: PyPI is reachable in the eval sandbox even though the
# --- HF Hub is not, so make sure 4-bit deps are present (needed to fit the T4).
import subprocess, sys
def _pip(*pkgs):
    try:
        subprocess.run([sys.executable, "-m", "pip", "install", "-q", *pkgs], check=False)
    except Exception as e:
        print("pip bootstrap skipped:", e, flush=True)
_pip("bitsandbytes>=0.43.0", "accelerate>=0.30.0", "sentencepiece", "tiktoken")

import re, json, time
import pandas as pd
import torch

START = time.time()
TIME_BUDGET = 27 * 60          # stop generating with margin before the 30-min hard limit
# On the platform the repo IS the working dir, so "." holds the weights.
# IOL_MODEL_DIR lets a local dry-run point at a downloaded snapshot instead.
MODEL_ID = os.environ.get("IOL_MODEL_DIR", ".")
MAX_NEW_TOKENS = 512

# ---------------------------------------------------------------------------
# Load the model in 4-bit (nf4) so 7B weights fit in 16 GB of T4 VRAM.
# Use float16 compute: the T4 is a Turing GPU with no native bfloat16.
# ---------------------------------------------------------------------------
from transformers import AutoTokenizer, AutoModelForCausalLM
try:
    from transformers import BitsAndBytesConfig
    _bnb = BitsAndBytesConfig(
        load_in_4bit=True,
        bnb_4bit_quant_type="nf4",
        bnb_4bit_use_double_quant=True,
        bnb_4bit_compute_dtype=torch.float16,
    )
    load_kwargs = dict(quantization_config=_bnb)
except Exception as e:
    print("bitsandbytes unavailable, falling back to fp16:", e, flush=True)
    load_kwargs = dict(torch_dtype=torch.float16)

try:
    tok = AutoTokenizer.from_pretrained(MODEL_ID)
except Exception as e:
    print("fast tokenizer failed, retrying slow:", e, flush=True)
    tok = AutoTokenizer.from_pretrained(MODEL_ID, use_fast=False)
model = AutoModelForCausalLM.from_pretrained(
    MODEL_ID, device_map="auto", **load_kwargs
).eval()
if tok.pad_token_id is None:
    tok.pad_token_id = tok.eos_token_id

df = pd.read_csv("/tmp/data/test.csv", dtype=str).fillna("")

# ---------------------------------------------------------------------------
# Helpers
# ---------------------------------------------------------------------------
def expected_n(query, context):
    # How many numbered items must this row's answer list contain?
    for pat in (r"(?m)^\s*(\d+)\s*[\.\)]", r"\((\d+)\)"):
        m = re.findall(pat, query)
        if m:
            return len(m)
    # matching tasks number their items in the context, not the query
    m = re.findall(r"(?m)^\s*(\d+)\s*[\.\)]", context)
    return len(m) if m else 1

def strip_num(s):
    # Remove a leading list marker ("1.", "2)", "(3)") but NEVER a bare number,
    # so numeric answers like "111" survive. Punctuation after the digit is required.
    return re.sub(r"^\s*(?:\(\d+\)|\d+\s*[\.\):])\s*", "", s).strip()

# A line that begins with a list marker of the same forms.
NUMLINE = r"^\s*(?:\(\d+\)|\d+\s*[\.\):-])"

def parse_output(text, n, task_type):
    ans_part, expl = text, ""
    m = re.search(r"(?is)\bEXPLANATION\b\s*:?", text)
    if m:
        ans_part = text[:m.start()]
        expl = text[m.end():].strip()
    m2 = re.search(r"(?is)\bANSWERS?\b\s*:?", ans_part)
    if m2:
        ans_part = ans_part[m2.end():]
    lines = [ln.strip() for ln in ans_part.splitlines() if ln.strip()]
    numbered = [ln for ln in lines if re.match(NUMLINE, ln)]
    use = numbered if numbered else lines
    answers = [strip_num(ln) for ln in use]
    # Fallback: model crammed items onto one comma-separated line (common for
    # matching/number tasks). Split it back out, but not for free-text tasks
    # where commas can legitimately appear inside an answer.
    if len(answers) < n and task_type in ("match_letters", "text_to_num", "num_to_text"):
        flat = []
        for ln in use:
            flat += re.split(r"\s*[,;]\s*", strip_num(ln))
        flat = [x for x in flat if x != ""]
        if len(flat) > len(answers):
            answers = flat
    if len(answers) < n:
        answers += [""] * (n - len(answers))
    expl = re.sub(r"\s+", " ", expl).strip()[:800]  # keep explanation single-line for CSV
    return answers[:n], expl

TASK_HINTS = {
    "translation": "Translate each item. Answer in the language the query asks for.",
    "fill_blanks": "Work out the rule from the paired forms, then give the missing form for each blank.",
    "match_letters": "For each numbered item, output ONLY the letter label of its correct match.",
    "text_to_num": "Convert each written number into digits.",
    "num_to_text": "Write each number out in words in the task language.",
}

def build_messages(r):
    hint = TASK_HINTS.get(r["task_type"].strip().lower(), "Answer every numbered item.")
    sys_prompt = (
        "You are an expert solver of International Linguistics Olympiad problems. "
        "Each problem is fully self-contained: reason only from the data shown, with no outside "
        "knowledge of the language. Infer the grammar, vocabulary, or number system from the "
        "given examples, then answer every numbered item.\n"
        "OUTPUT FORMAT (follow exactly):\n"
        "ANSWERS:\n"
        "1. <answer to item 1>\n"
        "2. <answer to item 2>\n"
        "(one line per item, numbered, in the query's order, no commentary between them)\n"
        "EXPLANATION:\n"
        "<2-4 short sentences, human-readable, describing the rule you found>"
    )
    user_prompt = r["context"].strip() + "\n\n" + r["query"].strip() + "\n\n" + hint
    return [
        {"role": "system", "content": sys_prompt},
        {"role": "user", "content": user_prompt},
    ]

# ---------------------------------------------------------------------------
# Run
# ---------------------------------------------------------------------------
out_rows = []
for _, r in df.iterrows():
    n = expected_n(r["query"], r["context"])
    if time.time() - START > TIME_BUDGET:
        # Out of time: still emit a valid, correctly-sized (empty) row.
        out_rows.append({"id": r["id"],
                         "pred": json.dumps([""] * n, ensure_ascii=False),
                         "explanation": ""})
        continue
    enc = tok.apply_chat_template(
        build_messages(r), add_generation_prompt=True,
        return_tensors="pt", return_dict=True,   # BatchEncoding incl. attention_mask
    ).to(model.device)
    input_len = enc["input_ids"].shape[-1]
    with torch.no_grad():
        gen = model.generate(
            **enc, max_new_tokens=MAX_NEW_TOKENS,
            do_sample=False, pad_token_id=tok.pad_token_id,
        )
    text = tok.decode(gen[0][input_len:], skip_special_tokens=True).strip()
    answers, expl = parse_output(text, n, r["task_type"].strip().lower())
    out_rows.append({"id": r["id"],
                     "pred": json.dumps(answers, ensure_ascii=False),
                     "explanation": expl})
    print(str(len(out_rows)) + "/" + str(len(df)) + " done", flush=True)

pd.DataFrame(out_rows, columns=["id", "pred", "explanation"]).to_csv(
    "submission.csv", index=False)
print("wrote submission.csv", flush=True)