Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
33 changes: 27 additions & 6 deletions common/audio/generate_speaker_predictions.py
Original file line number Diff line number Diff line change
@@ -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
Expand All @@ -7,15 +8,20 @@
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.

Args:
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.
Expand Down Expand Up @@ -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
Expand All @@ -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
}
19 changes: 12 additions & 7 deletions common/audio/speakers.py
Original file line number Diff line number Diff line change
Expand Up @@ -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.
Expand Down
3 changes: 2 additions & 1 deletion common/database/postgres_models.py
Original file line number Diff line number Diff line change
@@ -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
Expand All @@ -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
Expand Down
46 changes: 46 additions & 0 deletions tests/test_speakers.py
Original file line number Diff line number Diff line change
@@ -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,
}
]