File size: 8,446 Bytes
7c268e9
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
"""End-to-end streaming diarizer: waveform -> RTTM.

Two scatter-free graphs are executed per step (pre-encode + encoder); the
speaker-cache packing, streaming state update, log-mel front-end and
post-processing are all numpy and shared by host and AX650.
"""

import argparse
from pathlib import Path
from typing import Dict, Protocol

import numpy as np

from .feature import SAMPLE_RATE, log_mel_spectrogram, mel_filterbank
from .postprocess import PostProcessingParams, predlist_to_timestamps, timestamps_to_rttm_lines
from .state import SortformerConfig, init_state, iter_chunks, pre_encode_length, streaming_update


class GraphPair(Protocol):
    def pre_encode(self, chunk: np.ndarray) -> np.ndarray:
        """chunk [1, chunk_mel, 128] -> chunk embeddings [1, chunk_embs, 512]."""

    def encode(self, seq: np.ndarray, total_lengths: int) -> np.ndarray:
        """packed seq [1, state_len + chunk_embs, 512] -> preds [1, T, num_speakers]."""


def _provider_alias(ort, provider: str) -> str:
    aliases = {"cpu": "CPUExecutionProvider", "cuda": "CUDAExecutionProvider"}
    if provider == "auto":
        available = ort.get_available_providers()
        return "CUDAExecutionProvider" if "CUDAExecutionProvider" in available else "CPUExecutionProvider"
    return aliases.get(provider, provider)


class OnnxGraphPair:
    """ONNX Runtime implementation of the split graphs (host verification)."""

    def __init__(self, preencode_path: str, encoder_path: str, provider: str = "auto"):
        import onnxruntime as ort

        provider = _provider_alias(ort, provider)
        self.pre_session = ort.InferenceSession(str(preencode_path), providers=[provider])
        self.encoder_session = ort.InferenceSession(str(encoder_path), providers=[provider])
        # The encoder graph may be wider than the runtime state (e.g. a 390-wide
        # graph evaluated with a 40-frame FIFO); the host pads and masks the tail.
        encoder_inputs = {node.name: node for node in self.encoder_session.get_inputs()}
        self.seq_width = int(encoder_inputs["seq"].shape[1])

    def pre_encode(self, chunk: np.ndarray) -> np.ndarray:
        return self.pre_session.run(None, {"chunk": chunk.astype(np.float32)})[0]

    def encode(self, seq: np.ndarray, total_lengths: int) -> np.ndarray:
        feed = {"seq": seq.astype(np.float32), "total_lengths": np.array([total_lengths], dtype=np.float32)}
        return self.encoder_session.run(None, feed)[0]


class AxengineGraphPair:
    """axengine implementation of the split graphs (AX650 board)."""

    def __init__(self, preencode_path: str, encoder_path: str):
        import axengine

        self.pre_session = axengine.InferenceSession(str(preencode_path))
        self.encoder_session = axengine.InferenceSession(str(encoder_path))

    def pre_encode(self, chunk: np.ndarray) -> np.ndarray:
        return self.pre_session.run(None, {"chunk": chunk.astype(np.float32)})[0]

    def encode(self, seq: np.ndarray, total_lengths: int) -> np.ndarray:
        feed = {"seq": seq.astype(np.float32), "total_lengths": np.array([total_lengths], dtype=np.float32)}
        return self.encoder_session.run(None, feed)[0]


class StreamingDiarizer:
    """Sortformer 1.04 s streaming diarizer over the split graphs."""

    def __init__(self, graphs: GraphPair, config: SortformerConfig = None):
        self.graphs = graphs
        self.config = config or SortformerConfig()
        self.chunk_width = (
            self.config.chunk_left_context + self.config.chunk_len + self.config.chunk_right_context
        ) * self.config.subsampling_factor
        self.state_len = self.config.spkcache_len + self.config.fifo_len

    def process_features(self, features: np.ndarray, on_step=None) -> np.ndarray:
        cfg = self.config
        state = init_state(cfg)
        mel_dim = features.shape[1]
        predictions = []
        for step, (chunk, left_offset, right_offset) in enumerate(iter_chunks(cfg, features)):
            chunk_valid = chunk.shape[0]
            chunk_input = np.zeros((1, self.chunk_width, mel_dim), dtype=np.float32)
            chunk_input[0, :chunk_valid] = chunk

            chunk_embs_full = self.graphs.pre_encode(chunk_input)[0]
            chunk_embs_length = pre_encode_length(chunk_valid)
            chunk_embs = chunk_embs_full[:chunk_embs_length]

            spkcache_len = state.spkcache.shape[0]
            fifo_len = state.fifo.shape[0]
            total_lengths = spkcache_len + fifo_len + chunk_embs_length
            seq_width = getattr(self.graphs, "seq_width", self.state_len + chunk_embs_full.shape[0])
            seq = np.zeros((1, seq_width, cfg.fc_d_model), dtype=np.float32)
            seq[0, :spkcache_len] = state.spkcache
            seq[0, spkcache_len : spkcache_len + fifo_len] = state.fifo
            seq[0, spkcache_len + fifo_len : total_lengths] = chunk_embs

            if on_step is not None:
                on_step(step, chunk_input, seq, total_lengths)

            preds = self.graphs.encode(seq, total_lengths)[0]
            lc_enc = round(left_offset / cfg.subsampling_factor)
            rc_enc = int(np.ceil(right_offset / cfg.subsampling_factor))
            state, chunk_preds = streaming_update(cfg, state, chunk_embs, preds, lc_enc, rc_enc)
            predictions.append(chunk_preds)
        if not predictions:
            return np.zeros((0, cfg.num_speakers), dtype=np.float32)
        return np.concatenate(predictions, axis=0)

    def process_wav(self, waveform: np.ndarray, sample_rate: int = SAMPLE_RATE, on_step=None) -> np.ndarray:
        features = log_mel_spectrogram(waveform, sample_rate, mel_filter=mel_filterbank())
        return self.process_features(features, on_step=on_step)

    def rttm_lines(self, preds: np.ndarray, uri: str, params: PostProcessingParams = None, bypass: bool = False):
        timestamps = predlist_to_timestamps(preds, params=params, bypass_postprocessing=bypass)
        num_speakers = preds.shape[1] if preds.ndim == 2 else self.config.num_speakers
        return timestamps_to_rttm_lines(timestamps, uri, num_speakers)


def main() -> None:
    parser = argparse.ArgumentParser(description=__doc__)
    parser.add_argument("--preencode", default=None, help="preencode.onnx (host mode)")
    parser.add_argument("--encoder", default=None, help="encoder.onnx (host mode)")
    parser.add_argument("--preencode-axmodel", default=None, help="preencode axmodel (board mode)")
    parser.add_argument("--encoder-axmodel", default=None, help="encoder axmodel (board mode)")
    parser.add_argument("--wav", required=True)
    parser.add_argument("--rttm", default=None)
    parser.add_argument("--provider", default="auto")
    parser.add_argument("--chunk-len", type=int, default=6)
    parser.add_argument("--chunk-left-context", type=int, default=1)
    parser.add_argument("--chunk-right-context", type=int, default=7)
    parser.add_argument("--fifo-len", type=int, default=188)
    parser.add_argument("--spkcache-len", type=int, default=188)
    parser.add_argument("--spkcache-update-period", type=int, default=144)
    parser.add_argument("--bypass-postproc", action="store_true")
    args = parser.parse_args()

    if args.preencode and args.encoder:
        graphs = OnnxGraphPair(args.preencode, args.encoder, provider=args.provider)
    elif args.preencode_axmodel and args.encoder_axmodel:
        graphs = AxengineGraphPair(args.preencode_axmodel, args.encoder_axmodel)
    else:
        raise SystemExit("pass either --preencode/--encoder or the axmodel pair")

    config = SortformerConfig(
        chunk_len=args.chunk_len,
        chunk_left_context=args.chunk_left_context,
        chunk_right_context=args.chunk_right_context,
        fifo_len=args.fifo_len,
        spkcache_len=args.spkcache_len,
        spkcache_update_period=args.spkcache_update_period,
    )
    diarizer = StreamingDiarizer(graphs, config)

    import soundfile as sf

    waveform, sample_rate = sf.read(args.wav, dtype="float32", always_2d=True)
    waveform = waveform.mean(axis=1)
    preds = diarizer.process_wav(waveform, sample_rate)
    uri = Path(args.wav).stem
    lines = diarizer.rttm_lines(preds, uri, bypass=args.bypass_postproc)
    print(f"{uri}: {len(preds)} frames, {len(lines)} segments")
    if args.rttm:
        Path(args.rttm).write_text("\n".join(lines) + "\n")
        print(f"rttm -> {args.rttm}")


if __name__ == "__main__":
    main()