File size: 7,475 Bytes
8129b09
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
22b973c
 
8129b09
 
 
 
 
22b973c
 
8129b09
 
22b973c
 
 
8129b09
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1ec7af8
8129b09
 
 
1ec7af8
 
 
 
 
 
 
 
 
 
 
8129b09
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
#!/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")