Upload folder using huggingface_hub
Browse files- .gitattributes +5 -0
- .gitignore +7 -0
- README.md +94 -1
- bin/sortformer_ax650 +3 -0
- configuration.json +17 -0
- models/encoder.axmodel +3 -0
- models/encoder_fifo40.axmodel +3 -0
- models/model_meta.json +57 -0
- models/model_meta_fifo40.json +57 -0
- models/preencode.axmodel +3 -0
- python/example.py +63 -0
- python/sortformer_sdk/__init__.py +34 -0
- python/sortformer_sdk/diarize.py +179 -0
- python/sortformer_sdk/feature.py +128 -0
- python/sortformer_sdk/postprocess.py +181 -0
- python/sortformer_sdk/state.py +210 -0
- requirements.txt +4 -0
- run_ax650.sh +18 -0
- run_cpp_ax650.sh +26 -0
- samples/sample_meeting.wav +3 -0
.gitattributes
CHANGED
|
@@ -33,3 +33,8 @@ saved_model/**/* filter=lfs diff=lfs merge=lfs -text
|
|
| 33 |
*.zip filter=lfs diff=lfs merge=lfs -text
|
| 34 |
*.zst filter=lfs diff=lfs merge=lfs -text
|
| 35 |
*tfevents* filter=lfs diff=lfs merge=lfs -text
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 33 |
*.zip filter=lfs diff=lfs merge=lfs -text
|
| 34 |
*.zst filter=lfs diff=lfs merge=lfs -text
|
| 35 |
*tfevents* filter=lfs diff=lfs merge=lfs -text
|
| 36 |
+
bin/sortformer_ax650 filter=lfs diff=lfs merge=lfs -text
|
| 37 |
+
models/encoder.axmodel filter=lfs diff=lfs merge=lfs -text
|
| 38 |
+
models/encoder_fifo40.axmodel filter=lfs diff=lfs merge=lfs -text
|
| 39 |
+
models/preencode.axmodel filter=lfs diff=lfs merge=lfs -text
|
| 40 |
+
samples/sample_meeting.wav filter=lfs diff=lfs merge=lfs -text
|
.gitignore
ADDED
|
@@ -0,0 +1,7 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
__pycache__/
|
| 2 |
+
*.pyc
|
| 3 |
+
*.rttm
|
| 4 |
+
outputs/
|
| 5 |
+
.msc/
|
| 6 |
+
.mv/
|
| 7 |
+
._____temp/
|
README.md
CHANGED
|
@@ -1,3 +1,96 @@
|
|
| 1 |
---
|
| 2 |
-
license:
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 3 |
---
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
---
|
| 2 |
+
license: other
|
| 3 |
+
license_name: nvidia-open-model-license
|
| 4 |
+
license_link: https://www.nvidia.com/en-us/agreements/enterprise-software/nvidia-open-model-license/
|
| 5 |
+
language:
|
| 6 |
+
- en
|
| 7 |
+
pipeline_tag: audio-classification
|
| 8 |
+
tags:
|
| 9 |
+
- speaker-diarization
|
| 10 |
+
- sortformer
|
| 11 |
+
- axera
|
| 12 |
+
- ax650
|
| 13 |
+
- rttm
|
| 14 |
---
|
| 15 |
+
|
| 16 |
+
# Sortformer.AXERA
|
| 17 |
+
|
| 18 |
+
`nvidia/diar_streaming_sortformer_4spk-v2.1`(NVIDIA Open Model License)端到端流式说话人日志在 AX650N 的量化部署包:
|
| 19 |
+
输入 16 kHz 单声道会议音频,输出 RTTM 说话人时序标签。
|
| 20 |
+
|
| 21 |
+
- 模型: 2 张拆图 axmodel(U16 激活 / S8 权重)——preencode + encoder;encoder 提供 fifo188(精度优先)与 fifo40(速度优先)
|
| 22 |
+
- 依赖: Python 仅 numpy + soundfile + axengine;C++ 为预编译 aarch64 可执行文件
|
| 23 |
+
|
| 24 |
+
## 目录
|
| 25 |
+
|
| 26 |
+
```
|
| 27 |
+
├── models/ # preencode.axmodel / encoder.axmodel / encoder_fifo40.axmodel + model_meta
|
| 28 |
+
├── python/ # Python 推理入口(sortformer_sdk + example.py)
|
| 29 |
+
├── bin/ # C++ 可执行文件(aarch64):sortformer_ax650
|
| 30 |
+
├── samples/ # 演示音频(AMI 会议 120 s)
|
| 31 |
+
├── run_ax650.sh run_cpp_ax650.sh # 一键运行(Python / C++)
|
| 32 |
+
└── requirements.txt
|
| 33 |
+
```
|
| 34 |
+
|
| 35 |
+
## Python 运行
|
| 36 |
+
|
| 37 |
+
```bash
|
| 38 |
+
pip3 install numpy soundfile
|
| 39 |
+
# pyaxengine: https://github.com/AXERA-TECH/pyaxengine/releases/latest
|
| 40 |
+
pip3 install axengine-<version>-py3-none-any.whl
|
| 41 |
+
./run_ax650.sh # 自带样例 samples/sample_meeting.wav
|
| 42 |
+
./run_ax650.sh input.wav out.rttm --fast # fifo40 快速版
|
| 43 |
+
```
|
| 44 |
+
|
| 45 |
+
或:
|
| 46 |
+
|
| 47 |
+
```python
|
| 48 |
+
import sys; sys.path.insert(0, "python")
|
| 49 |
+
from sortformer_sdk import AxengineGraphPair, StreamingDiarizer, SortformerConfig
|
| 50 |
+
|
| 51 |
+
diarizer = StreamingDiarizer(
|
| 52 |
+
AxengineGraphPair("models/preencode.axmodel", "models/encoder.axmodel"),
|
| 53 |
+
SortformerConfig(),
|
| 54 |
+
)
|
| 55 |
+
preds = diarizer.process_wav(waveform_16k_float32, 16000) # (T, 4) @ 80 ms
|
| 56 |
+
lines = diarizer.rttm_lines(preds, uri="meeting")
|
| 57 |
+
```
|
| 58 |
+
|
| 59 |
+
## C++ 运行
|
| 60 |
+
|
| 61 |
+
```bash
|
| 62 |
+
export LD_LIBRARY_PATH=/soc/lib:${LD_LIBRARY_PATH:-}
|
| 63 |
+
./bin/sortformer_ax650 --preencode models/preencode.axmodel --encoder models/encoder.axmodel \
|
| 64 |
+
--wav input.wav --rttm out.rttm --threads 8
|
| 65 |
+
# 或一键:./run_cpp_ax650.sh input.wav out.rttm [--fast]
|
| 66 |
+
```
|
| 67 |
+
|
| 68 |
+
## 模型说明
|
| 69 |
+
|
| 70 |
+
- 输入: 16 kHz 单声道 wav;输出: RTTM(`SPEAKER <uri> 1 <start> <dur> <NA> <NA> speaker_k <NA>`)
|
| 71 |
+
- 流程: 主机 log-mel(128 mel / 25 ms 窗 / 10 ms 步长,与 NeMo 逐点对齐)→ preencode → 主机打包
|
| 72 |
+
`seq [1,390,512]` → encoder → 主机 `streaming_update`(静音画像 / top-k + AOSC 压缩 / FIFO 弹出)→ 后处理
|
| 73 |
+
- 拆图契约:spkcache 188 + fifo 188 + chunk 14 帧,图内无 Scatter/Where;80 ms/帧、最多 4 说话人、1.04 s 延迟
|
| 74 |
+
- 完整转换源码见 GitHub: [Sortformer.AXERA](https://github.com/ZY-2012/Sortformer.AXERA)
|
| 75 |
+
|
| 76 |
+
## 精度与性能(AX650N 实测)
|
| 77 |
+
|
| 78 |
+
ES2004a(1049 s,collar=0):
|
| 79 |
+
|
| 80 |
+
| 推理路径 | DER | RTF |
|
| 81 |
+
|---|---|---|
|
| 82 |
+
| NeMo FP32(参考) | 30.89% | — |
|
| 83 |
+
| C++ fifo188(精度优先) | 30.70% | 0.144 |
|
| 84 |
+
| C++ fifo40(速度优先) | 30.64% | 0.064 |
|
| 85 |
+
|
| 86 |
+
AMI-SDM(5 场加权)/ AliMeeting(4 场远场,TextGrid 参考),collar=0:
|
| 87 |
+
|
| 88 |
+
| 配置 | AMI-5 DER | Ali-4 DER |
|
| 89 |
+
|---|---|---|
|
| 90 |
+
| fifo188 | 33.77% | 21.71% |
|
| 91 |
+
| fifo40 | 34.34% | 22.02% |
|
| 92 |
+
|
| 93 |
+
## License
|
| 94 |
+
|
| 95 |
+
部署代码 Apache-2.0;上游 Sortformer 权重为
|
| 96 |
+
[NVIDIA Open Model License](https://www.nvidia.com/en-us/agreements/enterprise-software/nvidia-open-model-license/)。
|
bin/sortformer_ax650
ADDED
|
@@ -0,0 +1,3 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
version https://git-lfs.github.com/spec/v1
|
| 2 |
+
oid sha256:b54e33a4ff0ef1d28a4d96f3a0b010778715c4f19a16d40651abf08bf538ca65
|
| 3 |
+
size 144952
|
configuration.json
ADDED
|
@@ -0,0 +1,17 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
{
|
| 2 |
+
"model_name": "Sortformer-streaming-4spk-v2.1-ax650",
|
| 3 |
+
"task": "speaker-diarization",
|
| 4 |
+
"base_model": "nvidia/diar_streaming_sortformer_4spk-v2.1",
|
| 5 |
+
"chip": "AX650N",
|
| 6 |
+
"npu_mode": "NPU3",
|
| 7 |
+
"sample_rate": 16000,
|
| 8 |
+
"max_speakers": 4,
|
| 9 |
+
"frame_shift_ms": 80,
|
| 10 |
+
"latency_s": 1.04,
|
| 11 |
+
"pulsar2_version": "7.0",
|
| 12 |
+
"axmodel_files": [
|
| 13 |
+
"models/preencode.axmodel",
|
| 14 |
+
"models/encoder.axmodel",
|
| 15 |
+
"models/encoder_fifo40.axmodel"
|
| 16 |
+
]
|
| 17 |
+
}
|
models/encoder.axmodel
ADDED
|
@@ -0,0 +1,3 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
version https://git-lfs.github.com/spec/v1
|
| 2 |
+
oid sha256:5ce5abb7c3509c85ef76eda43e4858658e0c498de370a5baa2e423dcdef2ef96
|
| 3 |
+
size 139147096
|
models/encoder_fifo40.axmodel
ADDED
|
@@ -0,0 +1,3 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
version https://git-lfs.github.com/spec/v1
|
| 2 |
+
oid sha256:50124d50843561cb2a0653ed911afb2fae21fcc84f2e38c6fd408744ec621fa8
|
| 3 |
+
size 132804934
|
models/model_meta.json
ADDED
|
@@ -0,0 +1,57 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
{
|
| 2 |
+
"model": "nvidia/diar_streaming_sortformer_4spk-v2.1",
|
| 3 |
+
"contract": "split-preencode-encoder",
|
| 4 |
+
"simplified": true,
|
| 5 |
+
"config": {
|
| 6 |
+
"chunk_len": 6,
|
| 7 |
+
"chunk_left_context": 1,
|
| 8 |
+
"chunk_right_context": 7,
|
| 9 |
+
"fifo_len": 188,
|
| 10 |
+
"spkcache_len": 188,
|
| 11 |
+
"spkcache_update_period": 144
|
| 12 |
+
},
|
| 13 |
+
"preencode": {
|
| 14 |
+
"file": "preencode.onnx",
|
| 15 |
+
"inputs": {
|
| 16 |
+
"chunk": [
|
| 17 |
+
1,
|
| 18 |
+
112,
|
| 19 |
+
128
|
| 20 |
+
]
|
| 21 |
+
},
|
| 22 |
+
"outputs": {
|
| 23 |
+
"chunk_pre_encode_embs": [
|
| 24 |
+
1,
|
| 25 |
+
14,
|
| 26 |
+
512
|
| 27 |
+
]
|
| 28 |
+
},
|
| 29 |
+
"parity_min_cosine": 0.9999999403953552
|
| 30 |
+
},
|
| 31 |
+
"encoder": {
|
| 32 |
+
"file": "encoder.onnx",
|
| 33 |
+
"inputs": {
|
| 34 |
+
"seq": [
|
| 35 |
+
1,
|
| 36 |
+
390,
|
| 37 |
+
512
|
| 38 |
+
],
|
| 39 |
+
"total_lengths": [
|
| 40 |
+
1
|
| 41 |
+
]
|
| 42 |
+
},
|
| 43 |
+
"outputs": {
|
| 44 |
+
"spkcache_fifo_chunk_preds": [
|
| 45 |
+
1,
|
| 46 |
+
390,
|
| 47 |
+
4
|
| 48 |
+
]
|
| 49 |
+
},
|
| 50 |
+
"parity_min_cosine": 0.9999998807907104
|
| 51 |
+
},
|
| 52 |
+
"state_len": 376,
|
| 53 |
+
"chunk_mel_frames": 112,
|
| 54 |
+
"frame_duration_s": 0.08,
|
| 55 |
+
"sample_rate": 16000,
|
| 56 |
+
"mel_features": 128
|
| 57 |
+
}
|
models/model_meta_fifo40.json
ADDED
|
@@ -0,0 +1,57 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
{
|
| 2 |
+
"model": "nvidia/diar_streaming_sortformer_4spk-v2.1",
|
| 3 |
+
"contract": "split-preencode-encoder",
|
| 4 |
+
"simplified": true,
|
| 5 |
+
"config": {
|
| 6 |
+
"chunk_len": 6,
|
| 7 |
+
"chunk_left_context": 1,
|
| 8 |
+
"chunk_right_context": 7,
|
| 9 |
+
"fifo_len": 40,
|
| 10 |
+
"spkcache_len": 188,
|
| 11 |
+
"spkcache_update_period": 31
|
| 12 |
+
},
|
| 13 |
+
"preencode": {
|
| 14 |
+
"file": "preencode.onnx",
|
| 15 |
+
"inputs": {
|
| 16 |
+
"chunk": [
|
| 17 |
+
1,
|
| 18 |
+
112,
|
| 19 |
+
128
|
| 20 |
+
]
|
| 21 |
+
},
|
| 22 |
+
"outputs": {
|
| 23 |
+
"chunk_pre_encode_embs": [
|
| 24 |
+
1,
|
| 25 |
+
14,
|
| 26 |
+
512
|
| 27 |
+
]
|
| 28 |
+
},
|
| 29 |
+
"parity_min_cosine": 1.0
|
| 30 |
+
},
|
| 31 |
+
"encoder": {
|
| 32 |
+
"file": "encoder.onnx",
|
| 33 |
+
"inputs": {
|
| 34 |
+
"seq": [
|
| 35 |
+
1,
|
| 36 |
+
242,
|
| 37 |
+
512
|
| 38 |
+
],
|
| 39 |
+
"total_lengths": [
|
| 40 |
+
1
|
| 41 |
+
]
|
| 42 |
+
},
|
| 43 |
+
"outputs": {
|
| 44 |
+
"spkcache_fifo_chunk_preds": [
|
| 45 |
+
1,
|
| 46 |
+
242,
|
| 47 |
+
4
|
| 48 |
+
]
|
| 49 |
+
},
|
| 50 |
+
"parity_min_cosine": 0.9999999403953552
|
| 51 |
+
},
|
| 52 |
+
"state_len": 228,
|
| 53 |
+
"chunk_mel_frames": 112,
|
| 54 |
+
"frame_duration_s": 0.08,
|
| 55 |
+
"sample_rate": 16000,
|
| 56 |
+
"mel_features": 128
|
| 57 |
+
}
|
models/preencode.axmodel
ADDED
|
@@ -0,0 +1,3 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
version https://git-lfs.github.com/spec/v1
|
| 2 |
+
oid sha256:399ec91c5845102985614bc5f19556b43816f57d16b9030ab86be67b908a6701
|
| 3 |
+
size 2535538
|
python/example.py
ADDED
|
@@ -0,0 +1,63 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
#!/usr/bin/env python3
|
| 2 |
+
"""Sortformer streaming diarization on AX650N: 16 kHz wav -> RTTM.
|
| 3 |
+
|
| 4 |
+
python3 example.py meeting.wav out.rttm [--fast]
|
| 5 |
+
"""
|
| 6 |
+
|
| 7 |
+
import argparse
|
| 8 |
+
import sys
|
| 9 |
+
import time
|
| 10 |
+
from pathlib import Path
|
| 11 |
+
|
| 12 |
+
import numpy as np
|
| 13 |
+
import soundfile as sf
|
| 14 |
+
|
| 15 |
+
sys.path.insert(0, str(Path(__file__).resolve().parent))
|
| 16 |
+
|
| 17 |
+
from sortformer_sdk import AxengineGraphPair, SortformerConfig, StreamingDiarizer # noqa: E402
|
| 18 |
+
|
| 19 |
+
MODELS = Path(__file__).resolve().parents[1] / "models"
|
| 20 |
+
|
| 21 |
+
|
| 22 |
+
def main() -> None:
|
| 23 |
+
parser = argparse.ArgumentParser(description=__doc__)
|
| 24 |
+
parser.add_argument("wav")
|
| 25 |
+
parser.add_argument("rttm", nargs="?", default=None)
|
| 26 |
+
parser.add_argument("--fast", action="store_true", help="fifo40 encoder: RTF 0.064, DER +0.3~0.6pp")
|
| 27 |
+
parser.add_argument("--preencode", default=str(MODELS / "preencode.axmodel"))
|
| 28 |
+
parser.add_argument("--encoder", default=None)
|
| 29 |
+
parser.add_argument("--threads", type=int, default=8, help="reserved (numpy front-end is single-process)")
|
| 30 |
+
args = parser.parse_args()
|
| 31 |
+
|
| 32 |
+
encoder = args.encoder or str(MODELS / ("encoder_fifo40.axmodel" if args.fast else "encoder.axmodel"))
|
| 33 |
+
config = SortformerConfig(
|
| 34 |
+
fifo_len=40 if args.fast else 188,
|
| 35 |
+
spkcache_update_period=31 if args.fast else 144,
|
| 36 |
+
)
|
| 37 |
+
|
| 38 |
+
waveform, sample_rate = sf.read(args.wav, dtype="float32", always_2d=True)
|
| 39 |
+
waveform = waveform.mean(axis=1)
|
| 40 |
+
if sample_rate != 16000:
|
| 41 |
+
raise SystemExit(f"expected 16 kHz wav, got {sample_rate}")
|
| 42 |
+
|
| 43 |
+
diarizer = StreamingDiarizer(AxengineGraphPair(args.preencode, encoder), config)
|
| 44 |
+
start = time.perf_counter()
|
| 45 |
+
preds = diarizer.process_wav(waveform, sample_rate)
|
| 46 |
+
elapsed = time.perf_counter() - start
|
| 47 |
+
duration = waveform.shape[0] / sample_rate
|
| 48 |
+
|
| 49 |
+
uri = Path(args.wav).stem
|
| 50 |
+
lines = diarizer.rttm_lines(preds, uri)
|
| 51 |
+
if args.rttm:
|
| 52 |
+
Path(args.rttm).write_text("\n".join(lines) + "\n")
|
| 53 |
+
|
| 54 |
+
print(
|
| 55 |
+
f"{uri}: {len(preds)} frames, {len(lines)} segments, "
|
| 56 |
+
f"{elapsed:.1f}s / {duration:.1f}s audio (RTF {elapsed / duration:.4f})"
|
| 57 |
+
)
|
| 58 |
+
if args.rttm:
|
| 59 |
+
print(f"rttm -> {args.rttm}")
|
| 60 |
+
|
| 61 |
+
|
| 62 |
+
if __name__ == "__main__":
|
| 63 |
+
main()
|
python/sortformer_sdk/__init__.py
ADDED
|
@@ -0,0 +1,34 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
from .diarize import AxengineGraphPair, OnnxGraphPair, StreamingDiarizer
|
| 2 |
+
from .feature import log_mel_spectrogram, mel_filterbank
|
| 3 |
+
from .postprocess import (
|
| 4 |
+
PostProcessingParams,
|
| 5 |
+
generate_diarization_output_lines,
|
| 6 |
+
predlist_to_timestamps,
|
| 7 |
+
timestamps_to_rttm_lines,
|
| 8 |
+
)
|
| 9 |
+
from .state import (
|
| 10 |
+
SortformerConfig,
|
| 11 |
+
StreamingState,
|
| 12 |
+
init_state,
|
| 13 |
+
iter_chunks,
|
| 14 |
+
pre_encode_length,
|
| 15 |
+
streaming_update,
|
| 16 |
+
)
|
| 17 |
+
|
| 18 |
+
__all__ = [
|
| 19 |
+
"AxengineGraphPair",
|
| 20 |
+
"OnnxGraphPair",
|
| 21 |
+
"PostProcessingParams",
|
| 22 |
+
"SortformerConfig",
|
| 23 |
+
"StreamingDiarizer",
|
| 24 |
+
"StreamingState",
|
| 25 |
+
"generate_diarization_output_lines",
|
| 26 |
+
"init_state",
|
| 27 |
+
"iter_chunks",
|
| 28 |
+
"log_mel_spectrogram",
|
| 29 |
+
"mel_filterbank",
|
| 30 |
+
"pre_encode_length",
|
| 31 |
+
"predlist_to_timestamps",
|
| 32 |
+
"streaming_update",
|
| 33 |
+
"timestamps_to_rttm_lines",
|
| 34 |
+
]
|
python/sortformer_sdk/diarize.py
ADDED
|
@@ -0,0 +1,179 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""End-to-end streaming diarizer: waveform -> RTTM.
|
| 2 |
+
|
| 3 |
+
Two scatter-free graphs are executed per step (pre-encode + encoder); the
|
| 4 |
+
speaker-cache packing, streaming state update, log-mel front-end and
|
| 5 |
+
post-processing are all numpy and shared by host and AX650.
|
| 6 |
+
"""
|
| 7 |
+
|
| 8 |
+
import argparse
|
| 9 |
+
from pathlib import Path
|
| 10 |
+
from typing import Dict, Protocol
|
| 11 |
+
|
| 12 |
+
import numpy as np
|
| 13 |
+
|
| 14 |
+
from .feature import SAMPLE_RATE, log_mel_spectrogram, mel_filterbank
|
| 15 |
+
from .postprocess import PostProcessingParams, predlist_to_timestamps, timestamps_to_rttm_lines
|
| 16 |
+
from .state import SortformerConfig, init_state, iter_chunks, pre_encode_length, streaming_update
|
| 17 |
+
|
| 18 |
+
|
| 19 |
+
class GraphPair(Protocol):
|
| 20 |
+
def pre_encode(self, chunk: np.ndarray) -> np.ndarray:
|
| 21 |
+
"""chunk [1, chunk_mel, 128] -> chunk embeddings [1, chunk_embs, 512]."""
|
| 22 |
+
|
| 23 |
+
def encode(self, seq: np.ndarray, total_lengths: int) -> np.ndarray:
|
| 24 |
+
"""packed seq [1, state_len + chunk_embs, 512] -> preds [1, T, num_speakers]."""
|
| 25 |
+
|
| 26 |
+
|
| 27 |
+
def _provider_alias(ort, provider: str) -> str:
|
| 28 |
+
aliases = {"cpu": "CPUExecutionProvider", "cuda": "CUDAExecutionProvider"}
|
| 29 |
+
if provider == "auto":
|
| 30 |
+
available = ort.get_available_providers()
|
| 31 |
+
return "CUDAExecutionProvider" if "CUDAExecutionProvider" in available else "CPUExecutionProvider"
|
| 32 |
+
return aliases.get(provider, provider)
|
| 33 |
+
|
| 34 |
+
|
| 35 |
+
class OnnxGraphPair:
|
| 36 |
+
"""ONNX Runtime implementation of the split graphs (host verification)."""
|
| 37 |
+
|
| 38 |
+
def __init__(self, preencode_path: str, encoder_path: str, provider: str = "auto"):
|
| 39 |
+
import onnxruntime as ort
|
| 40 |
+
|
| 41 |
+
provider = _provider_alias(ort, provider)
|
| 42 |
+
self.pre_session = ort.InferenceSession(str(preencode_path), providers=[provider])
|
| 43 |
+
self.encoder_session = ort.InferenceSession(str(encoder_path), providers=[provider])
|
| 44 |
+
# The encoder graph may be wider than the runtime state (e.g. a 390-wide
|
| 45 |
+
# graph evaluated with a 40-frame FIFO); the host pads and masks the tail.
|
| 46 |
+
encoder_inputs = {node.name: node for node in self.encoder_session.get_inputs()}
|
| 47 |
+
self.seq_width = int(encoder_inputs["seq"].shape[1])
|
| 48 |
+
|
| 49 |
+
def pre_encode(self, chunk: np.ndarray) -> np.ndarray:
|
| 50 |
+
return self.pre_session.run(None, {"chunk": chunk.astype(np.float32)})[0]
|
| 51 |
+
|
| 52 |
+
def encode(self, seq: np.ndarray, total_lengths: int) -> np.ndarray:
|
| 53 |
+
feed = {"seq": seq.astype(np.float32), "total_lengths": np.array([total_lengths], dtype=np.float32)}
|
| 54 |
+
return self.encoder_session.run(None, feed)[0]
|
| 55 |
+
|
| 56 |
+
|
| 57 |
+
class AxengineGraphPair:
|
| 58 |
+
"""axengine implementation of the split graphs (AX650 board)."""
|
| 59 |
+
|
| 60 |
+
def __init__(self, preencode_path: str, encoder_path: str):
|
| 61 |
+
import axengine
|
| 62 |
+
|
| 63 |
+
self.pre_session = axengine.InferenceSession(str(preencode_path))
|
| 64 |
+
self.encoder_session = axengine.InferenceSession(str(encoder_path))
|
| 65 |
+
|
| 66 |
+
def pre_encode(self, chunk: np.ndarray) -> np.ndarray:
|
| 67 |
+
return self.pre_session.run(None, {"chunk": chunk.astype(np.float32)})[0]
|
| 68 |
+
|
| 69 |
+
def encode(self, seq: np.ndarray, total_lengths: int) -> np.ndarray:
|
| 70 |
+
feed = {"seq": seq.astype(np.float32), "total_lengths": np.array([total_lengths], dtype=np.float32)}
|
| 71 |
+
return self.encoder_session.run(None, feed)[0]
|
| 72 |
+
|
| 73 |
+
|
| 74 |
+
class StreamingDiarizer:
|
| 75 |
+
"""Sortformer 1.04 s streaming diarizer over the split graphs."""
|
| 76 |
+
|
| 77 |
+
def __init__(self, graphs: GraphPair, config: SortformerConfig = None):
|
| 78 |
+
self.graphs = graphs
|
| 79 |
+
self.config = config or SortformerConfig()
|
| 80 |
+
self.chunk_width = (
|
| 81 |
+
self.config.chunk_left_context + self.config.chunk_len + self.config.chunk_right_context
|
| 82 |
+
) * self.config.subsampling_factor
|
| 83 |
+
self.state_len = self.config.spkcache_len + self.config.fifo_len
|
| 84 |
+
|
| 85 |
+
def process_features(self, features: np.ndarray, on_step=None) -> np.ndarray:
|
| 86 |
+
cfg = self.config
|
| 87 |
+
state = init_state(cfg)
|
| 88 |
+
mel_dim = features.shape[1]
|
| 89 |
+
predictions = []
|
| 90 |
+
for step, (chunk, left_offset, right_offset) in enumerate(iter_chunks(cfg, features)):
|
| 91 |
+
chunk_valid = chunk.shape[0]
|
| 92 |
+
chunk_input = np.zeros((1, self.chunk_width, mel_dim), dtype=np.float32)
|
| 93 |
+
chunk_input[0, :chunk_valid] = chunk
|
| 94 |
+
|
| 95 |
+
chunk_embs_full = self.graphs.pre_encode(chunk_input)[0]
|
| 96 |
+
chunk_embs_length = pre_encode_length(chunk_valid)
|
| 97 |
+
chunk_embs = chunk_embs_full[:chunk_embs_length]
|
| 98 |
+
|
| 99 |
+
spkcache_len = state.spkcache.shape[0]
|
| 100 |
+
fifo_len = state.fifo.shape[0]
|
| 101 |
+
total_lengths = spkcache_len + fifo_len + chunk_embs_length
|
| 102 |
+
seq_width = getattr(self.graphs, "seq_width", self.state_len + chunk_embs_full.shape[0])
|
| 103 |
+
seq = np.zeros((1, seq_width, cfg.fc_d_model), dtype=np.float32)
|
| 104 |
+
seq[0, :spkcache_len] = state.spkcache
|
| 105 |
+
seq[0, spkcache_len : spkcache_len + fifo_len] = state.fifo
|
| 106 |
+
seq[0, spkcache_len + fifo_len : total_lengths] = chunk_embs
|
| 107 |
+
|
| 108 |
+
if on_step is not None:
|
| 109 |
+
on_step(step, chunk_input, seq, total_lengths)
|
| 110 |
+
|
| 111 |
+
preds = self.graphs.encode(seq, total_lengths)[0]
|
| 112 |
+
lc_enc = round(left_offset / cfg.subsampling_factor)
|
| 113 |
+
rc_enc = int(np.ceil(right_offset / cfg.subsampling_factor))
|
| 114 |
+
state, chunk_preds = streaming_update(cfg, state, chunk_embs, preds, lc_enc, rc_enc)
|
| 115 |
+
predictions.append(chunk_preds)
|
| 116 |
+
if not predictions:
|
| 117 |
+
return np.zeros((0, cfg.num_speakers), dtype=np.float32)
|
| 118 |
+
return np.concatenate(predictions, axis=0)
|
| 119 |
+
|
| 120 |
+
def process_wav(self, waveform: np.ndarray, sample_rate: int = SAMPLE_RATE, on_step=None) -> np.ndarray:
|
| 121 |
+
features = log_mel_spectrogram(waveform, sample_rate, mel_filter=mel_filterbank())
|
| 122 |
+
return self.process_features(features, on_step=on_step)
|
| 123 |
+
|
| 124 |
+
def rttm_lines(self, preds: np.ndarray, uri: str, params: PostProcessingParams = None, bypass: bool = False):
|
| 125 |
+
timestamps = predlist_to_timestamps(preds, params=params, bypass_postprocessing=bypass)
|
| 126 |
+
num_speakers = preds.shape[1] if preds.ndim == 2 else self.config.num_speakers
|
| 127 |
+
return timestamps_to_rttm_lines(timestamps, uri, num_speakers)
|
| 128 |
+
|
| 129 |
+
|
| 130 |
+
def main() -> None:
|
| 131 |
+
parser = argparse.ArgumentParser(description=__doc__)
|
| 132 |
+
parser.add_argument("--preencode", default=None, help="preencode.onnx (host mode)")
|
| 133 |
+
parser.add_argument("--encoder", default=None, help="encoder.onnx (host mode)")
|
| 134 |
+
parser.add_argument("--preencode-axmodel", default=None, help="preencode axmodel (board mode)")
|
| 135 |
+
parser.add_argument("--encoder-axmodel", default=None, help="encoder axmodel (board mode)")
|
| 136 |
+
parser.add_argument("--wav", required=True)
|
| 137 |
+
parser.add_argument("--rttm", default=None)
|
| 138 |
+
parser.add_argument("--provider", default="auto")
|
| 139 |
+
parser.add_argument("--chunk-len", type=int, default=6)
|
| 140 |
+
parser.add_argument("--chunk-left-context", type=int, default=1)
|
| 141 |
+
parser.add_argument("--chunk-right-context", type=int, default=7)
|
| 142 |
+
parser.add_argument("--fifo-len", type=int, default=188)
|
| 143 |
+
parser.add_argument("--spkcache-len", type=int, default=188)
|
| 144 |
+
parser.add_argument("--spkcache-update-period", type=int, default=144)
|
| 145 |
+
parser.add_argument("--bypass-postproc", action="store_true")
|
| 146 |
+
args = parser.parse_args()
|
| 147 |
+
|
| 148 |
+
if args.preencode and args.encoder:
|
| 149 |
+
graphs = OnnxGraphPair(args.preencode, args.encoder, provider=args.provider)
|
| 150 |
+
elif args.preencode_axmodel and args.encoder_axmodel:
|
| 151 |
+
graphs = AxengineGraphPair(args.preencode_axmodel, args.encoder_axmodel)
|
| 152 |
+
else:
|
| 153 |
+
raise SystemExit("pass either --preencode/--encoder or the axmodel pair")
|
| 154 |
+
|
| 155 |
+
config = SortformerConfig(
|
| 156 |
+
chunk_len=args.chunk_len,
|
| 157 |
+
chunk_left_context=args.chunk_left_context,
|
| 158 |
+
chunk_right_context=args.chunk_right_context,
|
| 159 |
+
fifo_len=args.fifo_len,
|
| 160 |
+
spkcache_len=args.spkcache_len,
|
| 161 |
+
spkcache_update_period=args.spkcache_update_period,
|
| 162 |
+
)
|
| 163 |
+
diarizer = StreamingDiarizer(graphs, config)
|
| 164 |
+
|
| 165 |
+
import soundfile as sf
|
| 166 |
+
|
| 167 |
+
waveform, sample_rate = sf.read(args.wav, dtype="float32", always_2d=True)
|
| 168 |
+
waveform = waveform.mean(axis=1)
|
| 169 |
+
preds = diarizer.process_wav(waveform, sample_rate)
|
| 170 |
+
uri = Path(args.wav).stem
|
| 171 |
+
lines = diarizer.rttm_lines(preds, uri, bypass=args.bypass_postproc)
|
| 172 |
+
print(f"{uri}: {len(preds)} frames, {len(lines)} segments")
|
| 173 |
+
if args.rttm:
|
| 174 |
+
Path(args.rttm).write_text("\n".join(lines) + "\n")
|
| 175 |
+
print(f"rttm -> {args.rttm}")
|
| 176 |
+
|
| 177 |
+
|
| 178 |
+
if __name__ == "__main__":
|
| 179 |
+
main()
|
python/sortformer_sdk/feature.py
ADDED
|
@@ -0,0 +1,128 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""NeMo-compatible log-mel front-end (numpy only, no torch/librosa at runtime).
|
| 2 |
+
|
| 3 |
+
Mirrors ``FilterbankFeatures`` (eval mode) from NeMo Speech 3.0 with the
|
| 4 |
+
Sortformer streaming configuration:
|
| 5 |
+
|
| 6 |
+
sample_rate=16000, n_fft=512, win_length=400 (hann, periodic=False),
|
| 7 |
+
hop_length=160, n_mels=128, preemph=0.97, mag_power=2.0,
|
| 8 |
+
log(x + 2**-24), normalize="NA", pad_to=16
|
| 9 |
+
|
| 10 |
+
``log_mel_spectrogram`` returns ``(num_frames, 128)`` float32 where
|
| 11 |
+
``num_frames = floor(num_samples / hop)`` rounded up to a multiple of 16.
|
| 12 |
+
"""
|
| 13 |
+
|
| 14 |
+
import numpy as np
|
| 15 |
+
|
| 16 |
+
SAMPLE_RATE = 16000
|
| 17 |
+
N_FFT = 512
|
| 18 |
+
WIN_LENGTH = 400
|
| 19 |
+
HOP_LENGTH = 160
|
| 20 |
+
N_MELS = 128
|
| 21 |
+
PREEMPH = 0.97
|
| 22 |
+
LOG_ZERO_GUARD = 2.0 ** -24
|
| 23 |
+
PAD_TO = 16
|
| 24 |
+
|
| 25 |
+
_MEL_CACHE = {}
|
| 26 |
+
|
| 27 |
+
|
| 28 |
+
def _hz_to_mel(freq):
|
| 29 |
+
# Slaney mel scale (librosa default, htk=False)
|
| 30 |
+
f_sp = 200.0 / 3.0
|
| 31 |
+
log_step = np.log(6.4) / 27.0
|
| 32 |
+
min_log_hz = 1000.0
|
| 33 |
+
min_log_mel = min_log_hz / f_sp
|
| 34 |
+
freq = np.asarray(freq, dtype=np.float64)
|
| 35 |
+
mel_linear = freq / f_sp
|
| 36 |
+
mel_log = min_log_mel + np.log(np.maximum(freq, min_log_hz) / min_log_hz) / log_step
|
| 37 |
+
return np.where(freq >= min_log_hz, mel_log, mel_linear)
|
| 38 |
+
|
| 39 |
+
|
| 40 |
+
def _mel_to_hz(mel):
|
| 41 |
+
f_sp = 200.0 / 3.0
|
| 42 |
+
log_step = np.log(6.4) / 27.0
|
| 43 |
+
min_log_hz = 1000.0
|
| 44 |
+
min_log_mel = min_log_hz / f_sp
|
| 45 |
+
mel = np.asarray(mel, dtype=np.float64)
|
| 46 |
+
hz_linear = f_sp * mel
|
| 47 |
+
hz_log = min_log_hz * np.exp(log_step * (mel - min_log_mel))
|
| 48 |
+
return np.where(mel >= min_log_mel, hz_log, hz_linear)
|
| 49 |
+
|
| 50 |
+
|
| 51 |
+
def mel_filterbank(sample_rate: int = SAMPLE_RATE, n_fft: int = N_FFT, n_mels: int = N_MELS):
|
| 52 |
+
"""librosa-compatible Slaney-normalized Slaney-scale mel filterbank."""
|
| 53 |
+
key = (sample_rate, n_fft, n_mels)
|
| 54 |
+
if key in _MEL_CACHE:
|
| 55 |
+
return _MEL_CACHE[key]
|
| 56 |
+
|
| 57 |
+
fmin, fmax = 0.0, sample_rate / 2.0
|
| 58 |
+
mels = np.linspace(_hz_to_mel(fmin), _hz_to_mel(fmax), n_mels + 2)
|
| 59 |
+
hz = _mel_to_hz(mels)
|
| 60 |
+
freqs = np.linspace(0.0, sample_rate / 2.0, 1 + n_fft // 2)
|
| 61 |
+
|
| 62 |
+
fdiff = np.diff(hz)
|
| 63 |
+
ramps = np.subtract.outer(hz, freqs)
|
| 64 |
+
lower = -ramps[np.arange(n_mels), :] / fdiff[np.arange(n_mels)][:, None]
|
| 65 |
+
upper = ramps[np.arange(2, n_mels + 2), :] / fdiff[np.arange(1, n_mels + 1)][:, None]
|
| 66 |
+
weights = np.maximum(0.0, np.minimum(lower, upper))
|
| 67 |
+
|
| 68 |
+
enorm = 2.0 / (hz[2 : n_mels + 2] - hz[:n_mels])
|
| 69 |
+
weights *= enorm[:, None]
|
| 70 |
+
result = weights.astype(np.float64)
|
| 71 |
+
_MEL_CACHE[key] = result
|
| 72 |
+
return result
|
| 73 |
+
|
| 74 |
+
|
| 75 |
+
def _hann_window(length: int):
|
| 76 |
+
return 0.5 - 0.5 * np.cos(2.0 * np.pi * np.arange(length) / (length - 1))
|
| 77 |
+
|
| 78 |
+
|
| 79 |
+
def log_mel_spectrogram(
|
| 80 |
+
waveform: np.ndarray,
|
| 81 |
+
sample_rate: int = SAMPLE_RATE,
|
| 82 |
+
*,
|
| 83 |
+
preemph: float = PREEMPH,
|
| 84 |
+
pad_to: int = PAD_TO,
|
| 85 |
+
mel_filter: np.ndarray = None,
|
| 86 |
+
) -> np.ndarray:
|
| 87 |
+
"""Compute the NeMo Sortformer streaming log-mel features for a mono waveform."""
|
| 88 |
+
if sample_rate != SAMPLE_RATE:
|
| 89 |
+
raise ValueError(f"expected {SAMPLE_RATE} Hz input, got {sample_rate}")
|
| 90 |
+
waveform = np.asarray(waveform, dtype=np.float32).reshape(-1)
|
| 91 |
+
|
| 92 |
+
if preemph is not None and waveform.size > 0:
|
| 93 |
+
emphasized = np.empty_like(waveform)
|
| 94 |
+
emphasized[0] = waveform[0]
|
| 95 |
+
emphasized[1:] = waveform[1:] - preemph * waveform[:-1]
|
| 96 |
+
waveform = emphasized
|
| 97 |
+
|
| 98 |
+
padded = np.pad(waveform, (N_FFT // 2, N_FFT // 2), mode="constant")
|
| 99 |
+
num_samples = waveform.shape[0]
|
| 100 |
+
num_frames = num_samples // HOP_LENGTH
|
| 101 |
+
if num_frames == 0:
|
| 102 |
+
return np.zeros((0, N_MELS), dtype=np.float32)
|
| 103 |
+
|
| 104 |
+
window = _hann_window(WIN_LENGTH).astype(np.float32)
|
| 105 |
+
left_pad = (N_FFT - WIN_LENGTH) // 2
|
| 106 |
+
frame_starts = np.arange(num_frames) * HOP_LENGTH
|
| 107 |
+
frames = np.lib.stride_tricks.as_strided(
|
| 108 |
+
padded,
|
| 109 |
+
shape=(num_frames, N_FFT),
|
| 110 |
+
strides=(padded.strides[0] * HOP_LENGTH, padded.strides[0]),
|
| 111 |
+
writeable=False,
|
| 112 |
+
).copy()
|
| 113 |
+
frames[:, :left_pad] = 0.0
|
| 114 |
+
frames[:, left_pad : left_pad + WIN_LENGTH] *= window
|
| 115 |
+
frames[:, left_pad + WIN_LENGTH :] = 0.0
|
| 116 |
+
|
| 117 |
+
spectrum = np.fft.rfft(frames, n=N_FFT, axis=1)
|
| 118 |
+
magnitude = np.abs(spectrum).astype(np.float32) ** 2.0
|
| 119 |
+
|
| 120 |
+
if mel_filter is None:
|
| 121 |
+
mel_filter = mel_filterbank()
|
| 122 |
+
mel = magnitude @ mel_filter.T.astype(np.float32)
|
| 123 |
+
log_mel = np.log(mel + LOG_ZERO_GUARD, dtype=np.float32)
|
| 124 |
+
|
| 125 |
+
if pad_to and log_mel.shape[0] % pad_to:
|
| 126 |
+
pad = pad_to - log_mel.shape[0] % pad_to
|
| 127 |
+
log_mel = np.pad(log_mel, ((0, pad), (0, 0)), mode="constant")
|
| 128 |
+
return log_mel.astype(np.float32)
|
python/sortformer_sdk/postprocess.py
ADDED
|
@@ -0,0 +1,181 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""Numpy port of NeMo's Sortformer/VAD post-processing (preds -> RTTM).
|
| 2 |
+
|
| 3 |
+
Mirrors ``predlist_to_timestamps`` / ``binarization_vectorized`` / ``filtering``
|
| 4 |
+
from ``nemo/collections/asr/parts/utils/vad_utils.py`` and
|
| 5 |
+
``generate_diarization_output_lines`` from ``speaker_utils.py`` (NeMo Speech 3.0).
|
| 6 |
+
"""
|
| 7 |
+
|
| 8 |
+
from dataclasses import dataclass
|
| 9 |
+
|
| 10 |
+
import numpy as np
|
| 11 |
+
|
| 12 |
+
FRAME_LENGTH_IN_SEC = 0.08
|
| 13 |
+
|
| 14 |
+
|
| 15 |
+
@dataclass
|
| 16 |
+
class PostProcessingParams:
|
| 17 |
+
onset: float = 0.5
|
| 18 |
+
offset: float = 0.5
|
| 19 |
+
pad_onset: float = 0.0
|
| 20 |
+
pad_offset: float = 0.0
|
| 21 |
+
min_duration_on: float = 0.0
|
| 22 |
+
min_duration_off: float = 0.0
|
| 23 |
+
filter_speech_first: float = 1.0
|
| 24 |
+
|
| 25 |
+
|
| 26 |
+
def merge_overlap_segment(segments: np.ndarray) -> np.ndarray:
|
| 27 |
+
if segments.shape[0] <= 1:
|
| 28 |
+
return segments
|
| 29 |
+
segments = segments[np.argsort(segments[:, 0])]
|
| 30 |
+
merge_boundary = segments[:-1, 1] >= segments[1:, 0]
|
| 31 |
+
head_padded = np.concatenate([[False], merge_boundary])
|
| 32 |
+
tail_padded = np.concatenate([merge_boundary, [False]])
|
| 33 |
+
head = segments[~head_padded, 0]
|
| 34 |
+
tail = segments[~tail_padded, 1]
|
| 35 |
+
return np.stack([head, tail], axis=1)
|
| 36 |
+
|
| 37 |
+
|
| 38 |
+
def filter_short_segments(segments: np.ndarray, threshold: float) -> np.ndarray:
|
| 39 |
+
return segments[segments[:, 1] - segments[:, 0] >= threshold]
|
| 40 |
+
|
| 41 |
+
|
| 42 |
+
def get_gap_segments(segments: np.ndarray) -> np.ndarray:
|
| 43 |
+
segments = segments[np.argsort(segments[:, 0])]
|
| 44 |
+
return np.column_stack((segments[:-1, 1], segments[1:, 0]))
|
| 45 |
+
|
| 46 |
+
|
| 47 |
+
def remove_segments(original_segments: np.ndarray, to_be_removed: np.ndarray) -> np.ndarray:
|
| 48 |
+
keep = np.ones(original_segments.shape[0], dtype=bool)
|
| 49 |
+
for segment in to_be_removed:
|
| 50 |
+
keep &= ~(original_segments == segment).all(axis=1)
|
| 51 |
+
return original_segments[keep]
|
| 52 |
+
|
| 53 |
+
|
| 54 |
+
def filtering(segments: np.ndarray, params: PostProcessingParams) -> np.ndarray:
|
| 55 |
+
if segments.shape[0] == 0:
|
| 56 |
+
return segments
|
| 57 |
+
|
| 58 |
+
def filter_speech():
|
| 59 |
+
if params.min_duration_on > 0:
|
| 60 |
+
return filter_short_segments(segments, params.min_duration_on)
|
| 61 |
+
return segments
|
| 62 |
+
|
| 63 |
+
def restore_short_gaps(current: np.ndarray) -> np.ndarray:
|
| 64 |
+
if params.min_duration_off <= 0 or current.shape[0] == 0:
|
| 65 |
+
return current
|
| 66 |
+
non_speech = get_gap_segments(current)
|
| 67 |
+
short_gaps = remove_segments(non_speech, filter_short_segments(non_speech, params.min_duration_off))
|
| 68 |
+
if short_gaps.shape[0] == 0:
|
| 69 |
+
return current
|
| 70 |
+
return merge_overlap_segment(np.concatenate([current, short_gaps], axis=0))
|
| 71 |
+
|
| 72 |
+
if params.filter_speech_first == 1.0:
|
| 73 |
+
segments = filter_speech()
|
| 74 |
+
segments = restore_short_gaps(segments)
|
| 75 |
+
else:
|
| 76 |
+
segments = restore_short_gaps(segments)
|
| 77 |
+
segments = filter_speech()
|
| 78 |
+
return segments
|
| 79 |
+
|
| 80 |
+
|
| 81 |
+
def binarization_vectorized(sequence: np.ndarray, params: PostProcessingParams) -> np.ndarray:
|
| 82 |
+
empty = np.empty((0, 2), dtype=np.float32)
|
| 83 |
+
num_frames = sequence.shape[0]
|
| 84 |
+
if num_frames == 0:
|
| 85 |
+
return empty
|
| 86 |
+
|
| 87 |
+
positions = np.arange(1, num_frames + 1)
|
| 88 |
+
onset, offset = params.onset, params.offset
|
| 89 |
+
|
| 90 |
+
if onset >= offset:
|
| 91 |
+
force_on = sequence > onset
|
| 92 |
+
force_off = sequence < offset
|
| 93 |
+
has_event = force_on | force_off
|
| 94 |
+
event_positions = np.where(has_event, positions, 0)
|
| 95 |
+
last_event_positions = np.maximum.accumulate(event_positions)
|
| 96 |
+
event_states = np.concatenate([[False], force_on])
|
| 97 |
+
above = event_states[last_event_positions]
|
| 98 |
+
else:
|
| 99 |
+
force_on = sequence >= offset
|
| 100 |
+
force_off = sequence <= onset
|
| 101 |
+
toggle = (sequence > onset) & (sequence < offset)
|
| 102 |
+
|
| 103 |
+
has_reset = force_on | force_off
|
| 104 |
+
reset_positions = np.where(has_reset, positions, 0)
|
| 105 |
+
last_reset_positions = np.maximum.accumulate(reset_positions)
|
| 106 |
+
|
| 107 |
+
reset_states = np.concatenate([[0], force_on.astype(np.int64)])
|
| 108 |
+
base_state = reset_states[last_reset_positions]
|
| 109 |
+
toggle_prefix = np.concatenate([[0], np.cumsum(toggle.astype(np.int64))])
|
| 110 |
+
toggles_since_reset = toggle_prefix[positions] - toggle_prefix[last_reset_positions]
|
| 111 |
+
above = np.logical_xor(base_state.astype(bool), (toggles_since_reset % 2).astype(bool))
|
| 112 |
+
|
| 113 |
+
padded = np.pad(above.astype(np.float32), (1, 1))
|
| 114 |
+
diff = padded[1:] - padded[:-1]
|
| 115 |
+
starts = np.where(diff > 0.5)[0]
|
| 116 |
+
ends = np.where(diff < -0.5)[0]
|
| 117 |
+
if starts.shape[0] == 0:
|
| 118 |
+
return empty
|
| 119 |
+
|
| 120 |
+
start_times = np.clip(starts.astype(np.float32) * FRAME_LENGTH_IN_SEC - params.pad_onset, 0.0, None)
|
| 121 |
+
end_times = ends.astype(np.float32) * FRAME_LENGTH_IN_SEC + params.pad_offset
|
| 122 |
+
valid = end_times > start_times
|
| 123 |
+
if not valid.any():
|
| 124 |
+
return empty
|
| 125 |
+
segments = np.stack([start_times[valid], end_times[valid]], axis=1)
|
| 126 |
+
if params.pad_onset > 0 or params.pad_offset > 0:
|
| 127 |
+
segments = merge_overlap_segment(segments)
|
| 128 |
+
return segments
|
| 129 |
+
|
| 130 |
+
|
| 131 |
+
def predlist_to_timestamps(
|
| 132 |
+
preds: np.ndarray,
|
| 133 |
+
offset: float = 0.0,
|
| 134 |
+
params: PostProcessingParams = None,
|
| 135 |
+
bypass_postprocessing: bool = False,
|
| 136 |
+
precision: int = 2,
|
| 137 |
+
):
|
| 138 |
+
"""Convert (num_frames, num_speakers) probabilities to per-speaker timestamps."""
|
| 139 |
+
if params is None:
|
| 140 |
+
params = PostProcessingParams()
|
| 141 |
+
if bypass_postprocessing:
|
| 142 |
+
params = PostProcessingParams(onset=0.5, offset=0.5)
|
| 143 |
+
|
| 144 |
+
timestamps = []
|
| 145 |
+
for spk in range(preds.shape[1]):
|
| 146 |
+
segments = binarization_vectorized(preds[:, spk], params)
|
| 147 |
+
if not bypass_postprocessing:
|
| 148 |
+
segments = filtering(segments, params)
|
| 149 |
+
if segments.shape[0] == 0:
|
| 150 |
+
timestamps.append([])
|
| 151 |
+
continue
|
| 152 |
+
segments = segments + offset
|
| 153 |
+
timestamps.append([[round(float(start), precision), round(float(end), precision)] for start, end in segments])
|
| 154 |
+
return timestamps
|
| 155 |
+
|
| 156 |
+
|
| 157 |
+
def generate_diarization_output_lines(timestamps, model_spk_num: int):
|
| 158 |
+
lines = []
|
| 159 |
+
for spk_idx in range(model_spk_num):
|
| 160 |
+
if not timestamps[spk_idx]:
|
| 161 |
+
continue
|
| 162 |
+
intervals = np.asarray(timestamps[spk_idx], dtype=np.float32).reshape(-1, 2)
|
| 163 |
+
for start, end in merge_overlap_segment(intervals):
|
| 164 |
+
lines.append(f"{start:.3f} {end:.3f} speaker_{int(spk_idx)}")
|
| 165 |
+
return lines
|
| 166 |
+
|
| 167 |
+
|
| 168 |
+
def timestamps_to_rttm_lines(timestamps, uri: str, model_spk_num: int):
|
| 169 |
+
lines = []
|
| 170 |
+
for spk_idx in range(model_spk_num):
|
| 171 |
+
intervals = timestamps[spk_idx]
|
| 172 |
+
if not intervals:
|
| 173 |
+
continue
|
| 174 |
+
merged = merge_overlap_segment(np.asarray(intervals, dtype=np.float32).reshape(-1, 2))
|
| 175 |
+
for start, end in merged:
|
| 176 |
+
duration = float(end) - float(start)
|
| 177 |
+
if duration > 0:
|
| 178 |
+
lines.append(
|
| 179 |
+
f"SPEAKER {uri} 1 {float(start):.3f} {duration:.3f} <NA> <NA> speaker_{int(spk_idx)} <NA>"
|
| 180 |
+
)
|
| 181 |
+
return lines
|
python/sortformer_sdk/state.py
ADDED
|
@@ -0,0 +1,210 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""Numpy port of NeMo's synchronous Sortformer streaming state machine.
|
| 2 |
+
|
| 3 |
+
Mirrors ``SortformerModules.streaming_update`` (eval, batch size 1, no speaker
|
| 4 |
+
permutation, learnable silence embedding disabled) from NeMo Speech 3.0:
|
| 5 |
+
|
| 6 |
+
nemo/collections/asr/modules/sortformer_modules.py
|
| 7 |
+
|
| 8 |
+
The module is dependency-free so it can run on the host (base env) and be
|
| 9 |
+
reused by the AX650 board SDK.
|
| 10 |
+
"""
|
| 11 |
+
|
| 12 |
+
import math
|
| 13 |
+
from dataclasses import dataclass, field
|
| 14 |
+
|
| 15 |
+
import numpy as np
|
| 16 |
+
|
| 17 |
+
NEG_INF = float("-inf")
|
| 18 |
+
|
| 19 |
+
|
| 20 |
+
@dataclass
|
| 21 |
+
class SortformerConfig:
|
| 22 |
+
num_speakers: int = 4
|
| 23 |
+
fc_d_model: int = 512
|
| 24 |
+
subsampling_factor: int = 8
|
| 25 |
+
chunk_len: int = 6
|
| 26 |
+
chunk_left_context: int = 1
|
| 27 |
+
chunk_right_context: int = 7
|
| 28 |
+
fifo_len: int = 188
|
| 29 |
+
spkcache_len: int = 188
|
| 30 |
+
spkcache_update_period: int = 144
|
| 31 |
+
spkcache_sil_frames_per_spk: int = 3
|
| 32 |
+
pred_score_threshold: float = 0.25
|
| 33 |
+
max_index: int = 99999
|
| 34 |
+
scores_boost_latest: float = 0.05
|
| 35 |
+
sil_threshold: float = 0.2
|
| 36 |
+
strong_boost_rate: float = 0.75
|
| 37 |
+
weak_boost_rate: float = 1.5
|
| 38 |
+
min_pos_scores_rate: float = 0.5
|
| 39 |
+
use_learnable_sil_emb: bool = False
|
| 40 |
+
|
| 41 |
+
|
| 42 |
+
@dataclass
|
| 43 |
+
class StreamingState:
|
| 44 |
+
spkcache: np.ndarray = field(default_factory=lambda: np.zeros((0, 512), dtype=np.float32))
|
| 45 |
+
spkcache_preds: np.ndarray = field(default_factory=lambda: np.zeros((0, 4), dtype=np.float32))
|
| 46 |
+
spkcache_compressed: bool = False
|
| 47 |
+
fifo: np.ndarray = field(default_factory=lambda: np.zeros((0, 512), dtype=np.float32))
|
| 48 |
+
fifo_preds: np.ndarray = field(default_factory=lambda: np.zeros((0, 4), dtype=np.float32))
|
| 49 |
+
mean_sil_emb: np.ndarray = field(default_factory=lambda: np.zeros(512, dtype=np.float32))
|
| 50 |
+
n_sil_frames: int = 0
|
| 51 |
+
|
| 52 |
+
|
| 53 |
+
def init_state(cfg: SortformerConfig) -> StreamingState:
|
| 54 |
+
state = StreamingState()
|
| 55 |
+
state.mean_sil_emb = np.zeros(cfg.fc_d_model, dtype=np.float32)
|
| 56 |
+
return state
|
| 57 |
+
|
| 58 |
+
|
| 59 |
+
def streaming_update(cfg: SortformerConfig, state: StreamingState, chunk: np.ndarray, preds: np.ndarray, lc: int, rc: int):
|
| 60 |
+
"""Update speaker cache / FIFO with one chunk; returns the chunk predictions.
|
| 61 |
+
|
| 62 |
+
``chunk`` has shape (lc + chunk_len + rc, emb_dim); ``preds`` has shape
|
| 63 |
+
(spkcache_len + fifo_len + chunk.shape[0], num_speakers) and covers the
|
| 64 |
+
speaker cache, FIFO and chunk regions.
|
| 65 |
+
"""
|
| 66 |
+
spkcache_len = state.spkcache.shape[0]
|
| 67 |
+
fifo_len = state.fifo.shape[0]
|
| 68 |
+
chunk_len = chunk.shape[0] - lc - rc
|
| 69 |
+
|
| 70 |
+
state.fifo_preds = preds[spkcache_len : spkcache_len + fifo_len]
|
| 71 |
+
chunk_body = chunk[lc : chunk_len + lc]
|
| 72 |
+
chunk_preds = preds[spkcache_len + fifo_len + lc : spkcache_len + fifo_len + chunk_len + lc]
|
| 73 |
+
|
| 74 |
+
state.fifo = np.concatenate([state.fifo, chunk_body], axis=0)
|
| 75 |
+
state.fifo_preds = np.concatenate([state.fifo_preds, chunk_preds], axis=0)
|
| 76 |
+
|
| 77 |
+
if fifo_len + chunk_len > cfg.fifo_len:
|
| 78 |
+
pop_out_len = cfg.spkcache_update_period
|
| 79 |
+
pop_out_len = max(pop_out_len, chunk_len - cfg.fifo_len + fifo_len)
|
| 80 |
+
pop_out_len = min(pop_out_len, fifo_len + chunk_len)
|
| 81 |
+
|
| 82 |
+
pop_out_embs = state.fifo[:pop_out_len]
|
| 83 |
+
pop_out_preds = state.fifo_preds[:pop_out_len]
|
| 84 |
+
if not cfg.use_learnable_sil_emb:
|
| 85 |
+
state.mean_sil_emb, state.n_sil_frames = _get_silence_profile(
|
| 86 |
+
cfg, state.mean_sil_emb, state.n_sil_frames, pop_out_embs, pop_out_preds
|
| 87 |
+
)
|
| 88 |
+
state.fifo = state.fifo[pop_out_len:]
|
| 89 |
+
state.fifo_preds = state.fifo_preds[pop_out_len:]
|
| 90 |
+
|
| 91 |
+
state.spkcache = np.concatenate([state.spkcache, pop_out_embs], axis=0)
|
| 92 |
+
if state.spkcache_compressed:
|
| 93 |
+
state.spkcache_preds = np.concatenate([state.spkcache_preds, pop_out_preds], axis=0)
|
| 94 |
+
else:
|
| 95 |
+
state.spkcache_preds = np.concatenate([preds[:spkcache_len], pop_out_preds], axis=0)
|
| 96 |
+
if state.spkcache.shape[0] > cfg.spkcache_len:
|
| 97 |
+
state.spkcache, state.spkcache_preds = _compress_spkcache(
|
| 98 |
+
cfg, state.spkcache, state.spkcache_preds, state.mean_sil_emb
|
| 99 |
+
)
|
| 100 |
+
state.spkcache_compressed = True
|
| 101 |
+
|
| 102 |
+
return state, chunk_preds
|
| 103 |
+
|
| 104 |
+
|
| 105 |
+
def _get_silence_profile(cfg, mean_sil_emb, n_sil_frames, emb_seq, preds):
|
| 106 |
+
is_sil = preds.sum(axis=1) < cfg.sil_threshold
|
| 107 |
+
sil_count = int(is_sil.sum())
|
| 108 |
+
if sil_count == 0:
|
| 109 |
+
return mean_sil_emb, n_sil_frames
|
| 110 |
+
sil_emb_sum = (emb_seq * is_sil[:, None]).sum(axis=0)
|
| 111 |
+
upd_n_sil_frames = n_sil_frames + sil_count
|
| 112 |
+
total_sil_sum = mean_sil_emb * n_sil_frames + sil_emb_sum
|
| 113 |
+
upd_mean_sil_emb = total_sil_sum / max(upd_n_sil_frames, 1)
|
| 114 |
+
return upd_mean_sil_emb.astype(np.float32), upd_n_sil_frames
|
| 115 |
+
|
| 116 |
+
|
| 117 |
+
def _get_log_pred_scores(cfg, preds):
|
| 118 |
+
log_probs = np.log(np.clip(preds, cfg.pred_score_threshold, None))
|
| 119 |
+
log_1_probs = np.log(np.clip(1.0 - preds, cfg.pred_score_threshold, None))
|
| 120 |
+
log_1_probs_sum = log_1_probs.sum(axis=1, keepdims=True)
|
| 121 |
+
return log_probs - log_1_probs + log_1_probs_sum - math.log(0.5)
|
| 122 |
+
|
| 123 |
+
|
| 124 |
+
def _disable_low_scores(cfg, preds, scores, min_pos_scores_per_spk):
|
| 125 |
+
is_speech = preds > 0.5
|
| 126 |
+
scores = np.where(is_speech, scores, NEG_INF)
|
| 127 |
+
is_pos = scores > 0
|
| 128 |
+
# NeMo sums over the frame dimension (torch: is_pos.sum(dim=1)); batch-free here -> axis=0.
|
| 129 |
+
is_nonpos_replace = (~is_pos) & is_speech & (is_pos.sum(axis=0, keepdims=True) >= min_pos_scores_per_spk)
|
| 130 |
+
return np.where(is_nonpos_replace, NEG_INF, scores)
|
| 131 |
+
|
| 132 |
+
|
| 133 |
+
def _boost_topk_scores(cfg, scores, n_boost_per_spk, scale_factor=1.0, offset=0.5):
|
| 134 |
+
n_frames, n_spk = scores.shape
|
| 135 |
+
n_boost_per_spk = min(n_boost_per_spk, n_frames)
|
| 136 |
+
if n_boost_per_spk <= 0:
|
| 137 |
+
return scores
|
| 138 |
+
for spk in range(n_spk):
|
| 139 |
+
column = scores[:, spk]
|
| 140 |
+
# Stable descending order: ties keep the smaller frame index (matches C++).
|
| 141 |
+
order = np.argsort(-column, kind="stable")[:n_boost_per_spk]
|
| 142 |
+
scores[order, spk] -= scale_factor * math.log(offset)
|
| 143 |
+
return scores
|
| 144 |
+
|
| 145 |
+
|
| 146 |
+
def _get_topk_indices(cfg, scores):
|
| 147 |
+
n_frames, n_spk = scores.shape
|
| 148 |
+
n_frames_no_sil = n_frames - cfg.spkcache_sil_frames_per_spk
|
| 149 |
+
scores_flatten = scores.T.reshape(-1) # speaker-major, matches permute(0, 2, 1).reshape()
|
| 150 |
+
# Stable descending order: ties keep the smaller flat index (matches C++).
|
| 151 |
+
order = np.argsort(-scores_flatten, kind="stable")
|
| 152 |
+
k = min(cfg.spkcache_len, scores_flatten.shape[0])
|
| 153 |
+
topk_indices = order[:k]
|
| 154 |
+
values = scores_flatten[topk_indices]
|
| 155 |
+
topk_indices = np.where(values != NEG_INF, topk_indices, cfg.max_index)
|
| 156 |
+
topk_indices_sorted = np.sort(topk_indices)
|
| 157 |
+
is_disabled = topk_indices_sorted == cfg.max_index
|
| 158 |
+
topk_indices_sorted = np.remainder(topk_indices_sorted, n_frames)
|
| 159 |
+
is_disabled = is_disabled | (topk_indices_sorted >= n_frames_no_sil)
|
| 160 |
+
topk_indices_sorted = np.where(is_disabled, 0, topk_indices_sorted)
|
| 161 |
+
return topk_indices_sorted, is_disabled
|
| 162 |
+
|
| 163 |
+
|
| 164 |
+
def _compress_spkcache(cfg, emb_seq, preds, mean_sil_emb):
|
| 165 |
+
n_frames, n_spk = preds.shape
|
| 166 |
+
spkcache_len_per_spk = cfg.spkcache_len // n_spk - cfg.spkcache_sil_frames_per_spk
|
| 167 |
+
strong_boost_per_spk = math.floor(spkcache_len_per_spk * cfg.strong_boost_rate)
|
| 168 |
+
weak_boost_per_spk = math.floor(spkcache_len_per_spk * cfg.weak_boost_rate)
|
| 169 |
+
min_pos_scores_per_spk = math.floor(spkcache_len_per_spk * cfg.min_pos_scores_rate)
|
| 170 |
+
|
| 171 |
+
scores = _get_log_pred_scores(cfg, preds)
|
| 172 |
+
scores = _disable_low_scores(cfg, preds, scores, min_pos_scores_per_spk)
|
| 173 |
+
if cfg.scores_boost_latest > 0:
|
| 174 |
+
scores[cfg.spkcache_len :, :] += cfg.scores_boost_latest
|
| 175 |
+
scores = _boost_topk_scores(cfg, scores, strong_boost_per_spk, scale_factor=2)
|
| 176 |
+
scores = _boost_topk_scores(cfg, scores, weak_boost_per_spk, scale_factor=1)
|
| 177 |
+
|
| 178 |
+
if cfg.spkcache_sil_frames_per_spk > 0:
|
| 179 |
+
pad = np.full((cfg.spkcache_sil_frames_per_spk, n_spk), np.inf, dtype=scores.dtype)
|
| 180 |
+
scores = np.concatenate([scores, pad], axis=0)
|
| 181 |
+
|
| 182 |
+
topk_indices, is_disabled = _get_topk_indices(cfg, scores)
|
| 183 |
+
spkcache = emb_seq[topk_indices]
|
| 184 |
+
spkcache = np.where(is_disabled[:, None], mean_sil_emb[None, :], spkcache)
|
| 185 |
+
spkcache_preds = preds[topk_indices]
|
| 186 |
+
spkcache_preds = np.where(is_disabled[:, None], 0.0, spkcache_preds)
|
| 187 |
+
return spkcache.astype(np.float32), spkcache_preds.astype(np.float32)
|
| 188 |
+
|
| 189 |
+
|
| 190 |
+
def iter_chunks(cfg: SortformerConfig, features: np.ndarray):
|
| 191 |
+
"""Yield (chunk_mel, left_offset, right_offset) following NeMo's ``streaming_feat_loader``."""
|
| 192 |
+
feat_len = features.shape[0]
|
| 193 |
+
start = 0
|
| 194 |
+
while start < feat_len:
|
| 195 |
+
left_offset = min(cfg.chunk_left_context * cfg.subsampling_factor, start)
|
| 196 |
+
end = min(start + cfg.chunk_len * cfg.subsampling_factor, feat_len)
|
| 197 |
+
right_offset = min(cfg.chunk_right_context * cfg.subsampling_factor, feat_len - end)
|
| 198 |
+
chunk = features[start - left_offset : end + right_offset]
|
| 199 |
+
yield chunk, left_offset, right_offset
|
| 200 |
+
start = end
|
| 201 |
+
|
| 202 |
+
|
| 203 |
+
def pre_encode_length(mel_frames: int, num_layers: int = 3) -> int:
|
| 204 |
+
"""Number of encoder frames after NeMo's dw_striding pre-encode (stride 2, kernel 3)."""
|
| 205 |
+
length = int(mel_frames)
|
| 206 |
+
for _ in range(num_layers):
|
| 207 |
+
if length <= 0:
|
| 208 |
+
return 0
|
| 209 |
+
length = (length - 1) // 2 + 1
|
| 210 |
+
return length
|
requirements.txt
ADDED
|
@@ -0,0 +1,4 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
numpy>=1.21
|
| 2 |
+
soundfile>=0.12
|
| 3 |
+
# 板端 NPU 运行库(板端镜像已内置;主机安装见 README)
|
| 4 |
+
# axengine
|
run_ax650.sh
ADDED
|
@@ -0,0 +1,18 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
#!/usr/bin/env bash
|
| 2 |
+
# Sortformer 一键说话人日志 (Python, axengine)
|
| 3 |
+
# 用法: ./run_ax650.sh [input.wav] [output.rttm] [--fast]
|
| 4 |
+
set -euo pipefail
|
| 5 |
+
cd "$(dirname "$0")"
|
| 6 |
+
|
| 7 |
+
export LD_LIBRARY_PATH=/soc/lib:${LD_LIBRARY_PATH:-}
|
| 8 |
+
|
| 9 |
+
WAV="${1:-samples/sample_meeting.wav}"
|
| 10 |
+
RTTM="${2:-$(basename "${WAV%.*}").rttm}"
|
| 11 |
+
EXTRA=("${@:3}")
|
| 12 |
+
|
| 13 |
+
PYTHON_BIN="${PYTHON_BIN:-python3}"
|
| 14 |
+
if [[ -x /root/miniforge3/bin/python ]]; then
|
| 15 |
+
PYTHON_BIN=/root/miniforge3/bin/python
|
| 16 |
+
fi
|
| 17 |
+
|
| 18 |
+
"$PYTHON_BIN" python/example.py "$WAV" "$RTTM" "${EXTRA[@]}"
|
run_cpp_ax650.sh
ADDED
|
@@ -0,0 +1,26 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
#!/usr/bin/env bash
|
| 2 |
+
# Sortformer 一键说话人日志 (C++, 预编译可执行文件)
|
| 3 |
+
# 用法: ./run_cpp_ax650.sh [input.wav] [output.rttm] [--fast]
|
| 4 |
+
set -euo pipefail
|
| 5 |
+
cd "$(dirname "$0")"
|
| 6 |
+
|
| 7 |
+
export LD_LIBRARY_PATH=/soc/lib:${LD_LIBRARY_PATH:-}
|
| 8 |
+
|
| 9 |
+
WAV="${1:-samples/sample_meeting.wav}"
|
| 10 |
+
RTTM="${2:-$(basename "${WAV%.*}").rttm}"
|
| 11 |
+
FAST=0
|
| 12 |
+
for arg in "$@"; do
|
| 13 |
+
[[ "$arg" == "--fast" ]] && FAST=1
|
| 14 |
+
done
|
| 15 |
+
|
| 16 |
+
if [[ "$FAST" == "1" ]]; then
|
| 17 |
+
ENCODER=models/encoder_fifo40.axmodel
|
| 18 |
+
EXTRA="--fifo-len 40 --spkcache-update-period 31"
|
| 19 |
+
else
|
| 20 |
+
ENCODER=models/encoder.axmodel
|
| 21 |
+
EXTRA=""
|
| 22 |
+
fi
|
| 23 |
+
|
| 24 |
+
./bin/sortformer_ax650 \
|
| 25 |
+
--preencode models/preencode.axmodel --encoder "$ENCODER" \
|
| 26 |
+
--wav "$WAV" --rttm "$RTTM" --threads 8 $EXTRA
|
samples/sample_meeting.wav
ADDED
|
@@ -0,0 +1,3 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
version https://git-lfs.github.com/spec/v1
|
| 2 |
+
oid sha256:f1d66664d87a7f8fe98c8cb44e61d26e8f97c8eaba1572234fcb2c38703f21ba
|
| 3 |
+
size 3840044
|