初始化项目,由ModelHub XC社区提供模型

Model: jbuaba/iolai-2026-qwen25-14b
Source: Original Platform
This commit is contained in:
ModelHub XC
2026-09-13 04:46:18 +08:00
commit 13827db9e9
21 changed files with 457381 additions and 0 deletions

26
solver/__init__.py Normal file
View File

@@ -0,0 +1,26 @@
from .items import count_answer_slots, pad_short_answers
from .matching import solve_matching
from .minimal import solve_row
from .model import (
GenStats,
ModelBundle,
apply_greedy_decoding,
generate,
generate_with_stats,
load_model,
)
from .normalize import safe_normalize_answers
__all__ = [
"GenStats",
"ModelBundle",
"apply_greedy_decoding",
"count_answer_slots",
"generate",
"generate_with_stats",
"load_model",
"pad_short_answers",
"safe_normalize_answers",
"solve_matching",
"solve_row",
]

96
solver/items.py Normal file
View File

@@ -0,0 +1,96 @@
from __future__ import annotations
import re
_LINE_NUMBER = re.compile(r"^[ \t]*(\d{1,3})[.)\]]", re.M)
_PAREN_NUMBER = re.compile(r"\((\d{1,3})\)")
_ITEM_RANGE = re.compile(r"\(?(\d{1,3})\s*(?:[-–—]|to)\s*(\d{1,3})\)?")
_LINE_LETTER = re.compile(r"^[ \t]*([A-Z])[.)\]]\s", re.M)
_PAREN_LETTER = re.compile(r"\(([A-Z])\)")
_NUMBERED_ANSWER = re.compile(r"^\s*\(?(\d{1,3})\)?[.):\]]\s*(.+)$")
def count_answer_slots(query: str, context: str = "") -> int:
"""How many answers this problem expects. Always >= 1."""
query = query or ""
line_nums = [int(m) for m in _LINE_NUMBER.findall(query)]
paren_nums = [int(m) for m in _PAREN_NUMBER.findall(query)]
range_count = 0
for start, end in _ITEM_RANGE.findall(query):
start_i, end_i = int(start), int(end)
if 0 < end_i - start_i < 60:
range_count = max(range_count, end_i - start_i + 1)
marker_count = max(len(set(line_nums)), len(set(paren_nums)))
if range_count and marker_count and range_count != marker_count:
return marker_count
marker_count = max(
marker_count,
len(set(_LINE_LETTER.findall(query))),
len(set(_PAREN_LETTER.findall(query))),
)
slot_count = max(range_count, marker_count)
if slot_count > 1:
return slot_count
lines = [line.strip() for line in query.splitlines() if line.strip()]
if len(lines) > 1:
head = lines[0]
body = lines[1:] if head.endswith((":", ".")) else lines
if body:
return len(body)
if context:
context_nums = len(set(int(m) for m in _LINE_NUMBER.findall(context)))
if context_nums > 1:
return context_nums
context_letters = len(set(_LINE_LETTER.findall(context)))
if context_letters > 1:
return context_letters
return max(slot_count, 1)
def source_fallbacks(query: str, slot_count: int) -> list[str]:
"""Per-item source text used when the model returns too few lines.
Empty predictions score zero on EM and chrF. Echoing the query item is
usually wrong on EM but recovers chrF on fill-blank / transcription tasks.
"""
query = query or ""
sources: list[str] = []
for line in query.splitlines():
stripped = line.strip()
if not stripped:
continue
match = _NUMBERED_ANSWER.match(stripped)
if match:
sources.append(match.group(2).strip())
if not sources:
lines = [line.strip() for line in query.splitlines() if line.strip()]
if len(lines) > 1 and lines[0].endswith((":", ".")):
sources = lines[1:]
sources = [s.split("|")[0].strip() if "|" in s else s for s in sources]
sources = [s for s in sources if s]
while len(sources) < slot_count:
sources.append(sources[-1] if sources else "?")
return sources[:slot_count]
def pad_short_answers(
answers: list[str],
slot_count: int,
fallbacks: list[str] | None = None,
) -> list[str]:
"""Pad undersized answer lists only. Never truncate — the grader keeps the first N."""
slot_count = max(1, int(slot_count))
padded = [str(a).strip() if a and str(a).strip() else "?" for a in answers]
while len(padded) < slot_count:
if fallbacks and len(padded) < len(fallbacks):
fill = str(fallbacks[len(padded)]).strip() or "?"
else:
fill = "?"
padded.append(fill)
return padded

173
solver/matching.py Normal file
View File

@@ -0,0 +1,173 @@
from __future__ import annotations
import re
from typing import Any
from .model import ModelBundle
_OPT_LINE = re.compile(r"^[ \t]*([A-Za-z])[.)]\s+(.+)$", re.M)
_ITEM_LINE = re.compile(r"^[ \t]*(\d{1,3})[.)]\s+(.+)$", re.M)
def parse_matching_block(
context: str,
) -> tuple[list[tuple[int, str]], list[tuple[str, str]]]:
items = [(int(a), b.strip()) for a, b in _ITEM_LINE.findall(context or "")]
opts = [(a.upper(), b.strip()) for a, b in _OPT_LINE.findall(context or "")]
seen_items: set[int] = set()
items = [x for x in items if not (x[0] in seen_items or seen_items.add(x[0]))]
seen_opts: set[str] = set()
opts = [x for x in opts if not (x[0] in seen_opts or seen_opts.add(x[0]))]
return items, opts
def best_assignment(score: list[list[float]]) -> list[int]:
n = len(score)
m = len(score[0]) if score else 0
if n == 0 or m == 0:
return []
try:
import numpy as np
from scipy.optimize import linear_sum_assignment
_, cols = linear_sum_assignment(-np.array(score, dtype=float))
return list(cols)
except Exception:
pass
used: set[int] = set()
out = [0] * n
order = sorted(
range(n),
key=lambda i: -(
max(score[i]) - sorted(score[i])[-2] if m > 1 else max(score[i])
),
)
for i in order:
j = max(
(jj for jj in range(m) if jj not in used),
key=lambda jj: score[i][jj],
default=0,
)
used.add(j)
out[i] = j
for _ in range(4):
improved = False
for a in range(n):
for b in range(a + 1, n):
cur = score[a][out[a]] + score[b][out[b]]
alt = score[a][out[b]] + score[b][out[a]]
if alt > cur + 1e-9:
out[a], out[b] = out[b], out[a]
improved = True
if not improved:
break
return out
def repair_letter_bijection(answers: list[str]) -> list[str]:
if len(answers) < 3:
return answers
if not all(re.fullmatch(r"[A-Z]", a or "") for a in answers):
return answers
n = len(answers)
universe = [chr(ord("A") + i) for i in range(n)]
if len(set(answers)) == n:
return answers
unused = [letter for letter in universe if letter not in set(answers)]
if not unused:
return answers
seen: set[str] = set()
out: list[str] = []
for answer in answers:
if answer in seen and unused:
out.append(unused.pop(0))
else:
seen.add(answer)
out.append(answer)
return out
def solve_matching(
bundle: ModelBundle,
row: Any,
slot_count: int,
*,
batch_size: int = 4,
) -> list[str] | None:
"""Score (item, option) next-token logprobs and take a 1-1 assignment.
Returns None on any failure so callers can fall back to greedy decode.
"""
import torch
items, opts = parse_matching_block(str(row.get("context", "") or ""))
if len(items) < 3 or len(opts) < 3 or len(items) != slot_count:
return None
letters = [opt[0] for opt in opts]
candidate_ids: list[list[int]] = []
for letter in letters:
ids: set[int] = set()
for form in (letter, " " + letter):
tokens = bundle.tok.encode(form, add_special_tokens=False)
if tokens:
ids.add(int(tokens[0]))
if not ids:
return None
candidate_ids.append(sorted(ids))
context = str(row.get("context", "") or "").strip()
prompts: list[str] = []
for number, item_text in items:
messages = [
{
"role": "system",
"content": (
"You match items to their correct counterparts in a "
"linguistics problem. Reply with one option letter only."
),
},
{
"role": "user",
"content": (
f"{context}\n\nWhich lettered option corresponds to item "
f"{number} ({item_text})? Reply with the option letter only."
),
},
]
prompts.append(
bundle.tok.apply_chat_template(
messages, tokenize=False, add_generation_prompt=True
)
)
score: list[list[float]] = []
try:
for start in range(0, len(prompts), batch_size):
chunk = prompts[start : start + batch_size]
encoded = bundle.tok(
chunk,
return_tensors="pt",
padding=True,
truncation=True,
max_length=6144,
)
encoded = {k: v.to(bundle.model.device) for k, v in encoded.items()}
with torch.no_grad():
logits = bundle.model(**encoded).logits[:, -1, :].float()
logprobs = torch.log_softmax(logits, dim=-1)
for batch_index in range(len(chunk)):
score.append(
[
max(float(logprobs[batch_index, token_id].item()) for token_id in ids)
for ids in candidate_ids
]
)
except Exception:
return None
columns = best_assignment(score)
if len(columns) != slot_count:
return None
return repair_letter_bijection([letters[col] for col in columns])

145
solver/minimal.py Normal file
View File

@@ -0,0 +1,145 @@
from __future__ import annotations
import re
from typing import Any, Mapping
from .items import count_answer_slots, pad_short_answers, source_fallbacks
from .matching import solve_matching
from .model import GenStats, ModelBundle, generate_with_stats
from .normalize import safe_normalize_answers
SYSTEM_PROMPT = (
"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."
)
MAX_NEW_TOKENS = 512
_LEADING_MARKER = re.compile(r"^\s*(?:\(?\d{1,3}\)?[.):\]]\s*|[-*•]\s+)")
_CODE_FENCE = re.compile(r"^```[a-zA-Z]*\s*$")
_PREAMBLE = 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|"
r"let me\b|first,|so,|therefore\b|thus\b)",
re.I,
)
_NUMBERED_LINE = re.compile(r"^\s*\(?(\d{1,3})\)?[.):\]]\s*(.+)$")
def clean_answer_text(text: str) -> str:
text = (text or "").strip()
text = _LEADING_MARKER.sub("", text)
text = text.strip().strip("`").strip()
if len(text) >= 2 and text[0] == text[-1] and text[0] in "\"'“”":
text = text[1:-1].strip()
return text.strip()
def extract_answer_lines(model_text: str, slot_count: int) -> list[str]:
labeled: dict[int, str] = {}
unlabeled: list[str] = []
for line in (model_text or "").splitlines():
if not line.strip() or _CODE_FENCE.match(line):
continue
if _PREAMBLE.match(line):
continue
numbered = _NUMBERED_LINE.match(line.strip())
if numbered:
label = int(numbered.group(1))
value = clean_answer_text(numbered.group(2))
if value and not _PREAMBLE.match(value):
labeled[label] = value
continue
cleaned = clean_answer_text(line)
if cleaned and not _PREAMBLE.match(cleaned):
unlabeled.append(cleaned)
if labeled:
slots: list[str | None] = [None] * slot_count
for label, value in labeled.items():
if 1 <= label <= slot_count:
slots[label - 1] = value
fill_from = 0
for index in range(slot_count):
if slots[index] is None and fill_from < len(unlabeled):
slots[index] = unlabeled[fill_from]
fill_from += 1
ordered = [s for s in slots if s is not None]
leftover = unlabeled[fill_from:]
return ordered + leftover
return unlabeled
def _greedy_answers(
row: Mapping[str, Any],
bundle: ModelBundle,
*,
max_new_tokens: int,
generate_fn,
) -> tuple[list[str], str, GenStats]:
context = str(row.get("context", "") or "").strip()
query = str(row.get("query", "") or "").strip()
messages = [
{"role": "system", "content": SYSTEM_PROMPT},
{"role": "user", "content": f"{context}\n\n{query}"},
]
if generate_fn is not None:
raw = generate_fn(bundle, messages, max_new_tokens)
stats = GenStats(
prompt_tokens=0,
new_tokens=0,
hit_max_new=False,
eos_limited=True,
)
else:
raw, stats = generate_with_stats(
bundle, messages, max_new_tokens=max_new_tokens
)
slot_count = count_answer_slots(query, context)
fallbacks = source_fallbacks(query, slot_count)
answers = pad_short_answers(
extract_answer_lines(raw, slot_count),
slot_count,
fallbacks,
)
return answers, raw, stats
def solve_row(
row: Mapping[str, Any],
bundle: ModelBundle,
*,
max_new_tokens: int = MAX_NEW_TOKENS,
generate_fn=None,
) -> tuple[list[str], str, GenStats]:
context = str(row.get("context", "") or "").strip()
query = str(row.get("query", "") or "").strip()
task_type = str(row.get("task_type", "") or "").strip().lower()
slot_count = count_answer_slots(query, context)
if task_type == "match_letters" and generate_fn is None:
try:
matched = solve_matching(bundle, row, slot_count)
if matched and len(matched) == slot_count:
stats = GenStats(
prompt_tokens=0,
new_tokens=0,
hit_max_new=False,
eos_limited=True,
)
return (
safe_normalize_answers(matched, task_type),
"",
stats,
)
except Exception:
pass
answers, raw, stats = _greedy_answers(
row, bundle, max_new_tokens=max_new_tokens, generate_fn=generate_fn
)
return safe_normalize_answers(answers, task_type), raw, stats

202
solver/model.py Normal file
View File

@@ -0,0 +1,202 @@
from __future__ import annotations
import os
from dataclasses import dataclass
from typing import Any
DEFAULT_MODEL_ID = "."
LOAD_MODE_AWQ = "awq"
LOAD_MODE_BNB = "bnb"
@dataclass
class ModelBundle:
tok: Any
model: Any
model_id: str
@dataclass(frozen=True)
class GenStats:
prompt_tokens: int
new_tokens: int
hit_max_new: bool
eos_limited: bool
def assert_gpu_resident(bundle: ModelBundle) -> None:
import torch
if not torch.cuda.is_available():
print("warn: CUDA unavailable", flush=True)
return
bad = []
for name, param in bundle.model.named_parameters():
if not str(param.device).startswith("cuda"):
bad.append(f"{name}:{param.device}")
if len(bad) >= 5:
break
if bad:
print(f"warn: non-CUDA parameters: {bad}", flush=True)
return
print(
f"gpu ok | VRAM {torch.cuda.memory_allocated() / 1e9:.2f} GB",
flush=True,
)
def load_model(
model_id: str | None = None,
*,
offline: bool | None = None,
load_mode: str | None = None,
) -> ModelBundle:
import torch
from transformers import AutoModelForCausalLM, AutoTokenizer
model_id = model_id or os.environ.get("IOL_MODEL_ID", DEFAULT_MODEL_ID)
load_mode = (load_mode or os.environ.get("IOL_LOAD", LOAD_MODE_AWQ)).strip().lower()
if offline is None:
offline = model_id == DEFAULT_MODEL_ID or os.environ.get("HF_HUB_OFFLINE") == "1"
if offline:
os.environ["HF_HUB_OFFLINE"] = "1"
os.environ["TRANSFORMERS_OFFLINE"] = "1"
else:
os.environ.pop("HF_HUB_OFFLINE", None)
os.environ.pop("TRANSFORMERS_OFFLINE", None)
os.environ.setdefault("TOKENIZERS_PARALLELISM", "false")
tok = AutoTokenizer.from_pretrained(model_id)
if tok.pad_token_id is None and tok.eos_token_id is not None:
tok.pad_token = tok.eos_token
dtype_kwargs = _dtype_kwargs(torch)
preferred: Any = {"": 0} if torch.cuda.is_available() else "auto"
def _load(device_map: Any):
if load_mode == LOAD_MODE_BNB:
from transformers import BitsAndBytesConfig
return AutoModelForCausalLM.from_pretrained(
model_id,
quantization_config=BitsAndBytesConfig(
load_in_4bit=True,
bnb_4bit_compute_dtype=torch.float16,
bnb_4bit_use_double_quant=True,
bnb_4bit_quant_type="nf4",
),
device_map=device_map,
).eval()
try:
return AutoModelForCausalLM.from_pretrained(
model_id,
device_map=device_map,
**dtype_kwargs,
).eval()
except ImportError as exc:
raise ImportError(
"AWQ load failed; install gptqmodel/autoawq or use IOL_LOAD=bnb"
) from exc
try:
model = _load(preferred)
except ImportError:
raise
except Exception as exc:
if preferred == "auto":
raise
print(f"warn: device_map retry auto ({exc})", flush=True)
model = _load("auto")
apply_greedy_decoding(model)
return ModelBundle(tok=tok, model=model, model_id=model_id)
def apply_greedy_decoding(model: Any) -> None:
try:
cfg = model.generation_config
cfg.do_sample = False
cfg.repetition_penalty = 1.0
for key in ("temperature", "top_p", "top_k", "typical_p"):
if hasattr(cfg, key):
setattr(cfg, key, None)
except Exception:
pass
def _dtype_kwargs(torch_mod) -> dict:
try:
import inspect
from transformers import AutoModelForCausalLM
if "dtype" in inspect.signature(AutoModelForCausalLM.from_pretrained).parameters:
return {"dtype": torch_mod.float16}
except Exception:
pass
return {"torch_dtype": torch_mod.float16}
def _prompt_tensors(tok: Any, model: Any, messages: list[dict[str, str]]):
text = tok.apply_chat_template(
messages,
tokenize=False,
add_generation_prompt=True,
)
enc = tok(
text,
return_tensors="pt",
truncation=True,
max_length=6144,
)
moved = {k: v.to(model.device) for k, v in enc.items()}
return moved, int(moved["input_ids"].shape[-1])
def _pad_token_id(bundle: ModelBundle) -> int | None:
if getattr(bundle.tok, "pad_token_id", None) is not None:
return int(bundle.tok.pad_token_id)
if getattr(bundle.tok, "eos_token_id", None) is not None:
return int(bundle.tok.eos_token_id)
return None
def generate(
bundle: ModelBundle,
messages: list[dict[str, str]],
*,
max_new_tokens: int = 512,
) -> str:
text, _ = generate_with_stats(bundle, messages, max_new_tokens=max_new_tokens)
return text
def generate_with_stats(
bundle: ModelBundle,
messages: list[dict[str, str]],
*,
max_new_tokens: int = 512,
) -> tuple[str, GenStats]:
import torch
apply_greedy_decoding(bundle.model)
inputs, prompt_len = _prompt_tensors(bundle.tok, bundle.model, messages)
kwargs: dict[str, Any] = {
"max_new_tokens": max_new_tokens,
"do_sample": False,
"repetition_penalty": 1.0,
}
pad_id = _pad_token_id(bundle)
if pad_id is not None:
kwargs["pad_token_id"] = pad_id
with torch.no_grad():
output = bundle.model.generate(**inputs, **kwargs)
new_tokens = int(output.shape[-1] - prompt_len)
hit_max = new_tokens >= max_new_tokens
text = bundle.tok.decode(output[0][prompt_len:], skip_special_tokens=True).strip()
return text, GenStats(
prompt_tokens=int(prompt_len),
new_tokens=new_tokens,
hit_max_new=hit_max,
eos_limited=not hit_max,
)

42
solver/normalize.py Normal file
View File

@@ -0,0 +1,42 @@
"""Per-line surface normalizers (Lipas/Hul). No arity force / pad / truncate."""
from __future__ import annotations
import re
def normalize_match_letter(ans: str) -> str:
ans = ans.strip()
m = re.fullmatch(r"[\(\[]?([A-Za-z])[\)\]]?[.)]?", ans)
if m:
return m.group(1).upper()
tokens = re.findall(r"\b([A-Za-z])\b", ans)
if tokens:
return tokens[-1].upper()
m = re.search(r"[A-Za-z]", ans)
return m.group(0).upper() if m else ans
def normalize_text_to_num(ans: str) -> str:
a = re.sub(r"(?i)^(answer|ans|result)\s*[:=]\s*", "", ans.strip()).strip()
if re.fullmatch(r"[\d\s+\-*/^=()]+", a.replace(",", "")):
a = a.replace(",", "").replace(" ", "")
if "=" in a and " = " not in a:
a = a.replace("=", " = ")
return a.strip()
m = re.search(r"\d+", a)
return m.group(0) if m and len(a) < 40 else a
def safe_normalize_answers(answers: list[str], task_type: str) -> list[str]:
"""Per-line only. Does not pad, truncate, or reorder."""
task_type = (task_type or "").strip().lower()
out: list[str] = []
for a in answers:
a = a.strip()
if task_type == "match_letters":
a = normalize_match_letter(a)
elif task_type == "text_to_num":
a = normalize_text_to_num(a)
out.append(a)
return out

99
solver/runtime.py Normal file
View File

@@ -0,0 +1,99 @@
"""Offline runtime bootstrap for Qwen3 under the Space's older transformers."""
from __future__ import annotations
import importlib
import importlib.metadata
import os
import subprocess
import sys
from pathlib import Path
RUNTIME_PACKAGES = {
"transformers": "4.51.3",
"tokenizers": "0.21.1",
"huggingface_hub": "0.30.2",
"autoawq": "0.2.9",
}
RUNTIME_WHEELS = (
"transformers-4.51.3-py3-none-any.whl",
"tokenizers-0.21.1-cp39-abi3-manylinux_2_17_x86_64.manylinux2014_x86_64.whl",
"huggingface_hub-0.30.2-py3-none-any.whl",
"autoawq-0.2.9-py3-none-any.whl",
)
def _wheelhouse() -> Path:
env = os.environ.get("IOL_WHEELHOUSE") or os.environ.get("QWEN3_WHEELHOUSE")
if env:
return Path(env)
return Path(__file__).resolve().parent.parent / "wheelhouse"
def _runtime_dir() -> Path:
return Path(os.environ.get("IOL_RUNTIME_DIR", "/tmp/iol_qwen3_runtime"))
def installed_versions() -> dict[str, str]:
versions: dict[str, str] = {}
for package in RUNTIME_PACKAGES:
try:
versions[package] = importlib.metadata.version(package)
except importlib.metadata.PackageNotFoundError:
versions[package] = "missing"
return versions
def ensure_runtime() -> dict[str, str]:
"""Install bundled wheels into /tmp and prefer them on sys.path.
No-op when wheelhouse is absent (local Qwen2.5 packs / unit tests).
"""
wheelhouse = _wheelhouse()
if not wheelhouse.is_dir():
return installed_versions()
wheel_paths = [wheelhouse / name for name in RUNTIME_WHEELS]
missing = [str(path) for path in wheel_paths if not path.is_file()]
if missing:
raise FileNotFoundError(f"Missing offline runtime wheels: {missing}")
runtime_dir = _runtime_dir()
marker = runtime_dir / ".iol-qwen3-runtime-v1"
if not marker.is_file():
runtime_dir.mkdir(parents=True, exist_ok=True)
subprocess.run(
[
sys.executable,
"-m",
"pip",
"install",
"--disable-pip-version-check",
"--no-index",
"--no-deps",
"--upgrade",
"--target",
str(runtime_dir),
*(str(path) for path in wheel_paths),
],
check=True,
timeout=180,
)
marker.write_text("offline Qwen3 runtime installed\n", encoding="utf-8")
runtime_path = str(runtime_dir)
if runtime_path in sys.path:
sys.path.remove(runtime_path)
sys.path.insert(0, runtime_path)
importlib.invalidate_caches()
versions = installed_versions()
mismatches = {
name: (versions[name], expected)
for name, expected in RUNTIME_PACKAGES.items()
if versions[name] != expected
}
if mismatches:
raise RuntimeError(f"Offline runtime version mismatch: {mismatches}")
print(f"offline runtime: {versions}", flush=True)
return versions