aboutsummaryrefslogtreecommitdiff
path: root/working/meeting-transcription-service/src/merge_transcript.py
diff options
context:
space:
mode:
Diffstat (limited to 'working/meeting-transcription-service/src/merge_transcript.py')
-rw-r--r--working/meeting-transcription-service/src/merge_transcript.py329
1 files changed, 329 insertions, 0 deletions
diff --git a/working/meeting-transcription-service/src/merge_transcript.py b/working/meeting-transcription-service/src/merge_transcript.py
new file mode 100644
index 0000000..13aec75
--- /dev/null
+++ b/working/meeting-transcription-service/src/merge_transcript.py
@@ -0,0 +1,329 @@
+#!/usr/bin/env python3
+"""Merge whisper word timings with speaker turns into transcript lines.
+
+Input is two JSON files: whisper-cli's ``-oj`` output (run at word level) and
+the diarizer's list of speaker turns. Output is one line per stretch of speech,
+``HH:MM:SS Speaker A: text``, the same shape the hosted services produced.
+
+Standard library only, so it runs under any Python 3.10+ without the venv.
+"""
+
+from __future__ import annotations
+
+import json
+import sys
+from dataclasses import dataclass
+from pathlib import Path
+
+DEFAULT_MAX_GAP_S = 3.0
+DEFAULT_SPEECH_MARGIN_S = 2.0
+DEFAULT_MAX_LOOP_S = 30.0
+
+
+@dataclass(frozen=True)
+class Unit:
+ """A piece of transcribed text with its start and end in seconds."""
+
+ start: float
+ end: float
+ text: str
+
+
+@dataclass(frozen=True)
+class Turn:
+ """A stretch of audio the diarizer attributes to one speaker."""
+
+ start: float
+ end: float
+ speaker: str
+
+
+def _timestamp(seconds: float) -> str:
+ """Render seconds as HH:MM:SS, floored."""
+ whole = int(seconds)
+ return f"{whole // 3600:02d}:{whole % 3600 // 60:02d}:{whole % 60:02d}"
+
+
+def _speaker_name(index: int) -> str:
+ """Name speakers A-Z in order of first speech, then by number."""
+ return chr(ord("A") + index) if index < 26 else str(index)
+
+
+def _speaker_for(unit: Unit, turns: list[Turn]) -> str:
+ """Pick the turn a unit belongs to.
+
+ The turn overlapping most of the unit wins. A unit overlapping nothing (a
+ zero-length word, or one whisper timed into a silence) goes to the turn
+ nearest its midpoint, because whisper's word timings drift by a few hundred
+ milliseconds and dropping the word would be worse than a near guess.
+ """
+ best = max(turns, key=lambda t: min(unit.end, t.end) - max(unit.start, t.start))
+ if min(unit.end, best.end) - max(unit.start, best.start) > 0:
+ return best.speaker
+ mid = (unit.start + unit.end) / 2
+
+ def distance(turn: Turn) -> float:
+ if turn.start <= mid <= turn.end:
+ return 0.0
+ return min(abs(mid - turn.start), abs(mid - turn.end))
+
+ return min(turns, key=distance).speaker
+
+
+def merge(units: list[Unit], turns: list[Turn], max_gap_s: float = DEFAULT_MAX_GAP_S) -> list[str]:
+ """Return transcript lines for ``units`` labelled by ``turns``.
+
+ Consecutive units from one speaker share a line. The line breaks when the
+ speaker changes, or when the speaker pauses longer than ``max_gap_s``, so a
+ long monologue still carries usable timestamps.
+
+ Raises:
+ ValueError: if there are no turns, no spoken words, or ``max_gap_s`` is negative.
+ """
+ if max_gap_s < 0:
+ raise ValueError("max_gap_s must not be negative")
+ spoken = sorted((u for u in units if u.text.strip()), key=lambda u: (u.start, u.end))
+ if not spoken:
+ raise ValueError("no speech: the transcription holds no words")
+ if not turns:
+ raise ValueError("no speaker turns: the diarization is empty")
+ ordered_turns = sorted(turns, key=lambda t: (t.start, t.end))
+
+ names: dict[str, str] = {}
+ lines: list[tuple[float, str, list[str]]] = []
+ previous_end = 0.0
+ for unit in spoken:
+ speaker = _speaker_for(unit, ordered_turns)
+ name = names.setdefault(speaker, _speaker_name(len(names)))
+ if lines and lines[-1][1] == name and unit.start - previous_end <= max_gap_s:
+ lines[-1][2].append(unit.text)
+ else:
+ lines.append((unit.start, name, [unit.text]))
+ previous_end = max(previous_end, unit.end)
+
+ return [
+ f"{_timestamp(start)} Speaker {name}: {' '.join(''.join(parts).split())}"
+ for start, name, parts in lines
+ ]
+
+
+def drop_outside_speech(
+ units: list[Unit], turns: list[Turn], margin_s: float = DEFAULT_SPEECH_MARGIN_S
+) -> tuple[list[Unit], int]:
+ """Return the units that belong to speech, and how many were dropped.
+
+ Whisper invents words ("Thank you.") when it is handed silence. The diarizer
+ marks where people actually spoke, so a unit that touches no turn and sits
+ more than ``margin_s`` from the nearest one is treated as invented. The margin
+ protects real words the diarizer clipped off the edge of a turn.
+
+ Raises:
+ ValueError: if there are no turns, or ``margin_s`` is negative.
+ """
+ if margin_s < 0:
+ raise ValueError("margin_s must not be negative")
+ if not turns:
+ raise ValueError("no speaker turns: the diarization is empty")
+
+ def gap(unit: Unit) -> float:
+ # Seconds between the unit and its nearest turn; zero when they touch.
+ # Rounded to the millisecond, the resolution of whisper's offsets.
+ nearest = min(max(turn.start - unit.end, unit.start - turn.end, 0.0) for turn in turns)
+ return round(nearest, 3)
+
+ kept = [unit for unit in units if gap(unit) <= margin_s]
+ return kept, len(units) - len(kept)
+
+
+def _unit(item: dict) -> Unit:
+ """A Unit from a whisper segment or token dict (offsets in milliseconds).
+
+ whisper-cli clamps a token's start to its segment's start without moving the
+ end, so some tokens arrive ending before they begin. Those become zero-length
+ at their start; left alone they corrupt both the overlap and the pause maths.
+ """
+ start = item["offsets"]["from"] / 1000
+ end = item["offsets"]["to"] / 1000
+ return Unit(start, max(start, end), item["text"])
+
+
+def find_repetition(
+ text: str, min_words: int = 3, max_words: int = 12, min_repeats: int = 4
+) -> tuple[str, int] | None:
+ """Find a phrase repeated back to back, whisper's hallucination signature.
+
+ Returns the phrase and its repeat count, or None. Four consecutive repeats of
+ a phrase of three or more words is the line: people say a thing two or three
+ times, and single-word runs ("yeah yeah yeah") are ordinary speech.
+ """
+ words = text.split()
+ keys = [w.lower().strip(".,!?;:\"'") for w in words]
+ for size in range(min_words, max_words + 1):
+ for i in range(len(keys) - size * min_repeats + 1):
+ phrase = keys[i : i + size]
+ if len(set(phrase)) < 2:
+ continue
+ count = 1
+ while keys[i + count * size : i + (count + 1) * size] == phrase:
+ count += 1
+ if count >= min_repeats:
+ return " ".join(words[i : i + size]), count
+ return None
+
+
+def collapse_repetitions(
+ units: list[Unit],
+ min_words: int = 3,
+ max_words: int = 12,
+ min_repeats: int = 4,
+ max_loop_s: float = DEFAULT_MAX_LOOP_S,
+) -> tuple[list[Unit], list[tuple[str, int]]]:
+ """Collapse whisper's repetition loops, keeping one copy of the phrase.
+
+ Whisper sometimes gets stuck and emits the same phrase over and over. A short
+ loop (up to ``max_loop_s`` of audio) costs a few seconds of speech, so it is
+ collapsed to a single occurrence and reported. A longer one means real speech
+ was lost for a stretch, and that is raised instead, so the job fails rather
+ than hand back a transcript with a hole in it.
+
+ Returns the surviving units and a list of (phrase, repeat count) for every
+ loop collapsed. Units are matched word by word, so a phrase spread over word
+ units and a phrase sitting in one segment unit are both found.
+
+ Raises:
+ ValueError: if a loop lasts longer than ``max_loop_s``, or the limit is negative.
+ """
+ if max_loop_s < 0:
+ raise ValueError("max_loop_s must not be negative")
+ units = list(units)
+ collapsed: list[tuple[str, int]] = []
+
+ def find() -> tuple[int, int, int] | None:
+ # (first word index, phrase size, repeat count) of the earliest loop, or None
+ words = [(w, ui) for ui, u in enumerate(units) for w in u.text.split()]
+ keys = [w.lower().strip(".,!?;:\"'") for w, _ in words]
+ best: tuple[int, int, int] | None = None
+ for size in range(min_words, max_words + 1):
+ for i in range(len(keys) - size * min_repeats + 1):
+ if best is not None and i >= best[0]:
+ break
+ phrase = keys[i : i + size]
+ if len(set(phrase)) < 2:
+ continue
+ count = 1
+ while keys[i + count * size : i + (count + 1) * size] == phrase:
+ count += 1
+ if count >= min_repeats:
+ best = (i, size, count)
+ break
+ return best
+
+ while True:
+ hit = find()
+ if hit is None:
+ return units, collapsed
+ i, size, count = hit
+ words = [(w, ui) for ui, u in enumerate(units) for w in u.text.split()]
+ phrase_text = " ".join(w for w, _ in words[i : i + size])
+ doomed = set(range(i + size, i + count * size)) # word indexes of the repeats
+ touched = {ui for wi, (_, ui) in enumerate(words) if wi in doomed or i <= wi < i + size}
+ span_start = min(units[ui].start for ui in touched)
+ span_end = max(units[ui].end for ui in touched)
+ duration = round(span_end - span_start, 3)
+ if duration > max_loop_s:
+ raise ValueError(
+ f"whisper looped: {phrase_text!r} repeats {count} times over {duration:.0f} s; "
+ "rerun whisper with -mc 0"
+ )
+ # Rebuild every touched unit from the words it keeps. A unit that held only
+ # repeats disappears; one that also held the first copy or later speech keeps
+ # those words. Every pass removes (count - 1) * size words, so this ends.
+ last_kept = max(ui for wi, (_, ui) in enumerate(words) if i <= wi < i + size)
+ rebuilt: list[Unit] = []
+ for ui, unit in enumerate(units):
+ if ui not in touched:
+ rebuilt.append(unit)
+ continue
+ keep = [w for wi, (w, wui) in enumerate(words) if wui == ui and wi not in doomed]
+ if not keep:
+ continue
+ # The kept copy takes over the time the loop occupied, so the merge does
+ # not read the removed stretch as a pause and break the line there.
+ end = max(unit.end, span_end) if ui == last_kept else unit.end
+ rebuilt.append(Unit(unit.start, end, " " + " ".join(keep)))
+ units = rebuilt
+ collapsed.append((phrase_text, count))
+
+
+def load_whisper_json(path: str | Path) -> list[Unit]:
+ """Read whisper-cli JSON output. Offsets there are in milliseconds.
+
+ With ``-ojf`` each segment carries its tokens and their offsets; those become
+ word-level units, which is what lets a speaker change land mid-segment. Plain
+ ``-oj`` output, or a segment with no tokens, falls back to the segment itself.
+
+ Raises:
+ ValueError: if the file is not whisper's JSON shape.
+ """
+ try:
+ data = json.loads(Path(path).read_text(encoding="utf-8"))
+ units: list[Unit] = []
+ for item in data["transcription"]:
+ tokens = item.get("tokens") or []
+ if not tokens:
+ units.append(_unit(item))
+ continue
+ for token in tokens:
+ text = token["text"]
+ if not text or text.startswith("[_"): # [_BEG_], [_TT_123], [_EOT_]
+ continue
+ unit = _unit(token)
+ if units and not text.startswith(" "):
+ # A sub-word piece or punctuation: it belongs to the word before it.
+ previous = units[-1]
+ units[-1] = Unit(previous.start, max(previous.end, unit.end), previous.text + text)
+ else:
+ units.append(unit)
+ return units
+ except OSError as err:
+ raise ValueError(f"{path}: cannot read whisper output ({err.strerror})") from err
+ except (json.JSONDecodeError, KeyError, TypeError) as err:
+ raise ValueError(f"{path}: not whisper-cli JSON output ({err!r})") from err
+
+
+def load_turns_json(path: str | Path) -> list[Turn]:
+ """Read the diarizer's turns: a list of {start, end, speaker}, in seconds.
+
+ Raises:
+ ValueError: if the file is not that shape.
+ """
+ try:
+ data = json.loads(Path(path).read_text(encoding="utf-8"))
+ if not isinstance(data, list):
+ raise TypeError("expected a list of turns")
+ return [Turn(float(t["start"]), float(t["end"]), str(t["speaker"])) for t in data]
+ except OSError as err:
+ raise ValueError(f"{path}: cannot read speaker turns ({err.strerror})") from err
+ except (json.JSONDecodeError, KeyError, TypeError, ValueError) as err:
+ raise ValueError(f"{path}: not a speaker-turns file ({err!r})") from err
+
+
+def main(argv: list[str]) -> int:
+ """CLI: ``merge_transcript.py WHISPER_JSON TURNS_JSON`` prints the transcript."""
+ if len(argv) != 2:
+ print("usage: merge_transcript.py WHISPER_JSON TURNS_JSON", file=sys.stderr)
+ return 2
+ try:
+ turns = load_turns_json(argv[1])
+ units, _dropped = drop_outside_speech(load_whisper_json(argv[0]), turns)
+ units, _collapsed = collapse_repetitions(units)
+ lines = merge(units, turns)
+ except ValueError as err:
+ print(f"Error: {err}", file=sys.stderr)
+ return 1
+ print("\n".join(lines))
+ return 0
+
+
+if __name__ == "__main__":
+ sys.exit(main(sys.argv[1:]))