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