初始化项目,由ModelHub XC社区提供模型
Model: jbuaba/iolai-2026-qwen25-14b Source: Original Platform
This commit is contained in:
26
solver/__init__.py
Normal file
26
solver/__init__.py
Normal 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
96
solver/items.py
Normal 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
173
solver/matching.py
Normal 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
145
solver/minimal.py
Normal 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
202
solver/model.py
Normal 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
42
solver/normalize.py
Normal 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
99
solver/runtime.py
Normal 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
|
||||
Reference in New Issue
Block a user