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