260 lines
10 KiB
Python
260 lines
10 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 = 40 # We want maximum possible samples
|
||
|
|
SAMPLE_TEMP = 0.40 # Sweet spot between 0.35 and 0.50
|
||
|
|
|
||
|
|
# 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+)")
|
||
|
|
_FENCE = re.compile(r"^```[a-zA-Z]*\s*$")
|
||
|
|
_CHATTY = re.compile(r"^\s*(?:here (?:are|is)\b|answers?\s*:?\s*$|explanation\b|note\b|okay\b|"
|
||
|
|
r"solution\b|reasoning\b|analysis\b|translations?\s*:?\s*$|the answers?\b|let me\b|first,|so,|therefore\b|thus\b)", re.I)
|
||
|
|
|
||
|
|
def clean_line(s):
|
||
|
|
s = s.strip()
|
||
|
|
s = _STRIP_PREFIX.sub("", s)
|
||
|
|
s = s.strip().strip("`").strip()
|
||
|
|
if len(s) >= 2 and s[0] == s[-1] and s[0] in "\"'“”": s = s[1:-1].strip()
|
||
|
|
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 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=128, do_sample=False, repetition_penalty=1.0)
|
||
|
|
text = tok.decode(out[0][inputs["input_ids"].shape[1]:], skip_special_tokens=True).strip()
|
||
|
|
ans = 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: ACCELERATED MULTI-SAMPLE VOTING ----
|
||
|
|
pass_num = 1
|
||
|
|
while time.time() < (DEADLINE - EXPLAIN_RESERVE) and pass_num < MAX_SAMPLES:
|
||
|
|
pass_num += 2 # We are generating 2 samples per pass
|
||
|
|
print(f"Starting Voting Passes {pass_num-1} & {pass_num} (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:
|
||
|
|
# THROUGHPUT MULTIPLIER: num_return_sequences=2 generates 2 samples at once using shared KV cache!
|
||
|
|
with torch.no_grad():
|
||
|
|
out = model.generate(
|
||
|
|
**inputs,
|
||
|
|
max_new_tokens=128,
|
||
|
|
do_sample=True,
|
||
|
|
temperature=SAMPLE_TEMP,
|
||
|
|
top_p=0.9,
|
||
|
|
repetition_penalty=1.0,
|
||
|
|
num_return_sequences=2
|
||
|
|
)
|
||
|
|
|
||
|
|
# Decode and parse both generated sequences
|
||
|
|
for seq_idx in range(2):
|
||
|
|
text = tok.decode(out[seq_idx][inputs["input_ids"].shape[1]:], skip_special_tokens=True).strip()
|
||
|
|
if text:
|
||
|
|
s_ans = 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 torch.cuda.OutOfMemoryError:
|
||
|
|
torch.cuda.empty_cache()
|
||
|
|
print(f"OOM on {i}, skipping to preserve memory.")
|
||
|
|
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 ~{pass_num} voting passes in {time.time()-T0:.0f}s elapsed.", flush=True)
|