1 Commits
Author SHA1 Message Date
admin 405fa6cdf4 Add post-diarization review checkpoint 2026-09-17 19:16:49 +02:00
2 changed files with 85 additions and 0 deletions
+43
View File
@@ -65,6 +65,7 @@ class MvpMeetingConfig:
diarization_runtime: str = "native" diarization_runtime: str = "native"
diarization_container_image: str | None = None diarization_container_image: str | None = None
diarization_container_args: Sequence[str] = () diarization_container_args: Sequence[str] = ()
stop_after_diarization: bool = False
@dataclass(frozen=True) @dataclass(frozen=True)
@@ -130,6 +131,10 @@ def regenerate_mvp_protocol(
) )
protocol_path = persist_protocol_generation(run_dir, result, context=context) protocol_path = persist_protocol_generation(run_dir, result, context=context)
except Exception as exc: except Exception as exc:
_record_protocol_generation_attempt(
run_dir, status="failed", runtime_seconds=time.perf_counter() - started,
failure={"type": type(exc).__name__, "message": str(exc)},
)
_emit( _emit(
progress_sink, progress_sink,
"failed", "failed",
@@ -138,11 +143,40 @@ def regenerate_mvp_protocol(
message=f"protocol_generation: {type(exc).__name__}: {exc}", message=f"protocol_generation: {type(exc).__name__}: {exc}",
) )
raise raise
_record_protocol_generation_attempt(
run_dir, status="completed", runtime_seconds=time.perf_counter() - started
)
_emit(progress_sink, "protocol_generation", "completed", started) _emit(progress_sink, "protocol_generation", "completed", started)
_emit(progress_sink, "completed", "completed", started) _emit(progress_sink, "completed", "completed", started)
return MvpRunResult(0, run_dir, protocol_path) return MvpRunResult(0, run_dir, protocol_path)
def _record_protocol_generation_attempt(
run_dir: Path,
*,
status: str,
runtime_seconds: float,
failure: dict[str, str] | None = None,
) -> None:
"""Record downstream attempts without changing completed upstream artifacts."""
metadata_path = run_dir / "run_metadata.json"
if not metadata_path.is_file():
return
metadata = json.loads(metadata_path.read_text(encoding="utf-8"))
attempts = metadata.setdefault("protocol_generation_attempts", [])
if not isinstance(attempts, list):
attempts = metadata["protocol_generation_attempts"] = []
attempt: dict[str, Any] = {
"status": status,
"runtime_seconds": round(runtime_seconds, 3),
}
if failure is not None:
attempt["failure"] = failure
attempts.append(attempt)
metadata["status"] = "completed" if status == "completed" else "awaiting_speaker_review"
_write_json(metadata_path, metadata)
def create_unique_run_dir( def create_unique_run_dir(
output_root: Path, output_root: Path,
meeting_name: str, meeting_name: str,
@@ -419,6 +453,15 @@ def run_mvp_meeting(
) )
_emit(progress_sink, "diarization", "completed", overall_started) _emit(progress_sink, "diarization", "completed", overall_started)
if config.stop_after_diarization:
metadata["status"] = "awaiting_speaker_review"
metadata["speaker_review"] = {
"status": "awaiting_optional_mapping",
"speaker_count": metadata["diarization"].get("speaker_count"),
}
_emit(progress_sink, "completed", "completed", overall_started)
return MvpRunResult(0, run_dir, None)
current_stage = "protocol_generation" current_stage = "protocol_generation"
stage_started = time.perf_counter() stage_started = time.perf_counter()
_emit(progress_sink, "protocol_generation", "started", overall_started) _emit(progress_sink, "protocol_generation", "started", overall_started)
+42
View File
@@ -9,6 +9,7 @@ from unittest.mock import patch
from scripts import run_mvp_meeting as cli from scripts import run_mvp_meeting as cli
from src.meeting_lab.audio import PreparedAudio from src.meeting_lab.audio import PreparedAudio
from src.meeting_lab.diarization.backend import DiarizationResult
from src.meeting_lab.llm.ollama import OllamaGeneration from src.meeting_lab.llm.ollama import OllamaGeneration
from src.meeting_lab.protocol.generate_direct_protocol import generate_direct_protocol from src.meeting_lab.protocol.generate_direct_protocol import generate_direct_protocol
from src.meeting_lab.models.meeting_context import load_meeting_context from src.meeting_lab.models.meeting_context import load_meeting_context
@@ -82,6 +83,27 @@ def fake_prepare(source, destination, **kwargs):
return PreparedAudio(source, source.suffix.removeprefix("."), destination, "ffmpeg", "ffmpeg") return PreparedAudio(source, source.suffix.removeprefix("."), destination, "ffmpeg", "ffmpeg")
def fake_diarize(audio, output, mode, **kwargs):
output.mkdir(parents=True, exist_ok=True)
metadata = output / "metadata.json"
ordinary = output / "diarization.rttm"
exclusive = output / "exclusive_diarization.rttm"
turns = output / "turns.json"
exclusive_turns = output / "exclusive_turns.json"
details = {"speaker_count": 2, "runtime_seconds": 0.1}
metadata.write_text(json.dumps(details), encoding="utf-8")
ordinary.write_text("", encoding="utf-8")
exclusive.write_text("", encoding="utf-8")
turns.write_text("[]", encoding="utf-8")
exclusive_turns.write_text(
json.dumps([
{"start": 0, "end": 0.5, "speaker_id": "SPEAKER_00"},
{"start": 0.5, "end": 1, "speaker_id": "SPEAKER_01"},
]), encoding="utf-8"
)
return DiarizationResult(output, metadata, ordinary, exclusive, turns, exclusive_turns, details)
def fake_protocol(transcript, context, **kwargs): def fake_protocol(transcript, context, **kwargs):
rendered_context = load_meeting_context(context) rendered_context = load_meeting_context(context)
assert rendered_context.meeting_id == "programmatic-test" assert rendered_context.meeting_id == "programmatic-test"
@@ -107,6 +129,26 @@ class MvpApiTests(unittest.TestCase):
model="test:model", model="test:model",
) )
def test_post_diarization_checkpoint_skips_protocol_generation(self):
with tempfile.TemporaryDirectory() as directory:
root = Path(directory)
config = replace(self.config(root), diarization="cpu", stop_after_diarization=True)
with (
patch.object(mvp_api, "prepare_audio", side_effect=fake_prepare),
patch.object(mvp_api, "transcribe_audio", side_effect=fake_transcribe),
patch.object(mvp_api, "diarize_audio", side_effect=fake_diarize),
patch.object(mvp_api, "generate_direct_protocol") as generate,
):
result = mvp_api.run_mvp_meeting(config, meeting_context=context_data())
self.assertEqual(result.exit_code, 0)
self.assertIsNone(result.protocol_path)
generate.assert_not_called()
metadata = json.loads((result.run_dir / "run_metadata.json").read_text())
self.assertEqual(metadata["status"], "awaiting_speaker_review")
self.assertEqual(metadata["speaker_review"]["speaker_count"], 2)
self.assertTrue((result.run_dir / "diarization/transcript_diarized.json").is_file())
def test_programmatic_context_is_persisted_without_source_yaml(self): def test_programmatic_context_is_persisted_without_source_yaml(self):
with tempfile.TemporaryDirectory() as directory: with tempfile.TemporaryDirectory() as directory:
root = Path(directory) root = Path(directory)