#!/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 "" 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\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\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")