DeepFilterNet3 AXERA 量化部署包: axmodels (GRU 状态携带) + Python SDK + C++ 可执行文件
Browse files- .gitattributes +3 -0
- .gitignore +8 -0
- README.md +98 -0
- assets/enhanced_sample.wav +3 -0
- assets/noisy_snr0.wav +3 -0
- axmodels/df_dec.axmodel +3 -0
- axmodels/enc.axmodel +3 -0
- axmodels/erb_dec.axmodel +3 -0
- bin/df3_ax +3 -0
- configuration.json +20 -0
- deepfilternet3_ax/__init__.py +14 -0
- deepfilternet3_ax/__main__.py +58 -0
- deepfilternet3_ax/dsp.py +233 -0
- deepfilternet3_ax/enhance.py +92 -0
- deepfilternet3_ax/session.py +123 -0
- inference.py +52 -0
- requirements.txt +4 -0
- run.sh +16 -0
.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"
|