Download python/inflect_ax_tts/backend.py from AXERA-TECH/inflect_micro_v2: direct link, hf CLI and curl.
- Browser
- Download file 3.62 kB
-
https://huggingface.co/AXERA-TECH/inflect_micro_v2/resolve/main/python/inflect_ax_tts/backend.py
- Command line
-
hf download hf://AXERA-TECH/inflect_micro_v2/python/inflect_ax_tts/backend.py
-
curl -L -o backend.py https://huggingface.co/AXERA-TECH/inflect_micro_v2/resolve/main/python/inflect_ax_tts/backend.py
3.62 kB
| """Inference backend wrapper: pyaxengine (default) or onnxruntime (host stand-in). | |
| Both expose an onnxruntime-like API (run / get_inputs / get_outputs). The | |
| backend is selected from the model file suffix unless overridden: | |
| *.axmodel -> axengine.InferenceSession (pyaxengine, on-device or host sim) | |
| *.onnx -> onnxruntime.InferenceSession (host numeric stand-in only) | |
| Encoder input dtype note: the compiled AXMODELs declare their token inputs as | |
| S32 at runtime (COMPILE_NOTES §3 folds the S64 input_processors into the | |
| model; SIMULATE §1 confirmed S32 binaries), while the ONNX graphs take int64. | |
| The wrapper therefore adapts the feed dtype from the session's input metadata. | |
| """ | |
| from __future__ import annotations | |
| from pathlib import Path | |
| import numpy as np | |
| class ModelSession: | |
| """Minimal name-addressable inference session over axengine/onnxruntime.""" | |
| def __init__(self, model_path: str | Path, backend: str = "auto") -> None: | |
| self.path = Path(model_path) | |
| if not self.path.is_file(): | |
| raise FileNotFoundError(self.path) | |
| if backend == "auto": | |
| backend = "axengine" if self.path.suffix == ".axmodel" else "onnxruntime" | |
| self.backend = backend | |
| if backend == "axengine": | |
| try: | |
| from axengine import InferenceSession | |
| except ImportError as exc: | |
| raise ImportError( | |
| "pyaxengine is required for .axmodel files " | |
| "(pip install axengine on an AXera host/BSP environment)" | |
| ) from exc | |
| self._sess = InferenceSession(str(self.path)) | |
| elif backend == "onnxruntime": | |
| try: | |
| import onnxruntime as ort | |
| except ImportError as exc: | |
| raise ImportError( | |
| "onnxruntime is required for .onnx files " | |
| "(host numeric stand-in only)" | |
| ) from exc | |
| self._sess = ort.InferenceSession( | |
| str(self.path), providers=["CPUExecutionProvider"] | |
| ) | |
| else: | |
| raise ValueError(f"unknown backend {backend!r}") | |
| # -- metadata ----------------------------------------------------------- | |
| def _inputs(self): | |
| return self._sess.get_inputs() | |
| def output_names(self) -> list[str]: | |
| return [o.name for o in self._sess.get_outputs()] | |
| def input_dtype(self, name: str) -> np.dtype: | |
| """Best-effort numpy dtype of a named input. | |
| onnxruntime exposes `.type` strings ('tensor(int64)'); pyaxengine | |
| exposes numpy-like `.dtype`. Fallback: int32 for .axmodel (see module | |
| docstring), int64 for .onnx. | |
| """ | |
| for meta in self._inputs(): | |
| if meta.name != name: | |
| continue | |
| dt = getattr(meta, "dtype", None) | |
| if dt is not None: | |
| try: | |
| return np.dtype(dt) | |
| except TypeError: | |
| pass | |
| type_str = str(getattr(meta, "type", "") or "") | |
| if "int32" in type_str: | |
| return np.dtype(np.int32) | |
| if "int64" in type_str: | |
| return np.dtype(np.int64) | |
| if "float" in type_str: | |
| return np.dtype(np.float32) | |
| return np.dtype(np.int32 if self.path.suffix == ".axmodel" else np.int64) | |
| # -- inference ---------------------------------------------------------- | |
| def run(self, feeds: dict[str, np.ndarray]) -> dict[str, np.ndarray]: | |
| results = self._sess.run(None, feeds) | |
| return dict(zip(self.output_names(), results)) | |