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/pytorch/infer.py from AXERA-TECH/Laya: direct link, hf CLI and curl.
- Browser
- Download file 5.45 kB
-
https://huggingface.co/AXERA-TECH/Laya/resolve/main/python/pytorch/infer.py
- Command line
-
hf download hf://AXERA-TECH/Laya/python/pytorch/infer.py
-
curl -L -o infer.py https://huggingface.co/AXERA-TECH/Laya/resolve/main/python/pytorch/infer.py
5.45 kB
| #!/usr/bin/env python3 | |
| """Run Laya with the original PyTorch checkpoint using the AX650 request schema.""" | |
| import argparse | |
| import json | |
| import os | |
| import sys | |
| import time | |
| from pathlib import Path | |
| from typing import Any, Dict, Optional, Tuple | |
| # Avoid importing TensorFlow through Transformers. Some TensorFlow installations can | |
| # delay or deadlock Laya model construction, and TensorFlow is not used here. | |
| os.environ.setdefault("USE_TF", "0") | |
| MODEL_SPECS: Dict[str, Tuple[str, Optional[str]]] = { | |
| "english": ("convaiinnovations/laya", None), | |
| "multilingual": ("convaiinnovations/laya", "multilingual"), | |
| "typed-decisions": ("convaiinnovations/laya", "typed-decisions"), | |
| } | |
| MODEL_NAMES = { | |
| "english": "laya-english", | |
| "multilingual": "laya-multilingual", | |
| "typed-decisions": "laya-typed-decisions", | |
| } | |
| def parse_args() -> argparse.Namespace: | |
| parser = argparse.ArgumentParser( | |
| description=( | |
| "Run an original Laya PyTorch checkpoint with the same state/questions " | |
| "JSON schema used by the packaged AX650 runtime." | |
| ) | |
| ) | |
| parser.add_argument( | |
| "--variant", | |
| required=True, | |
| choices=tuple(MODEL_SPECS), | |
| help="Checkpoint to load.", | |
| ) | |
| parser.add_argument( | |
| "--input", | |
| type=Path, | |
| help="Request JSON file. Omit for resident JSON Lines mode on stdin.", | |
| ) | |
| parser.add_argument( | |
| "--device", | |
| default="auto", | |
| help="PyTorch device: auto, cpu, cuda, cuda:0, or mps (default: auto).", | |
| ) | |
| return parser.parse_args() | |
| def load_agent(variant: str, device: str): | |
| try: | |
| import laya | |
| except ImportError as exc: | |
| raise SystemExit( | |
| "The Python dependencies are missing. Run: " | |
| "python -m pip install -r python/requirements.txt" | |
| ) from exc | |
| repo, subfolder = MODEL_SPECS[variant] | |
| selected_device = None if device == "auto" else device | |
| agent = laya.load( | |
| repo, | |
| subfolder=subfolder, | |
| device=selected_device, | |
| ) | |
| # Match the fixed sequence and option budgets used by the packaged AXModels. | |
| # The original upstream checkpoints support larger contexts. | |
| agent.cfg["max_len"] = 256 | |
| agent.cfg["head_max_len"] = 128 | |
| return agent | |
| 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") | |
| question_type = question.get("type") | |
| if question_type not in {"choice", "score", "noul"}: | |
| raise ValueError( | |
| f"question {question_id!r} has unsupported type {question_type!r}" | |
| ) | |
| if question_type in {"choice", "score"}: | |
| criteria = question.get("criteria") | |
| if not isinstance(criteria, (dict, list)): | |
| raise ValueError( | |
| f"question {question_id!r}.criteria must be an object or list" | |
| ) | |
| if not 2 <= len(criteria) <= 4: | |
| raise ValueError( | |
| f"question {question_id!r} must contain 2 to 4 criteria" | |
| ) | |
| return request | |
| def predict(agent, variant: str, request: Dict[str, Any]) -> Dict[str, Any]: | |
| request = validate_request(request) | |
| started = time.perf_counter() | |
| result = agent.predict(request["state"], request["questions"]) | |
| latency_ms = (time.perf_counter() - started) * 1000.0 | |
| # Keep the primary result fields aligned with `axllm run`. Python adds its | |
| # backend and wall-clock timing under `python_runtime`. | |
| result["model"] = MODEL_NAMES[variant] | |
| result["python_runtime"] = { | |
| "backend": "pytorch", | |
| "device": str(agent.device), | |
| "latency_ms": round(latency_ms, 3), | |
| "sequence_length": 256, | |
| "max_options": 4, | |
| } | |
| return result | |
| def run_file(agent, variant: str, input_path: Path) -> None: | |
| request = json.loads(input_path.read_text(encoding="utf-8")) | |
| result = predict(agent, variant, request) | |
| print(json.dumps(result, indent=2, ensure_ascii=False)) | |
| def run_json_lines(agent, variant: str) -> None: | |
| for line_number, line in enumerate(sys.stdin, start=1): | |
| line = line.strip() | |
| if not line: | |
| continue | |
| if line == "/exit": | |
| return | |
| try: | |
| request = json.loads(line) | |
| result = predict(agent, variant, request) | |
| print(json.dumps(result, ensure_ascii=False), flush=True) | |
| except Exception as exc: # Keep the resident process available after a bad request. | |
| error = { | |
| "error": str(exc), | |
| "line": line_number, | |
| } | |
| print(json.dumps(error, ensure_ascii=False), flush=True) | |
| def main() -> None: | |
| args = parse_args() | |
| agent = load_agent(args.variant, args.device) | |
| if args.input is not None: | |
| run_file(agent, args.variant, args.input) | |
| else: | |
| run_json_lines(agent, args.variant) | |
| if __name__ == "__main__": | |
| main() | |