laya-mobile / examples /laya_sequence.py
charioteer's picture
Core ML fp16 and ONNX conversions of convaiinnovations/laya@55cf4c4e
d8b837f verified
Raw History Blame Contribute Delete
3.65 kB
"""Build the Laya input sequence for one `choice` question with only the `tokenizers` package.
This follows `laya.common.build_sequence` (laya 0.3.22) for one question:
[CLS] "choice question: <instructions>" [SEP]
[MASK] " <label>: <description>" ... (one block per option, 48 tokens at most each)
[SEP] <state> [SEP]
The whole sequence is at most 512 tokens. The question part (instructions and
options) gets at most 192 tokens. The state fills the rest and is cut.
"""
from __future__ import annotations
import json
from typing import TYPE_CHECKING
if TYPE_CHECKING: # only for the type hints, so tests/check.py runs without tokenizers
from tokenizers import Tokenizer
CLS, SEP, PAD, MASK = 50281, 50282, 50283, 50284
MAX_LEN, HEAD_MAX_LEN, OPTION_MAX = 512, 192, 48
def encode(tok: Tokenizer, text: str) -> list[int]:
return tok.encode(text.replace("[MASK]", " "), add_special_tokens=False).ids
def build(tok: Tokenizer, instructions: str, options: dict[str, str], state) -> tuple[list[int], list[int]]:
"""Token ids and the position of each option's [MASK] marker."""
head = encode(tok, f"choice question: {instructions}")
opts = [[MASK] + encode(tok, f" {label}: {text}" if text else f" {label}")[:OPTION_MAX]
for label, text in options.items()]
budget = HEAD_MAX_LEN - sum(len(o) for o in opts)
if budget < 16: # many options: every option gets the same share of the question budget
per = max(4, (HEAD_MAX_LEN - 16) // len(opts))
opts = [o[:per] for o in opts]
budget = HEAD_MAX_LEN - sum(len(o) for o in opts)
ids = [CLS] + head[:max(8, budget)] + [SEP]
markers = []
for o in opts:
markers.append(len(ids))
ids += o
ids.append(SEP)
room = max(0, MAX_LEN - len(ids) - 1)
text = state if isinstance(state, str) else json.dumps(state, ensure_ascii=False)
state_ids = encode(tok, text)
# A list state is a conversation, newest last: keep its end. Any other state keeps its start.
state_ids = state_ids[max(0, len(state_ids) - room):] if isinstance(state, list) else state_ids[:room]
ids = (ids + state_ids + [SEP])[:MAX_LEN]
return ids, [m for m in markers if m < MAX_LEN]
def temperature(config_path, k: int) -> float:
"""The calibrated temperature of `rl_agent_config.json` for a choice question with k options.
laya clamps every temperature to [0.5, 5.0]. Only the `choice:11+` value (0.1006) is outside.
"""
with open(config_path, encoding="utf-8") as f:
config = json.load(f)
size = "2" if k <= 2 else "3-5" if k <= 5 else "6-10" if k <= 10 else "11+"
t = config["temperature_by_options"].get(f"choice:{size}", config["temperature"][0])
return min(max(t, 0.5), 5.0)
def probabilities(logits, t: float) -> list[float]:
"""softmax(logits / t), computed in float64."""
import numpy as np
z = np.asarray(logits, np.float64) / t
p = np.exp(z - z.max())
return (p / p.sum()).tolist()
# The example question of this repository: a customer-support router with five teams.
# tests/vectors.json holds what the `laya` package returns for it.
INSTRUCTIONS = "Which team should handle this message?"
OPTIONS = {
"billing": "charges, invoices, refunds, payment methods, and subscription prices",
"technical": "bugs, error messages, crashes, setup, and how to use a feature",
"account": "sign-in problems, passwords, profile details, and closing an account",
"shipping": "delivery status, tracking numbers, addresses, and lost or damaged parcels",
"other": "anything else, such as feedback, partnerships, or press requests",
}