HY-2012 commited on
Commit
26f00f8
·
verified ·
1 Parent(s): 3d54d8c

DeepFilterNet3 AXERA 量化部署包: axmodels (GRU 状态携带) + Python SDK + C++ 可执行文件

Browse files
.gitattributes CHANGED
@@ -34,3 +34,6 @@ saved_model/**/* filter=lfs diff=lfs merge=lfs -text
34
  *.zip filter=lfs diff=lfs merge=lfs -text
35
  *.zst filter=lfs diff=lfs merge=lfs -text
36
  *tfevents* filter=lfs diff=lfs merge=lfs -text
 
 
 
 
34
  *.zip filter=lfs diff=lfs merge=lfs -text
35
  *.zst filter=lfs diff=lfs merge=lfs -text
36
  *tfevents* filter=lfs diff=lfs merge=lfs -text
37
+ assets/enhanced_sample.wav filter=lfs diff=lfs merge=lfs -text
38
+ assets/noisy_snr0.wav filter=lfs diff=lfs merge=lfs -text
39
+ bin/df3_ax filter=lfs diff=lfs merge=lfs -text
.gitignore ADDED
@@ -0,0 +1,8 @@
 
 
 
 
 
 
 
 
 
1
+ # HF 仓库不需要的产物
2
+ __pycache__/
3
+ *.pyc
4
+ outputs/
5
+ *.log
6
+ .msc
7
+ .mv
8
+ ._____temp
README.md CHANGED
@@ -1,3 +1,101 @@
1
  ---
 
 
 
2
  license: mit
 
 
 
 
 
 
 
3
  ---
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
  ---
2
+ language:
3
+ - zh
4
+ - en
5
  license: mit
6
+ pipeline_tag: audio-to-audio
7
+ tags:
8
+ - speech-enhancement
9
+ - noise-reduction
10
+ - axera
11
+ - ax650
12
+ - deepfilternet
13
  ---
14
+
15
+ # DeepFilterNet3.AXERA
16
+
17
+ DeepFilterNet3 (48kHz 全频带实时语音降噪) 的 AXERA 平台量化部署包 (当前支持 AX650N)。
18
+
19
+ - 模型: enc (U16) / erb_dec (U16) / df_dec (全 FP32) 3 个子图, GRU 状态跨块携带版
20
+ - 依赖: 仅 numpy + axengine (Python); C++ 为预编译 aarch64 可执行文件
21
+
22
+ ## 目录
23
+
24
+ ```
25
+ ├── axmodels/ # enc.axmodel / erb_dec.axmodel / df_dec.axmodel
26
+ ├── inference.py # Python 推理入口 (依赖仅 numpy + axengine)
27
+ ├── requirements.txt
28
+ ├── assets/ # 演示音频 (带噪输入 + 增强输出样例)
29
+ ├── bin/ # C++ 可执行文件 (aarch64)
30
+ └── run.sh # 一键运行
31
+ ```
32
+
33
+ ## Python 运行
34
+
35
+ ```bash
36
+ pip3 install numpy
37
+ # pyaxengine: https://github.com/AXERA-TECH/pyaxengine/releases/latest
38
+ pip3 install axengine-<version>-py3-none-any.whl
39
+ ./run.sh input.wav output.wav
40
+ ```
41
+
42
+ 或:
43
+
44
+ ```python
45
+ from inference import enhance_audio
46
+ out = enhance_audio(audio_float32_48k, model_dir="axmodels")
47
+ ```
48
+
49
+ ## C++ 运行
50
+
51
+ ```bash
52
+ LD_LIBRARY_PATH=<ax_runtime>/lib ./bin/df3_ax \
53
+ axmodels/enc.axmodel axmodels/erb_dec.axmodel axmodels/df_dec.axmodel \
54
+ input.wav output.wav
55
+ ```
56
+
57
+ ## 模型说明
58
+
59
+ - 输入: 48kHz 单声道 (PCM16 wav 或 float32 数组)
60
+ - 3 子图流水线: STFT(fft 960/hop 480) -> ERB+unit_norm 特征 -> enc (GRU, 状态携带)
61
+ -> erb_dec (m 掩码) + df_dec (coefs) -> ERB 增益 + DF 滤波(5 阶, 2 帧前瞻) -> ISTFT
62
+ - 静态 T=99 帧/块, GRU 状态跨块传递 (实时流式语义, 与官方 tract 部署一致)
63
+ - 完整转换源码见 GitHub: [DeepFilterNet3.AXERA](https://github.com/ZY-2012/DeepFilterNet3.AXERA)
64
+
65
+ ## RTF(AX650N 实测, noisy_snr0.wav 10.6s)
66
+
67
+ | 推理路径 | 耗时 | RTF | 峰值内存 |
68
+ |------|------|------|------|
69
+ | C++ | 1.50 s | 0.142 | — |
70
+ | Python | 1.44 s | 0.136 | 28.9 MB |
71
+
72
+ > RTF = 推理耗时 / 音频时长(不含模型加载,RTF < 1.0 即可实时)
73
+
74
+ ## 示例结果(AX650N 实测)
75
+
76
+ | 音频 | 输入 RMS | 输出 RMS | 效果 |
77
+ |------|---------|---------|------|
78
+ | `assets/noisy_snr0.wav`(SNR 0dB 带噪语音) | 0.0640 | 0.0489 | 噪声移除、语音保留 |
79
+ | `assets/enhanced_sample.wav`(上者降噪输出) | — | — | 试听对比用 |
80
+
81
+ > 其他实测(不在本仓):干净语音 RMS 0.0287→0.0270(几乎无损伤);
82
+ > 纯噪声 0.0336→0.0021(约 -24dB 压制)
83
+
84
+ ## 精度(AX650N 实测)
85
+
86
+ | 指标 | 数值 |
87
+ |------|------|
88
+ | 逐张量 cosine vs ONNX(enc/erb_dec) | ≥ 0.9946(U16 量化) |
89
+ | df_dec 输出 coefs | 0.99999(全 FP32) |
90
+ | 板端端到端 vs 官方 torch 参考 | cosine 0.9857(量化损失仅 0.0005) |
91
+
92
+ ## 参考
93
+
94
+ - [DeepFilterNet](https://github.com/Rikorose/DeepFilterNet)
95
+ - [DeepFilterNet3.AXERA](https://github.com/ZY-2012/DeepFilterNet3.AXERA)(完整转换源码)
96
+ - [Magnetar](https://github.com/AXERA-TECH/Magnetar) — AXERA 模型部署 agent 工具
97
+
98
+ ## 技术讨论
99
+
100
+ - GitHub issues
101
+ - QQ 群: 139953715
assets/enhanced_sample.wav ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:87cdbffb148bfac4f52def56a19d0f3b7bd90afac9fddf4963f4e5cbe6c249f9
3
+ size 1017226
assets/noisy_snr0.wav ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:e1e08601f3b7ceb2f36d45c86343e38cd8927a73d7ad5526d6c4687c33aa7186
3
+ size 1017226
axmodels/df_dec.axmodel ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:3dd647f3df37c3442d933ae4aef33a24b8c5d5d072aed782c772a4d747bbe5e8
3
+ size 4193578
axmodels/enc.axmodel ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:87a1a741bffa304511adb0143e54876adda2e2286e702cee435fd62d2f613cd9
3
+ size 1252300
axmodels/erb_dec.axmodel ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:d64fcca474a08263e4f21ff59bbc3e6281413ad359eb2cdaa99e6b3d6575834e
3
+ size 2333248
bin/df3_ax ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:3464f3e384dad5b8b38e9531183fd499b29b3e78625e75d5e640e6816699e963
3
+ size 262360
configuration.json ADDED
@@ -0,0 +1,20 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "model_name": "DeepFilterNet3",
3
+ "supported_chips": ["AX650N"],
4
+ "quantization": {
5
+ "enc": "U16 (MinMax + SmoothQuant 0.5)",
6
+ "erb_dec": "U16 (MinMax + SmoothQuant 0.5)",
7
+ "df_dec": "FP32 (输出层精度保底)"
8
+ },
9
+ "sample_rate": 48000,
10
+ "fft_size": 960,
11
+ "hop_size": 480,
12
+ "chunk_frames": 99,
13
+ "gru_stateful": true,
14
+ "pulsar2_version": "7.0",
15
+ "axmodel_files": [
16
+ "axmodels/enc.axmodel",
17
+ "axmodels/erb_dec.axmodel",
18
+ "axmodels/df_dec.axmodel"
19
+ ]
20
+ }
deepfilternet3_ax/__init__.py ADDED
@@ -0,0 +1,14 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """DeepFilterNet3 NPU 推理 SDK (Python, 仅依赖 numpy + axengine)。
2
+
3
+ 用法:
4
+ from deepfilternet3_ax import DeepFilterNet3
5
+ enh = DeepFilterNet3("models/")
6
+ out = enh.enhance(audio_float32_48k) # (N,) -> (N,)
7
+
8
+ 命令行:
9
+ python -m deepfilternet3_ax input.wav -o output.wav [--model-dir models/]
10
+ """
11
+ from .enhance import DeepFilterNet3, enhance_audio
12
+
13
+ __version__ = "1.0.0"
14
+ __all__ = ["DeepFilterNet3", "enhance_audio"]
deepfilternet3_ax/__main__.py ADDED
@@ -0,0 +1,58 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """命令行降噪: python -m deepfilternet3_ax input.wav -o output.wav
2
+
3
+ wav 读写用标准库 wave (PCM16), 核心依赖仅 numpy + axengine。
4
+ """
5
+ import argparse
6
+ import time
7
+ import wave
8
+ from pathlib import Path
9
+
10
+ import numpy as np
11
+
12
+
13
+ def read_wav(path: Path):
14
+ with wave.open(str(path), "rb") as w:
15
+ assert w.getnchannels() == 1, "仅支持单声道"
16
+ sr = w.getframerate()
17
+ data = np.frombuffer(w.readframes(w.getnframes()), dtype=np.int16)
18
+ return (data.astype(np.float32) / 32768.0), sr
19
+
20
+
21
+ def write_wav(path: Path, audio: np.ndarray, sr: int):
22
+ pcm = np.clip(audio, -1.0, 1.0)
23
+ pcm = (pcm * 32767.0).astype(np.int16)
24
+ with wave.open(str(path), "wb") as w:
25
+ w.setnchannels(1)
26
+ w.setsampwidth(2)
27
+ w.setframerate(sr)
28
+ w.writeframes(pcm.tobytes())
29
+
30
+
31
+ def main():
32
+ parser = argparse.ArgumentParser(description="DeepFilterNet3 降噪 (AXERA 平台)")
33
+ parser.add_argument("input", type=Path, help="输入 48kHz 单声道 wav")
34
+ parser.add_argument("-o", "--output", type=Path, default=None,
35
+ help="输出 wav (默认 <输入名>_enhanced.wav)")
36
+ parser.add_argument("--model-dir", type=Path, default=Path("models"),
37
+ help="axmodel 目录 (默认 ./models)")
38
+ args = parser.parse_args()
39
+
40
+ from .enhance import DeepFilterNet3, SR
41
+
42
+ audio, sr = read_wav(args.input)
43
+ if sr != SR:
44
+ raise SystemExit(f"输入采样率 {sr}Hz 不支持, 请转成 48kHz 单声道")
45
+
46
+ enh = DeepFilterNet3(args.model_dir)
47
+ t0 = time.time()
48
+ out = enh.enhance(audio)
49
+ dt = time.time() - t0
50
+ rtf = dt / (len(audio) / SR)
51
+
52
+ out_path = args.output or args.input.with_name(args.input.stem + "_enhanced.wav")
53
+ write_wav(out_path, out, SR)
54
+ print(f"enhanced -> {out_path} ({dt:.2f}s, RTF={rtf:.3f})")
55
+
56
+
57
+ if __name__ == "__main__":
58
+ main()
deepfilternet3_ax/dsp.py ADDED
@@ -0,0 +1,233 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """DeepFilterNet3 numpy DSP 前后处理(SDK 核心, 仅依赖 numpy)。
2
+
3
+ 对齐来源:
4
+ - 特征/STFT: libDF/src/lib.rs (DFState analysis/feat_erb/feat_cplx, erb_fb, norm alpha)
5
+ - 模型语义: df/deepfilternet3.py DfNet.forward (pad_feat/pad_spec crop2/pad2)
6
+ - DF 滤波: df/modules.py DfOp.forward_real_unfold (spec_pad(5, lookahead=2))
7
+ - 流式语义: libDF/src/tract.rs (buf_x/buf_y 环形缓冲, 输出帧 = t-2)
8
+
9
+ 两种对齐变体:
10
+ batch — 复刻 DfNet.forward: 特征 crop2/pad2 后进网络, m/coefs(t) 应用于 spec'(t)
11
+ streaming— 复刻 tract: 原始特征进网络, coefs(t) 滤波 spec(t-2+o), m(t) 应用于 spec(t)
12
+ 数值对分 (vs enhance() 输出) 决定最终采用哪种。
13
+ """
14
+
15
+ import numpy as np
16
+
17
+ SR = 48000
18
+ FFT_SIZE = 960
19
+ HOP = 480
20
+ NB_ERB = 32
21
+ NB_DF = 96
22
+ DF_ORDER = 5
23
+ DF_LOOKAHEAD = 2
24
+ MIN_NB_ERB_FREQS = 2
25
+ FREQ_SIZE = FFT_SIZE // 2 + 1 # 481
26
+
27
+ MEAN_NORM_INIT = (-60.0, -90.0)
28
+ UNIT_NORM_INIT = (0.001, 0.0001)
29
+
30
+
31
+ def _freq2erb(f):
32
+ return 9.265 * np.log(1 + f / (24.7 * 9.265))
33
+
34
+
35
+ def _erb2freq(e):
36
+ return 24.7 * 9.265 * (np.exp(e / 9.265) - 1)
37
+
38
+
39
+ def erb_bandwidths(sr=SR, fft_size=FFT_SIZE, nb_bands=NB_ERB,
40
+ min_nb_freqs=MIN_NB_ERB_FREQS):
41
+ """libDF erb_fb(): 每 ERB 带频点宽度 (总宽 = fft_size/2+1)。"""
42
+ nyq = sr / 2
43
+ freq_width = sr / fft_size
44
+ erb_low = _freq2erb(0.0)
45
+ erb_high = _freq2erb(nyq)
46
+ step = (erb_high - erb_low) / nb_bands
47
+ erb = []
48
+ prev_freq = 0
49
+ freq_over = 0
50
+ for i in range(1, nb_bands + 1):
51
+ f = _erb2freq(erb_low + i * step)
52
+ fb = int(round(f / freq_width))
53
+ nb_f = fb - prev_freq - freq_over
54
+ if nb_f < min_nb_freqs:
55
+ freq_over = min_nb_freqs - nb_f
56
+ nb_f = min_nb_freqs
57
+ else:
58
+ freq_over = 0
59
+ erb.append(nb_f)
60
+ prev_freq = fb
61
+ erb[nb_bands - 1] += 1
62
+ too_large = int(np.sum(erb)) - (fft_size // 2 + 1)
63
+ if too_large > 0:
64
+ erb[nb_bands - 1] -= too_large
65
+ assert int(np.sum(erb)) == fft_size // 2 + 1
66
+ return np.asarray(erb, dtype=np.int64)
67
+
68
+
69
+ def norm_alpha(sr=SR, hop=HOP, tau=1.0):
70
+ """calc_norm_alpha: exp(-dt/tau) 保留 3 位小数 (libDF)。"""
71
+ alpha = np.exp(-(hop / sr) / tau)
72
+ a = 1.0
73
+ precision = 3
74
+ while a >= 1.0:
75
+ a = round(float(alpha) * 10 ** precision) / 10 ** precision
76
+ precision += 1
77
+ return a
78
+
79
+
80
+ class DFStateNp:
81
+ """libDF DFState 的 numpy 复刻 (流式 STFT/ISTFT + stateful 特征)。"""
82
+
83
+ def __init__(self):
84
+ i = np.arange(FFT_SIZE)
85
+ self.window = np.sin(0.5 * np.pi * np.sin(0.5 * np.pi * (i + 0.5) / (FFT_SIZE / 2)) ** 2)
86
+ self.window = self.window.astype(np.float32)
87
+ self.wnorm = np.float32(1.0 / (FFT_SIZE ** 2 / (2 * HOP)))
88
+ self.erb_w = erb_bandwidths()
89
+ self.alpha = norm_alpha()
90
+ self.band_idx = np.searchsorted(np.cumsum(self.erb_w), np.arange(FREQ_SIZE)) # (481,)
91
+ self.reset()
92
+
93
+ def reset(self):
94
+ self.analysis_mem = np.zeros(FFT_SIZE - HOP, np.float32)
95
+ self.synthesis_mem = np.zeros(FFT_SIZE - HOP, np.float32)
96
+ self.mean_norm_state = np.linspace(*MEAN_NORM_INIT, NB_ERB).astype(np.float32)
97
+ self.unit_norm_state = np.linspace(*UNIT_NORM_INIT, NB_DF).astype(np.float32)
98
+
99
+ def analysis(self, frame):
100
+ buf = np.concatenate([self.analysis_mem, frame]) * self.window
101
+ self.analysis_mem = frame.copy()
102
+ return (np.fft.rfft(buf).astype(np.complex64) * self.wnorm)
103
+
104
+ def synthesis(self, spec):
105
+ # libDF realfft 逆变换不归一化 (wnorm 在 analysis 侧补偿), np.fft 自带 1/n, 需乘回
106
+ x = np.fft.irfft(spec, n=FFT_SIZE) * FFT_SIZE * self.window
107
+ out = x[: HOP] + self.synthesis_mem
108
+ self.synthesis_mem = x[HOP:].copy()
109
+ return out.astype(np.float32)
110
+
111
+ def feat_erb(self, spec):
112
+ p = spec.real ** 2 + spec.imag ** 2
113
+ e = np.add.reduceat(p, np.append([0], np.cumsum(self.erb_w)[:-1])) / self.erb_w
114
+ e = np.log10(e + 1e-10) * 10.0
115
+ s = self.mean_norm_state
116
+ s[:] = e * (1 - self.alpha) + s * self.alpha
117
+ return ((e - s) / 40.0).astype(np.float32)
118
+
119
+ def feat_cplx(self, spec):
120
+ x = spec[:NB_DF].copy()
121
+ s = self.unit_norm_state
122
+ s[:] = np.abs(x) * (1 - self.alpha) + s * self.alpha
123
+ return (x / np.sqrt(s)).astype(np.complex64)
124
+
125
+
126
+ def frames_from_audio(audio, state):
127
+ """audio (N,) float32 -> [spec (481,) complex64] 流式帧序列。
128
+
129
+ 帧数 = floor(N/hop), 与 libDF python binding 一致 (末尾不足一帧零填充)。
130
+ """
131
+ n_frames = len(audio) // HOP
132
+ specs = []
133
+ for f in range(n_frames):
134
+ s = f * HOP
135
+ frame = audio[s: s + HOP]
136
+ if len(frame) < HOP:
137
+ frame = np.pad(frame, (0, HOP - len(frame)))
138
+ specs.append(state.analysis(frame))
139
+ return specs
140
+
141
+
142
+ def _crop2_pad2(x):
143
+ """DfNet.forward pad_feat/pad_spec: T 维头部 crop 2 帧 + 尾部补 2 零帧。"""
144
+ return np.concatenate([x[2:], np.zeros_like(x[:2])], axis=0)
145
+
146
+
147
+ def _df_filter(spec_in, coefs, T_axis=0):
148
+ """spec_pad(5, lookahead=2): 前/后各 pad 2 帧, 逐帧复数乘加。
149
+
150
+ spec_in: (T, NB_DF) complex64; coefs: (T, NB_DF, DF_ORDER, 2) float32
151
+ -> (T, NB_DF) complex64: out(t) = Σ_o coefs(t,o) · spec_in(t-2+o)
152
+ """
153
+ T = spec_in.shape[T_axis]
154
+ padded = np.pad(spec_in, ((2, 2), (0, 0)))
155
+ out = np.zeros((T, NB_DF), np.complex64)
156
+ c = coefs[..., 0] + 1j * coefs[..., 1] # (T, NB_DF, O)
157
+ for o in range(DF_ORDER):
158
+ out += padded[o: o + T] * c[..., o]
159
+ return out
160
+
161
+
162
+ def pipeline(specs, enc_fn, erb_dec_fn, df_dec_fn, mode="streaming", T=99):
163
+ """完整频谱域流水线 (不含 STFT/ISTFT)。
164
+
165
+ specs: (T_all, 481) complex64
166
+ 返回 (T_all, 481) complex64 增强谱。
167
+ mode: "batch" 复刻 DfNet.forward crop 语义; "streaming" 复刻 tract 环形缓冲语义
168
+ """
169
+ st = DFStateNp()
170
+ fb = np.stack([st.feat_erb(s) for s in specs]) # (T_all, 32)
171
+ fc = np.stack([st.feat_cplx(s) for s in specs]) # (T_all, 96) complex
172
+ feat_erb = fb[:, None, :] # (T_all, 1, 32)
173
+ feat_spec = np.stack([fc.real, fc.imag], axis=0) # (2, T_all, 96)
174
+
175
+ if mode == "batch":
176
+ feat_erb_in = _crop2_pad2(feat_erb)
177
+ feat_spec_in = _crop2_pad2(feat_spec)
178
+ spec_work = _crop2_pad2(specs)
179
+ m, coefs = _chunked(feat_erb_in, feat_spec_in, enc_fn, erb_dec_fn, df_dec_fn, T)
180
+ spec_m = spec_work * m[:, st.band_idx]
181
+ spec_e = np.empty_like(spec_work)
182
+ spec_e[:, :NB_DF] = _df_filter(spec_work[:, :NB_DF], coefs)
183
+ spec_e[:, NB_DF:] = spec_m[:, NB_DF:]
184
+ return spec_e
185
+ else: # streaming
186
+ m, coefs = _chunked(feat_erb, feat_spec, enc_fn, erb_dec_fn, df_dec_fn, T)
187
+ spec_m = specs * m[:, st.band_idx]
188
+ spec_e = np.empty_like(specs)
189
+ spec_e[:, :NB_DF] = _df_filter(specs[:, :NB_DF], coefs)
190
+ spec_e[:, NB_DF:] = spec_m[:, NB_DF:]
191
+ return spec_e
192
+
193
+
194
+ def _chunked(feat_erb, feat_spec, enc_fn, erb_dec_fn, df_dec_fn, T=99):
195
+ """分块跑 3 子图 (GRU 每块 h0=0)。
196
+
197
+ enc_fn(fb(1,1,T,32), fs(1,2,T,96)) -> dict(e0,e1,e2,e3,emb,c0,lsnr)
198
+ erb_dec_fn(emb,e3,e2,e1,e0) -> m (1,1,T,32)
199
+ df_dec_fn(emb,c0) -> coefs (1,T,96,10)
200
+ -> m_all (T_all,32), coefs_all (T_all,96,5,2)
201
+ """
202
+ T_all = feat_erb.shape[0]
203
+ n_ch = (T_all + T - 1) // T
204
+ ms, cs = [], []
205
+ for c in range(n_ch):
206
+ sl = slice(c * T, min((c + 1) * T, T_all))
207
+ t = sl.stop - sl.start
208
+ fb = np.zeros((1, 1, T, 32), np.float32)
209
+ fs = np.zeros((1, 2, T, 96), np.float32)
210
+ fb[:, 0, :t, :] = feat_erb[sl, 0] # (t,32)
211
+ fs[0, :, :t, :] = feat_spec[:, sl] # (2,t,96)
212
+ o = enc_fn(fb, fs)
213
+ m = erb_dec_fn(o["emb"], o["e3"], o["e2"], o["e1"], o["e0"])[0, 0]
214
+ coefs = df_dec_fn(o["emb"], o["c0"])[0]
215
+ coefs = coefs.reshape(T, NB_DF, DF_ORDER, 2)
216
+ ms.append(m[:t])
217
+ cs.append(coefs[:t])
218
+ return np.concatenate(ms), np.concatenate(cs)
219
+
220
+
221
+ def enhance_offline(audio, enc_fn, erb_dec_fn, df_dec_fn, mode="streaming", T=99):
222
+ """离线增强 (对齐 enhance()): 输出与输入等长。
223
+
224
+ audio: (N,) float32 48kHz 单声道
225
+ """
226
+ st = DFStateNp()
227
+ audio_pad = np.concatenate([audio, np.zeros(FFT_SIZE, np.float32)])
228
+ specs = np.stack(frames_from_audio(audio_pad, st)) # (T_all, 481)
229
+ spec_e = pipeline(specs, enc_fn, erb_dec_fn, df_dec_fn, mode, T)
230
+ st.reset()
231
+ out = np.concatenate([st.synthesis(s) for s in spec_e]).astype(np.float32)
232
+ d = FFT_SIZE - HOP
233
+ return out[d: len(audio) + d]
deepfilternet3_ax/enhance.py ADDED
@@ -0,0 +1,92 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """DeepFilterNet3 端到端增强入口: 48kHz 单声道 float32 音频 -> 降噪后同形音频。
2
+
3
+ 仅依赖 numpy + axengine (NPU-only, 无 onnxruntime/torch 回退)。
4
+ 前后处理 (STFT/ERB 特征/掩码/DF 滤波/ISTFT) 为 libDF 的 numpy 复刻,
5
+ 已与官方 libDF + torch 参考对分 (cosine=1.0)。
6
+ """
7
+ from __future__ import annotations
8
+
9
+ from pathlib import Path
10
+
11
+ import numpy as np
12
+
13
+ from . import dsp
14
+ from .session import DfNet3Ax
15
+
16
+ SR = 48000
17
+
18
+
19
+ class DeepFilterNet3:
20
+ """DeepFilterNet3 NPU 降噪器。
21
+
22
+ Attributes:
23
+ chunk_frames: 每块静态时间帧数 (99, 对应 1s @48k 流式 STFT)
24
+ """
25
+
26
+ chunk_frames = dsp.DFStateNp.__module__ and 99
27
+
28
+ def __init__(self, model_dir: str | Path):
29
+ self.net = DfNet3Ax(model_dir)
30
+ self.df_state = dsp.DFStateNp()
31
+
32
+ # ---- NPU 分块推理 (对齐 cache/validate_dsp2.py 的验证语义) ----
33
+ def _run_chunks(self, feat_erb, feat_spec, t_all):
34
+ T = self.chunk_frames
35
+ ms, cs = [], []
36
+ h_enc = None # GRU 状态跨块携带 (stateful axmodel 时生效)
37
+ h_df = None
38
+ for c in range((t_all + T - 1) // T):
39
+ sl = slice(c * T, min((c + 1) * T, t_all))
40
+ t = sl.stop - sl.start
41
+ fb = np.zeros((1, 1, T, 32), np.float32)
42
+ fs = np.zeros((1, 2, T, 96), np.float32)
43
+ fb[:, 0, :t, :] = feat_erb[sl]
44
+ fs[0, :, :t, :] = feat_spec[:, sl]
45
+ o, h_enc = self.net.run_enc(fb, fs, h_enc)
46
+ m = self.net.run_erb_dec(o["emb"], o["e3"], o["e2"], o["e1"], o["e0"])
47
+ coefs, h_df = self.net.run_df_dec(o["emb"], o["c0"], h_df)
48
+ ms.append(m[0, 0][:t])
49
+ cs.append(coefs[0][:t])
50
+ m = np.concatenate(ms)
51
+ coefs = np.concatenate(cs).reshape(-1, 96, 5, 2)
52
+ return m, coefs
53
+
54
+ def enhance(self, audio: np.ndarray, sr: int = SR) -> np.ndarray:
55
+ """降噪。audio: (N,) float32 48kHz 单声道 -> (N,) float32。
56
+
57
+ 输入非 48kHz 时建议调用方自行重采样。
58
+ """
59
+ if sr != SR:
60
+ raise ValueError(f"仅支持 48kHz 输入, got {sr}Hz (请先重采样)")
61
+ audio = np.asarray(audio, dtype=np.float32)
62
+ if audio.ndim != 1:
63
+ raise ValueError(f"audio 需为单声道一维数组, got shape {audio.shape}")
64
+
65
+ st = dsp.DFStateNp()
66
+ audio_pad = np.concatenate([audio, np.zeros(dsp.FFT_SIZE, np.float32)])
67
+ specs = np.stack(dsp.frames_from_audio(audio_pad, st))
68
+
69
+ # 特征 (全序列 stateful EMA, 与参考一致)
70
+ fb = np.stack([st.feat_erb(s) for s in specs]) # (T_all, 32)
71
+ fc = np.stack([st.feat_cplx(s) for s in specs]) # (T_all, 96) complex
72
+ feat_erb = fb # (T_all, 32)
73
+ feat_spec = np.stack([fc.real, fc.imag], axis=0) # (2, T_all, 96)
74
+
75
+ m, coefs = self._run_chunks(feat_erb, feat_spec, len(specs))
76
+
77
+ # 后处理: ERB 掩码 + DF 滤波 (流式语义, 与参考一致)
78
+ gains = m[:, st.band_idx]
79
+ spec_m = specs * gains
80
+ spec_e = np.empty_like(specs)
81
+ spec_e[:, :96] = dsp._df_filter(specs[:, :96], coefs)
82
+ spec_e[:, 96:] = spec_m[:, 96:]
83
+
84
+ st.reset()
85
+ out = np.concatenate([st.synthesis(s) for s in spec_e]).astype(np.float32)
86
+ d = dsp.FFT_SIZE - dsp.HOP
87
+ return out[d: len(audio) + d]
88
+
89
+
90
+ def enhance_audio(audio: np.ndarray, model_dir: str | Path, sr: int = SR) -> np.ndarray:
91
+ """便捷入口: 一次调用完成加载与增强 (重复处理建议复用 DeepFilterNet3)。"""
92
+ return DeepFilterNet3(model_dir).enhance(audio, sr=sr)
deepfilternet3_ax/session.py ADDED
@@ -0,0 +1,123 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """3 个 axmodel 的 axengine 会话封装 (enc/erb_dec/df_dec 流水线)。"""
2
+ from __future__ import annotations
3
+
4
+ from dataclasses import dataclass
5
+ from pathlib import Path
6
+ from typing import Any
7
+
8
+ import numpy as np
9
+
10
+ T = 99 # 静态时间帧 (1s @ 48kHz 流式 STFT)
11
+
12
+
13
+ @dataclass(frozen=True)
14
+ class TensorInfo:
15
+ name: str
16
+ shape: tuple[int, ...]
17
+ dtype: np.dtype
18
+
19
+
20
+ def _numpy_dtype(value: Any) -> np.dtype:
21
+ text = str(value).lower()
22
+ mapping = (
23
+ (("tensor(float)", "float32", "fp32", "f32"), np.float32),
24
+ (("tensor(float16)", "float16", "fp16", "f16"), np.float16),
25
+ (("tensor(int64)", "int64", "s64"), np.int64),
26
+ (("tensor(int32)", "int32", "s32"), np.int32),
27
+ (("tensor(uint16)", "uint16", "u16"), np.uint16),
28
+ (("tensor(uint8)", "uint8", "u8"), np.uint8),
29
+ )
30
+ for aliases, dtype in mapping:
31
+ if any(alias in text for alias in aliases):
32
+ return np.dtype(dtype)
33
+ raise ValueError(f"unsupported runtime tensor dtype: {value}")
34
+
35
+
36
+ def _tensor_info(value: Any) -> TensorInfo:
37
+ shape = getattr(value, "shape", None)
38
+ if shape is None:
39
+ shape = getattr(value, "dims", None)
40
+ if shape is None or any(dim is None for dim in shape):
41
+ raise ValueError(f"dynamic or missing tensor shape for {value.name}: {shape}")
42
+ dtype = getattr(value, "dtype", None)
43
+ if dtype is None:
44
+ dtype = getattr(value, "type", None)
45
+ return TensorInfo(
46
+ name=value.name,
47
+ shape=tuple(int(dim) for dim in shape),
48
+ dtype=_numpy_dtype(dtype),
49
+ )
50
+
51
+
52
+ class _AxSession:
53
+ """单个 axmodel 的 axengine 会话 (参考 AXERA 官方推理工程封装)。"""
54
+
55
+ def __init__(self, model_path: str | Path):
56
+ import axengine
57
+ self.path = Path(model_path)
58
+ if not self.path.is_file():
59
+ raise FileNotFoundError(self.path)
60
+ self._session = axengine.InferenceSession(str(self.path))
61
+ self.inputs = [_tensor_info(v) for v in self._session.get_inputs()]
62
+ self.outputs = [_tensor_info(v) for v in self._session.get_outputs()]
63
+ self.input_by_name = {v.name: v for v in self.inputs}
64
+
65
+ def run(self, feed: dict[str, np.ndarray]) -> dict[str, np.ndarray]:
66
+ missing = [v.name for v in self.inputs if v.name not in feed]
67
+ if missing:
68
+ raise KeyError(f"missing inputs for {self.path.name}: {missing}")
69
+ prepared = {
70
+ name: np.ascontiguousarray(np.asarray(feed[name], dtype=meta.dtype))
71
+ for name, meta in self.input_by_name.items()
72
+ }
73
+ values = self._session.run(None, prepared)
74
+ if isinstance(values, dict):
75
+ return {name: np.asarray(v) for name, v in values.items()}
76
+ if not isinstance(values, (list, tuple)):
77
+ values = [values]
78
+ if len(values) != len(self.outputs):
79
+ raise RuntimeError(
80
+ f"unexpected output count from {self.path.name}: "
81
+ f"{len(values)} != {len(self.outputs)}")
82
+ return {meta.name: np.asarray(v) for meta, v in zip(self.outputs, values)}
83
+
84
+
85
+ class DfNet3Ax:
86
+ """DeepFilterNet3 三子图 NPU 流水线 (与 dsp.pipeline 配合)。
87
+
88
+ enc: feat_erb(1,1,T,32)+feat_spec(1,2,T,96) -> e0..e3, emb, c0, lsnr
89
+ erb_dec: emb,e3,e2,e1,e0 -> m(1,1,T,32)
90
+ df_dec: emb,c0 -> coefs(1,T,96,10)
91
+ """
92
+
93
+ def __init__(self, model_dir: str | Path):
94
+ model_dir = Path(model_dir)
95
+ self.enc = _AxSession(model_dir / "enc.axmodel")
96
+ self.erb_dec = _AxSession(model_dir / "erb_dec.axmodel")
97
+ self.df_dec = _AxSession(model_dir / "df_dec.axmodel")
98
+ # 状态携带模式检测: axmodel 是否有 h0 输入
99
+ self.stateful = all("h0" in s.input_by_name for s in (self.enc, self.df_dec))
100
+
101
+ def run_enc(self, feat_erb: np.ndarray, feat_spec: np.ndarray,
102
+ h: np.ndarray | None = None) -> tuple[dict[str, np.ndarray], np.ndarray | None]:
103
+ feed = {"feat_erb": feat_erb, "feat_spec": feat_spec}
104
+ if self.stateful:
105
+ if h is None:
106
+ h = np.zeros((1, 1, 256), np.float32)
107
+ feed["h0"] = h
108
+ out = self.enc.run(feed)
109
+ return out, np.asarray(out["h_out"])
110
+ return self.enc.run(feed), None
111
+
112
+ def run_erb_dec(self, emb, e3, e2, e1, e0) -> np.ndarray:
113
+ return self.erb_dec.run({"emb": emb, "e3": e3, "e2": e2, "e1": e1, "e0": e0})["m"]
114
+
115
+ def run_df_dec(self, emb, c0, h: np.ndarray | None = None) -> tuple[np.ndarray, np.ndarray | None]:
116
+ feed = {"emb": emb, "c0": c0}
117
+ if self.stateful:
118
+ if h is None:
119
+ h = np.zeros((2, 1, 256), np.float32)
120
+ feed["h0"] = h
121
+ out = self.df_dec.run(feed)
122
+ return np.asarray(out["coefs"]), np.asarray(out["h_out"])
123
+ return self.df_dec.run(feed)["coefs"], None
inference.py ADDED
@@ -0,0 +1,52 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """DeepFilterNet3 推理入口 (AXERA 平台, 当前 AX650N) (依赖仅 numpy + axengine)。
2
+
3
+ 用法:
4
+ python3 inference.py input.wav [-o output.wav] [--model-dir axmodels]
5
+ """
6
+ import argparse
7
+ import wave
8
+ from pathlib import Path
9
+
10
+ import numpy as np
11
+
12
+ SR = 48000
13
+
14
+
15
+ def read_wav(path):
16
+ with wave.open(str(path), "rb") as w:
17
+ assert w.getnchannels() == 1, "仅支持单声道"
18
+ sr = w.getframerate()
19
+ data = np.frombuffer(w.readframes(w.getnframes()), dtype=np.int16)
20
+ return (data.astype(np.float32) / 32768.0), sr
21
+
22
+
23
+ def write_wav(path, audio, sr):
24
+ pcm = (np.clip(audio, -1.0, 1.0) * 32767.0).astype(np.int16)
25
+ with wave.open(str(path), "wb") as w:
26
+ w.setnchannels(1)
27
+ w.setsampwidth(2)
28
+ w.setframerate(sr)
29
+ w.writeframes(pcm.tobytes())
30
+
31
+
32
+ def main():
33
+ parser = argparse.ArgumentParser(description=__doc__)
34
+ parser.add_argument("input", type=Path)
35
+ parser.add_argument("-o", "--output", type=Path, default=None)
36
+ parser.add_argument("--model-dir", type=Path, default=Path("axmodels"))
37
+ args = parser.parse_args()
38
+
39
+ from deepfilternet3_ax import DeepFilterNet3
40
+
41
+ audio, sr = read_wav(args.input)
42
+ if sr != SR:
43
+ raise SystemExit(f"仅支持 48kHz 输入, got {sr}Hz")
44
+ enh = DeepFilterNet3(args.model_dir)
45
+ out = enh.enhance(audio)
46
+ out_path = args.output or args.input.with_name(args.input.stem + "_enhanced.wav")
47
+ write_wav(out_path, out, SR)
48
+ print(f"enhanced -> {out_path}")
49
+
50
+
51
+ if __name__ == "__main__":
52
+ main()
requirements.txt ADDED
@@ -0,0 +1,4 @@
 
 
 
 
 
1
+ numpy>=1.24
2
+ # pyaxengine (板端 axengine, 从官方 release 下载 wheel 安装):
3
+ # https://github.com/AXERA-TECH/pyaxengine/releases/latest
4
+ # pip3 install axengine-<version>-py3-none-any.whl
run.sh ADDED
@@ -0,0 +1,16 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ #!/usr/bin/env bash
2
+ # DeepFilterNet3 一键降噪 (Python)
3
+ # 用法: ./run.sh input.wav [output.wav]
4
+ set -euo pipefail
5
+ cd "$(dirname "$0")"
6
+
7
+ INPUT="${1:?用法: ./run.sh input.wav [output.wav]}"
8
+ OUTPUT="${2:-${INPUT%.wav}_enhanced.wav}"
9
+
10
+ PYTHON_BIN="${PYTHON_BIN:-python3}"
11
+ if [[ -x /root/miniforge3/bin/python ]]; then
12
+ PYTHON_BIN=/root/miniforge3/bin/python
13
+ fi
14
+
15
+ PYTHONPATH="$(pwd)${PYTHONPATH:+:$PYTHONPATH}" \
16
+ "$PYTHON_BIN" inference.py "$INPUT" -o "$OUTPUT" --model-dir "$(pwd)/axmodels"