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)