| |
| """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) |
| |
| |
| 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") |
|
|
| |
| 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 |
|
|
| |
| 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") |
|
|
| |
| 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") |
|
|