初始化项目,由ModelHub XC社区提供模型
Model: newendiesports/iol-2026-qwen14b-entry Source: Original Platform
This commit is contained in:
251
script.py
Normal file
251
script.py
Normal file
@@ -0,0 +1,251 @@
|
||||
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)
|
||||
Reference in New Issue
Block a user