vlm-twin-spec-decoding / code /proto_compare.py
LeoMaxwell's picture
add code, data, results, patches, report
ee3e28a verified
Raw History Blame Contribute Delete
1.44 kB
import json, statistics as st
def load(p):
d={}
for line in open(p):
r=json.loads(line); d[r["pid"]]=r
return d
def analyze(tag, spec_f, van1_f, van2_f):
sp, v1, v2 = load(spec_f), load(van1_f), load(van2_f)
pids = sorted(set(sp)&set(v1)&set(v2))
n=len(pids)
sp_m = st.mean([sp[p]["tok_s"] for p in pids])
v1_m = st.mean([v1[p]["tok_s"] for p in pids])
v2_m = st.mean([v2[p]["tok_s"] for p in pids])
ratios_paired = sorted(sp[p]["tok_s"]/((v1[p]["tok_s"]+v2[p]["tok_s"])/2) for p in pids)
med = st.median(ratios_paired)
q1 = ratios_paired[n//4]; q3 = ratios_paired[(3*n)//4]
a = sp_m/v1_m
b = st.mean([sp[p]["tok_s"]/v1[p]["tok_s"] for p in pids])
c = (sum(sp[p]["gen_len"] for p in pids)/sum(sp[p]["t"] for p in pids)) / (sum(v1[p]["gen_len"] for p in pids)/sum(v1[p]["t"] for p in pids))
d = sp_m/min(v1_m, v2_m)
print(f"{tag}: n={n} ours_median={med:.3f} [IQR {q1:.3f}-{q3:.3f}] | litA_meanratio={a:.3f} litB_meanofratios={b:.3f} litC_totaltok={c:.3f} litD_vs_slower_van={d:.3f}")
print(f" abs tok/s: spec {sp_m:.1f} van1 {v1_m:.1f} van2 {v2_m:.1f}")
B="/testessfs10/users/zeyu.zhang/wangyu_ssd/stage1/"
for g in (4,6,8):
analyze(f"strict g{g}", B+f"nb_spec_g{g}.jsonl", B+f"nb_van1_g{g}.jsonl", B+f"nb_van2_g{g}.jsonl")
for g in (6,8):
analyze(f"relax g{g}", B+f"rxb_spec_g{g}.jsonl", B+f"rxb_van1_g{g}.jsonl", B+f"rxb_van2_g{g}.jsonl")