103 lines
4.2 KiB
Python
103 lines
4.2 KiB
Python
"""Custom Quizbowl pipelines for QANTA 2026 Arena.
|
|
Registers `quizbowl-tossup` and `quizbowl-bonus`. The build script rewrites the
|
|
two flags below per model (LLM vs VLM). Every name passed to register_pipeline
|
|
is imported from transformers, per the submission requirements.
|
|
"""
|
|
import json, re
|
|
from transformers.pipelines import PIPELINE_REGISTRY
|
|
from transformers import (
|
|
Pipeline,
|
|
AutoModelForCausalLM, AutoModelForImageTextToText, AutoProcessor,
|
|
)
|
|
|
|
PT_MODEL = AutoModelForCausalLM # build script -> AutoModelForImageTextToText for VLMs
|
|
IS_VLM = False # build script -> True for VLMs
|
|
|
|
|
|
def _to_pil(images):
|
|
from PIL import Image
|
|
out = []
|
|
for im in images or []:
|
|
out.append(Image.open(im).convert("RGB") if isinstance(im, str) else im)
|
|
return out
|
|
|
|
|
|
def _extract(text):
|
|
text = re.sub(r"<think>.*?</think>", "", text, flags=re.S) # strip reasoning traces
|
|
m = re.search(r"\{.*\}", text, re.S)
|
|
if m:
|
|
try:
|
|
return json.loads(m.group())
|
|
except Exception:
|
|
pass
|
|
return {}
|
|
|
|
|
|
class _Base(Pipeline):
|
|
def _sanitize_parameters(self, **kw):
|
|
return {}, {}, {}
|
|
|
|
def preprocess(self, inputs):
|
|
return inputs
|
|
|
|
def _forward(self, inputs):
|
|
return {"text": self._gen(self._prompt(inputs), inputs.get("images"))}
|
|
|
|
def _proc(self):
|
|
if not hasattr(self, "_p"):
|
|
self._p = AutoProcessor.from_pretrained(self.model.name_or_path, trust_remote_code=True)
|
|
return self._p
|
|
|
|
def _gen(self, prompt, images=None):
|
|
if IS_VLM and images:
|
|
proc = self._proc()
|
|
msgs = [{"role": "user", "content": [{"type": "image"} for _ in images] +
|
|
[{"type": "text", "text": prompt}]}]
|
|
chat = proc.apply_chat_template(msgs, add_generation_prompt=True)
|
|
inp = proc(text=chat, images=_to_pil(images), return_tensors="pt").to(self.model.device)
|
|
ids = self.model.generate(**inp, max_new_tokens=200, do_sample=False)
|
|
return proc.batch_decode(ids[:, inp["input_ids"].shape[1]:], skip_special_tokens=True)[0]
|
|
tok = self.tokenizer
|
|
if getattr(tok, "chat_template", None):
|
|
text = tok.apply_chat_template([{"role": "user", "content": prompt}],
|
|
tokenize=False, add_generation_prompt=True)
|
|
else:
|
|
text = prompt
|
|
inp = tok(text, return_tensors="pt").to(self.model.device)
|
|
ids = self.model.generate(**inp, max_new_tokens=200, do_sample=False,
|
|
pad_token_id=tok.eos_token_id)
|
|
return tok.decode(ids[0, inp["input_ids"].shape[1]:], skip_special_tokens=True)
|
|
|
|
|
|
class QuizbowlTossupPipeline(_Base):
|
|
def _prompt(self, inputs):
|
|
return ("You are a quizbowl expert. From the (possibly partial) question, give your best guess.\n"
|
|
'Respond ONLY as JSON: {"answer": "<concise answer>", "confidence": <0-1 float>, '
|
|
'"buzz": <true|false>}. Set buzz=true only if confident enough to interrupt.\n'
|
|
"QUESTION: " + inputs["question_text"])
|
|
|
|
def postprocess(self, mo):
|
|
d = _extract(mo["text"])
|
|
c = min(max(float(d.get("confidence", 0.5) or 0.5), 0.0), 1.0)
|
|
return {"answer": str(d.get("answer", "")).strip(),
|
|
"confidence": c,
|
|
"buzz": bool(d.get("buzz", c >= 0.7))}
|
|
|
|
|
|
class QuizbowlBonusPipeline(_Base):
|
|
def _prompt(self, inputs):
|
|
return ("You are a quizbowl expert answering one bonus part.\n"
|
|
'Respond ONLY as JSON: {"answer": "<concise>", "confidence": <0-1 float>, '
|
|
'"explanation": "<= 30 words"}.\n'
|
|
"LEADIN: " + inputs.get("leadin", "") + "\nPART: " + inputs["part"])
|
|
|
|
def postprocess(self, mo):
|
|
d = _extract(mo["text"])
|
|
c = min(max(float(d.get("confidence", 0.5) or 0.5), 0.0), 1.0)
|
|
exp = " ".join(str(d.get("explanation", "")).split()[:30])
|
|
return {"answer": str(d.get("answer", "")).strip(), "confidence": c, "explanation": exp}
|
|
|
|
|
|
PIPELINE_REGISTRY.register_pipeline("quizbowl-tossup", pipeline_class=QuizbowlTossupPipeline, pt_model=PT_MODEL)
|
|
PIPELINE_REGISTRY.register_pipeline("quizbowl-bonus", pipeline_class=QuizbowlBonusPipeline, pt_model=PT_MODEL)
|