Download python/sortformer_sdk/diarize.py from HY-2012/Sortformer.AXERA: direct link, hf CLI and curl.
- Browser
- Download file 8.45 kB
-
https://huggingface.co/HY-2012/Sortformer.AXERA/resolve/main/python/sortformer_sdk/diarize.py
- Command line
-
hf download hf://HY-2012/Sortformer.AXERA/python/sortformer_sdk/diarize.py
-
curl -L -o diarize.py https://huggingface.co/HY-2012/Sortformer.AXERA/resolve/main/python/sortformer_sdk/diarize.py
8.45 kB
| """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() | |