HY-2012 commited on
Commit
7c268e9
·
verified ·
1 Parent(s): 3f67ccd

Upload folder using huggingface_hub

Browse files
.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: mit
 
 
 
 
 
 
 
 
 
 
 
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