Add diarization and reusable MVP meeting pipeline
This commit is contained in:
@@ -0,0 +1,360 @@
|
||||
"""Reusable audio-to-direct-protocol MVP orchestration."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import re
|
||||
import shutil
|
||||
import sys
|
||||
import time
|
||||
from dataclasses import dataclass
|
||||
from datetime import datetime
|
||||
from pathlib import Path
|
||||
from typing import Any, Callable, Mapping, Sequence
|
||||
|
||||
from src.meeting_lab.diarization import (
|
||||
DEFAULT_MODEL as DEFAULT_DIARIZATION_MODEL,
|
||||
diarize_audio,
|
||||
write_diarized_transcript,
|
||||
)
|
||||
from src.meeting_lab.llm.ollama import DEFAULT_ENDPOINT
|
||||
from src.meeting_lab.models.meeting_context import (
|
||||
MeetingContext,
|
||||
create_meeting_context,
|
||||
load_meeting_context,
|
||||
validate_meeting_context,
|
||||
write_meeting_context,
|
||||
)
|
||||
from src.meeting_lab.progress import ProgressEvent, ProgressSink, ProgressStatus
|
||||
from src.meeting_lab.protocol.generate_direct_protocol import (
|
||||
DEFAULT_MODEL,
|
||||
DirectProtocolResult,
|
||||
generate_direct_protocol,
|
||||
load_compact_transcript,
|
||||
)
|
||||
from src.meeting_lab.transcription.whisper import transcribe_audio
|
||||
|
||||
|
||||
DEFAULT_OUTPUT_ROOT = Path("meeting_data/runs")
|
||||
ContextInput = MeetingContext | Mapping[str, Any]
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class MvpMeetingConfig:
|
||||
audio_file: Path
|
||||
whisper_model: Path
|
||||
whisper_executable: str = "whisper-cli"
|
||||
context_file: Path | None = None
|
||||
output_root: Path = DEFAULT_OUTPUT_ROOT
|
||||
language: str = "de"
|
||||
threads: str | int = "auto"
|
||||
model: str = DEFAULT_MODEL
|
||||
ollama_endpoint: str = DEFAULT_ENDPOINT
|
||||
diarization: str = "off"
|
||||
diarization_runtime: str = "native"
|
||||
diarization_container_image: str | None = None
|
||||
diarization_container_args: Sequence[str] = ()
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class MvpRunResult:
|
||||
exit_code: int
|
||||
run_dir: Path | None
|
||||
protocol_path: Path | None
|
||||
|
||||
|
||||
def create_unique_run_dir(
|
||||
output_root: Path,
|
||||
meeting_name: str,
|
||||
now: Callable[[], datetime] = datetime.now,
|
||||
) -> Path:
|
||||
safe_name = re.sub(r"[^A-Za-z0-9_.-]+", "_", meeting_name).strip("._-") or "meeting"
|
||||
base = output_root / f"{safe_name}_{now().strftime('%Y%m%d_%H%M%S')}"
|
||||
candidate = base
|
||||
suffix = 1
|
||||
while candidate.exists():
|
||||
candidate = output_root / f"{base.name}_{suffix:02d}"
|
||||
suffix += 1
|
||||
candidate.mkdir(parents=True)
|
||||
return candidate
|
||||
|
||||
|
||||
def _write_json(path: Path, value: Any) -> None:
|
||||
path.write_text(json.dumps(value, ensure_ascii=False, indent=2) + "\n", encoding="utf-8")
|
||||
|
||||
|
||||
def _effective_context(value: ContextInput | None) -> MeetingContext | None:
|
||||
if value is None:
|
||||
return None
|
||||
if isinstance(value, MeetingContext):
|
||||
validate_meeting_context(value.data)
|
||||
return value
|
||||
if isinstance(value, Mapping):
|
||||
return create_meeting_context(dict(value), source_file=Path("<programmatic>"))
|
||||
raise TypeError("meeting_context must be MeetingContext, mapping, or None.")
|
||||
|
||||
|
||||
def _validate_inputs(
|
||||
config: MvpMeetingConfig, meeting_context: MeetingContext | None
|
||||
) -> None:
|
||||
if not config.audio_file.is_file():
|
||||
raise FileNotFoundError(f"Audio file does not exist: {config.audio_file}")
|
||||
if not config.whisper_model.is_file():
|
||||
raise FileNotFoundError(f"Whisper model does not exist: {config.whisper_model}")
|
||||
if meeting_context is not None and config.context_file is not None:
|
||||
raise ValueError("Use either a context file or a programmatic Meeting Context, not both.")
|
||||
if meeting_context is not None:
|
||||
validate_meeting_context(meeting_context.data)
|
||||
elif config.context_file is not None:
|
||||
if not config.context_file.is_file():
|
||||
raise FileNotFoundError(
|
||||
f"Meeting Context file does not exist: {config.context_file}"
|
||||
)
|
||||
load_meeting_context(config.context_file)
|
||||
if config.diarization not in ("off", "auto", "gpu", "cpu"):
|
||||
raise ValueError(f"Unsupported diarization mode: {config.diarization}")
|
||||
if config.diarization_runtime not in ("native", "container"):
|
||||
raise ValueError(
|
||||
f"Unsupported diarization runtime: {config.diarization_runtime}"
|
||||
)
|
||||
if (
|
||||
config.diarization != "off"
|
||||
and config.diarization_runtime == "container"
|
||||
and not config.diarization_container_image
|
||||
):
|
||||
raise ValueError("A diarization container image is required.")
|
||||
|
||||
|
||||
def _emit(
|
||||
sink: ProgressSink | None,
|
||||
stage: str,
|
||||
status: ProgressStatus,
|
||||
overall_started: float,
|
||||
*,
|
||||
message: str | None = None,
|
||||
) -> None:
|
||||
if sink is not None:
|
||||
sink(
|
||||
ProgressEvent(
|
||||
stage=stage,
|
||||
status=status,
|
||||
elapsed_seconds=time.perf_counter() - overall_started,
|
||||
message=message,
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
def _persist_protocol(run_dir: Path, result: DirectProtocolResult) -> Path:
|
||||
protocol_dir = run_dir / "protocol"
|
||||
protocol_dir.mkdir(exist_ok=True)
|
||||
(protocol_dir / "exact_prompt.txt").write_text(result.exact_prompt, encoding="utf-8")
|
||||
_write_json(protocol_dir / "raw_response.json", result.raw_response)
|
||||
_write_json(protocol_dir / "runtime_metadata.json", result.runtime_metadata)
|
||||
protocol_path = run_dir / "protocol.md"
|
||||
protocol_path.write_text(result.protocol_text, encoding="utf-8")
|
||||
return protocol_path
|
||||
|
||||
|
||||
def run_mvp_meeting(
|
||||
config: MvpMeetingConfig,
|
||||
*,
|
||||
meeting_context: ContextInput | None = None,
|
||||
progress_sink: ProgressSink | None = None,
|
||||
) -> MvpRunResult:
|
||||
"""Run the existing MVP directly, without subprocess or GUI dependencies."""
|
||||
overall_started = time.perf_counter()
|
||||
validation_started = time.perf_counter()
|
||||
_emit(progress_sink, "preparing", "started", overall_started)
|
||||
try:
|
||||
effective_context = _effective_context(meeting_context)
|
||||
_validate_inputs(config, effective_context)
|
||||
except Exception as exc:
|
||||
_emit(
|
||||
progress_sink,
|
||||
"failed",
|
||||
"failed",
|
||||
overall_started,
|
||||
message=f"preparing: {type(exc).__name__}: {exc}",
|
||||
)
|
||||
print(f"Error: {type(exc).__name__}: {exc}", file=sys.stderr)
|
||||
return MvpRunResult(2, None, None)
|
||||
|
||||
validation_runtime = time.perf_counter() - validation_started
|
||||
run_dir = create_unique_run_dir(config.output_root, config.audio_file.stem)
|
||||
timestamp = datetime.now().astimezone().isoformat(timespec="seconds")
|
||||
transcript_path = run_dir / "transcript" / "transcript.json"
|
||||
protocol_path = run_dir / "protocol.md"
|
||||
stage_runtimes: dict[str, float | None] = {
|
||||
"validation": round(validation_runtime, 3),
|
||||
"setup": None,
|
||||
"whisper": None,
|
||||
"transcript_validation": None,
|
||||
"protocol": None,
|
||||
}
|
||||
if config.diarization != "off":
|
||||
stage_runtimes["diarization"] = None
|
||||
stage_runtimes["diarization_alignment"] = None
|
||||
metadata: dict[str, Any] = {
|
||||
"run_id": run_dir.name,
|
||||
"timestamp": timestamp,
|
||||
"input_audio": str(config.audio_file.resolve()),
|
||||
"transcript_output": str(transcript_path.resolve()),
|
||||
"protocol_output": str(protocol_path.resolve()),
|
||||
"whisper_model": str(config.whisper_model.resolve()),
|
||||
"model": config.model,
|
||||
"ollama_endpoint": config.ollama_endpoint,
|
||||
"status": "running",
|
||||
"stage_runtimes_seconds": stage_runtimes,
|
||||
"total_runtime_seconds": None,
|
||||
"failure": None,
|
||||
"diarization": {
|
||||
"enabled": config.diarization != "off",
|
||||
"backend": "pyannote.audio" if config.diarization != "off" else None,
|
||||
"model": DEFAULT_DIARIZATION_MODEL if config.diarization != "off" else None,
|
||||
"requested_device_mode": config.diarization,
|
||||
"runtime": config.diarization_runtime if config.diarization != "off" else None,
|
||||
"metadata_path": None,
|
||||
"transcript_diarized": None,
|
||||
},
|
||||
}
|
||||
current_stage = "preparing"
|
||||
stage_started = time.perf_counter()
|
||||
|
||||
try:
|
||||
audio_dir = run_dir / "audio"
|
||||
transcript_dir = run_dir / "transcript"
|
||||
context_dir = run_dir / "context"
|
||||
protocol_dir = run_dir / "protocol"
|
||||
audio_dir.mkdir()
|
||||
transcript_dir.mkdir()
|
||||
context_dir.mkdir()
|
||||
protocol_dir.mkdir()
|
||||
_write_json(
|
||||
audio_dir / "input_manifest.json",
|
||||
{
|
||||
"source_file": str(config.audio_file.resolve()),
|
||||
"filename": config.audio_file.name,
|
||||
"size_bytes": config.audio_file.stat().st_size,
|
||||
},
|
||||
)
|
||||
|
||||
preserved_context: Path | None = None
|
||||
if effective_context is not None:
|
||||
preserved_context = context_dir / "meeting_context.yaml"
|
||||
write_meeting_context(effective_context, preserved_context)
|
||||
elif config.context_file is not None:
|
||||
preserved_context = context_dir / "meeting_context.yaml"
|
||||
shutil.copy2(config.context_file, preserved_context)
|
||||
stage_runtimes["setup"] = round(time.perf_counter() - stage_started, 3)
|
||||
_emit(progress_sink, "preparing", "completed", overall_started)
|
||||
|
||||
current_stage = "transcription"
|
||||
stage_started = time.perf_counter()
|
||||
_emit(progress_sink, "transcription", "started", overall_started)
|
||||
transcription = transcribe_audio(
|
||||
config.audio_file,
|
||||
config.whisper_model,
|
||||
transcript_dir,
|
||||
config.language,
|
||||
executable=config.whisper_executable,
|
||||
threads=config.threads,
|
||||
)
|
||||
stage_runtimes["whisper"] = round(time.perf_counter() - stage_started, 3)
|
||||
_emit(progress_sink, "transcription", "completed", overall_started)
|
||||
|
||||
stage_started = time.perf_counter()
|
||||
load_compact_transcript(transcription.transcript_json)
|
||||
stage_runtimes["transcript_validation"] = round(
|
||||
time.perf_counter() - stage_started, 3
|
||||
)
|
||||
|
||||
protocol_transcript = transcription.transcript_json
|
||||
if config.diarization != "off":
|
||||
current_stage = "diarization"
|
||||
stage_started = time.perf_counter()
|
||||
_emit(progress_sink, "diarization", "started", overall_started)
|
||||
diarization_dir = run_dir / "diarization"
|
||||
diarization = diarize_audio(
|
||||
config.audio_file,
|
||||
diarization_dir,
|
||||
config.diarization,
|
||||
runtime=config.diarization_runtime,
|
||||
container_image=config.diarization_container_image,
|
||||
container_args=config.diarization_container_args,
|
||||
)
|
||||
stage_runtimes["diarization"] = round(
|
||||
time.perf_counter() - stage_started, 3
|
||||
)
|
||||
metadata["diarization"].update(
|
||||
{
|
||||
"actual_device": diarization.metadata.get("actual_device"),
|
||||
"device_name": diarization.metadata.get("device_name"),
|
||||
"runtime_seconds": diarization.metadata.get("runtime_seconds"),
|
||||
"speaker_count": diarization.metadata.get("speaker_count"),
|
||||
"metadata_path": str(diarization.metadata_path.resolve()),
|
||||
}
|
||||
)
|
||||
|
||||
stage_started = time.perf_counter()
|
||||
protocol_transcript, diarized_text = write_diarized_transcript(
|
||||
transcription.transcript_json,
|
||||
diarization.exclusive_turns_json,
|
||||
diarization_dir,
|
||||
)
|
||||
load_compact_transcript(protocol_transcript)
|
||||
stage_runtimes["diarization_alignment"] = round(
|
||||
time.perf_counter() - stage_started, 3
|
||||
)
|
||||
metadata["diarization"].update(
|
||||
{
|
||||
"transcript_diarized": str(protocol_transcript.resolve()),
|
||||
"transcript_diarized_text": str(diarized_text.resolve()),
|
||||
}
|
||||
)
|
||||
_emit(progress_sink, "diarization", "completed", overall_started)
|
||||
|
||||
current_stage = "protocol_generation"
|
||||
stage_started = time.perf_counter()
|
||||
_emit(progress_sink, "protocol_generation", "started", overall_started)
|
||||
result = generate_direct_protocol(
|
||||
protocol_transcript,
|
||||
preserved_context,
|
||||
model=config.model,
|
||||
endpoint=config.ollama_endpoint,
|
||||
)
|
||||
stage_runtimes["protocol"] = round(time.perf_counter() - stage_started, 3)
|
||||
protocol_path = _persist_protocol(run_dir, result)
|
||||
_emit(progress_sink, "protocol_generation", "completed", overall_started)
|
||||
metadata["status"] = "completed"
|
||||
_emit(progress_sink, "completed", "completed", overall_started)
|
||||
except Exception as exc:
|
||||
metadata_stage = {
|
||||
"preparing": "setup",
|
||||
"transcription": "whisper",
|
||||
"diarization": "diarization",
|
||||
"protocol_generation": "protocol",
|
||||
}.get(current_stage, current_stage)
|
||||
runtime_key = metadata_stage
|
||||
if runtime_key in stage_runtimes and stage_runtimes[runtime_key] is None:
|
||||
stage_runtimes[runtime_key] = round(time.perf_counter() - stage_started, 3)
|
||||
metadata["status"] = "failed"
|
||||
metadata["failure"] = {
|
||||
"stage": metadata_stage,
|
||||
"type": type(exc).__name__,
|
||||
"message": str(exc),
|
||||
}
|
||||
protocol_path = None
|
||||
_emit(
|
||||
progress_sink,
|
||||
"failed",
|
||||
"failed",
|
||||
overall_started,
|
||||
message=f"{current_stage}: {type(exc).__name__}: {exc}",
|
||||
)
|
||||
print(f"Error: {type(exc).__name__}: {exc}", file=sys.stderr)
|
||||
finally:
|
||||
metadata["total_runtime_seconds"] = round(time.perf_counter() - overall_started, 3)
|
||||
_write_json(run_dir / "run_metadata.json", metadata)
|
||||
|
||||
exit_code = 0 if metadata["status"] == "completed" else 2
|
||||
return MvpRunResult(exit_code, run_dir, protocol_path)
|
||||
Reference in New Issue
Block a user