Skip to content
Merged
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
1 change: 1 addition & 0 deletions config/pipeline_configs/PyannoteApi.yaml
Original file line number Diff line number Diff line change
Expand Up @@ -3,3 +3,4 @@ PyannoteApiPipeline:
out_dir: ./pyannoteapi
timeout: 3600 # 1 hour
request_buffer: 30
model: precision-3
1 change: 1 addition & 0 deletions config/pipeline_configs/PyannoteOrchestration.yaml
Original file line number Diff line number Diff line change
Expand Up @@ -3,3 +3,4 @@ PyannoteOrchestrationPipeline:
out_dir: ./pyannote-orchestration
timeout: 3600 # 1 hour
request_buffer: 30
model: precision-3
1 change: 1 addition & 0 deletions config/pipeline_configs/PyannoteTranscription.yaml
Original file line number Diff line number Diff line change
Expand Up @@ -3,3 +3,4 @@ PyannoteTranscriptionPipeline:
out_dir: ./pyannote-transcription
timeout: 3600 # 1 hour
request_buffer: 30
model: precision-3
2 changes: 2 additions & 0 deletions src/openbench/engine/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -14,6 +14,7 @@
from .openai_engine import OpenAIApi
from .pyannote_engine import (
PyannoteAIApi,
PyannoteAIModel,
PyannoteApiDiarizationOutput,
PyannoteApiOrchestrationOutput,
PyannoteApiSegment,
Expand Down Expand Up @@ -44,6 +45,7 @@
"ElevenLabsApiResponse",
"OpenAIApi",
"PyannoteAIApi",
"PyannoteAIModel",
"PyannoteApiDiarizationOutput",
"PyannoteApiOrchestrationOutput",
"PyannoteApiSegment",
Expand Down
11 changes: 10 additions & 1 deletion src/openbench/engine/pyannote_engine.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand All @@ -18,6 +19,7 @@

__all__ = [
"PyannoteAIApi",
"PyannoteAIModel",
"PyannoteApiDiarizationOutput",
"PyannoteApiOrchestrationOutput",
"PyannoteApiSegment",
Expand All @@ -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."""
Expand Down Expand Up @@ -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"
Expand All @@ -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"):
Expand Down Expand Up @@ -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
Expand Down
10 changes: 9 additions & 1 deletion src/openbench/pipeline/diarization/pyannote_api.py
Original file line number Diff line number Diff line change
Expand Up @@ -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

Expand All @@ -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")
Expand All @@ -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(
Expand Down
10 changes: 9 additions & 1 deletion src/openbench/pipeline/orchestration/orchestration_pyannote.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand All @@ -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
Expand All @@ -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,
)

Expand Down
9 changes: 6 additions & 3 deletions src/openbench/pipeline/pipeline_aliases.py
Original file line number Diff line number Diff line change
Expand Up @@ -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(
Expand Down Expand Up @@ -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(
Expand Down Expand Up @@ -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 #################
Expand Down
10 changes: 9 additions & 1 deletion src/openbench/pipeline/transcription/transcription_pyannote.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand All @@ -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
Expand All @@ -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,
)

Expand Down
Loading