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
22 changes: 22 additions & 0 deletions src/openbench/pipeline/diarization/speakerkit.py
Original file line number Diff line number Diff line change
Expand Up @@ -33,6 +33,20 @@ class SpeakerKitPipelineConfig(DiarizationPipelineConfig):
cli_path: str = Field(..., description="The absolute path to the SpeakerKit CLI")
model_path: str | None = Field(None, description="The absolute path to the SpeakerKit model directory")
engine: Literal["pyannote", "sortformer"] = Field("pyannote", description="The engine to use")
sortformer_model_version: str | None = Field(
None,
description=(
"Sortformer model version passed as `--sortformer-model-version` (e.g. `v2-1` or `nemotron-3-diarization`). "
"Only applicable when `engine` is `sortformer`; when unset the CLI default is used."
),
)
sortformer_model_variant: str | None = Field(
None,
description=(
"Sortformer model variant passed as `--sortformer-model-variant` (e.g. `384_94MB` or `684_74MB`). "
"Only applicable when `engine` is `sortformer`; when unset the CLI default is used."
),
)

@property
def is_sortformer(self) -> bool:
Expand All @@ -53,6 +67,14 @@ def generate_cli_args(self, inputs: SpeakerKitInput) -> list[str]:
"--verbose",
]

if self.is_sortformer:
# speakerkitpro-cli >= 3.1.6 defaults to Nemotron 3 Diarization; aliases pin the model explicitly so
# they keep evaluating the same model regardless of the CLI default.
if self.sortformer_model_version is not None:
cmd.extend(["--sortformer-model-version", self.sortformer_model_version])
if self.sortformer_model_variant is not None:
cmd.extend(["--sortformer-model-variant", self.sortformer_model_variant])

if self.model_path is not None:
cmd.extend(["--model-path", self.model_path])

Expand Down
24 changes: 23 additions & 1 deletion src/openbench/pipeline/pipeline_aliases.py
Original file line number Diff line number Diff line change
Expand Up @@ -130,13 +130,35 @@ def register_pipeline_aliases() -> None:
"out_dir": "./speakerkit-sortformer-report",
"cli_path": os.getenv("SPEAKERKIT_CLI_PATH"),
"engine": "sortformer",
# Pinned explicitly: speakerkitpro-cli >= 3.1.6 defaults to Nemotron 3 Diarization.
"sortformer_model_version": "v2-1",
"sortformer_model_variant": "384_94MB",
},
description=(
"SpeakerKit speaker diarization pipeline using Sortformer model compressed to 94MB. Requires CLI installation and API key. "
"SpeakerKit speaker diarization pipeline using the Sortformer v2 model compressed to 94MB "
"(`--sortformer-model-version v2-1 --sortformer-model-variant 384_94MB`). Requires CLI installation and API key. "
"Set `SPEAKERKIT_CLI_PATH` and `SPEAKERKIT_API_KEY` env vars. For access to the CLI binary contact speakerkitpro@argmaxinc.com."
),
)

PipelineRegistry.register_alias(
"speakerkit-nemotron-3-diarization",
SpeakerKitPipeline,
default_config={
"out_dir": "./speakerkit-nemotron-3-diarization-report",
"cli_path": os.getenv("SPEAKERKIT_CLI_PATH"),
"engine": "sortformer",
"sortformer_model_version": "nemotron-3-diarization",
"sortformer_model_variant": "684_74MB",
},
description=(
"SpeakerKit speaker diarization pipeline using NVIDIA Nemotron 3 Diarization (Sortformer v3, 74MB variant; "
"`--sortformer-model-version nemotron-3-diarization --sortformer-model-variant 684_74MB`). Requires CLI installation "
"and API key. Set `SPEAKERKIT_CLI_PATH` and `SPEAKERKIT_API_KEY` env vars. "
"For access to the CLI binary contact speakerkitpro@argmaxinc.com."
),
)

PipelineRegistry.register_alias(
"argmax-oss-diarization",
ArgmaxOpenSourceDiarizationPipeline,
Expand Down
51 changes: 51 additions & 0 deletions tests/pipeline/test_speakerkit_cli_args.py
Original file line number Diff line number Diff line change
Expand Up @@ -29,3 +29,54 @@ def test_cli_args_sortformer_num_speakers_and_api_key(monkeypatch):
assert cmd[cmd.index("--model-path") + 1] == "/models"
assert cmd[cmd.index("--num-speakers") + 1] == "3"
assert cmd[cmd.index("--api-key") + 1] == "secret"


def test_cli_args_sortformer_model_version_and_variant(monkeypatch):
monkeypatch.delenv("SPEAKERKIT_API_KEY", raising=False)
cmd = _config(
engine="sortformer",
sortformer_model_version="nemotron-3-diarization",
sortformer_model_variant="684_74MB",
).generate_cli_args(_inputs())
assert cmd[cmd.index("--sortformer-model-version") + 1] == "nemotron-3-diarization"
assert cmd[cmd.index("--sortformer-model-variant") + 1] == "684_74MB"


def test_cli_args_sortformer_model_flags_omitted_when_unset(monkeypatch):
monkeypatch.delenv("SPEAKERKIT_API_KEY", raising=False)
cmd = _config(engine="sortformer").generate_cli_args(_inputs())
assert "--sortformer-model-version" not in cmd
assert "--sortformer-model-variant" not in cmd


def test_cli_args_sortformer_model_flags_ignored_for_pyannote(monkeypatch):
monkeypatch.delenv("SPEAKERKIT_API_KEY", raising=False)
cmd = _config(
engine="pyannote", sortformer_model_version="v2-1", sortformer_model_variant="384_94MB"
).generate_cli_args(_inputs())
assert "--sortformer-model-version" not in cmd
assert "--sortformer-model-variant" not in cmd


def _alias_cli_args(alias: str, monkeypatch) -> list[str]:
import openbench.pipeline # noqa: F401 - importing the package registers the aliases
from openbench.pipeline.pipeline_registry import PipelineRegistry

monkeypatch.delenv("SPEAKERKIT_API_KEY", raising=False)
config = dict(PipelineRegistry.get_alias_info(alias).default_config)
config["cli_path"] = "/opt/speakerkitpro-cli"
return SpeakerKitPipelineConfig(**config).generate_cli_args(_inputs())


def test_sortformer_compressed_alias_pins_v2_model(monkeypatch):
cmd = _alias_cli_args("speakerkit-sortformer-compressed", monkeypatch)
assert cmd[cmd.index("--diarizer") + 1] == "sortformer"
assert cmd[cmd.index("--sortformer-model-version") + 1] == "v2-1"
assert cmd[cmd.index("--sortformer-model-variant") + 1] == "384_94MB"


def test_nemotron_3_diarization_alias_pins_v3_model(monkeypatch):
cmd = _alias_cli_args("speakerkit-nemotron-3-diarization", monkeypatch)
assert cmd[cmd.index("--diarizer") + 1] == "sortformer"
assert cmd[cmd.index("--sortformer-model-version") + 1] == "nemotron-3-diarization"
assert cmd[cmd.index("--sortformer-model-variant") + 1] == "684_74MB"
Loading