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