File size: 2,773 Bytes
a330cfa | 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 | """
Minimal ONNX inference sample for the exported Smartwatch LM.
Requirements:
pip install numpy onnxruntime tokenizers
Usage:
python onnx_sample.py "How many steps today?"
"""
from __future__ import annotations
import argparse
import sys
from pathlib import Path
import numpy as np
import onnxruntime as ort
from tokenizers import Tokenizer
from reply_utils import build_prompt, process_model_output
MODEL_DIR = Path(__file__).parent
ONNX_PATH = MODEL_DIR / "smartwatch_lm_merged.onnx"
TOKENIZER_PATH = MODEL_DIR / "tokenizer.json"
BLOCK_SIZE = 256
MAX_NEW_TOKENS = 40
TEMPERATURE = 0.5
TOP_K = 40
EOS_TOKEN_ID = 0
def sample_next_token(logits: np.ndarray, temperature: float, top_k: int) -> int:
scaled = logits / max(temperature, 1e-8)
scaled = scaled - scaled.max()
probs = np.exp(scaled)
probs /= probs.sum()
k = min(top_k, probs.size)
top_idx = np.argpartition(probs, -k)[-k:]
top_probs = probs[top_idx]
top_probs /= top_probs.sum()
return int(np.random.choice(top_idx, p=top_probs))
def generate(session: ort.InferenceSession, tokenizer: Tokenizer, prompt: str) -> str:
ids = tokenizer.encode(prompt).ids
prompt_len = len(ids)
for step in range(MAX_NEW_TOKENS):
seq = ids[-BLOCK_SIZE:]
x = np.array([seq], dtype=np.int64)
logits = session.run(None, {"input_ids": x})[0]
next_logits = logits[0, -1, :]
next_id = sample_next_token(next_logits, TEMPERATURE, TOP_K)
ids.append(next_id)
if step > 2 and next_id == EOS_TOKEN_ID:
break
return tokenizer.decode(ids[prompt_len:])
def main() -> None:
parser = argparse.ArgumentParser(description="ONNX chat sample")
parser.add_argument("message", nargs="?", default="How many steps today?")
args = parser.parse_args()
if not ONNX_PATH.is_file():
print(f"Missing {ONNX_PATH.name}", file=sys.stderr)
sys.exit(1)
if not TOKENIZER_PATH.is_file():
print(f"Missing {TOKENIZER_PATH.name}", file=sys.stderr)
sys.exit(1)
session = ort.InferenceSession(str(ONNX_PATH), providers=["CPUExecutionProvider"])
tokenizer = Tokenizer.from_file(str(TOKENIZER_PATH))
prompt = build_prompt([], args.message)
continuation = generate(session, tokenizer, prompt)
slot_data = {
"STEPS_TODAY": "4,231",
"STEP_GOAL": "10,000",
"STEPS_REMAINING": "5,769",
}
raw, parsed, display = process_model_output(prompt, continuation, slot_data)
print(f"user> {args.message}")
print(f"raw> {raw}")
print(f"intent> {parsed.intent}")
print(f"bot> {display}")
if __name__ == "__main__":
main()
|