File size: 3,189 Bytes
f722fb4
819e690
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
"""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)