onw / check_ref.py
ryugyosoft's picture
onw 0.2: renamed from npue; onw command; Qwen3.5 (dense) / Gemma 4 / vision support; LM head segment
49c3379 verified
Raw History Blame Contribute Delete
3.02 kB
"""onw engine vs the original HF model (PyTorch bf16, CPU): first-token logits and greedy tokens.
usage: python check_ref.py HF_DIR ENGINE_DIR [DEVICE=CPU] [n_tokens=30]"""
import os, sys, time
import numpy as np
import openvino as ov
import torch
from transformers import AutoModelForCausalLM, AutoTokenizer
from onw.runtime import SegmentedModel
PROMPT = "日本の首都について2文で説明してください。"
def main():
hf, eng = sys.argv[1], sys.argv[2]
dev = sys.argv[3] if len(sys.argv) > 3 else "CPU"
n = int(sys.argv[4]) if len(sys.argv) > 4 else 30
tok = AutoTokenizer.from_pretrained(hf)
kw = {"enable_thinking": False} if "enable_thinking" in (tok.chat_template or "") else {}
ids = tok.apply_chat_template([{"role": "user", "content": PROMPT}], add_generation_prompt=True, return_tensors="pt", **kw)
ids = ids["input_ids"] if hasattr(ids, "keys") else ids
cache = os.path.join(eng, f"ref_{n}.npz") # the reference is slow for big models: keep it
if os.path.exists(cache):
z = np.load(cache)
ref_lg, ref_out = z["lg"], z["out"].tolist()
else:
try:
ref = AutoModelForCausalLM.from_pretrained(hf, dtype=torch.bfloat16)
except Exception: # multimodal checkpoints (e.g. Qwen3.6): text-only use
from transformers import AutoModelForImageTextToText
ref = AutoModelForImageTextToText.from_pretrained(hf, dtype=torch.bfloat16)
with torch.no_grad():
ref_lg = ref(ids).logits[0, -1].float().numpy()
ref_out = ref.generate(ids, max_new_tokens=n, do_sample=False)[0, ids.shape[1]:].tolist()
del ref
np.savez(cache, lg=ref_lg, out=np.array(ref_out))
core = ov.Core()
cfg = {"INFERENCE_PRECISION_HINT": "f32"} if dev == "CPU" else {"CACHE_DIR": eng + "/npu_cache"}
m = SegmentedModel(core, eng, dev, cfg)
t0 = time.time()
toks, s0, lg = ids[0].tolist(), 0, None
while s0 < len(toks):
k = min(m.S, len(toks) - s0)
lg = m.step(toks[s0:s0 + k])[-1].astype(np.float32)
s0 += k
tp = time.time() - t0
first, out = lg, []
t0 = time.time()
while len(out) < n:
t = int(lg.argmax())
out.append(t)
if t == tok.eos_token_id:
break
lg = m.step([t])[-1].astype(np.float32)
td = time.time() - t0
k = next((i for i in range(min(len(out), len(ref_out))) if out[i] != ref_out[i]), None)
print(f"prompt {ids.shape[1]} tokens: prefill {tp*1000:.0f} ms, decode {len(out)/max(td, 1e-9):.1f} tok/s")
print(f"first-token logits rel err {np.linalg.norm(first - ref_lg) / np.linalg.norm(ref_lg):.4f}, "
f"top1 ref {ref_lg.argmax()} engine {first.argmax()}; greedy identical: {out == ref_out[:len(out)]}"
+ ("" if k is None else f" (first difference at {k})"))
print("ref :", repr(tok.decode(ref_out)))
print("engine:", repr(tok.decode(out)))
if __name__ == "__main__":
main()