diff --git a/common/audio/generate_speaker_predictions.py b/common/audio/generate_speaker_predictions.py index ad08d1e2..a06308ee 100644 --- a/common/audio/generate_speaker_predictions.py +++ b/common/audio/generate_speaker_predictions.py @@ -1,4 +1,5 @@ import logging +from typing import TypedDict from common.format_transcript import transcript_as_speaker_and_utterance from common.llm.client import FastOrBestLLM, create_default_chatbot @@ -7,7 +8,12 @@ logger = logging.getLogger(__name__) -async def generate_speaker_predictions(dialogue_entries: list[DialogueEntry]) -> dict[str, str]: +class SpeakerPredictionResult(TypedDict): + predicted_name: str + confidence: float | None + + +async def generate_speaker_predictions(dialogue_entries: list[DialogueEntry]) -> dict[str, SpeakerPredictionResult]: """ Generate speaker name predictions based on dialogue entries. @@ -15,7 +21,7 @@ async def generate_speaker_predictions(dialogue_entries: list[DialogueEntry]) -> dialogue_entries: List of DialogueEntry objects containing speaker and text Returns: - Dictionary mapping original speaker labels to predicted names + Dictionary mapping original speaker labels to predicted names and confidence """ # Create a system message that explains the task system_message = """You are an expert at analysing conversation transcripts and identifying speakers. @@ -44,9 +50,18 @@ async def generate_speaker_predictions(dialogue_entries: list[DialogueEntry]) -> if not speaker_prediction.predictions: logger.warning("No predictions found, returning original speaker labels") - return {entry["speaker"]: entry["speaker"] for entry in dialogue_entries} + return { + entry["speaker"]: SpeakerPredictionResult(predicted_name=entry["speaker"], confidence=None) + for entry in dialogue_entries + } - return {pred.original_speaker: pred.predicted_name for pred in speaker_prediction.predictions} + return { + pred.original_speaker: SpeakerPredictionResult( + predicted_name=pred.predicted_name, + confidence=pred.confidence, + ) + for pred in speaker_prediction.predictions + } except Exception as e: # noqa: BLE001 # flagged by ruff - investigate when we have time. error_message = str(e) # Check for content filter errors from Azure OpenAI @@ -66,8 +81,14 @@ async def generate_speaker_predictions(dialogue_entries: list[DialogueEntry]) -> ) # Return original speaker labels - return {entry["speaker"]: entry["speaker"] for entry in dialogue_entries} + return { + entry["speaker"]: SpeakerPredictionResult(predicted_name=entry["speaker"], confidence=None) + for entry in dialogue_entries + } else: # For other errors, log and return original speaker labels logger.error("Error predicting speaker names: %s", error_message) - return {entry["speaker"]: entry["speaker"] for entry in dialogue_entries} + return { + entry["speaker"]: SpeakerPredictionResult(predicted_name=entry["speaker"], confidence=None) + for entry in dialogue_entries + } diff --git a/common/audio/speakers.py b/common/audio/speakers.py index edc998d9..c855eee2 100644 --- a/common/audio/speakers.py +++ b/common/audio/speakers.py @@ -126,14 +126,19 @@ async def process_speakers_and_dialogue_entries( # Step 5: Update entries with predicted names predicted_entries = [] for entry in labelled_dialogue_entries: - predicted_entries.append( - DialogueEntry( - speaker=speaker_predictions.get(entry["speaker"], entry["speaker"]), - text=entry["text"], - start_time=entry["start_time"], - end_time=entry["end_time"], - ) + prediction = speaker_predictions.get( + entry["speaker"], + {"predicted_name": entry["speaker"], "confidence": None}, ) + dialogue_entry = DialogueEntry( + speaker=prediction["predicted_name"], + text=entry["text"], + start_time=entry["start_time"], + end_time=entry["end_time"], + ) + if prediction["confidence"] is not None: + dialogue_entry["speaker_confidence"] = prediction["confidence"] + predicted_entries.append(dialogue_entry) return predicted_entries except Exception as e: # noqa: BLE001 # flagged by ruff - investigate when we have time. diff --git a/common/database/postgres_models.py b/common/database/postgres_models.py index 7aa5f60c..b5f628b1 100644 --- a/common/database/postgres_models.py +++ b/common/database/postgres_models.py @@ -1,6 +1,6 @@ from datetime import datetime from enum import StrEnum, auto -from typing import TypedDict +from typing import NotRequired, TypedDict from uuid import UUID, uuid4 from sqlalchemy import TIMESTAMP, Column @@ -15,6 +15,7 @@ class DialogueEntry(TypedDict): text: str start_time: float end_time: float + speaker_confidence: NotRequired[float] # Create factory functions for columns to avoid reusing column objects diff --git a/tests/test_speakers.py b/tests/test_speakers.py new file mode 100644 index 00000000..c4613f5f --- /dev/null +++ b/tests/test_speakers.py @@ -0,0 +1,46 @@ +from unittest.mock import AsyncMock, patch + +import pytest + +from common.audio.speakers import process_speakers_and_dialogue_entries + + +@pytest.mark.asyncio +async def test_process_speakers_persists_prediction_confidence(): + dialogue_entries = [ + { + "speaker": "raw-speaker-a", + "text": "Hello from Alice", + "start_time": 0.0, + "end_time": 1.0, + }, + { + "speaker": "raw-speaker-a", + "text": "Continuing", + "start_time": 1.0, + "end_time": 2.0, + }, + ] + + with patch( + "common.audio.speakers.generate_speaker_predictions", + new=AsyncMock( + return_value={ + "Unknown speaker 0": { + "predicted_name": "Alice", + "confidence": 0.91, + } + } + ), + ): + result = await process_speakers_and_dialogue_entries(dialogue_entries) + + assert result == [ + { + "speaker": "Alice", + "text": "Hello from Alice Continuing", + "start_time": 0.0, + "end_time": 2.0, + "speaker_confidence": 0.91, + } + ]