diff options
Diffstat (limited to 'working/meeting-transcription-service/src/diarize.py')
| -rw-r--r-- | working/meeting-transcription-service/src/diarize.py | 114 |
1 files changed, 114 insertions, 0 deletions
diff --git a/working/meeting-transcription-service/src/diarize.py b/working/meeting-transcription-service/src/diarize.py new file mode 100644 index 0000000..a6ae85e --- /dev/null +++ b/working/meeting-transcription-service/src/diarize.py @@ -0,0 +1,114 @@ +#!/usr/bin/env python3 +"""Run pyannote speaker diarization on one audio file and write the turns as JSON. + +Usage: diarize.py AUDIO OUT_JSON [--speakers N | --min-speakers N --max-speakers N] + +Output is a list of {"start", "end", "speaker"} in seconds, which is what +merge_transcript.py reads. HF_TOKEN is only needed the first time, to download +the gated model; after that the cached copy loads offline. + +pyannote and torch are imported inside run(), so the pure helpers here can be +tested without the multi-gigabyte environment. +""" + +from __future__ import annotations + +import argparse +import json +import os +import sys +from collections.abc import Iterable +from pathlib import Path +from typing import Any + +MODEL = "pyannote/speaker-diarization-community-1" + + +def turns_from_tracks(tracks: Iterable[tuple[Any, Any, str]]) -> list[dict[str, Any]]: + """Convert pyannote (segment, track, label) triples into sorted turn dicts. + + Times are rounded to milliseconds. Segments with no length are dropped. + """ + turns = [ + {"start": round(float(seg.start), 3), "end": round(float(seg.end), 3), "speaker": str(label)} + for seg, _track, label in tracks + if float(seg.end) > float(seg.start) + ] + return sorted(turns, key=lambda t: (t["start"], t["end"])) + + +def _positive_int(value: str) -> int: + number = int(value) + if number < 1: + raise argparse.ArgumentTypeError("must be 1 or more") + return number + + +def parse_args(argv: list[str]) -> argparse.Namespace: + """Parse the command line. Exits with usage on bad or contradictory counts.""" + parser = argparse.ArgumentParser(description="Speaker diarization with pyannote.") + parser.add_argument("audio") + parser.add_argument("out") + parser.add_argument("--speakers", type=_positive_int, help="exact number of speakers") + parser.add_argument("--min-speakers", type=_positive_int) + parser.add_argument("--max-speakers", type=_positive_int) + args = parser.parse_args(argv) + if args.speakers is not None and (args.min_speakers or args.max_speakers): + parser.error("--speakers cannot be combined with --min-speakers/--max-speakers") + if args.min_speakers and args.max_speakers and args.min_speakers > args.max_speakers: + parser.error("--min-speakers cannot exceed --max-speakers") + return args + + +def pipeline_kwargs(args: argparse.Namespace) -> dict[str, int]: + """Only the speaker-count options that were actually given.""" + options = { + "num_speakers": args.speakers, + "min_speakers": args.min_speakers, + "max_speakers": args.max_speakers, + } + return {name: value for name, value in options.items() if value is not None} + + +def run(args: argparse.Namespace) -> list[dict[str, Any]]: + """Load the pipeline, diarize the audio, and return the turns.""" + # Heavy, and only installed in the service venv, so imported here on purpose. + from pyannote.audio import Pipeline # pyright: ignore[reportMissingImports] + + pipeline = Pipeline.from_pretrained(MODEL, token=os.environ.get("HF_TOKEN") or None) + if pipeline is None: + raise RuntimeError(f"could not load {MODEL}: check HF_TOKEN and that its terms are accepted") + output = pipeline(args.audio, **pipeline_kwargs(args)) + # The exclusive variant never overlaps two speakers, which is what a + # word-by-word merge wants. Older pipelines return the annotation itself. + annotation = getattr(output, "exclusive_speaker_diarization", None) + if annotation is None: + annotation = getattr(output, "speaker_diarization", output) + return turns_from_tracks(annotation.itertracks(yield_label=True)) + + +def main(argv: list[str]) -> int: + """CLI entry point. Writes OUT_JSON atomically; non-zero on any failure.""" + args = parse_args(argv) + if not Path(args.audio).is_file(): + print(f"Error: audio file not found: {args.audio}", file=sys.stderr) + return 1 + try: + turns = run(args) + except Exception as err: # noqa: BLE001 - report any model failure and exit non-zero + print(f"Error: diarization failed: {err}", file=sys.stderr) + return 1 + if not turns: + print("Error: diarization found no speech", file=sys.stderr) + return 1 + out = Path(args.out) + partial = out.with_name(out.name + ".partial") + partial.write_text(json.dumps(turns), encoding="utf-8") + partial.replace(out) + speakers = len({t["speaker"] for t in turns}) + print(f"{len(turns)} turns, {speakers} speakers -> {out}", file=sys.stderr) + return 0 + + +if __name__ == "__main__": + sys.exit(main(sys.argv[1:])) |
