IOL-AI-v3 / probe_pipeline.py
ShayanShamsi's picture
Upload probe_pipeline.py with huggingface_hub
22b973c verified
Raw
History Blame Contribute Delete
7.48 kB
#!/usr/bin/env python3
"""Solving pipeline used by the probe submission (vLLM or transformers stack).
Strategy: budgeted thinking pass -> JSON answers; incremental writes;
stage-0 fast pass first on the transformers path so nothing is ever blank.
"""
import csv
import json
import time
from iol_common import direct_prompt, infer_labels, parse_items, strip_think
def _rows_to_probs(rows):
probs = []
for r in rows:
labels = infer_labels(r.get("context", ""), r.get("query", ""))
probs.append({
"id": r["id"], "context": r.get("context", ""),
"query": r.get("query", ""), "task_type": (r.get("task_type") or "").strip(),
"labels": labels, "n_items": len(labels),
})
return probs
def _write(out_csv, results, order):
with open(out_csv, "w", newline="", encoding="utf-8") as f:
w = csv.DictWriter(f, fieldnames=["id", "pred", "explanation"])
w.writeheader()
for pid in order:
r = results[pid]
w.writerow({"id": pid, "pred": json.dumps(r["pred"], ensure_ascii=False),
"explanation": r.get("explanation", "")})
def _explanation(sample_text):
tail = strip_think(sample_text).strip()
return (tail[:500] + "...") if len(tail) > 500 else tail
def solve_all_vllm(llm, rows, out_csv, start, deadline_s, log):
from vllm import SamplingParams
probs = _rows_to_probs(rows)
order = [p["id"] for p in probs]
results = {p["id"]: {"pred": [""] * p["n_items"], "explanation": ""} for p in probs}
_write(out_csv, results, order)
tok = llm.get_tokenizer()
remaining = deadline_s - (time.time() - start)
# scale the thinking budget to the time left and the row count, assuming
# a conservative ~90 tok/s aggregate on the T4
est_tokens = max(remaining - 120, 60) * 90
think_budget = int(min(3072, max(768, est_tokens / max(len(probs), 1))))
log(f"vllm solve: {len(probs)} rows, think_budget={think_budget}, remaining={remaining:.0f}s")
def render(p, thinking):
try:
return tok.apply_chat_template(
[{"role": "user", "content": p}], tokenize=False,
add_generation_prompt=True, enable_thinking=thinking)
except TypeError:
return tok.apply_chat_template(
[{"role": "user", "content": p}], tokenize=False, add_generation_prompt=True)
prompts = [render(direct_prompt(p["context"], p["query"], p["task_type"], p["labels"]), True)
for p in probs]
sp = SamplingParams(temperature=0.0, max_tokens=think_budget + 512)
t = time.time()
outs = llm.generate(prompts, sp)
texts = [o.outputs[0].text for o in outs]
log(f"vllm gen done in {time.time()-t:.0f}s, "
f"{sum(len(o.outputs[0].token_ids) for o in outs)} tokens out")
# budget forcing: close unfinished thinking and squeeze the answer out
todo = [i for i, tx in enumerate(texts) if "</think>" not in tx]
if todo and deadline_s - (time.time() - start) > 90:
log(f"budget-forcing {len(todo)} unclosed samples")
cont_prompts = [prompts[i] + texts[i] +
"\n\nOkay, time is up — I must answer now.\n</think>\n\n"
for i in todo]
cont = llm.generate(cont_prompts, SamplingParams(temperature=0.0, max_tokens=512))
for i, o in zip(todo, cont):
texts[i] = texts[i] + "\n</think>\n\n" + o.outputs[0].text
for p, text in zip(probs, texts):
results[p["id"]] = {"pred": parse_items(text, p["labels"]),
"explanation": _explanation(text)}
_write(out_csv, results, order)
log("vllm solve: submission written")
def solve_all_transformers(model_dir, rows, out_csv, start, deadline_s, log):
import torch
from transformers import AutoModelForCausalLM, AutoTokenizer
probs = _rows_to_probs(rows)
order = [p["id"] for p in probs]
results = {p["id"]: {"pred": [""] * p["n_items"], "explanation": ""} for p in probs}
for pid in order[:54]:
results[pid]["explanation"] = "[diag] tf solver started, model loading"
tok = AutoTokenizer.from_pretrained(model_dir)
if tok.pad_token_id is None:
tok.pad_token = tok.eos_token
tok.padding_side = "left"
model = AutoModelForCausalLM.from_pretrained(
model_dir, torch_dtype=torch.float16).eval()
model.to("cuda" if torch.cuda.is_available() else "cpu")
log(f"transformers model loaded from {model_dir}")
for pid in order[:63]:
results[pid]["explanation"] = "[diag] tf model loaded, generating"
_write(out_csv, results, order)
def render(p, thinking):
try:
return tok.apply_chat_template(
[{"role": "user", "content": p}], tokenize=False,
add_generation_prompt=True, enable_thinking=thinking)
except TypeError:
return tok.apply_chat_template(
[{"role": "user", "content": p}], tokenize=False, add_generation_prompt=True)
def gen_batch(prompt_texts, max_new, batch=4):
res = []
for s in range(0, len(prompt_texts), batch):
chunk = prompt_texts[s:s + batch]
enc = tok(chunk, return_tensors="pt", padding=True, truncation=True,
max_length=7000).to(model.device)
with torch.no_grad():
out = model.generate(**enc, max_new_tokens=max_new, do_sample=False,
pad_token_id=tok.pad_token_id)
res.extend(tok.batch_decode(out[:, enc["input_ids"].shape[1]:],
skip_special_tokens=True))
return res
# Stage 0: fast no-think pass, incremental writes so a timeout keeps partials
prompts0 = [render(direct_prompt(p["context"], p["query"], p["task_type"], p["labels"]),
False) for p in probs]
t = time.time()
B = 4
for s in range(0, len(prompts0), B):
if time.time() - start > deadline_s - 60:
log("tf stage0: deadline, stopping")
break
outs0 = gen_batch(prompts0[s:s + B], 400, batch=B)
for p, text in zip(probs[s:s + B], outs0):
results[p["id"]] = {"pred": parse_items(text, p["labels"]),
"explanation": _explanation(text)}
_write(out_csv, results, order)
log(f"tf stage0 {min(s+B, len(probs))}/{len(probs)}")
log(f"tf stage0 done in {time.time()-t:.0f}s")
# Stage 1: thinking pass per problem while time remains
for p in probs:
if time.time() - start > deadline_s - 90:
log("tf stage1: deadline, stopping")
break
try:
text = gen_batch([render(direct_prompt(p["context"], p["query"], p["task_type"],
p["labels"]), True)], 1800, batch=1)[0]
pred = parse_items(text, p["labels"])
if any(x.strip() for x in pred):
old = results[p["id"]]["pred"]
merged = [a if a.strip() else (old[i] if i < len(old) else "")
for i, a in enumerate(pred)]
results[p["id"]] = {"pred": merged, "explanation": _explanation(text)}
_write(out_csv, results, order)
except Exception as e:
log(f"tf stage1 {p['id']}: {e!r}")
_write(out_csv, results, order)
log("tf solve: submission written")