初始化项目,由ModelHub XC社区提供模型
Model: divaspoudel/iol-ai-challenge Source: Original Platform
This commit is contained in:
315
script.py
Normal file
315
script.py
Normal file
@@ -0,0 +1,315 @@
|
||||
import os
|
||||
os.environ["HF_HUB_OFFLINE"] = "1"
|
||||
os.environ["TRANSFORMERS_OFFLINE"] = "1"
|
||||
|
||||
import re
|
||||
import json
|
||||
import time
|
||||
import sys
|
||||
import subprocess
|
||||
|
||||
try:
|
||||
import bitsandbytes # noqa: F401
|
||||
except ImportError:
|
||||
subprocess.run([sys.executable, "-m", "pip", "install", "-q", "bitsandbytes"], check=True)
|
||||
import bitsandbytes
|
||||
|
||||
import torch
|
||||
import pandas as pd
|
||||
from transformers import AutoTokenizer, AutoModelForCausalLM, BitsAndBytesConfig
|
||||
|
||||
MODEL_ID = "."
|
||||
TIME_LIMIT_S = 28 * 60
|
||||
START = time.time()
|
||||
|
||||
# Try 8-bit first (much better reasoning fidelity than 4-bit). Falls back to 4-bit
|
||||
# only if it doesn't fit in memory. Precision matters a lot for symbolic pattern work.
|
||||
USE_8BIT_FIRST = True
|
||||
|
||||
|
||||
def time_left():
|
||||
return TIME_LIMIT_S - (time.time() - START)
|
||||
|
||||
|
||||
def load_model():
|
||||
if USE_8BIT_FIRST:
|
||||
try:
|
||||
bnb_config = BitsAndBytesConfig(load_in_8bit=True)
|
||||
tok = AutoTokenizer.from_pretrained(MODEL_ID, local_files_only=True)
|
||||
model = AutoModelForCausalLM.from_pretrained(
|
||||
MODEL_ID,
|
||||
quantization_config=bnb_config,
|
||||
device_map="auto",
|
||||
local_files_only=True,
|
||||
).eval()
|
||||
print("Loaded in 8-bit.", flush=True)
|
||||
return tok, model
|
||||
except Exception as e:
|
||||
print(f"8-bit load failed ({e}), falling back to 4-bit.", flush=True)
|
||||
|
||||
bnb_config = BitsAndBytesConfig(
|
||||
load_in_4bit=True,
|
||||
bnb_4bit_quant_type="nf4",
|
||||
bnb_4bit_compute_dtype=torch.float16,
|
||||
bnb_4bit_use_double_quant=True,
|
||||
)
|
||||
tok = AutoTokenizer.from_pretrained(MODEL_ID, local_files_only=True)
|
||||
model = AutoModelForCausalLM.from_pretrained(
|
||||
MODEL_ID,
|
||||
quantization_config=bnb_config,
|
||||
device_map="auto",
|
||||
torch_dtype=torch.float16,
|
||||
local_files_only=True,
|
||||
).eval()
|
||||
return tok, model
|
||||
|
||||
|
||||
def generate(messages, max_new_tokens=900, do_sample=False, temperature=0.7):
|
||||
ids = tok.apply_chat_template(
|
||||
messages, add_generation_prompt=True, return_tensors="pt"
|
||||
).to(model.device)
|
||||
gen_kwargs = dict(
|
||||
max_new_tokens=max_new_tokens,
|
||||
pad_token_id=tok.eos_token_id,
|
||||
)
|
||||
if do_sample:
|
||||
gen_kwargs.update(do_sample=True, temperature=temperature, top_p=0.9)
|
||||
else:
|
||||
gen_kwargs.update(do_sample=False, temperature=None, top_p=None)
|
||||
|
||||
with torch.no_grad():
|
||||
out = model.generate(ids, **gen_kwargs)
|
||||
return tok.decode(out[0][ids.shape[-1]:], skip_special_tokens=True).strip()
|
||||
|
||||
|
||||
def count_items(query: str) -> int:
|
||||
nums = re.findall(r"(?m)^\s*(\d+)\s*[.)]\s", query)
|
||||
if nums:
|
||||
return max(int(n) for n in nums)
|
||||
lines = [l for l in query.splitlines() if l.strip()]
|
||||
return max(1, len(lines))
|
||||
|
||||
|
||||
def extract_json_object(text: str):
|
||||
"""Pull the LAST valid JSON object out of a text blob (final-answer JSON should
|
||||
come after the reasoning, so we scan from the end)."""
|
||||
text = text.strip()
|
||||
try:
|
||||
return json.loads(text)
|
||||
except Exception:
|
||||
pass
|
||||
# find all {...} candidates, try from the last one backwards
|
||||
candidates = list(re.finditer(r"\{.*?\}", text, re.DOTALL))
|
||||
for m in reversed(candidates):
|
||||
try:
|
||||
return json.loads(m.group(0))
|
||||
except Exception:
|
||||
continue
|
||||
# try greedy full-span match as last resort
|
||||
m = re.search(r"\{.*\}", text, re.DOTALL)
|
||||
if m:
|
||||
try:
|
||||
return json.loads(m.group(0))
|
||||
except Exception:
|
||||
pass
|
||||
return None
|
||||
|
||||
|
||||
def extract_answers_fallback(text: str, n: int):
|
||||
lines = [l.strip() for l in text.splitlines() if l.strip()]
|
||||
cleaned = []
|
||||
for l in lines:
|
||||
l = re.sub(r"^\s*\d+[.)]\s*", "", l)
|
||||
l = l.strip(" -\t")
|
||||
if l:
|
||||
cleaned.append(l)
|
||||
if len(cleaned) >= n:
|
||||
return cleaned[-n:] # prefer the tail: likely the final-answer block
|
||||
parts = [p.strip() for p in text.replace("\n", ",").split(",") if p.strip()]
|
||||
if len(parts) >= n:
|
||||
return parts[-n:]
|
||||
cleaned = cleaned or parts or [text.strip()]
|
||||
while len(cleaned) < n:
|
||||
cleaned.append(cleaned[-1] if cleaned else "")
|
||||
return cleaned[:n]
|
||||
|
||||
|
||||
def fit_answers(answers, n):
|
||||
answers = list(answers)
|
||||
if len(answers) < n:
|
||||
answers = answers + [answers[-1] if answers else ""] * (n - len(answers))
|
||||
return answers[:n]
|
||||
|
||||
|
||||
# Single combined system prompt: reasoning + mandatory self-verification + strict
|
||||
# final JSON block. This replaces the old two-call reasoning->extraction pipeline,
|
||||
# roughly halving generation cost per row while forcing the model to check its own
|
||||
# rule against every given example (the single highest-leverage habit for IOL-style
|
||||
# problems) before committing to an answer.
|
||||
SOLVE_SYSTEM = (
|
||||
"You are an expert at International Linguistics Olympiad (IOL) problems. You will "
|
||||
"be given example data (context) in an unfamiliar language and a query asking you "
|
||||
"to translate, match, fill blanks, or convert forms.\n\n"
|
||||
"Follow this exact procedure:\n"
|
||||
"1. REASONING: Identify morphemes, word order, and correspondences using ONLY the "
|
||||
"given examples. Do not assume rules from natural languages you know unless the "
|
||||
"data supports them.\n"
|
||||
"2. VERIFY: Explicitly re-derive every example in the context using your proposed "
|
||||
"rules and confirm each one matches. If any mismatch, revise your rule before "
|
||||
"continuing.\n"
|
||||
"3. ANSWER: Apply the verified rule to the query.\n"
|
||||
"4. Output a final line starting with FINAL_JSON: followed by a single JSON object "
|
||||
"with exactly two keys:\n"
|
||||
' "answers": a JSON list of strings, one per numbered item in the query, in the '
|
||||
"same order as the query items.\n"
|
||||
' "explanation": 3-6 short bullet points (use \\n between them) summarizing the '
|
||||
"key rules used.\n"
|
||||
"The FINAL_JSON line must be the last thing you output, and must contain nothing "
|
||||
"else on that line besides FINAL_JSON: {...}."
|
||||
)
|
||||
|
||||
|
||||
def solve_once(problem_text: str, n_items: int, max_new_tokens: int, do_sample=False, temperature=0.7):
|
||||
raw = generate(
|
||||
[
|
||||
{"role": "system", "content": SOLVE_SYSTEM},
|
||||
{"role": "user", "content": problem_text + f"\n\n(The query has {n_items} numbered item(s).)"},
|
||||
],
|
||||
max_new_tokens=max_new_tokens,
|
||||
do_sample=do_sample,
|
||||
temperature=temperature,
|
||||
)
|
||||
|
||||
obj = None
|
||||
m = re.search(r"FINAL_JSON:\s*(\{.*\})", raw, re.DOTALL)
|
||||
if m:
|
||||
try:
|
||||
obj = json.loads(m.group(1))
|
||||
except Exception:
|
||||
obj = extract_json_object(m.group(1))
|
||||
if obj is None:
|
||||
obj = extract_json_object(raw)
|
||||
|
||||
answers, explanation = None, ""
|
||||
if isinstance(obj, dict):
|
||||
a = obj.get("answers")
|
||||
if isinstance(a, list):
|
||||
answers = [str(x) for x in a]
|
||||
explanation = str(obj.get("explanation", "")).strip()
|
||||
|
||||
if answers is None:
|
||||
answers = extract_answers_fallback(raw, n_items)
|
||||
if not explanation:
|
||||
snippet = raw.strip().replace("\n", " ")
|
||||
explanation = (snippet[:400] + "...") if len(snippet) > 400 else snippet
|
||||
|
||||
return fit_answers(answers, n_items), explanation, raw
|
||||
|
||||
|
||||
def majority_vote(all_answers, n_items):
|
||||
"""Per-position majority vote across N sampled runs."""
|
||||
final = []
|
||||
for i in range(n_items):
|
||||
votes = {}
|
||||
for ans_list in all_answers:
|
||||
v = ans_list[i] if i < len(ans_list) else ""
|
||||
votes[v] = votes.get(v, 0) + 1
|
||||
final.append(max(votes.items(), key=lambda kv: kv[1])[0])
|
||||
return final
|
||||
|
||||
|
||||
def solve_row(context: str, query: str, n_items: int, row_budget_s: float):
|
||||
problem_text = f"{context.strip()}\n\n{query.strip()}"
|
||||
|
||||
# Scale tokens to problem size; cap so one bad row can't eat the whole budget.
|
||||
approx_len = len(context) + len(query)
|
||||
max_new = 700 if approx_len < 1500 else 1000
|
||||
max_new = min(max_new, 1100)
|
||||
|
||||
row_start = time.time()
|
||||
|
||||
def row_time_left():
|
||||
return row_budget_s - (time.time() - row_start)
|
||||
|
||||
# First (greedy, most reliable) attempt.
|
||||
answers, explanation, _ = solve_once(problem_text, n_items, max_new, do_sample=False)
|
||||
|
||||
# If time allows, take 2 more sampled attempts and majority-vote for robustness.
|
||||
# This costs roughly the same as the old broken two-pass pipeline, but spends the
|
||||
# extra budget on getting the ANSWER right instead of just reformatting it.
|
||||
if row_time_left() > max_new * 0.03: # crude per-token-time heuristic guard
|
||||
try:
|
||||
extra_runs = [answers]
|
||||
for t in (0.5, 0.8):
|
||||
if row_time_left() < 20:
|
||||
break
|
||||
a2, _, _ = solve_once(problem_text, n_items, max_new, do_sample=True, temperature=t)
|
||||
extra_runs.append(a2)
|
||||
if len(extra_runs) > 1:
|
||||
answers = majority_vote(extra_runs, n_items)
|
||||
except Exception:
|
||||
pass # keep the greedy answer if sampling fails
|
||||
|
||||
return answers, explanation
|
||||
|
||||
|
||||
def main():
|
||||
global tok, model
|
||||
tok, model = load_model()
|
||||
|
||||
df = pd.read_csv("/tmp/data/test.csv", dtype=str).fillna("")
|
||||
n_rows = len(df)
|
||||
out_rows = []
|
||||
|
||||
for i, r in df.iterrows():
|
||||
rid = r["id"]
|
||||
context, query = r.get("context", ""), r.get("query", "")
|
||||
n_items = count_items(query)
|
||||
|
||||
rows_remaining = n_rows - i
|
||||
# Reserve a small safety margin (30s) for writing the CSV at the end.
|
||||
safe_time_left = max(time_left() - 30, 0)
|
||||
row_budget = safe_time_left / max(rows_remaining, 1)
|
||||
|
||||
if safe_time_left < 20:
|
||||
# Truly out of time: still try a fast, cheap single greedy pass rather
|
||||
# than silently submitting blanks — a bad guess beats an empty string
|
||||
# on chrF and exact-match both.
|
||||
try:
|
||||
answers, explanation, _ = solve_once(
|
||||
f"{context.strip()}\n\n{query.strip()}", n_items, max_new_tokens=300
|
||||
)
|
||||
except Exception as e:
|
||||
answers = [""] * n_items
|
||||
explanation = f"(error: {e})"
|
||||
else:
|
||||
try:
|
||||
answers, explanation = solve_row(context, query, n_items, row_budget)
|
||||
except Exception as e:
|
||||
answers = [""] * n_items
|
||||
explanation = f"(error: {e})"
|
||||
|
||||
out_rows.append(
|
||||
{
|
||||
"id": rid,
|
||||
"pred": json.dumps(answers, ensure_ascii=False),
|
||||
"explanation": explanation,
|
||||
}
|
||||
)
|
||||
print(
|
||||
f"[{i + 1}/{n_rows}] id={rid} n_items={n_items} "
|
||||
f"row_budget={row_budget:.0f}s time_left={time_left():.0f}s",
|
||||
flush=True,
|
||||
)
|
||||
|
||||
# Checkpoint after every row so a hard timeout still leaves a submission
|
||||
# instead of losing everything if the process gets killed mid-loop.
|
||||
pd.DataFrame(out_rows, columns=["id", "pred", "explanation"]).to_csv(
|
||||
"submission.csv", index=False
|
||||
)
|
||||
|
||||
print("wrote submission.csv", flush=True)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
Reference in New Issue
Block a user