252 lines
9.6 KiB
Python
252 lines
9.6 KiB
Python
import os
|
||
import json
|
||
import pandas as pd
|
||
import torch
|
||
import re
|
||
import time
|
||
import unicodedata
|
||
from collections import Counter, defaultdict
|
||
from transformers import AutoTokenizer, AutoModelForCausalLM
|
||
|
||
# Offline Environment Setup
|
||
os.environ["HF_HUB_OFFLINE"] = "1"
|
||
os.environ["TRANSFORMERS_OFFLINE"] = "1"
|
||
MODEL_ID = "."
|
||
|
||
# Time Budgets
|
||
T0 = time.time()
|
||
DEADLINE = T0 + 1650 # 27.5 minutes max limit
|
||
EXPLAIN_RESERVE = 240 # Save 4 minutes at the end for Jury Explanations
|
||
|
||
TEST_CSV = "/tmp/data/test.csv"
|
||
OUT_CSV = "submission.csv"
|
||
MAX_SAMPLES = 25 # Increased to 25 to maximize test-time compute
|
||
SAMPLE_TEMP = 0.5
|
||
|
||
# 1. Startup Placeholder Creation
|
||
if os.path.exists(TEST_CSV):
|
||
try:
|
||
init_df = pd.read_csv(TEST_CSV, dtype=str).fillna("")
|
||
init_rows = [{"id": str(r["id"]), "pred": json.dumps(["?"], ensure_ascii=False), "explanation": "Startup."} for _, r in init_df.iterrows()]
|
||
pd.DataFrame(init_rows).to_csv(OUT_CSV, index=False, encoding="utf-8-sig")
|
||
except Exception:
|
||
pass
|
||
|
||
print("Loading tokenizer and 14B AWQ model offline...", flush=True)
|
||
tok = AutoTokenizer.from_pretrained(MODEL_ID)
|
||
if tok.pad_token is None:
|
||
tok.pad_token = tok.eos_token
|
||
|
||
model = AutoModelForCausalLM.from_pretrained(
|
||
MODEL_ID,
|
||
torch_dtype=torch.float16,
|
||
device_map={"": 0}
|
||
).eval()
|
||
|
||
df = pd.read_csv(TEST_CSV, dtype=str).fillna("")
|
||
ids = [str(x) for x in df["id"].tolist()]
|
||
|
||
# ---- ALIGNMENT & FALLBACK LOGIC ----
|
||
_LINE_NUM = re.compile(r"^[ \t]*(\d{1,3})[.)\]]", re.M)
|
||
_PAREN_NUM = re.compile(r"\((\d{1,3})\)")
|
||
_LINE_LET = re.compile(r"^[ \t]*([A-Z])[.)\]]\s", re.M)
|
||
|
||
def detect_n_items(query, context=""):
|
||
q = query or ""
|
||
line = [int(m) for m in _LINE_NUM.findall(q)]
|
||
par = [int(m) for m in _PAREN_NUM.findall(q)]
|
||
cand = max(len(set(line)), len(set(par)))
|
||
cand = max(cand, len(set(_LINE_LET.findall(q))))
|
||
n = max(cand, 1)
|
||
if n > 1: return n
|
||
lines = [l.strip() for l in q.splitlines() if l.strip()]
|
||
if len(lines) > 1 and lines[0].endswith((":", ".")):
|
||
body = lines[1:]
|
||
if body: return len(body)
|
||
if context:
|
||
cn = len(set(int(m) for m in _LINE_NUM.findall(context)))
|
||
if cn > 1: return cn
|
||
return max(n, 1)
|
||
|
||
def extract_item_sources(query, n):
|
||
q = query or ""
|
||
out = []
|
||
for ln in q.splitlines():
|
||
s = ln.strip()
|
||
if not s: continue
|
||
m = re.match(r"^\(?(\d{1,3})\)?[.):\\]]\\s*(.+)$", s)
|
||
if m: out.append(m.group(2).strip())
|
||
if not out:
|
||
lines = [l.strip() for l in q.splitlines() if l.strip()]
|
||
if len(lines) > 1 and lines[0].endswith((":",".")): out = lines[1:]
|
||
out = [o.split("|")[0].strip() if "|" in o else o for o in out]
|
||
out = [o for o in out if o]
|
||
while len(out) < n: out.append(out[-1] if out else "?")
|
||
return out[:n]
|
||
|
||
ns = [detect_n_items(r.get("query", ""), r.get("context", "")) for _, r in df.iterrows()]
|
||
srcs = {i: extract_item_sources(r.get("query", ""), n) for i, (_, r), n in zip(ids, df.iterrows(), ns)}
|
||
|
||
# ---- PARSING & CLEANING ----
|
||
_STRIP_PREFIX = re.compile(r"^\s*(?:\(?\d{1,3}\)?[.):\\]]\\s*|[-*•]\\s+)")
|
||
def clean_line(s):
|
||
# Normalize Unicode characters
|
||
s = unicodedata.normalize("NFC", s.strip())
|
||
# Standardize spaces and hyphens
|
||
s = s.replace('\xa0', ' ').replace('–', '-').replace('—', '-')
|
||
# Standardize curly quotes to standard ASCII quotes
|
||
s = s.replace('’', "'").replace('‘', "'").replace('“', '"').replace('”', '"')
|
||
|
||
# Strip question numbers/bullets
|
||
s = _STRIP_PREFIX.sub("", s)
|
||
s = s.strip().strip("`").strip()
|
||
|
||
# Strip outer wrapping quotes if the model wrapped answers
|
||
if len(s) >= 2 and s[0] == s[-1] and s[0] in "\"'“”":
|
||
s = s[1:-1].strip()
|
||
|
||
# Collapse internal multiple spaces
|
||
s = re.sub(r"\s+", " ", s)
|
||
return s.strip()
|
||
|
||
def fit_to_n(items, n, fb=None):
|
||
items = [i for i in items if i and i.strip()]
|
||
if len(items) > n: items = items[-n:]
|
||
while len(items) < n:
|
||
if fb and len(items) < len(fb): items.append(fb[len(items)])
|
||
else: items.append(items[-1] if items else "?")
|
||
return items[:n]
|
||
|
||
def parse_raw_lines(text, n, fb=None):
|
||
lines = [clean_line(ln) for ln in (text or "").splitlines() if clean_line(ln)]
|
||
return fit_to_n(lines, n, fb)
|
||
|
||
# ---- VOTING ENGINE ----
|
||
def norm(s):
|
||
s = unicodedata.normalize("NFC", (s or "").strip().lower())
|
||
s = re.sub(r"\\s+", " ", s)
|
||
return s.strip(" .!?;:,")
|
||
|
||
def vote(cands, anchor=None):
|
||
cands = [c for c in cands if c and c.strip()]
|
||
if anchor is None: anchor = cands[0] if cands else "?"
|
||
if len(cands) < 3: return anchor
|
||
|
||
groups = defaultdict(list)
|
||
for c in cands: groups[norm(c)].append(c)
|
||
|
||
anchor_support = len(groups.get(norm(anchor), []))
|
||
best_key, best_n = None, 0
|
||
for k, v in groups.items():
|
||
if len(v) > best_n:
|
||
best_key, best_n = k, len(v)
|
||
|
||
if best_key is not None and best_n >= 2 and best_n > anchor_support:
|
||
return Counter(groups[best_key]).most_common(1)[0][0]
|
||
return anchor
|
||
|
||
def write_submission(preds, expl):
|
||
rows = [{"id": i, "pred": json.dumps(preds[i], ensure_ascii=False), "explanation": expl.get(i, "")} for i in ids]
|
||
pd.DataFrame(rows).to_csv(OUT_CSV, index=False, encoding="utf-8-sig")
|
||
|
||
# ---- SYSTEM PROMPTS ----
|
||
SYS_TRIVIAL = (
|
||
"You solve International Linguistics Olympiad problems. Answer every numbered "
|
||
"item. Put each answer on its own line, in order, with no numbering and no extra text."
|
||
)
|
||
|
||
EXPLAIN_SYSTEM = (
|
||
"You explain International Linguistics Olympiad solutions to a human judge. "
|
||
"State the key rules of the language: morphemes, word order, sound changes. Concise (2-4 sentences)."
|
||
)
|
||
|
||
preds = {i: list(srcs[i]) for i in ids}
|
||
expl = {i: "Deducted rules from context." for i in ids}
|
||
samples = {i: [] for i in ids}
|
||
|
||
# ---- PASS 1: GREEDY ANCHOR ----
|
||
print("Starting Pass 1 (Greedy Anchor)...", flush=True)
|
||
for i, n, (_, r) in zip(ids, ns, df.iterrows()):
|
||
if time.time() > DEADLINE: break
|
||
|
||
msgs = [
|
||
{"role": "system", "content": SYS_TRIVIAL},
|
||
{"role": "user", "content": f"{r['context'].strip()}\n\n{r['query'].strip()}"}
|
||
]
|
||
prompt_text = tok.apply_chat_template(msgs, tokenize=False, add_generation_prompt=True)
|
||
inputs = tok(prompt_text, return_tensors="pt", max_length=6144, truncation=True).to(model.device)
|
||
|
||
try:
|
||
with torch.no_grad():
|
||
out = model.generate(**inputs, max_new_tokens=256, do_sample=False, repetition_penalty=1.0)
|
||
text = tok.decode(out[0][inputs["input_ids"].shape[1]:], skip_special_tokens=True).strip()
|
||
ans = parse_raw_lines(text, n, srcs[i])
|
||
|
||
preds[i] = ans
|
||
samples[i].append(ans)
|
||
except Exception as e:
|
||
print(f"Error on {i}: {e}")
|
||
|
||
write_submission(preds, expl)
|
||
print(f"Pass 1 Complete. Time elapsed: {time.time()-T0:.1f}s", flush=True)
|
||
|
||
# ---- PASSES 2-N: SELF-CONSISTENCY VOTING ----
|
||
for pass_num in range(MAX_SAMPLES):
|
||
if time.time() > (DEADLINE - EXPLAIN_RESERVE): break
|
||
print(f"Starting Voting Pass {pass_num+2} (Time left: {DEADLINE - time.time():.1f}s)...", flush=True)
|
||
|
||
for i, n, (_, r) in zip(ids, ns, df.iterrows()):
|
||
if time.time() > (DEADLINE - EXPLAIN_RESERVE): break
|
||
|
||
msgs = [
|
||
{"role": "system", "content": SYS_TRIVIAL},
|
||
{"role": "user", "content": f"{r['context'].strip()}\n\n{r['query'].strip()}"}
|
||
]
|
||
prompt_text = tok.apply_chat_template(msgs, tokenize=False, add_generation_prompt=True)
|
||
inputs = tok(prompt_text, return_tensors="pt", max_length=6144, truncation=True).to(model.device)
|
||
|
||
try:
|
||
with torch.no_grad():
|
||
out = model.generate(**inputs, max_new_tokens=256, do_sample=True, temperature=SAMPLE_TEMP, top_p=0.9, repetition_penalty=1.0)
|
||
text = tok.decode(out[0][inputs["input_ids"].shape[1]:], skip_special_tokens=True).strip()
|
||
if text:
|
||
s_ans = parse_raw_lines(text, n, srcs[i])
|
||
samples[i].append(s_ans)
|
||
|
||
# Apply Voting Immediately
|
||
if len(samples[i]) >= 3:
|
||
g = samples[i][0]
|
||
preds[i] = fit_to_n([vote([s[k] for s in samples[i] if k < len(s)], anchor=g[k] if k < len(g) else None) for k in range(n)], n, srcs[i])
|
||
except Exception:
|
||
pass
|
||
|
||
write_submission(preds, expl)
|
||
|
||
# ---- ISOLATED JURY EXPLANATION PASS ----
|
||
if time.time() < DEADLINE:
|
||
print(f"Starting Isolated Explanation Pass (Time left: {DEADLINE - time.time():.1f}s)...", flush=True)
|
||
for i, n, (_, r) in zip(ids, ns, df.iterrows()):
|
||
if time.time() > DEADLINE: break
|
||
|
||
ans_str = "\n".join(f"- {a}" for a in preds[i])
|
||
msgs = [
|
||
{"role": "system", "content": EXPLAIN_SYSTEM},
|
||
{"role": "user", "content": f"{r['context'].strip()}\n\n{r['query'].strip()}\n\nAnswers given:\n{ans_str}\n\nBriefly explain the linguistic rules."}
|
||
]
|
||
|
||
prompt_text = tok.apply_chat_template(msgs, tokenize=False, add_generation_prompt=True)
|
||
inputs = tok(prompt_text, return_tensors="pt", max_length=6144, truncation=True).to(model.device)
|
||
|
||
try:
|
||
with torch.no_grad():
|
||
out = model.generate(**inputs, max_new_tokens=150, do_sample=False, repetition_penalty=1.0)
|
||
text = tok.decode(out[0][inputs["input_ids"].shape[1]:], skip_special_tokens=True).strip()
|
||
e = re.sub(r"\\s+", " ", text)
|
||
if e: expl[i] = e[:1000]
|
||
except Exception:
|
||
pass
|
||
|
||
write_submission(preds, expl)
|
||
|
||
print(f"DONE. Completed execution in {time.time()-T0:.0f}s elapsed.", flush=True)
|