rnnoise / python /rnnoise_sdk /inference_npu.py
inoryQwQ's picture
merge AX650 + AX620E(SenseVoice 风格双芯布局)
f722fb4 verified
Raw
History Blame Contribute Delete
3.19 kB
"""RNNoise AX620E 推理会话(NPU 专用发布版:无 onnxruntime/torch 回退)。"""
import numpy as np
from . import dsp
DEFAULT_PROVIDER = "AxEngineExecutionProvider"
INPUT_NAMES = ["features", "conv1_mem", "conv2_mem",
"gru1_s", "gru2_s", "gru3_s"]
INPUT_SHAPES = {"features": (1, 65), "conv1_mem": (1, 130),
"conv2_mem": (1, 256), "gru1_s": (1, 384),
"gru2_s": (1, 384), "gru3_s": (1, 384)}
OUTPUT_NAMES = ["gains", "vad", "conv1_mem_new", "conv2_mem_new",
"gru1_s_new", "gru2_s_new", "gru3_s_new"]
_STATE_INPUTS = ["conv1_mem", "conv2_mem", "gru1_s", "gru2_s", "gru3_s"]
class RNNoiseDenoiser:
"""48k 单声道实时降噪器(AX 芯片端到端,无 CPU 回退)。"""
def __init__(self, model_path, providers=None):
try:
import axengine as axe
except ImportError as exc:
raise RuntimeError(
"SDK 为 NPU 专用发布版,仅支持在 AX 芯片上运行;请先安装 "
"requirements.txt 并在板端执行(无 onnxruntime/torch 回退)"
) from exc
self.session = axe.InferenceSession(
model_path, providers=providers or [DEFAULT_PROVIDER])
self.backend = "axengine"
self.input_names = [i.name for i in self.session.get_inputs()]
self.output_names = [o.name for o in self.session.get_outputs()]
self.reset()
def reset(self):
self.st = dsp.RNNoiseState()
self._states = {k: np.zeros(INPUT_SHAPES[k], dtype=np.float32)
for k in _STATE_INPUTS}
def process_frame(self, frame):
frame = np.asarray(frame, dtype=np.float32).reshape(-1)
if frame.size != dsp.FRAME_SIZE:
raise ValueError(f"帧长必须为 {dsp.FRAME_SIZE},实际 {frame.size}")
ana = dsp.analyze_frame(self.st, frame)
if ana["silence"]:
out, _ = dsp.synthesize_frame(self.st, ana, None, 0.0)
return out, 0.0
feeds = {
"features": np.ascontiguousarray(
ana["features"][None, :].astype(np.float32)),
}
feeds.update({k: np.ascontiguousarray(v)
for k, v in self._states.items()})
outs = self.session.run(None, feeds)
out_idx = {n: i for i, n in enumerate(self.output_names)}
gains = outs[out_idx["gains"]]
vad = outs[out_idx["vad"]]
for k in _STATE_INPUTS:
self._states[k] = np.asarray(
outs[out_idx[k + "_new"]], dtype=np.float32)
out, vad = dsp.synthesize_frame(self.st, ana, gains, vad)
return out, float(vad)
def process(self, pcm):
pcm = np.asarray(pcm, dtype=np.float32)
if pcm.ndim == 1:
n = pcm.size // dsp.FRAME_SIZE
frames = pcm[:n * dsp.FRAME_SIZE].reshape(n, dsp.FRAME_SIZE)
else:
frames = pcm.reshape(-1, dsp.FRAME_SIZE)
outs, vads = [], []
for fr in frames:
o, v = self.process_frame(fr)
outs.append(o)
vads.append(v)
return np.concatenate(outs), np.asarray(vads, dtype=np.float32)