vlm-twin-spec-decoding / code /stage1_accept.py
LeoMaxwell's picture
add code, data, results, patches, report
ee3e28a verified
Raw History Blame Contribute Delete
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()