Download onnx_sample.py from prathamkode/smartwatch-lm-0.2: direct link, hf CLI and curl.
- Browser
- Download file 2.77 kB
-
https://huggingface.co/prathamkode/smartwatch-lm-0.2/resolve/main/onnx_sample.py
- Command line
-
hf download hf://prathamkode/smartwatch-lm-0.2/onnx_sample.py
-
curl -L -o onnx_sample.py https://huggingface.co/prathamkode/smartwatch-lm-0.2/resolve/main/onnx_sample.py
2.77 kB
| """ | |
| 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() | |