""" 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()