Python runtime 对齐发布规范:NPU 默认、依赖补全、前 5 帧置零
#5
by inoryQwQ - opened
- README.md +4 -0
- requirements.txt +2 -0
- scripts/openwakeword_ax.py +13 -1
- scripts/runtime.py +2 -1
README.md
CHANGED
|
@@ -53,6 +53,10 @@ pip install axengine-x.x.x-py3-none-any.whl
|
|
| 53 |
pip install -r requirements.txt
|
| 54 |
```
|
| 55 |
|
|
|
|
|
|
|
|
|
|
|
|
|
| 56 |
AX650:
|
| 57 |
|
| 58 |
```bash
|
|
|
|
| 53 |
pip install -r requirements.txt
|
| 54 |
```
|
| 55 |
|
| 56 |
+
`requirements.txt` 已包含 axengine(pyaxengine)wheel 依赖,直接
|
| 57 |
+
`pip install -r requirements.txt` 即可。`--backend onnx`(CPU 回退)仅用于
|
| 58 |
+
开发/验证,发布运行请使用默认 `axengine`(NPU)。
|
| 59 |
+
|
| 60 |
AX650:
|
| 61 |
|
| 62 |
```bash
|
requirements.txt
CHANGED
|
@@ -1 +1,3 @@
|
|
| 1 |
numpy>=1.24
|
|
|
|
|
|
|
|
|
| 1 |
numpy>=1.24
|
| 2 |
+
# AX 芯片 NPU 推理引擎(pyaxengine;GitHub Releases 包名为 axengine,板端必需)
|
| 3 |
+
axengine @ https://github.com/AXERA-TECH/pyaxengine/releases/download/0.1.3.rc3/axengine-0.1.3-py3-none-any.whl
|
scripts/openwakeword_ax.py
CHANGED
|
@@ -162,6 +162,7 @@ def infer_clip(
|
|
| 162 |
mel_buffer = np.ones((76, 32), dtype=np.float32)
|
| 163 |
feature_buffer = np.zeros((34, 96), dtype=np.float32)
|
| 164 |
scores: dict[str, list[list[float]]] = {name: [] for name in CLASSIFIERS}
|
|
|
|
| 165 |
|
| 166 |
for start in range(0, samples.size, 1280):
|
| 167 |
chunk = samples[start : start + 1280]
|
|
@@ -206,6 +207,11 @@ def infer_clip(
|
|
| 206 |
classifier.run({classifier.inputs[0].name: classifier_input}),
|
| 207 |
)
|
| 208 |
scores[name].append(np.asarray(output).reshape(-1).tolist())
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 209 |
|
| 210 |
max_scores = {
|
| 211 |
name: np.max(np.asarray(values, dtype=np.float64), axis=0).tolist()
|
|
@@ -260,7 +266,13 @@ def main(
|
|
| 260 |
target_hardware: str = "AX650",
|
| 261 |
) -> None:
|
| 262 |
parser = argparse.ArgumentParser()
|
| 263 |
-
parser.add_argument(
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 264 |
parser.add_argument("--mel-backend", choices=("numpy", "model"), default="numpy")
|
| 265 |
parser.add_argument("--mel-weights", type=Path, default=DEFAULT_MEL_WEIGHTS)
|
| 266 |
parser.add_argument("--models-dir", type=Path, default=default_models_dir)
|
|
|
|
| 162 |
mel_buffer = np.ones((76, 32), dtype=np.float32)
|
| 163 |
feature_buffer = np.zeros((34, 96), dtype=np.float32)
|
| 164 |
scores: dict[str, list[list[float]]] = {name: [] for name in CLASSIFIERS}
|
| 165 |
+
frame_index = 0
|
| 166 |
|
| 167 |
for start in range(0, samples.size, 1280):
|
| 168 |
chunk = samples[start : start + 1280]
|
|
|
|
| 207 |
classifier.run({classifier.inputs[0].name: classifier_input}),
|
| 208 |
)
|
| 209 |
scores[name].append(np.asarray(output).reshape(-1).tolist())
|
| 210 |
+
# 对齐 openwakeword 官方行为:前 5 帧为初始化窗口,得分置零
|
| 211 |
+
if frame_index < 5:
|
| 212 |
+
for name in CLASSIFIERS:
|
| 213 |
+
scores[name][-1] = [0.0] * len(scores[name][-1])
|
| 214 |
+
frame_index += 1
|
| 215 |
|
| 216 |
max_scores = {
|
| 217 |
name: np.max(np.asarray(values, dtype=np.float64), axis=0).tolist()
|
|
|
|
| 266 |
target_hardware: str = "AX650",
|
| 267 |
) -> None:
|
| 268 |
parser = argparse.ArgumentParser()
|
| 269 |
+
parser.add_argument(
|
| 270 |
+
"--backend",
|
| 271 |
+
choices=("axengine", "onnx"),
|
| 272 |
+
default="axengine",
|
| 273 |
+
help="推理后端:axengine=NPU(发布默认,仅板端可用);"
|
| 274 |
+
"onnx=CPU 仅用于开发/验证,不属于发布运行路径",
|
| 275 |
+
)
|
| 276 |
parser.add_argument("--mel-backend", choices=("numpy", "model"), default="numpy")
|
| 277 |
parser.add_argument("--mel-weights", type=Path, default=DEFAULT_MEL_WEIGHTS)
|
| 278 |
parser.add_argument("--models-dir", type=Path, default=default_models_dir)
|
scripts/runtime.py
CHANGED
|
@@ -65,7 +65,8 @@ class InferenceSession:
|
|
| 65 |
raise RuntimeError(
|
| 66 |
"axengine is unavailable; run this backend on an AXERA board"
|
| 67 |
) from error
|
| 68 |
-
self._session = axengine.InferenceSession(
|
|
|
|
| 69 |
elif backend == "onnx":
|
| 70 |
try:
|
| 71 |
import onnxruntime as ort
|
|
|
|
| 65 |
raise RuntimeError(
|
| 66 |
"axengine is unavailable; run this backend on an AXERA board"
|
| 67 |
) from error
|
| 68 |
+
self._session = axengine.InferenceSession(
|
| 69 |
+
str(self.path), providers=["AxEngineExecutionProvider"])
|
| 70 |
elif backend == "onnx":
|
| 71 |
try:
|
| 72 |
import onnxruntime as ort
|