#!/usr/bin/env python3 """sherpa-onnx speaker-diarization engine for scripts/diarize.mjs. Kept deliberately dumb: decode audio, run pyannote segmentation + a speaker embedding model, cluster, print raw turns as JSON on stdout. All policy (filenames, provenance shape, atomic writes) lives in the .mjs wrapper, so swapping this engine for pyannote later means replacing one file. stdout is JSON ONLY. Progress goes to stderr so the wrapper can stream it. WINDOWING, AND WHY IT EXISTS ---------------------------- Long recordings used to be killed by the kernel partway through. Two independent memory sinks, both measured rather than guessed: 1. sherpa's clustering holds a pairwise distance matrix over speech-segment embeddings -- O(n^2) in SEGMENT count. At 30k segments that is 6.7 GiB and at 40k it is 11.9 GiB, which brackets the 10.6 GB and 9.6 GB peaks observed before this. Note it is segment count, not duration: speaker-turn density varies 40x across this corpus, which is why a sparse 7h42m file survived while a dense 6h12m one did not. 2. the whole-file decode below buffers the ENTIRE decoded stream as one bytes object -- 1.84 GB for an 8h file, and Python's incremental buffer growth peaks near twice that. Windowing fixes both at once, which is why it is the real answer and not a mitigation: per-window n falls by the window count so the matrix falls by its SQUARE, and decoding one window at a time removes the whole-file buffer entirely. A 45-minute window holds 165 MiB against 1.72 GiB for a whole 8h file. The hard part is that speaker 2 in window 1 and speaker 5 in window 2 may be the same person, and a diarization that cannot say so is not worth having. So each window's local speakers are reduced to an embedding centroid and the centroids are clustered globally, with the same threshold, to recover identity across seams. SHORT FILES TAKE THE EXACT ORIGINAL CODE PATH -- same decode call, same sd.process, same speaker ids. That is a hard requirement, not tidiness: it is what keeps every sidecar already on disk reproducible. Usage: diarize-sherpa.py --seg --emb [--threshold F] [--threads N] [--ffmpeg PATH] [--ffprobe PATH] [--window-minutes M] [--window-after-minutes M] Output: {"turns":[{"start":s,"end":s,"speaker":i}], "audioSeconds":f, "version":"...", "sampleRate":n, "windowing":{...}|null} """ import argparse, json, subprocess, sys, time # How much of a speaker's audio to embed when building its per-window centroid. # Whole turns are used up to a cap: TitaNet needs a second or two to say anything # useful, and past ~10s per excerpt more audio stops changing the embedding. EMB_MAX_SEGMENTS = 5 EMB_MAX_SEC = 10.0 EMB_MIN_SEC = 0.5 # Two consecutive turns of the same GLOBAL speaker separated by less than this # are joined. Matches the segmentation default for min_duration_off, so a seam # does not leave a visible split that the same audio processed whole would not # have had. MERGE_GAP_SEC = 0.5 def decode_whole_file(np, ffmpeg, audio, sample_rate): """The ORIGINAL decode, unchanged. See the module docstring: short files must take this path byte-for-byte so existing sidecars stay reproducible.""" proc = subprocess.run( [ffmpeg, "-v", "error", "-i", audio, "-f", "f32le", "-ac", "1", "-ar", str(sample_rate), "-"], stdout=subprocess.PIPE, stderr=subprocess.PIPE, ) if proc.returncode != 0: sys.stderr.write(proc.stderr.decode("utf8", "replace")) print(f"diarize-sherpa: ffmpeg failed ({proc.returncode})", file=sys.stderr) return None return np.frombuffer(proc.stdout, dtype=np.float32) def decode_window(np, ffmpeg, audio, sample_rate, start_sec, dur_sec): """Decode ONE window into a PREALLOCATED array. The allocation is the point. subprocess.run(stdout=PIPE) accumulates the whole stream and grows its buffer as it goes, so peak memory is roughly twice the final size; here the exact sample count is known in advance from the window length, so ffmpeg writes straight into one array that is never resized. Returns the array trimmed to what actually arrived -- the last window is short, and a truncated file is shorter than its container claims. """ n = int(round(dur_sec * sample_rate)) buf = np.empty(n, dtype=np.float32) view = memoryview(buf).cast("B") args = [ffmpeg, "-v", "error", # -ss BEFORE -i is an input seek: ffmpeg does not decode and discard # the preceding hours, which is what makes a late window as cheap as # an early one. Verified sample-exact on this corpus's files. "-ss", f"{start_sec:.6f}", "-t", f"{dur_sec:.6f}", "-i", audio, "-f", "f32le", "-ac", "1", "-ar", str(sample_rate), "-"] proc = subprocess.Popen(args, stdout=subprocess.PIPE, stderr=subprocess.PIPE) filled = 0 try: while filled < view.nbytes: got = proc.stdout.readinto(view[filled:]) if not got: break filled += got # ffmpeg can round a window up by a frame. Drained, not buffered. while proc.stdout.read(65536): pass finally: proc.stdout.close() err = proc.stderr.read().decode("utf8", "replace") proc.stderr.close() rc = proc.wait() if rc != 0: sys.stderr.write(err) raise RuntimeError(f"ffmpeg failed ({rc}) decoding window at {start_sec:.1f}s") return buf[: filled // 4] def probe_duration(ffprobe, audio): """Container duration in seconds, or None when it cannot be measured. None is load-bearing: it sends the run down the whole-file path, i.e. exactly what this script did before windowing existed. Guessing a duration and windowing on it would be the one way to make an unreadable file produce a WRONG answer rather than the old one. """ try: p = subprocess.run( [ffprobe, "-v", "error", "-show_entries", "format=duration", "-of", "default=noprint_wrappers=1:nokey=1", audio], stdout=subprocess.PIPE, stderr=subprocess.PIPE, ) if p.returncode != 0: return None sec = float(p.stdout.decode("utf8", "replace").strip()) return sec if sec > 0 else None except Exception: return None def plan_windows(total_sec, window_sec, overlap_sec): """[(start, dur, own_start, own_end)] covering total_sec. Windows OVERLAP so the segmentation model has context either side of a seam, but each second of the recording is OWNED by exactly one window -- the ownership boundaries meet in the middle of the overlap. A turn is kept by the window that owns its midpoint, so the overlap buys context without producing the same speech twice. """ step = window_sec - overlap_sec spans = [] start = 0.0 while True: dur = min(window_sec, total_sec - start) spans.append((start, dur)) if start + dur >= total_sec - 1e-6: break start += step out = [] for i, (s, d) in enumerate(spans): own_s = s if i == 0 else s + overlap_sec / 2.0 own_e = s + d if i == len(spans) - 1 else s + d - overlap_sec / 2.0 out.append((s, d, own_s, own_e)) return out def speaker_centroid(np, extractor, samples, sample_rate, segments): """One L2-normalized embedding for a local speaker, or None. Averaged over the speaker's LONGEST turns rather than computed from one: a single turn can be a two-word interjection, and an identity built from that is what makes a speaker split across a seam. """ picked = sorted(segments, key=lambda s: s[1] - s[0], reverse=True) picked = picked[:EMB_MAX_SEGMENTS] vecs = [] for start, end, _ in picked: end = min(end, start + EMB_MAX_SEC) if end - start < EMB_MIN_SEC: continue a = int(start * sample_rate) b = min(int(end * sample_rate), samples.size) if b - a < int(EMB_MIN_SEC * sample_rate): continue stream = extractor.create_stream() stream.accept_waveform(sample_rate, samples[a:b]) stream.input_finished() if not extractor.is_ready(stream): continue v = np.asarray(extractor.compute(stream), dtype=np.float32) norm = float(np.linalg.norm(v)) if norm > 0: vecs.append(v / norm) if not vecs: return None mean = np.mean(np.stack(vecs), axis=0) norm = float(np.linalg.norm(mean)) return mean / norm if norm > 0 else None def merge_and_renumber(turns): """Absolute-time turns -> final output. Two jobs, both seam repair. Renumbering by first appearance, because global cluster labels come out of the clusterer in no useful order and the rest of this app reads speaker 0 as "the first voice heard". Then unioning each speaker's own intervals, so a speaker who talked straight through a seam is not reported as having stopped and started. PER SPEAKER, NOT PAIRWISE DOWN THE SORTED LIST. Merging only CONSECUTIVE turns looks equivalent and is not: two different local speakers can map to the same global one, and their turns then interleave with — or overlap — other people's. A pairwise pass leaves those unjoined, and worse leaves the same speaker overlapping THEMSELVES, which double-counts their talk time in every downstream consumer. Measured on a 6h12m VOD before this was fixed: 35 self-overlaps and 24 sub-gap splits that a pairwise merge could not see. Cross-speaker overlap is left alone — that is genuine overlapping speech and the whole-file path emits it too. """ order = {} for t in sorted(turns, key=lambda t: (t["start"], t["end"])): if t["speaker"] not in order: order[t["speaker"]] = len(order) by_speaker = {} for t in turns: by_speaker.setdefault(order[t["speaker"]], []).append( (t["start"], t["end"]) ) merged = [] for speaker, spans in by_speaker.items(): spans.sort() cur_start, cur_end = spans[0] for start, end in spans[1:]: if start - cur_end <= MERGE_GAP_SEC: cur_end = max(cur_end, end) else: merged.append( {"start": cur_start, "end": cur_end, "speaker": speaker} ) cur_start, cur_end = start, end merged.append({"start": cur_start, "end": cur_end, "speaker": speaker}) merged.sort(key=lambda t: (t["start"], t["end"], t["speaker"])) return merged def run_windowed(np, sherpa_onnx, sd, args, total_sec, window_sec, overlap_sec): """Window -> local diarization -> per-speaker centroid -> global clustering.""" sample_rate = sd.sample_rate windows = plan_windows(total_sec, window_sec, overlap_sec) print( f"diarize-sherpa: {total_sec/3600:.2f} h -> {len(windows)} window(s) of " f"{window_sec/60:.0f} min ({overlap_sec:.0f}s overlap), windowed mode", file=sys.stderr, flush=True, ) extractor = sherpa_onnx.SpeakerEmbeddingExtractor( sherpa_onnx.SpeakerEmbeddingExtractorConfig( model=args.emb, num_threads=args.threads ) ) kept = [] # {"start","end","key"} in ABSOLUTE seconds centroids = [] # parallel to keys keys = [] # (window index, local speaker) decoded_sec = 0.0 for wi, (ws, dur, own_s, own_e) in enumerate(windows): samples = decode_window(np, args.ffmpeg, args.audio, sample_rate, ws, dur) if samples.size == 0: print(f"diarize-sherpa: window {wi+1} decoded zero samples; skipping", file=sys.stderr, flush=True) continue decoded_sec += samples.size / sample_rate t0 = time.time() result = sd.process(samples).sort_by_start_time() segs = [(s.start, s.end, s.speaker) for s in result] # Keep the turns this window OWNS, by midpoint. The last window owns its # own right edge, or the tail of the recording would be dropped. is_last = wi == len(windows) - 1 owned_by_speaker = {} for start, end, spk in segs: mid = start + (end - start) / 2.0 + ws if mid < own_s or (mid >= own_e and not is_last): continue owned_by_speaker.setdefault(spk, []).append((start, end, spk)) kept.append({"start": ws + start, "end": ws + end, "key": (wi, spk)}) # Centroids ONLY for speakers with owned turns: a local speaker whose # every turn belongs to the neighbouring window would otherwise add a # cluster that names nothing. for spk, spk_segs in owned_by_speaker.items(): c = speaker_centroid(np, extractor, samples, sample_rate, spk_segs) if c is None: continue centroids.append(c) keys.append((wi, spk)) elapsed = time.time() - t0 print( f"diarize-sherpa: window {wi+1}/{len(windows)} " f"({ws/60:.0f}-{(ws+dur)/60:.0f} min): {len(segs)} turns, " f"{len(owned_by_speaker)} local speaker(s) in {elapsed:.1f}s", file=sys.stderr, flush=True, ) del samples, result, segs if not kept: return [], decoded_sec, len(windows) # Cross-window identity. Same clusterer and same threshold as within a # window, so "these two are the same person" means the same thing at both # scales. if centroids: labels = sherpa_onnx.FastClustering( sherpa_onnx.FastClusteringConfig( num_clusters=args.num_speakers, threshold=args.threshold ) )(np.stack(centroids)) global_of = {key: int(label) for key, label in zip(keys, labels)} else: global_of = {} # A local speaker with no usable centroid (all its turns too short to embed) # keeps an identity of its own rather than being merged into someone else or # dropped. Unmatched, but present and honest. spare = (max(global_of.values()) + 1) if global_of else 0 for t in kept: if t["key"] not in global_of: global_of[t["key"]] = spare spare += 1 t["speaker"] = global_of[t["key"]] del t["key"] return merge_and_renumber(kept), decoded_sec, len(windows) def main(): ap = argparse.ArgumentParser() ap.add_argument("audio") ap.add_argument("--seg", required=True) ap.add_argument("--emb", required=True) ap.add_argument("--threshold", type=float, default=0.5) ap.add_argument("--threads", type=int, default=4) ap.add_argument("--num-speakers", type=int, default=-1) ap.add_argument("--min-duration-on", type=float, default=0.3) ap.add_argument("--min-duration-off", type=float, default=0.5) ap.add_argument("--ffmpeg", default="ffmpeg") ap.add_argument("--ffprobe", default="ffprobe") # Window length. 0 disables windowing entirely and forces the original # whole-file path, whatever the duration. ap.add_argument("--window-minutes", type=float, default=45.0) # Files at or under this take the original whole-file path. Well above any # duration that has ever been a problem here, so windowing stays the # exception rather than quietly becoming the default. ap.add_argument("--window-after-minutes", type=float, default=90.0) ap.add_argument("--window-overlap-seconds", type=float, default=30.0) args = ap.parse_args() try: import numpy as np import sherpa_onnx except ImportError as e: print(f"diarize-sherpa: missing dependency: {e}", file=sys.stderr) return 3 cfg = sherpa_onnx.OfflineSpeakerDiarizationConfig( segmentation=sherpa_onnx.OfflineSpeakerSegmentationModelConfig( pyannote=sherpa_onnx.OfflineSpeakerSegmentationPyannoteModelConfig( model=args.seg ), num_threads=args.threads, ), embedding=sherpa_onnx.SpeakerEmbeddingExtractorConfig( model=args.emb, num_threads=args.threads ), clustering=sherpa_onnx.FastClusteringConfig( num_clusters=args.num_speakers, threshold=args.threshold ), min_duration_on=args.min_duration_on, min_duration_off=args.min_duration_off, ) if not cfg.validate(): print("diarize-sherpa: invalid engine config", file=sys.stderr) return 2 sd = sherpa_onnx.OfflineSpeakerDiarization(cfg) window_sec = args.window_minutes * 60.0 overlap_sec = max(0.0, min(args.window_overlap_seconds, window_sec / 2.0)) # Measured BEFORE decoding anything, because deciding after a whole-file # decode would already have paid the memory this exists to avoid. total_sec = probe_duration(args.ffprobe, args.audio) if window_sec > 0 else None windowed = ( window_sec > 0 and total_sec is not None and total_sec > args.window_after_minutes * 60.0 ) if windowed: # Timed WHOLE, decode included: in this mode decoding is interleaved with # the work rather than being a step before it, so there is no honest way # to separate them. t0 = time.time() try: turns, audio_seconds, window_count = run_windowed( np, sherpa_onnx, sd, args, total_sec, window_sec, overlap_sec ) except RuntimeError as e: print(f"diarize-sherpa: {e}", file=sys.stderr) return 4 elapsed = time.time() - t0 windowing = { "windows": window_count, "windowSeconds": round(window_sec, 3), "overlapSeconds": round(overlap_sec, 3), } else: # Decode to the rate the models expect. ffmpeg reads any container, which # is what lets this run against a persisted source video as well as # audio.mp3. samples = decode_whole_file(np, args.ffmpeg, args.audio, sd.sample_rate) if samples is None: return 4 if samples.size == 0: print("diarize-sherpa: decoded zero samples", file=sys.stderr) return 5 audio_seconds = samples.size / sd.sample_rate print( f"diarize-sherpa: {audio_seconds/60:.1f} min decoded, diarizing…", file=sys.stderr, flush=True, ) # Timed from HERE, excluding the decode, exactly as before windowing # existed — the s/audio-hour figures already recorded in plans/ are # engine time, and a number that quietly changed meaning is worse than no # number. t0 = time.time() result = sd.process(samples).sort_by_start_time() elapsed = time.time() - t0 turns = [ {"start": round(s.start, 3), "end": round(s.end, 3), "speaker": s.speaker} for s in result ] windowing = None if windowed: turns = [ {"start": round(t["start"], 3), "end": round(t["end"], 3), "speaker": int(t["speaker"])} for t in turns ] speakers = len({t["speaker"] for t in turns}) print( f"diarize-sherpa: {len(turns)} turns, {speakers} speakers in " f"{elapsed:.1f}s ({elapsed/max(audio_seconds/3600,1e-9):.0f} s/audio-hour)", file=sys.stderr, flush=True, ) json.dump( { "turns": turns, "audioSeconds": round(audio_seconds, 3), "sampleRate": sd.sample_rate, "version": getattr(sherpa_onnx, "__version__", None), "windowing": windowing, }, sys.stdout, ) return 0 if __name__ == "__main__": sys.exit(main())