初始化项目,由ModelHub XC社区提供模型
Model: nttruong1007/qb-granite31-8b Source: Original Platform
This commit is contained in:
102
pipeline.py
Normal file
102
pipeline.py
Normal file
@@ -0,0 +1,102 @@
|
||||
"""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)
|
||||
Reference in New Issue
Block a user