Add diarization and reusable MVP meeting pipeline

This commit is contained in:
2026-08-23 19:29:47 +02:00
parent f2d21c1faf
commit 8dab928763
18 changed files with 1697 additions and 180 deletions
+124 -12
View File
@@ -6,7 +6,9 @@ from pathlib import Path
from unittest.mock import Mock, patch
from scripts import run_mvp_meeting
from src.meeting_lab.orchestration import mvp as mvp_api
from src.meeting_lab.protocol.generate_direct_protocol import DirectProtocolResult
from src.meeting_lab.diarization.backend import DiarizationResult
from src.meeting_lab.transcription.whisper import TranscriptionError, TranscriptionResult
@@ -50,7 +52,15 @@ def fake_transcribe(
metadata = output_dir / "runtime_metadata.json"
raw.write_text('{"transcription": []}\n', encoding="utf-8")
transcript.write_text(
json.dumps({"text": "Besprechungstext.", "segments": []}) + "\n",
json.dumps(
{
"text": "Besprechungstext.",
"segments": [
{"id": 0, "start": 0.0, "end": 1.0, "text": "Besprechungstext."}
],
}
)
+ "\n",
encoding="utf-8",
)
text.write_text("Besprechungstext.\n", encoding="utf-8")
@@ -58,6 +68,44 @@ def fake_transcribe(
return TranscriptionResult(output_dir, raw, transcript, text, metadata, 1.25)
def fake_diarize(audio_path, output_dir, device_mode, **kwargs):
output_dir.mkdir(parents=True, exist_ok=True)
metadata = {
"backend": "pyannote.audio",
"model": "pyannote/speaker-diarization-community-1",
"requested_device_mode": device_mode,
"actual_device": "cuda",
"device_name": "Fake GPU",
"runtime_seconds": 2.5,
"speaker_count": 1,
"credentials_persisted": False,
}
paths = {
"metadata": output_dir / "metadata.json",
"ordinary": output_dir / "diarization.rttm",
"exclusive": output_dir / "exclusive_diarization.rttm",
"turns": output_dir / "turns.json",
"exclusive_turns": output_dir / "exclusive_turns.json",
}
paths["metadata"].write_text(json.dumps(metadata), encoding="utf-8")
paths["ordinary"].write_text("", encoding="utf-8")
paths["exclusive"].write_text("", encoding="utf-8")
paths["turns"].write_text("[]", encoding="utf-8")
paths["exclusive_turns"].write_text(
json.dumps([{"start": 0, "end": 10, "speaker_id": "SPEAKER_00"}]),
encoding="utf-8",
)
return DiarizationResult(
output_dir,
paths["metadata"],
paths["ordinary"],
paths["exclusive"],
paths["turns"],
paths["exclusive_turns"],
metadata,
)
class MvpOrchestratorTests(unittest.TestCase):
def create_inputs(self, root: Path) -> tuple[Path, Path, Path]:
audio = root / "team meeting.wav"
@@ -88,9 +136,9 @@ class MvpOrchestratorTests(unittest.TestCase):
root = Path(directory)
args = self.args(root)
with (
patch.object(run_mvp_meeting, "transcribe_audio", side_effect=fake_transcribe) as whisper,
patch.object(mvp_api, "transcribe_audio", side_effect=fake_transcribe) as whisper,
patch.object(
run_mvp_meeting,
mvp_api,
"generate_direct_protocol",
return_value=protocol_result(),
) as protocol,
@@ -129,9 +177,9 @@ class MvpOrchestratorTests(unittest.TestCase):
["--whisper-executable", "/tools/whisper-cli", "--threads", "4"],
)
with (
patch.object(run_mvp_meeting, "transcribe_audio", side_effect=fake_transcribe) as whisper,
patch.object(mvp_api, "transcribe_audio", side_effect=fake_transcribe) as whisper,
patch.object(
run_mvp_meeting,
mvp_api,
"generate_direct_protocol",
return_value=protocol_result(),
) as protocol,
@@ -148,17 +196,81 @@ class MvpOrchestratorTests(unittest.TestCase):
protocol.call_args.kwargs["endpoint"], "http://ollama.test:11434"
)
def test_diarization_is_off_by_default_and_preserves_protocol_input(self) -> None:
with tempfile.TemporaryDirectory() as directory:
root = Path(directory)
args = self.args(root)
with (
patch.object(mvp_api, "transcribe_audio", side_effect=fake_transcribe),
patch.object(mvp_api, "diarize_audio") as diarization,
patch.object(
mvp_api,
"generate_direct_protocol",
return_value=protocol_result(),
) as protocol,
):
code, run_dir, _ = run_mvp_meeting.run(args)
self.assertEqual(code, 0)
diarization.assert_not_called()
self.assertEqual(protocol.call_args.args[0], run_dir / "transcript/transcript.json")
self.assertFalse((run_dir / "diarization").exists())
def test_diarization_cli_propagates_and_uses_derived_protocol_input(self) -> None:
with tempfile.TemporaryDirectory() as directory:
root = Path(directory)
args = self.args(
root,
[
"--diarization", "gpu",
"--diarization-runtime", "container",
"--diarization-container-image", "rocm/test",
"--diarization-container-arg=--device=/dev/kfd",
],
)
with (
patch.object(
mvp_api, "transcribe_audio", side_effect=fake_transcribe
),
patch.object(
mvp_api, "diarize_audio", side_effect=fake_diarize
) as diarization,
patch.object(
mvp_api,
"generate_direct_protocol",
return_value=protocol_result(),
) as protocol,
):
code, run_dir, _ = run_mvp_meeting.run(args)
self.assertEqual(code, 0)
self.assertEqual(diarization.call_args.args[2], "gpu")
self.assertEqual(diarization.call_args.kwargs["runtime"], "container")
self.assertEqual(
diarization.call_args.kwargs["container_args"], ("--device=/dev/kfd",)
)
derived = run_dir / "diarization/transcript_diarized.json"
self.assertEqual(protocol.call_args.args[0], derived)
self.assertIn("SPEAKER_00", derived.read_text(encoding="utf-8"))
self.assertEqual(
json.loads((run_dir / "transcript/transcript.json").read_text())["text"],
"Besprechungstext.",
)
run_metadata = json.loads((run_dir / "run_metadata.json").read_text())
self.assertTrue(run_metadata["diarization"]["enabled"])
self.assertNotIn("HF_TOKEN", json.dumps(run_metadata))
def test_whisper_failure_is_recorded_and_protocol_is_not_called(self) -> None:
with tempfile.TemporaryDirectory() as directory:
root = Path(directory)
args = self.args(root)
with (
patch.object(
run_mvp_meeting,
mvp_api,
"transcribe_audio",
side_effect=TranscriptionError("whisper stopped"),
),
patch.object(run_mvp_meeting, "generate_direct_protocol") as protocol,
patch.object(mvp_api, "generate_direct_protocol") as protocol,
):
code, run_dir, protocol_path = run_mvp_meeting.run(args)
@@ -176,9 +288,9 @@ class MvpOrchestratorTests(unittest.TestCase):
root = Path(directory)
args = self.args(root)
with (
patch.object(run_mvp_meeting, "transcribe_audio", side_effect=fake_transcribe),
patch.object(mvp_api, "transcribe_audio", side_effect=fake_transcribe),
patch.object(
run_mvp_meeting,
mvp_api,
"generate_direct_protocol",
side_effect=ValueError("generation stopped"),
),
@@ -209,9 +321,9 @@ class MvpOrchestratorTests(unittest.TestCase):
root = Path(directory)
args = self.args(root)
with (
patch.object(run_mvp_meeting, "transcribe_audio", side_effect=fake_transcribe),
patch.object(mvp_api, "transcribe_audio", side_effect=fake_transcribe),
patch.object(
run_mvp_meeting,
mvp_api,
"generate_direct_protocol",
return_value=protocol_result(),
),
@@ -237,7 +349,7 @@ class MvpOrchestratorTests(unittest.TestCase):
"--output-root", str(args.output_root),
]
with patch.object(
run_mvp_meeting,
mvp_api,
"transcribe_audio",
side_effect=TranscriptionError("failed"),
):