File size: 4,903 Bytes
ee3e28a
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
"""Stage 1 gate: does bnb-NF4 acceptance transfer to the AWQ draft?
Teacher-force stored vanilla trajectories (res05_*_thinking.jsonl) through the AWQ
engine: one request per sample with prompt_token_ids = prompt + gen_ids and
prompt_logprobs=1; accept[i] = stored token has rank 1 in the AWQ distribution at
its position. Same sim_speedup as stage0_analyze. Gate: overall accept >= 0.92.
"""
import argparse, io, json
from PIL import Image
import pyarrow.parquet as pq


def sim_rounds(bits, gamma):
    p, rounds, T = 0, 0, len(bits)
    while p < T:
        run = 0
        while run < gamma and p + run < T and bits[p + run] == "1":
            run += 1
        p += run + 1
        rounds += 1
    return rounds


def main():
    ap = argparse.ArgumentParser()
    ap.add_argument("--model", required=True, help="AWQ checkpoint")
    ap.add_argument("--tok-model", default="", help="processor source (default: --model)")
    ap.add_argument("--parquet", required=True)
    ap.add_argument("--in-jsonl", required=True)
    ap.add_argument("--out", required=True)
    ap.add_argument("--n", type=int, default=0, help="0 = all")
    ap.add_argument("--gpu-mem-util", type=float, default=0.30)
    ap.add_argument("--max-model-len", type=int, default=6144)
    ap.add_argument("--img-max-side", type=int, default=1024)
    args = ap.parse_args()

    data = [json.loads(l) for l in open(args.in_jsonl) if "pid" in l]
    if args.n:
        data = data[: args.n]
    tbl = pq.read_table(args.parquet, columns=["pid", "query", "decoded_image"])
    rowmap = {r["pid"]: r for r in tbl.to_pylist()}

    from stage0_gonogo import build_inputs
    from transformers import AutoProcessor
    proc = AutoProcessor.from_pretrained(args.tok_model or args.model)

    from vllm import LLM, SamplingParams
    llm = LLM(model=args.model, gpu_memory_utilization=args.gpu_mem_util,
              max_model_len=args.max_model_len, enable_prefix_caching=False,
              disable_log_stats=True, limit_mm_per_prompt={"image": 1})
    sp = SamplingParams(temperature=0, max_tokens=1, prompt_logprobs=1, detokenize=False)

    done = set()
    try:
        done = {json.loads(l)["pid"] for l in open(args.out) if "pid" in l}
    except FileNotFoundError:
        pass
    fout = open(args.out, "a")
    for idx, rec in enumerate(data):
        pid = rec["pid"]
        if pid in done:
            continue
        try:
            row = rowmap[pid]
            img = Image.open(io.BytesIO(row["decoded_image"]["bytes"])).convert("RGB")
            if max(img.size) > args.img_max_side:
                sc = args.img_max_side / max(img.size)
                img = img.resize((int(img.width * sc), int(img.height * sc)))
            enc = build_inputs(proc, row["query"], img)
            pids_ = enc["input_ids"][0].tolist()
            if len(pids_) != rec["in_len"]:
                raise RuntimeError(f"prompt rebuild mismatch {len(pids_)} vs {rec['in_len']}")
            gen = [int(t) for t in rec["gen_ids"]]
            full = pids_ + gen
            # vLLM re-expands the image region, so the engine-internal prompt is
            # longer than ours; gen tokens are the tail -- index from the end.
            out = llm.generate(
                [{"prompt_token_ids": full, "multi_modal_data": {"image": img},
                  "multi_modal_uuids": {"image": [f"acc-{pid}"]}}],
                sp, use_tqdm=False)[0]
            plp = out.prompt_logprobs
            ptk = list(out.prompt_token_ids)
            assert plp is not None and ptk[-len(gen):] == gen, \
                f"tail misalign plp={len(plp) if plp else None} ptk={len(ptk)}"
            bits = []
            for i in range(len(gen)):
                lp = plp[-(len(gen) - i)][gen[i]]
                bits.append("1" if lp.rank == 1 else "0")
            bits = "".join(bits)
            fout.write(json.dumps(dict(pid=pid, gen_len=len(gen), accept=bits)) + "\n")
            fout.flush()
            print(f"[{idx+1}/{len(data)}] {pid} gen={len(gen)} "
                  f"acc={bits.count('1')/len(bits):.3f}", flush=True)
        except Exception as e:
            import traceback
            print(f"[err] {pid}: {e}\n{traceback.format_exc()}", flush=True)
    fout.close()

    recs = [json.loads(l) for l in open(args.out) if "pid" in l]
    T = sum(r["gen_len"] for r in recs)
    acc = sum(r["accept"].count("1") for r in recs) / T
    print(f"[agg] n={len(recs)} tok={T} AWQ_accept={acc:.4f} (bnb was 0.929) "
          f"gate={'PASS' if acc >= 0.92 else 'FAIL'}", flush=True)
    for gamma in (4, 6, 8):
        rounds = sum(sim_rounds(r["accept"], gamma) for r in recs)
        for cost in (0.30, 0.37, 0.42):
            print(f"[sim] gamma={gamma} c={cost:.2f} strict_speedup="
                  f"{T / (rounds * (gamma * cost + 1.0)):.3f}", flush=True)
    print("[done]", flush=True)


if __name__ == "__main__":
    main()