Download code/stage1_accept.py from LeoMaxwell/vlm-twin-spec-decoding: direct link, hf CLI and curl.
- Browser
- Download file 4.9 kB
-
https://huggingface.co/LeoMaxwell/vlm-twin-spec-decoding/resolve/main/code/stage1_accept.py
- Command line
-
hf download hf://LeoMaxwell/vlm-twin-spec-decoding/code/stage1_accept.py
-
curl -L -o stage1_accept.py https://huggingface.co/LeoMaxwell/vlm-twin-spec-decoding/resolve/main/code/stage1_accept.py
4.9 kB
| """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() | |