Files
iol-2026-qwen14b-entry/script.py
ModelHub XC 17e3da015b 初始化项目,由ModelHub XC社区提供模型
Model: newendiesports/iol-2026-qwen14b-entry
Source: Original Platform
2026-10-04 13:26:33 +08:00

252 lines
9.6 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

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)