diff --git a/config/pipeline_configs/PyannoteApi.yaml b/config/pipeline_configs/PyannoteApi.yaml index 9e363a0..e739456 100644 --- a/config/pipeline_configs/PyannoteApi.yaml +++ b/config/pipeline_configs/PyannoteApi.yaml @@ -3,3 +3,4 @@ PyannoteApiPipeline: out_dir: ./pyannoteapi timeout: 3600 # 1 hour request_buffer: 30 + model: precision-3 diff --git a/config/pipeline_configs/PyannoteOrchestration.yaml b/config/pipeline_configs/PyannoteOrchestration.yaml index 53b3373..5219a05 100644 --- a/config/pipeline_configs/PyannoteOrchestration.yaml +++ b/config/pipeline_configs/PyannoteOrchestration.yaml @@ -3,3 +3,4 @@ PyannoteOrchestrationPipeline: out_dir: ./pyannote-orchestration timeout: 3600 # 1 hour request_buffer: 30 + model: precision-3 diff --git a/config/pipeline_configs/PyannoteTranscription.yaml b/config/pipeline_configs/PyannoteTranscription.yaml index 99e7995..fb8fa78 100644 --- a/config/pipeline_configs/PyannoteTranscription.yaml +++ b/config/pipeline_configs/PyannoteTranscription.yaml @@ -3,3 +3,4 @@ PyannoteTranscriptionPipeline: out_dir: ./pyannote-transcription timeout: 3600 # 1 hour request_buffer: 30 + model: precision-3 diff --git a/src/openbench/engine/__init__.py b/src/openbench/engine/__init__.py index 677ebf6..ed3e4ac 100644 --- a/src/openbench/engine/__init__.py +++ b/src/openbench/engine/__init__.py @@ -14,6 +14,7 @@ from .openai_engine import OpenAIApi from .pyannote_engine import ( PyannoteAIApi, + PyannoteAIModel, PyannoteApiDiarizationOutput, PyannoteApiOrchestrationOutput, PyannoteApiSegment, @@ -44,6 +45,7 @@ "ElevenLabsApiResponse", "OpenAIApi", "PyannoteAIApi", + "PyannoteAIModel", "PyannoteApiDiarizationOutput", "PyannoteApiOrchestrationOutput", "PyannoteApiSegment", diff --git a/src/openbench/engine/pyannote_engine.py b/src/openbench/engine/pyannote_engine.py index 2d3456c..c4f4fb8 100644 --- a/src/openbench/engine/pyannote_engine.py +++ b/src/openbench/engine/pyannote_engine.py @@ -7,6 +7,7 @@ import time from datetime import datetime from pathlib import Path +from typing import Literal import requests from argmaxtools.utils import get_logger @@ -18,6 +19,7 @@ __all__ = [ "PyannoteAIApi", + "PyannoteAIModel", "PyannoteApiDiarizationOutput", "PyannoteApiOrchestrationOutput", "PyannoteApiSegment", @@ -27,6 +29,9 @@ logger = get_logger(__name__) +# Diarization models exposed by https://docs.pyannote.ai/api-reference/diarize +PyannoteAIModel = Literal["precision-3", "precision-2", "community-1"] + def to_camel(string: str) -> str: """Convert snake_case to camelCase.""" @@ -168,6 +173,8 @@ class PyannoteAIApi: timeout: Timeout for job polling in seconds request_buffer: Buffer for request rate limiting transcription: Whether to enable transcription (STT) in addition to diarization + model: Diarization model to request. Sent explicitly on every job so results do not + change when pyannoteAI rotates the API default. """ diarization_url = "https://api.pyannote.ai/v1/diarize" @@ -179,10 +186,12 @@ def __init__( timeout: int = 1800, request_buffer: int = 30, transcription: bool = False, + model: PyannoteAIModel = "precision-3", ) -> None: self.timeout = timeout self.request_buffer = request_buffer self.transcription = transcription + self.model = model # Check that the API key is set if not os.getenv("PYANNOTE_TOKEN"): @@ -244,7 +253,7 @@ def diarize( Returns: The response from the diarization endpoint """ - data = {"url": audio_url} + data = {"url": audio_url, "model": self.model} if num_speakers is not None: data["numSpeakers"] = num_speakers diff --git a/src/openbench/pipeline/diarization/pyannote_api.py b/src/openbench/pipeline/diarization/pyannote_api.py index c9fd830..df89dfc 100644 --- a/src/openbench/pipeline/diarization/pyannote_api.py +++ b/src/openbench/pipeline/diarization/pyannote_api.py @@ -8,7 +8,7 @@ from pydantic import Field from ...dataset import DiarizationSample -from ...engine import PyannoteAIApi, PyannoteApiDiarizationOutput +from ...engine import PyannoteAIApi, PyannoteAIModel, PyannoteApiDiarizationOutput from ..base import Pipeline, PipelineType, register_pipeline from .common import DiarizationOutput, DiarizationPipelineConfig @@ -27,6 +27,13 @@ class PyannoteApiConfig(DiarizationPipelineConfig): default=30, description="Buffer for the request rate limit", ) + model: PyannoteAIModel = Field( + default="precision-3", + description=( + "PyannoteAI diarization model. Pinned explicitly so results do not change " + "when pyannoteAI rotates the API default." + ), + ) TEMP_AUDIO_DIR = Path("audio_temp") @@ -43,6 +50,7 @@ def build_pipeline( api = PyannoteAIApi( timeout=self.config.timeout, request_buffer=self.config.request_buffer, + model=self.config.model, transcription=False, ) return lambda input_sample: api( diff --git a/src/openbench/pipeline/orchestration/orchestration_pyannote.py b/src/openbench/pipeline/orchestration/orchestration_pyannote.py index b9b3c14..2f737f5 100644 --- a/src/openbench/pipeline/orchestration/orchestration_pyannote.py +++ b/src/openbench/pipeline/orchestration/orchestration_pyannote.py @@ -10,7 +10,7 @@ from pydantic import Field from ...dataset import OrchestrationSample -from ...engine import PyannoteAIApi, PyannoteApiOrchestrationOutput +from ...engine import PyannoteAIApi, PyannoteAIModel, PyannoteApiOrchestrationOutput from ...pipeline_prediction import Transcript from ..base import Pipeline, PipelineType, register_pipeline from .common import OrchestrationConfig, OrchestrationOutput @@ -34,6 +34,13 @@ class PyannoteOrchestrationPipelineConfig(OrchestrationConfig): default=30, description="Buffer for the request rate limit", ) + model: PyannoteAIModel = Field( + default="precision-3", + description=( + "PyannoteAI diarization model. Pinned explicitly so results do not change " + "when pyannoteAI rotates the API default." + ), + ) @register_pipeline @@ -54,6 +61,7 @@ def build_pipeline( api = PyannoteAIApi( timeout=self.config.timeout, request_buffer=self.config.request_buffer, + model=self.config.model, transcription=True, ) diff --git a/src/openbench/pipeline/pipeline_aliases.py b/src/openbench/pipeline/pipeline_aliases.py index 4783dd8..2567d8d 100644 --- a/src/openbench/pipeline/pipeline_aliases.py +++ b/src/openbench/pipeline/pipeline_aliases.py @@ -107,8 +107,9 @@ def register_pipeline_aliases() -> None: "out_dir": "./pyannoteapi", "timeout": 3600, "request_buffer": 30, + "model": "precision-3", }, - description="Pyannote API speaker diarization pipeline. Requires API key from https://www.pyannote.ai/. Set `PYANNOTE_API_KEY` env var.", + description="PyannoteAI API speaker diarization pipeline using the precision-3 model. Requires `PYANNOTE_TOKEN` env var from https://www.pyannote.ai/.", ) PipelineRegistry.register_alias( @@ -401,8 +402,9 @@ def register_pipeline_aliases() -> None: "out_dir": "./pyannote_orchestration_results", "timeout": 3600, "request_buffer": 30, + "model": "precision-3", }, - description="PyannoteAI orchestration pipeline (diarization + transcription). Uses the precision-2 model with Nvidia Parakeet STT. Requires `PYANNOTE_TOKEN` env var from https://www.pyannote.ai/.", + description="PyannoteAI orchestration pipeline (diarization + transcription). Uses the precision-3 model with Nvidia Parakeet STT. Requires `PYANNOTE_TOKEN` env var from https://www.pyannote.ai/.", ) PipelineRegistry.register_alias( @@ -738,8 +740,9 @@ def register_pipeline_aliases() -> None: "out_dir": "./pyannote_transcription_results", "timeout": 3600, "request_buffer": 30, + "model": "precision-3", }, - description="PyannoteAI transcription pipeline (ignores speaker attribution). Uses the precision-2 model with Nvidia Parakeet STT. Requires `PYANNOTE_TOKEN` env var from https://www.pyannote.ai/.", + description="PyannoteAI transcription pipeline (ignores speaker attribution). Uses the precision-3 model with Nvidia Parakeet STT. Requires `PYANNOTE_TOKEN` env var from https://www.pyannote.ai/.", ) ################# SPEECH GENERATION PIPELINES ################# diff --git a/src/openbench/pipeline/transcription/transcription_pyannote.py b/src/openbench/pipeline/transcription/transcription_pyannote.py index ab956e5..a8e254a 100644 --- a/src/openbench/pipeline/transcription/transcription_pyannote.py +++ b/src/openbench/pipeline/transcription/transcription_pyannote.py @@ -10,7 +10,7 @@ from pydantic import Field from ...dataset import TranscriptionSample -from ...engine import PyannoteAIApi, PyannoteApiOrchestrationOutput +from ...engine import PyannoteAIApi, PyannoteAIModel, PyannoteApiOrchestrationOutput from ...pipeline_prediction import Transcript from ..base import Pipeline, PipelineType, register_pipeline from .common import TranscriptionConfig, TranscriptionOutput @@ -34,6 +34,13 @@ class PyannoteTranscriptionPipelineConfig(TranscriptionConfig): default=30, description="Buffer for the request rate limit", ) + model: PyannoteAIModel = Field( + default="precision-3", + description=( + "PyannoteAI diarization model. Pinned explicitly so results do not change " + "when pyannoteAI rotates the API default." + ), + ) @register_pipeline @@ -55,6 +62,7 @@ def build_pipeline( api = PyannoteAIApi( timeout=self.config.timeout, request_buffer=self.config.request_buffer, + model=self.config.model, transcription=True, )