| """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]) |
| |
| |
| 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() |
|
|