Files
iol-qwen2.5-14b-sft-awq/script.py

665 lines
24 KiB
Python
Raw Permalink Normal View History

#!/usr/bin/env python3
"""Offline Qwen2.5-14B-AWQ + book-RAG submission for the IOL-AI challenge."""
from __future__ import annotations
import argparse
import csv
import io
import json
import os
import re
import sys
from dataclasses import dataclass
from pathlib import Path
from typing import Any
# The evaluation container has no internet access. Fail locally instead of waiting
# for network retries if a model/tokenizer file was not included in the repository.
os.environ.setdefault("HF_HUB_OFFLINE", "1")
os.environ.setdefault("TRANSFORMERS_OFFLINE", "1")
os.environ.setdefault("HF_DATASETS_OFFLINE", "1")
os.environ.setdefault("TOKENIZERS_PARALLELISM", "false")
os.environ.setdefault("PYTORCH_CUDA_ALLOC_CONF", "expandable_segments:True")
REPO_DIR = Path(__file__).resolve().parent
RESOURCE_DIR = REPO_DIR / "rag_resources"
DEFAULT_INPUT = Path(os.environ.get("IOL_TEST_CSV", "/tmp/data/test.csv"))
DEFAULT_OUTPUT = Path(os.environ.get("IOL_SUBMISSION_CSV", "submission.csv"))
from rag_resources.retriever import BookRetriever # noqa: E402
@dataclass
class ParseResult:
thinking_trace: str
final_text: str
answers: list[str]
valid: bool
error: str = ""
def _load_config() -> dict[str, Any]:
with (RESOURCE_DIR / "config.json").open("r", encoding="utf-8") as handle:
config = json.load(handle)
integer_overrides = {
"IOL_TOP_METHODS": "top_methods",
"IOL_TOP_EXAMPLES": "top_examples",
"IOL_RAG_MAX_CHARS": "rag_max_chars",
"IOL_MAX_NEW_TOKENS": "max_new_tokens",
"IOL_EXPLANATION_MAX_NEW_TOKENS": "explanation_max_new_tokens",
"IOL_MAX_ATTEMPTS": "max_attempts",
"IOL_SEED": "seed",
}
for environment_name, config_name in integer_overrides.items():
if environment_name in os.environ:
config[config_name] = int(os.environ[environment_name])
if "IOL_CHAR_TFIDF_WEIGHT" in os.environ:
config["char_tfidf_weight"] = float(os.environ["IOL_CHAR_TFIDF_WEIGHT"])
if "IOL_ENABLE_EXPLANATIONS" in os.environ:
config["enable_explanations"] = os.environ[
"IOL_ENABLE_EXPLANATIONS"
].strip().casefold() in {"1", "true", "yes", "on"}
return config
def _clean_answer_line(line: str) -> str:
value = line.strip()
value = re.sub(r"^(?:[-*•]\s+)", "", value)
value = re.sub(r"^(?:\(?\d+\)?|\(?[A-Za-z]\)?)[.):]\s+", "", value)
value = value.strip()
pairs = {'"': '"', "'": "'", "`": "`", "“": "”", "‘": "’"}
if len(value) >= 2 and value[0] in pairs and value[-1] == pairs[value[0]]:
value = value[1:-1].strip()
return value
def parse_model_response(thinking_trace: str, final_text: str) -> ParseResult:
"""Parse a separate reasoning trace and a strict FINAL ANSWERS block."""
marker_matches = list(
re.finditer(r"(?im)^\s*FINAL\s+ANSWERS\s*:\s*", final_text)
)
if not marker_matches:
return ParseResult(
thinking_trace=thinking_trace.strip(),
final_text=final_text.strip(),
answers=[],
valid=False,
error="missing FINAL ANSWERS: marker",
)
marker = marker_matches[-1]
reasoning_outside_think = final_text[:marker.start()].strip()
combined_trace = "\n\n".join(
part for part in (thinking_trace.strip(), reasoning_outside_think) if part
)
answer_block = final_text[marker.end():].strip()
answer_block = re.sub(r"^```(?:text)?\s*", "", answer_block, flags=re.I)
answer_block = re.sub(r"\s*```\s*$", "", answer_block)
if not answer_block:
return ParseResult(
thinking_trace=combined_trace,
final_text=final_text.strip(),
answers=[],
valid=False,
error="empty FINAL ANSWERS block",
)
# Be tolerant if the model emits a JSON list even though bare lines were asked for.
answers: list[str] = []
parsed_json_list = False
if answer_block.startswith("["):
try:
decoded = json.loads(answer_block)
if isinstance(decoded, list):
parsed_json_list = True
answers = [str(value).strip() for value in decoded if str(value).strip()]
except json.JSONDecodeError:
pass
if not answers and not parsed_json_list:
for line in answer_block.splitlines():
stripped = line.strip()
if not stripped or stripped.startswith("```"):
continue
if re.match(r"(?i)^(?:explanation|reasoning|notes?)\s*:", stripped):
break
cleaned = _clean_answer_line(stripped)
if cleaned:
answers.append(cleaned)
if not answers:
return ParseResult(
thinking_trace=combined_trace,
final_text=final_text.strip(),
answers=[],
valid=False,
error="no non-empty answers in FINAL ANSWERS block",
)
return ParseResult(
thinking_trace=combined_trace,
final_text=final_text.strip(),
answers=answers,
valid=True,
)
def _best_effort_answers(final_texts: list[str]) -> list[str]:
"""Keep submission shape valid if all strict retries fail."""
for text in reversed(final_texts):
candidates = []
for line in text.splitlines():
cleaned = _clean_answer_line(line)
if cleaned and not re.match(
r"(?i)^(?:final answers|explanation|reasoning)\s*:?'?",
cleaned,
):
candidates.append(cleaned)
if candidates:
return [candidates[-1]]
return [""]
def _version_tuple(version: str) -> tuple[int, int, int]:
numbers = [int(value) for value in re.findall(r"\d+", version)[:3]]
return tuple((numbers + [0, 0, 0])[:3]) # type: ignore[return-value]
def _keep_last_token_hidden_state(_module: Any, inputs: tuple[Any, ...]) -> tuple[Any, ...] | None:
"""Avoid materializing full-sequence vocabulary logits during generation."""
if not inputs:
return None
hidden_states = inputs[0]
if hidden_states.ndim == 3 and hidden_states.shape[1] > 1:
return (hidden_states[:, -1:, :],) + inputs[1:]
return None
def load_model_and_tokenizer() -> tuple[Any, Any, Any]:
"""Load the pre-quantized AWQ model shipped in this repository."""
try:
import torch
import transformers
from transformers import AutoModelForCausalLM, AutoTokenizer
except ImportError as exc:
raise RuntimeError(
"PyTorch, Transformers, Accelerate, AutoAWQ, and safetensors must be "
"available in the evaluation image."
) from exc
if _version_tuple(transformers.__version__) < (4, 37, 0):
raise RuntimeError(
f"Qwen2.5 requires transformers>=4.37.0; found {transformers.__version__}."
)
model_config_path = REPO_DIR / "config.json"
if not model_config_path.exists():
raise RuntimeError(
"Qwen2.5-14B-Instruct-AWQ files are missing. Put the complete model and "
"tokenizer snapshot in the same repository directory as script.py."
)
with model_config_path.open("r", encoding="utf-8") as handle:
model_metadata = json.load(handle)
quantization_method = str(
model_metadata.get("quantization_config", {}).get("quant_method", "")
).casefold()
if model_metadata.get("model_type") != "qwen2" or quantization_method != "awq":
raise RuntimeError(
"This script expects the official Qwen2.5-14B-Instruct-AWQ snapshot "
"(model_type=qwen2, quant_method=awq)."
)
tokenizer = AutoTokenizer.from_pretrained(
str(REPO_DIR), local_files_only=True, trust_remote_code=False
)
if tokenizer.pad_token_id is None:
tokenizer.pad_token_id = tokenizer.eos_token_id
model = AutoModelForCausalLM.from_pretrained(
str(REPO_DIR),
local_files_only=True,
trust_remote_code=False,
device_map="auto",
torch_dtype=torch.float16,
low_cpu_mem_usage=True,
).eval()
# Transformers 4.44.1 projects every prompt token through Qwen2's large
# vocabulary head and then upcasts all logits to FP32. Generation only uses
# the final-token logits, so trim the token dimension immediately before
# lm_head to avoid a multi-gigabyte prefill allocation on the T4.
model.lm_head.register_forward_pre_hook(_keep_last_token_hidden_state)
input_device = model.get_input_embeddings().weight.device
return model, tokenizer, input_device
def _problem_prompt(row: dict[str, Any], rag_text: str, retry_note: str = "") -> str:
work_language = str(row.get("work_lang", "")).strip()
task_language = str(row.get("task_lang", "")).strip()
language_lines = []
if work_language:
language_lines.append(f"Working language: {work_language}")
if task_language:
language_lines.append(f"Problem language: {task_language}")
language_metadata = "\n".join(language_lines)
sections = [rag_text] if rag_text else []
sections.append("CURRENT PROBLEM")
if language_metadata:
sections.append(language_metadata)
sections.extend(
[
"CONTEXT:\n" + str(row.get("context", "")),
"QUERY:\n" + str(row.get("query", "")),
]
)
if retry_note:
sections.append(retry_note)
return "\n\n".join(sections)
def _encode_fitted_prompt(
tokenizer: Any,
system_prompt: str,
row: dict[str, Any],
rag_text: str,
retry_note: str,
context_window: int,
requested_new_tokens: int,
) -> tuple[dict[str, Any], int]:
"""Trim retrieved material, never the current problem, to fit the context."""
fitted_rag = rag_text
while True:
messages = [
{"role": "system", "content": system_prompt},
{"role": "user", "content": _problem_prompt(row, fitted_rag, retry_note)},
]
rendered = tokenizer.apply_chat_template(
messages,
tokenize=False,
add_generation_prompt=True,
)
encoded = tokenizer(rendered, return_tensors="pt", add_special_tokens=False)
prompt_tokens = int(encoded["input_ids"].shape[-1])
available = context_window - prompt_tokens
if available >= requested_new_tokens:
return encoded, requested_new_tokens
if fitted_rag:
if len(fitted_rag) <= 600:
fitted_rag = ""
else:
new_length = max(0, int(len(fitted_rag) * 0.72))
fitted_rag = fitted_rag[:new_length].rsplit("\n", 1)[0].rstrip()
fitted_rag += "\n[retrieved material shortened to fit the model context]"
continue
if available >= 256:
return encoded, available
raise RuntimeError(
f"The current problem and system prompt use {prompt_tokens} tokens, "
f"leaving only {available} tokens in the {context_window}-token context."
)
def generate_once(
model: Any,
tokenizer: Any,
input_device: Any,
system_prompt: str,
row: dict[str, Any],
rag_text: str,
retry_note: str,
config: dict[str, Any],
seed: int,
) -> ParseResult:
import torch
encoded, max_new_tokens = _encode_fitted_prompt(
tokenizer=tokenizer,
system_prompt=system_prompt,
row=row,
rag_text=rag_text,
retry_note=retry_note,
context_window=int(config["model_context_tokens"]),
requested_new_tokens=int(config["max_new_tokens"]),
)
encoded = {name: value.to(input_device) for name, value in encoded.items()}
torch.manual_seed(seed)
if torch.cuda.is_available():
torch.cuda.manual_seed_all(seed)
# These values match the generation configuration shipped with Qwen2.5
# Instruct checkpoints while keeping an explicit challenge-time token cap.
with torch.inference_mode():
output = model.generate(
**encoded,
max_new_tokens=max_new_tokens,
do_sample=True,
temperature=float(config["temperature"]),
top_p=float(config["top_p"]),
top_k=int(config["top_k"]),
repetition_penalty=float(config["repetition_penalty"]),
use_cache=True,
pad_token_id=tokenizer.pad_token_id,
)
prompt_length = int(encoded["input_ids"].shape[-1])
generated_ids = output[0, prompt_length:].tolist()
response_text = tokenizer.decode(generated_ids, skip_special_tokens=True).strip()
return parse_model_response("", response_text)
def solve_row(
row: dict[str, Any],
row_index: int,
retriever: BookRetriever,
model: Any,
tokenizer: Any,
input_device: Any,
system_prompt: str,
config: dict[str, Any],
) -> ParseResult:
retrieval = retriever.retrieve(
row,
top_methods=int(config["top_methods"]),
top_examples=int(config["top_examples"]),
char_tfidf_weight=float(config["char_tfidf_weight"]),
)
rag_text = retriever.format_for_prompt(
retrieval, max_chars=int(config["rag_max_chars"])
)
failed_results: list[ParseResult] = []
last_error = ""
max_attempts = int(config["max_attempts"])
for attempt in range(max_attempts):
retry_note = ""
if attempt:
retry_note = (
"FORMAT RETRY. The previous sample was rejected by the answer parser "
f"because it had {last_error}. Generate a fresh solution. Even if "
"uncertain, end with exactly FINAL ANSWERS: followed by at least one "
"non-empty bare answer line and nothing else in that block."
)
result = generate_once(
model=model,
tokenizer=tokenizer,
input_device=input_device,
system_prompt=system_prompt,
row=row,
rag_text=rag_text,
retry_note=retry_note,
config=config,
seed=int(config["seed"]) + row_index * max_attempts + attempt,
)
if result.valid:
if attempt:
print(f" parser accepted retry {attempt + 1}/{max_attempts}", flush=True)
return result
last_error = result.error
failed_results.append(result)
print(
f" parser rejected attempt {attempt + 1}/{max_attempts}: {last_error}",
flush=True,
)
print(" WARNING: all strict retries failed; using a best-effort final line", flush=True)
last_result = failed_results[-1]
return ParseResult(
thinking_trace=last_result.thinking_trace,
final_text=last_result.final_text,
answers=_best_effort_answers(
[failed_result.final_text for failed_result in failed_results]
),
valid=False,
error=f"all {max_attempts} strict parsing attempts failed",
)
def generate_explanation(
solution: ParseResult,
model: Any,
tokenizer: Any,
input_device: Any,
config: dict[str, Any],
) -> str:
"""Summarize the model's accepted reasoning for the optional jury track."""
import torch
source_reasoning = solution.thinking_trace.strip() or solution.final_text.strip()
answer_text = "\n".join(solution.answers)
messages = [
{
"role": "system",
"content": (
"Write a short, human-readable explanation of the solution in a few "
"concise bullet points. State the linguistic rule or pattern found, "
"the key evidence, and how the final answers follow. Do not reproduce "
"the raw reasoning trace, trial and error, or meta-commentary. Do not "
"change the final answers. Output only the explanation."
),
},
{
"role": "user",
"content": (
f"MODEL REASONING:\n{source_reasoning}\n\n"
f"FINAL ANSWERS:\n{answer_text}"
),
},
]
rendered = tokenizer.apply_chat_template(
messages, tokenize=False, add_generation_prompt=True
)
encoded = tokenizer(rendered, return_tensors="pt", add_special_tokens=False)
encoded = {name: value.to(input_device) for name, value in encoded.items()}
prompt_tokens = int(encoded["input_ids"].shape[-1])
available = int(config["model_context_tokens"]) - prompt_tokens
max_new_tokens = min(int(config["explanation_max_new_tokens"]), available)
if max_new_tokens < 32:
return (
"- Inferred the relevant linguistic patterns from the supplied examples.\n"
"- Applied those patterns to produce the listed answers."
)
with torch.inference_mode():
output = model.generate(
**encoded,
max_new_tokens=max_new_tokens,
do_sample=False,
use_cache=True,
pad_token_id=tokenizer.pad_token_id,
)
generated_ids = output[0, prompt_tokens:].tolist()
explanation = tokenizer.decode(
generated_ids, skip_special_tokens=True
).strip()
if explanation:
return explanation
return (
"- Inferred the relevant linguistic patterns from the supplied examples.\n"
"- Applied those patterns to produce the listed answers."
)
def read_test_rows(path: Path) -> list[dict[str, Any]]:
with path.open("r", encoding="utf-8-sig", newline="") as handle:
return list(csv.DictReader(handle))
def _write_submission_stream(handle: Any, rows: list[dict[str, str]]) -> None:
writer = csv.DictWriter(handle, fieldnames=["id", "pred", "explanation"])
writer.writeheader()
writer.writerows(rows)
def write_submission(path: Path, rows: list[dict[str, str]]) -> None:
path.parent.mkdir(parents=True, exist_ok=True)
temporary_path = path.with_name(f".{path.name}.tmp")
with temporary_path.open("w", encoding="utf-8", newline="") as handle:
_write_submission_stream(handle, rows)
os.replace(temporary_path, path)
def run_self_test() -> None:
import numpy as np
hidden_states = np.arange(24).reshape(1, 3, 8)
trimmed = _keep_last_token_hidden_state(None, (hidden_states,))
assert trimmed is not None and trimmed[0].shape == (1, 1, 8)
assert np.array_equal(trimmed[0], hidden_states[:, -1:, :])
assert _keep_last_token_hidden_state(None, (hidden_states[:, -1:, :],)) is None
parsed = parse_model_response(
"A useful analysis.",
"One last check.\nFINAL ANSWERS:\n1. čha\n2) multi word form",
)
assert parsed.valid
assert parsed.answers == ["čha", "multi word form"]
assert "One last check." in parsed.thinking_trace
assert not parse_model_response("", "FINAL ANSWERS:\n").valid
assert not parse_model_response("", "FINAL ANSWERS:\n[]").valid
assert not parse_model_response("", 'FINAL ANSWERS:\n[""]').valid
assert not parse_model_response("", "The answer is x.").valid
retriever = BookRetriever(RESOURCE_DIR)
canary_row = {
"context": "The following words are number expressions in an unknown language.",
"query": "Determine the rule and write the number 25 in words.",
"answer": "PRIVATE_ANSWER_CANARY",
}
result = retriever.retrieve(
canary_row, top_methods=2, top_examples=2, char_tfidf_weight=3.0
)
prompt = retriever.format_for_prompt(result, max_chars=8000)
assert result["methods"] and result["examples"]
assert "PRIVATE_ANSWER_CANARY" not in prompt
assert all(example.get("string_only_usable") for example in result["examples"])
exact_example = retriever.examples[0]
exact_result = retriever.retrieve(
{"context": exact_example["context"], "query": exact_example["query"]},
top_methods=1,
top_examples=1,
char_tfidf_weight=3.0,
)
assert exact_result["examples"][0]["id"] == exact_example["id"]
submission_buffer = io.StringIO(newline="")
_write_submission_stream(
submission_buffer,
[
{
"id": "007",
"pred": json.dumps(["čha", "multi word"], ensure_ascii=False),
"explanation": "- Identified the relevant pattern.",
}
],
)
submission_buffer.seek(0)
assert list(csv.DictReader(submission_buffer)) == [
{
"id": "007",
"pred": '["čha", "multi word"]',
"explanation": "- Identified the relevant pattern.",
}
]
incremental_rows = [
{"id": "001", "pred": "[]", "explanation": ""},
{"id": "002", "pred": "[]", "explanation": ""},
]
incremental_rows[0]["pred"] = json.dumps(["answer"], ensure_ascii=False)
incremental_buffer = io.StringIO(newline="")
_write_submission_stream(incremental_buffer, incremental_rows)
incremental_buffer.seek(0)
assert list(csv.DictReader(incremental_buffer)) == incremental_rows
print(
"Self-test passed: last-token logits hook, parser, empty-answer rejection, "
"hybrid book-only retrieval, exact-match retrieval, and submission CSV "
"serialization."
)
def main() -> None:
if hasattr(sys.stdout, "reconfigure"):
sys.stdout.reconfigure(encoding="utf-8")
parser = argparse.ArgumentParser(description=__doc__)
parser.add_argument("--input", type=Path, default=DEFAULT_INPUT)
parser.add_argument("--output", type=Path, default=DEFAULT_OUTPUT)
parser.add_argument(
"--explanations",
choices=("on", "off"),
default=None,
help="Generate optional jury-track explanations (default: config setting)",
)
parser.add_argument(
"--self-test", action="store_true", help="Test parser/retrieval without loading Qwen"
)
args = parser.parse_args()
if args.self_test:
run_self_test()
return
if not args.input.exists():
raise SystemExit(f"Test file not found: {args.input}")
config = _load_config()
explanations_enabled = (
args.explanations == "on"
if args.explanations is not None
else bool(config.get("enable_explanations", False))
)
system_prompt = (RESOURCE_DIR / "system_prompt.txt").read_text(encoding="utf-8").strip()
test_rows = read_test_rows(args.input)
retriever = BookRetriever(RESOURCE_DIR)
# Do not create an all-empty submission before the model is known to load.
# Otherwise a missing/incompatible shard or startup OOM is silently scored as
# zero EM and zero chrF instead of being reported as a runtime failure.
model, tokenizer, input_device = load_model_and_tokenizer()
submission_rows: list[dict[str, str]] = [
{
"id": str(row.get("id", "")),
"pred": "[]",
"explanation": "",
}
for row in test_rows
]
write_submission(args.output, submission_rows)
print(
f"Initialized incremental submission with {len(submission_rows)} rows at "
f"{args.output}; explanations={'on' if explanations_enabled else 'off'}",
flush=True,
)
for index, row in enumerate(test_rows):
row_id = str(row.get("id", ""))
print(f"[{index + 1}/{len(test_rows)}] solving id={row_id}", flush=True)
solution = solve_row(
row=row,
row_index=index,
retriever=retriever,
model=model,
tokenizer=tokenizer,
input_device=input_device,
system_prompt=system_prompt,
config=config,
)
explanation = ""
if explanations_enabled:
explanation = generate_explanation(
solution=solution,
model=model,
tokenizer=tokenizer,
input_device=input_device,
config=config,
)
submission_rows[index] = {
"id": row_id,
"pred": json.dumps(solution.answers, ensure_ascii=False),
"explanation": explanation,
}
write_submission(args.output, submission_rows)
print(f"Checkpointed prediction {index + 1}/{len(test_rows)}", flush=True)
print(f"Wrote {len(submission_rows)} predictions to {args.output}", flush=True)
if __name__ == "__main__":
main()