Add diarization and reusable MVP meeting pipeline
This commit is contained in:
+124
-12
@@ -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"),
|
||||
):
|
||||
|
||||
Reference in New Issue
Block a user