Text Classification
Transformers
English
multilingual
laya
typed-decisions
non-autoregressive
axera
ax650
Instructions to use AXERA-TECH/Laya with libraries, inference providers, notebooks, and local apps. Follow these links to get started.
- Libraries
- Transformers
How to use AXERA-TECH/Laya with Transformers:
# Use a pipeline as a high-level helper from transformers import pipeline pipe = pipeline("text-classification", model="AXERA-TECH/Laya")# Load model directly from transformers import AutoModel model = AutoModel.from_pretrained("AXERA-TECH/Laya", device_map="auto") - Notebooks
- Google Colab
- Kaggle
Download python/ax650/infer.py from AXERA-TECH/Laya: direct link, hf CLI and curl.
- Browser
- Download file 12.6 kB
-
https://huggingface.co/AXERA-TECH/Laya/resolve/main/python/ax650/infer.py
- Command line
-
hf download hf://AXERA-TECH/Laya/python/ax650/infer.py
-
curl -L -o infer.py https://huggingface.co/AXERA-TECH/Laya/resolve/main/python/ax650/infer.py
12.6 kB
| #!/usr/bin/env python3 | |
| """Run a packaged Laya AXModel through PyAXEngine on an AX650 board.""" | |
| import argparse | |
| import json | |
| import sys | |
| import time | |
| from pathlib import Path | |
| from typing import Any, Dict, Iterable, List, Tuple | |
| import numpy as np | |
| from transformers import AutoTokenizer | |
| QTYPES = {"choice": 0, "score": 1, "noul": 2} | |
| QTYPE_NAMES = {value: key for key, value in QTYPES.items()} | |
| def softmax(values: np.ndarray) -> np.ndarray: | |
| values = values.astype(np.float64) | |
| values -= values.max() | |
| result = np.exp(values) | |
| return result / result.sum() | |
| def confidence_from_probs(probabilities: np.ndarray, count: int) -> float: | |
| if count < 2: | |
| return 1.0 | |
| values = probabilities[:count] | |
| entropy = -(values * np.log(np.clip(values, 1e-12, 1.0))).sum() | |
| return float(np.clip(1.0 - entropy / np.log(count), 0.0, 1.0)) | |
| def temp_bucket(qtype: int, count: int) -> str: | |
| size = "2" if count <= 2 else "3-5" if count <= 5 else "6-10" if count <= 10 else "11+" | |
| return f"{QTYPE_NAMES[int(qtype)]}:{size}" | |
| def render(value: Any) -> str: | |
| if isinstance(value, str): | |
| return value | |
| return json.dumps(value, ensure_ascii=False, separators=(", ", ": "), default=str) | |
| def internal_question(question: Dict[str, Any]) -> Dict[str, Any]: | |
| qtype = question["type"] | |
| criteria = question.get("criteria") | |
| if qtype == "choice" and isinstance(criteria, list): | |
| criteria = {item: None for item in criteria} | |
| instructions = question["instructions"] | |
| if not isinstance(instructions, str): | |
| instructions = json.dumps(instructions) | |
| return {"t": qtype, "ins": instructions, "crit": criteria} | |
| def render_options(question: Dict[str, Any]) -> List[str]: | |
| qtype = question["t"] | |
| criteria = question.get("crit") | |
| if qtype == "choice": | |
| if not isinstance(criteria, dict): | |
| raise ValueError("choice.criteria must be an object or list") | |
| return [ | |
| key if value is None or value == "" else f"{key}: {render(value)}" | |
| for key, value in criteria.items() | |
| ] | |
| if qtype == "score": | |
| if not isinstance(criteria, list): | |
| raise ValueError("score.criteria must be a list") | |
| return [f"level {index}: {render(value)}" for index, value in enumerate(criteria)] | |
| criteria = criteria or {} | |
| false_value = criteria.get("false") | |
| true_value = criteria.get("true") | |
| return [ | |
| "false: " + (render(false_value) if false_value not in (None, "") else "no, the statement does not hold"), | |
| "true: " + (render(true_value) if true_value not in (None, "") else "yes, the statement holds"), | |
| ] | |
| def build_sequence( | |
| tokenizer, | |
| state: Any, | |
| question: Dict[str, Any], | |
| max_len: int, | |
| head_max_len: int, | |
| ) -> Tuple[List[int], List[int]]: | |
| mask_token = tokenizer.mask_token | |
| options = render_options(question) | |
| instructions = str(question["ins"]).replace(mask_token, " ") | |
| head_ids = tokenizer( | |
| f"{question['t']} question: {instructions}", | |
| add_special_tokens=False, | |
| )["input_ids"] | |
| option_ids = [ | |
| [tokenizer.mask_token_id] | |
| + tokenizer( | |
| " " + option.replace(mask_token, " "), | |
| add_special_tokens=False, | |
| )["input_ids"][:48] | |
| for option in options | |
| ] | |
| budget = head_max_len - sum(len(option) for option in option_ids) | |
| if budget < 16: | |
| per_option = max(4, (head_max_len - 16) // max(1, len(option_ids))) | |
| option_ids = [option[:per_option] for option in option_ids] | |
| budget = head_max_len - sum(len(option) for option in option_ids) | |
| head_ids = head_ids[: max(8, budget)] | |
| ids = [tokenizer.cls_token_id] + head_ids + [tokenizer.sep_token_id] | |
| markers = [] | |
| for option in option_ids: | |
| markers.append(len(ids)) | |
| ids.extend(option) | |
| ids.append(tokenizer.sep_token_id) | |
| room = max(0, max_len - len(ids) - 1) | |
| state_text = state if isinstance(state, str) else json.dumps(state, ensure_ascii=False) | |
| state_ids = tokenizer( | |
| state_text.replace(mask_token, " "), | |
| add_special_tokens=False, | |
| )["input_ids"][:room] | |
| ids = ids + state_ids + [tokenizer.sep_token_id] | |
| return ids[:max_len], [marker for marker in markers if marker < max_len] | |
| def encode_question( | |
| tokenizer, | |
| state: Any, | |
| question: Dict[str, Any], | |
| seq_len: int, | |
| num_options: int, | |
| head_max_len: int, | |
| ) -> Dict[str, np.ndarray]: | |
| internal = internal_question(question) | |
| options = render_options(internal) | |
| if not 2 <= len(options) <= num_options: | |
| raise ValueError( | |
| f"question has {len(options)} options; this graph supports 2..{num_options}" | |
| ) | |
| ids, markers = build_sequence(tokenizer, state, internal, seq_len, head_max_len) | |
| if len(markers) != len(options): | |
| raise ValueError("an option marker was truncated; shorten the question") | |
| pad = seq_len - len(ids) | |
| return { | |
| "input_ids": np.asarray( | |
| [ids + [tokenizer.pad_token_id] * pad], dtype=np.int32 | |
| ), | |
| "attention_mask": np.asarray( | |
| [[1] * len(ids) + [0] * pad], dtype=np.int32 | |
| ), | |
| "marker_pos": np.asarray( | |
| [markers + [0] * (num_options - len(markers))], dtype=np.int32 | |
| ), | |
| "marker_mask": np.asarray( | |
| [[1] * len(markers) + [0] * (num_options - len(markers))], | |
| dtype=np.int32, | |
| ), | |
| "qtype": np.asarray([QTYPES[internal["t"]]], dtype=np.int32), | |
| } | |
| def validate_request(request: Any) -> Dict[str, Any]: | |
| if not isinstance(request, dict): | |
| raise ValueError("request must be a JSON object") | |
| if "state" not in request: | |
| raise ValueError("request is missing required field: state") | |
| questions = request.get("questions") | |
| if not isinstance(questions, dict) or not questions: | |
| raise ValueError("request.questions must be a non-empty object") | |
| for question_id, question in questions.items(): | |
| if not isinstance(question, dict): | |
| raise ValueError(f"question {question_id!r} must be an object") | |
| if question.get("type") not in QTYPES: | |
| raise ValueError( | |
| f"question {question_id!r} has unsupported type {question.get('type')!r}" | |
| ) | |
| if "instructions" not in question: | |
| raise ValueError(f"question {question_id!r} is missing instructions") | |
| return request | |
| def format_answer( | |
| question: Dict[str, Any], | |
| probabilities: np.ndarray, | |
| confidence: float, | |
| act_probability: float, | |
| latency_ms: float, | |
| ) -> Dict[str, Any]: | |
| common = { | |
| "type": question["type"], | |
| "confidence": float(confidence), | |
| "action": {"act_probability": float(act_probability)}, | |
| "python_latency_ms": float(latency_ms), | |
| } | |
| if question["type"] == "choice": | |
| labels = list(question["criteria"]) | |
| return { | |
| "type": "choice", | |
| "choice": labels[int(probabilities.argmax())], | |
| "probabilities": { | |
| label: float(value) for label, value in zip(labels, probabilities) | |
| }, | |
| **{key: common[key] for key in ("confidence", "action", "python_latency_ms")}, | |
| } | |
| if question["type"] == "score": | |
| score = float(np.dot(np.arange(len(probabilities)), probabilities)) | |
| return { | |
| "type": "score", | |
| "probabilities": { | |
| str(index): float(value) | |
| for index, value in enumerate(probabilities) | |
| }, | |
| "legend": { | |
| str(index): value | |
| for index, value in enumerate(question["criteria"]) | |
| }, | |
| "score": score, | |
| **{key: common[key] for key in ("confidence", "action", "python_latency_ms")}, | |
| } | |
| return { | |
| "type": "noul", | |
| "noul": float(probabilities[1]), | |
| **{key: common[key] for key in ("confidence", "action", "python_latency_ms")}, | |
| } | |
| class LayaAx650: | |
| def __init__(self, model_dir: Path): | |
| try: | |
| from axengine import InferenceSession, get_available_providers | |
| except ImportError as exc: | |
| raise RuntimeError( | |
| "PyAXEngine is missing. Install an axengine wheel from " | |
| "https://github.com/AXERA-TECH/pyaxengine/releases" | |
| ) from exc | |
| self.model_dir = model_dir.resolve() | |
| config_path = self.model_dir / "config.json" | |
| if not config_path.is_file(): | |
| raise FileNotFoundError(config_path) | |
| self.config = json.loads(config_path.read_text(encoding="utf-8")) | |
| self.seq_len = int(self.config.get("sequence_length", 256)) | |
| self.num_options = int(self.config.get("num_options", 4)) | |
| self.head_max_len = int(self.config.get("head_max_len", 128)) | |
| self.temperatures = self.config.get("temperature", [1.0, 1.0, 1.0]) | |
| self.temperatures_by_options = self.config.get( | |
| "temperature_by_options", {} | |
| ) | |
| tokenizer_dir = self.model_dir / self.config.get( | |
| "tokenizer_dir", "tokenizer" | |
| ) | |
| self.tokenizer = AutoTokenizer.from_pretrained(tokenizer_dir) | |
| model_path = self.model_dir / self.config.get( | |
| "filename_axmodel", "model.axmodel" | |
| ) | |
| if not model_path.is_file(): | |
| raise FileNotFoundError(model_path) | |
| provider = "AxEngineExecutionProvider" | |
| providers = get_available_providers() | |
| if provider not in providers: | |
| raise RuntimeError(f"{provider} is unavailable; found {providers}") | |
| self.session = InferenceSession(str(model_path), providers=[provider]) | |
| def predict(self, request: Dict[str, Any]) -> Dict[str, Any]: | |
| request = validate_request(request) | |
| answers = {} | |
| total_latency_ms = 0.0 | |
| for name, question in request["questions"].items(): | |
| inputs = encode_question( | |
| self.tokenizer, | |
| request["state"], | |
| question, | |
| self.seq_len, | |
| self.num_options, | |
| self.head_max_len, | |
| ) | |
| started = time.perf_counter() | |
| logits, act_logits = self.session.run(None, inputs) | |
| latency_ms = (time.perf_counter() - started) * 1000.0 | |
| total_latency_ms += latency_ms | |
| count = int(inputs["marker_mask"].sum()) | |
| qtype = QTYPES[question["type"]] | |
| scale = self.temperatures_by_options.get( | |
| temp_bucket(qtype, count), self.temperatures[qtype] | |
| ) | |
| probabilities = softmax( | |
| logits[0, :count] / max(1e-3, float(scale)) | |
| ) | |
| act_probability = softmax(act_logits[0])[0] | |
| confidence = ( | |
| max(float(probabilities[1]), 1.0 - float(probabilities[1])) | |
| if question["type"] == "noul" | |
| else confidence_from_probs(probabilities, count) | |
| ) | |
| answers[name] = format_answer( | |
| question, | |
| probabilities, | |
| confidence, | |
| float(act_probability), | |
| latency_ms, | |
| ) | |
| return { | |
| "model": self.config.get("model_name", self.model_dir.name), | |
| "backend": "pyaxengine", | |
| "answers": answers, | |
| "total_python_latency_ms": total_latency_ms, | |
| } | |
| def parse_args() -> argparse.Namespace: | |
| parser = argparse.ArgumentParser(description=__doc__) | |
| parser.add_argument( | |
| "model_dir", | |
| type=Path, | |
| help="Packaged checkpoint directory, such as multilingual.", | |
| ) | |
| parser.add_argument( | |
| "--input", | |
| type=Path, | |
| help="Request JSON file. Omit for resident JSON Lines mode on stdin.", | |
| ) | |
| return parser.parse_args() | |
| def read_requests(input_path: Path) -> Iterable[Tuple[Dict[str, Any], bool]]: | |
| if input_path is not None: | |
| yield json.loads(input_path.read_text(encoding="utf-8")), True | |
| return | |
| for line in sys.stdin: | |
| line = line.strip() | |
| if not line: | |
| continue | |
| if line == "/exit": | |
| return | |
| yield json.loads(line), False | |
| def main() -> None: | |
| args = parse_args() | |
| runner = LayaAx650(args.model_dir) | |
| for request, pretty in read_requests(args.input): | |
| result = runner.predict(request) | |
| print( | |
| json.dumps(result, indent=2 if pretty else None, ensure_ascii=False), | |
| flush=True, | |
| ) | |
| if __name__ == "__main__": | |
| main() | |