Skip to content

Commit 9ebe0e9

Browse files
authored
Add Argmax Sortformer (#96)
* feat: Argmax Sortformer diarization pipelines * refactor: sortformer orchestration * fix: sortformer usage with speakerkit * refactor: add verbose to whisperkitpro * refactor: only add --diarization-mode for sortformer * refactor: CLI args * fix: missing is_sortformer * chore: ignore .sh
1 parent a529f39 commit 9ebe0e9

5 files changed

Lines changed: 143 additions & 59 deletions

File tree

‎.gitignore‎

Lines changed: 3 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -45,4 +45,6 @@ inference_outputs/
4545
miscellaneous/
4646

4747
# Default openbench-cli output directory
48-
downloaded_datasets/
48+
downloaded_datasets/
49+
50+
*.sh

‎src/openbench/engine/whisperkitpro_engine.py‎

Lines changed: 19 additions & 8 deletions
Original file line numberDiff line numberDiff line change
@@ -92,24 +92,28 @@ class WhisperKitProConfig(BaseModel):
9292
description="The compute units to use for the audio encoder. Default is CPU_AND_NE.",
9393
)
9494
text_decoder_compute_units: ct.ComputeUnit = Field(
95-
ct.ComputeUnit.CPU_AND_GPU,
95+
ct.ComputeUnit.CPU_AND_NE,
9696
description="The compute units to use for the text decoder. Default is CPU_AND_GPU.",
9797
)
9898
diarization: bool = Field(
9999
False,
100100
description="Whether to perform diarization",
101101
)
102-
orchestration_strategy: Literal["word", "segment"] = Field(
103-
"segment",
104-
description="The orchestration strategy to use either `word` or `segment`",
102+
diarization_mode: Literal["realtime", "prerecorded"] = Field(
103+
"prerecorded",
104+
description="Sortformer streaming mode: `realtime` (1.04s latency) or `prerecorded` (9.84s latency). This is only applicable when `engine` is `sortformer`.",
105+
)
106+
orchestration_strategy: Literal["segment", "subsegment"] = Field(
107+
"subsegment",
108+
description="The orchestration strategy to use either `segment` or `subsegment`",
105109
)
106110
speaker_models_path: str | None = Field(
107111
None,
108112
description="The path to the speaker models directory",
109113
)
110-
clusterer_version: Literal["pyannote3", "pyannote4"] = Field(
111-
"pyannote4",
112-
description="The version of the clusterer to use",
114+
engine: Literal["pyannote", "sortformer"] = Field(
115+
"pyannote",
116+
description="The engine to use. If `sortformer` the diarization model used is Sortformer, otherwise it is pyannote.",
113117
)
114118
use_exclusive_reconciliation: bool = Field(
115119
False,
@@ -167,6 +171,7 @@ def generate_cli_args(self, model_path: Path | None = None) -> list[str]:
167171
COMPUTE_UNITS_MAPPER[self.text_decoder_compute_units],
168172
"--fast-load",
169173
str(self.fast_load).lower(),
174+
"--verbose",
170175
]
171176
)
172177

@@ -180,9 +185,15 @@ def generate_cli_args(self, model_path: Path | None = None) -> list[str]:
180185
if self.diarization:
181186
args.extend(["--diarization"])
182187
args.extend(["--orchestration-strategy", self.orchestration_strategy])
188+
183189
# Add rttm path
184190
args.extend(["--rttm-path", self.rttm_path])
185-
args.extend(["--clusterer-version", self.clusterer_version])
191+
args.extend(["--engine", self.engine])
192+
193+
# Only add diarization mode if using Sortformer
194+
if self.engine == "sortformer":
195+
args.extend(["--diarization-mode", self.diarization_mode])
196+
186197
# If speaker models path is provided use it
187198
if self.speaker_models_path:
188199
args.extend(["--speaker-models-path", self.speaker_models_path])

‎src/openbench/pipeline/diarization/speakerkit.py‎

Lines changed: 43 additions & 25 deletions
Original file line numberDiff line numberDiff line change
@@ -23,48 +23,68 @@
2323
TEMP_AUDIO_DIR = Path("audio_temp")
2424

2525

26-
class SpeakerKitPipelineConfig(DiarizationPipelineConfig):
27-
cli_path: str = Field(..., description="The absolute path to the SpeakerKit CLI")
28-
clusterer_version: Literal["pyannote3", "pyannote4"] = Field(
29-
"pyannote4", description="The version of the clusterer to use"
30-
)
31-
model_path: str | None = Field(None, description="The absolute path to the SpeakerKit model")
32-
33-
3426
class SpeakerKitInput(TypedDict):
3527
audio_path: Path
3628
output_path: Path
3729
num_speakers: int | None
3830

3931

40-
class SpeakerKitCli:
41-
def __init__(self, config: SpeakerKitPipelineConfig):
42-
self.cli_path = config.cli_path
43-
self.model_path = config.model_path
44-
self.clusterer_version = config.clusterer_version
32+
class SpeakerKitPipelineConfig(DiarizationPipelineConfig):
33+
cli_path: str = Field(..., description="The absolute path to the SpeakerKit CLI")
34+
model_path: str | None = Field(None, description="The absolute path to the SpeakerKit model directory")
35+
engine: Literal["pyannote", "sortformer"] = Field("pyannote", description="The engine to use")
4536

46-
def __call__(self, speakerkit_input: SpeakerKitInput) -> tuple[Path, float]:
37+
@property
38+
def is_sortformer(self) -> bool:
39+
return self.engine == "sortformer"
40+
41+
def generate_cli_args(self, inputs: SpeakerKitInput) -> list[str]:
4742
cmd = [
4843
self.cli_path,
4944
"diarize",
5045
"--audio-path",
51-
str(speakerkit_input["audio_path"]),
46+
str(inputs["audio_path"]),
5247
"--rttm-path",
53-
str(speakerkit_input["output_path"]),
54-
"--clusterer-version",
55-
self.clusterer_version,
48+
str(inputs["output_path"]),
49+
"--engine",
50+
self.engine,
5651
"--verbose",
5752
]
5853

59-
if self.model_path:
54+
if self.model_path is not None:
6055
cmd.extend(["--model-path", self.model_path])
6156

62-
if speakerkit_input["num_speakers"] is not None:
63-
cmd.extend(["--num-speakers", str(speakerkit_input["num_speakers"])])
57+
if inputs["num_speakers"] is not None:
58+
cmd.extend(["--num-speakers", str(inputs["num_speakers"])])
6459

6560
if "SPEAKERKIT_API_KEY" in os.environ:
6661
cmd.extend(["--api-key", os.environ["SPEAKERKIT_API_KEY"]])
6762

63+
return cmd
64+
65+
def parse_stdout(self, stdout: str) -> float:
66+
# Default pattern for pyannote models
67+
pattern = r"Model Load Time:\s+\d+\.\d+\s+ms\nTotal Time:\s+(\d+\.\d+)\s+ms"
68+
divisor = 1000.0
69+
70+
# if model is sortfomer we override the pattern and divisor
71+
if self.is_sortformer:
72+
pattern = r"Prediction time:\s+(\d+\.\d+)\s+seconds"
73+
divisor = 1.0
74+
75+
matches = re.search(pattern, stdout)
76+
if matches is None:
77+
raise ValueError(f"Could not parse prediction time from stdout: {stdout!r}")
78+
return float(matches.group(1)) / divisor
79+
80+
81+
class SpeakerKitCli:
82+
def __init__(self, config: SpeakerKitPipelineConfig):
83+
self.config = config
84+
85+
def __call__(self, speakerkit_input: SpeakerKitInput) -> tuple[Path, float]:
86+
cmd = self.config.generate_cli_args(speakerkit_input)
87+
6888
try:
6989
result = subprocess.run(cmd, check=True, capture_output=True, text=True)
7090
logger.debug(f"Diarization CLI stdout:\n{result.stdout}")
@@ -81,11 +101,9 @@ def __call__(self, speakerkit_input: SpeakerKitInput) -> tuple[Path, float]:
81101
speakerkit_input["audio_path"].unlink()
82102

83103
# Parse stdout and take the total time it took to diarize
84-
pattern = r"Model Load Time:\s+\d+\.\d+\s+ms\nTotal Time:\s+(\d+\.\d+)\s+ms"
85-
matches = re.search(pattern, result.stdout)
86-
total_time = float(matches.group(1))
104+
total_time = self.config.parse_stdout(result.stdout)
87105

88-
return speakerkit_input["output_path"], total_time / 1000
106+
return speakerkit_input["output_path"], total_time
89107

90108

91109
@register_pipeline

‎src/openbench/pipeline/orchestration/orchestration_whisperkitpro.py‎

Lines changed: 12 additions & 7 deletions
Original file line numberDiff line numberDiff line change
@@ -69,13 +69,17 @@ class WhisperKitProOrchestrationConfig(OrchestrationConfig):
6969
ComputeUnit.CPU_AND_NE,
7070
description="The compute units to use for the text decoder. Default is CPU_AND_NE.",
7171
)
72-
orchestration_strategy: Literal["word", "segment"] = Field(
73-
"segment",
74-
description="The orchestration strategy to use either `word` or `segment`",
72+
orchestration_strategy: Literal["segment", "subsegment"] = Field(
73+
"subsegment",
74+
description="The orchestration strategy to use either `segment` or `subsegment`",
7575
)
76-
clusterer_version: Literal["pyannote3", "pyannote4"] = Field(
77-
"pyannote4",
78-
description="The version of the clusterer to use",
76+
engine: Literal["pyannote", "sortformer"] = Field(
77+
"pyannote",
78+
description="The engine to use. If `sortformer` the diarization model used is Sortformer, otherwise it is pyannote.",
79+
)
80+
diarization_mode: Literal["realtime", "prerecorded"] = Field(
81+
"prerecorded",
82+
description="Sortformer streaming mode: `realtime` (1.04s latency) or `prerecorded` (9.84s latency). This is only applicable when `engine` is `sortformer`.",
7983
)
8084
use_exclusive_reconciliation: bool = Field(
8185
False,
@@ -107,7 +111,8 @@ def build_pipeline(self) -> WhisperKitPro:
107111
chunking_strategy="vad",
108112
diarization=True,
109113
orchestration_strategy=self.config.orchestration_strategy,
110-
clusterer_version_string=self.config.clusterer_version,
114+
engine=self.config.engine,
115+
diarization_mode=self.config.diarization_mode,
111116
use_exclusive_reconciliation=self.config.use_exclusive_reconciliation,
112117
fast_load=self.config.fast_load,
113118
)

‎src/openbench/pipeline/pipeline_aliases.py‎

Lines changed: 66 additions & 18 deletions
Original file line numberDiff line numberDiff line change
@@ -112,9 +112,23 @@ def register_pipeline_aliases() -> None:
112112
default_config={
113113
"out_dir": "./speakerkit-report",
114114
"cli_path": os.getenv("SPEAKERKIT_CLI_PATH"),
115-
"clusterer_version": "pyannote4",
115+
"engine": "pyannote",
116116
},
117-
description="SpeakerKit speaker diarization pipeline. 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",
117+
description="SpeakerKit speaker diarization pipeline using community-1 model from pyannote. 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",
118+
)
119+
120+
PipelineRegistry.register_alias(
121+
"speakerkit-sortformer-compressed",
122+
SpeakerKitPipeline,
123+
default_config={
124+
"out_dir": "./speakerkit-sortformer-report",
125+
"cli_path": os.getenv("SPEAKERKIT_CLI_PATH"),
126+
"engine": "sortformer",
127+
},
128+
description=(
129+
"SpeakerKit speaker diarization pipeline using Sortformer model compressed to 94MB. Requires CLI installation and API key. "
130+
"Set `SPEAKERKIT_CLI_PATH` and `SPEAKERKIT_API_KEY` env vars. For access to the CLI binary contact speakerkitpro@argmaxinc.com."
131+
),
118132
)
119133

120134
PipelineRegistry.register_alias(
@@ -203,8 +217,8 @@ def register_pipeline_aliases() -> None:
203217
"repo_id": "argmaxinc/whisperkit-pro",
204218
"model_variant": "openai_whisper-tiny",
205219
"cli_path": os.getenv("WHISPERKITPRO_CLI_PATH"),
206-
"orchestration_strategy": "segment",
207-
"clusterer_version_string": "pyannote4",
220+
"orchestration_strategy": "subsegment",
221+
"engine": "pyannote",
208222
"use_exclusive_reconciliation": True,
209223
},
210224
description="WhisperKitPro orchestration pipeline using the tiny version of the model. Requires `WHISPERKITPRO_CLI_PATH` env var and depending on your permissions also `WHISPERKITPRO_API_KEY` env var.",
@@ -217,8 +231,8 @@ def register_pipeline_aliases() -> None:
217231
"repo_id": "argmaxinc/whisperkit-pro",
218232
"model_variant": "openai_whisper-large-v3",
219233
"cli_path": os.getenv("WHISPERKITPRO_CLI_PATH"),
220-
"orchestration_strategy": "segment",
221-
"clusterer_version_string": "pyannote4",
234+
"orchestration_strategy": "subsegment",
235+
"engine": "pyannote",
222236
"use_exclusive_reconciliation": True,
223237
},
224238
description="WhisperKitPro orchestration pipeline using the large-v3 version of the model. Requires `WHISPERKITPRO_CLI_PATH` env var and depending on your permissions also `WHISPERKITPRO_API_KEY` env var.",
@@ -231,8 +245,8 @@ def register_pipeline_aliases() -> None:
231245
"repo_id": "argmaxinc/whisperkit-pro",
232246
"model_variant": "openai_whisper-large-v3-v20240930",
233247
"cli_path": os.getenv("WHISPERKITPRO_CLI_PATH"),
234-
"orchestration_strategy": "segment",
235-
"clusterer_version_string": "pyannote4",
248+
"orchestration_strategy": "subsegment",
249+
"engine": "pyannote",
236250
"use_exclusive_reconciliation": True,
237251
},
238252
description="WhisperKitPro orchestration pipeline using the large-v3-v20240930 version of the model (which is the same as large-v3-turbo from OpenAI). Requires `WHISPERKITPRO_CLI_PATH` env var and depending on your permissions also `WHISPERKITPRO_API_KEY` env var.",
@@ -245,8 +259,8 @@ def register_pipeline_aliases() -> None:
245259
"repo_id": "argmaxinc/whisperkit-pro",
246260
"model_variant": "openai_whisper-large-v3-v20240930_626MB",
247261
"cli_path": os.getenv("WHISPERKITPRO_CLI_PATH"),
248-
"orchestration_strategy": "segment",
249-
"clusterer_version_string": "pyannote4",
262+
"orchestration_strategy": "subsegment",
263+
"engine": "pyannote",
250264
"use_exclusive_reconciliation": True,
251265
},
252266
description="WhisperKitPro orchestration pipeline using the large-v3-v20240930 version of the model compressed to 626MB (which is the same as large-v3-turbo from OpenAI). Requires `WHISPERKITPRO_CLI_PATH` env var and depending on your permissions also `WHISPERKITPRO_API_KEY` env var.",
@@ -259,8 +273,8 @@ def register_pipeline_aliases() -> None:
259273
"repo_id": "argmaxinc/parakeetkit-pro",
260274
"model_variant": "nvidia_parakeet-v2",
261275
"cli_path": os.getenv("WHISPERKITPRO_CLI_PATH"),
262-
"orchestration_strategy": "segment",
263-
"clusterer_version_string": "pyannote4",
276+
"orchestration_strategy": "subsegment",
277+
"engine": "pyannote",
264278
"use_exclusive_reconciliation": True,
265279
},
266280
description="WhisperKitPro orchestration pipeline using the parakeet-v2 version of the model. Requires `WHISPERKITPRO_CLI_PATH` env var and depending on your permissions also `WHISPERKITPRO_API_KEY` env var.",
@@ -273,22 +287,39 @@ def register_pipeline_aliases() -> None:
273287
"repo_id": "argmaxinc/parakeetkit-pro",
274288
"model_variant": "nvidia_parakeet-v2_476MB",
275289
"cli_path": os.getenv("WHISPERKITPRO_CLI_PATH"),
276-
"orchestration_strategy": "segment",
277-
"clusterer_version_string": "pyannote4",
290+
"orchestration_strategy": "subsegment",
291+
"engine": "pyannote",
278292
"use_exclusive_reconciliation": True,
279293
},
280294
description="WhisperKitPro orchestration pipeline using the parakeet-v2 version of the model compressed to 476MB. Requires `WHISPERKITPRO_CLI_PATH` env var and depending on your permissions also `WHISPERKITPRO_API_KEY` env var.",
281295
)
282296

297+
PipelineRegistry.register_alias(
298+
"whisperkitpro-orchestration-parakeet-v2-compressed-sortformer-compressed",
299+
WhisperKitProOrchestrationPipeline,
300+
default_config={
301+
"repo_id": "argmaxinc/parakeetkit-pro",
302+
"model_variant": "nvidia_parakeet-v2_476MB",
303+
"cli_path": os.getenv("WHISPERKITPRO_CLI_PATH"),
304+
"orchestration_strategy": "subsegment",
305+
"engine": "sortformer",
306+
"diarization_mode": "prerecorded",
307+
},
308+
description=(
309+
"WhisperKitPro orchestration pipeline using the parakeet-v2 version of the model compressed to 476MB and using Sortformer for diarization. "
310+
"Requires `WHISPERKITPRO_CLI_PATH` env var and depending on your permissions also `WHISPERKITPRO_API_KEY` env var."
311+
),
312+
)
313+
283314
PipelineRegistry.register_alias(
284315
"whisperkitpro-orchestration-parakeet-v3",
285316
WhisperKitProOrchestrationPipeline,
286317
default_config={
287318
"repo_id": "argmaxinc/parakeetkit-pro",
288319
"model_variant": "nvidia_parakeet-v3",
289320
"cli_path": os.getenv("WHISPERKITPRO_CLI_PATH"),
290-
"orchestration_strategy": "segment",
291-
"clusterer_version_string": "pyannote4",
321+
"orchestration_strategy": "subsegment",
322+
"engine": "pyannote",
292323
"use_exclusive_reconciliation": True,
293324
},
294325
description="WhisperKitPro orchestration pipeline using the parakeet-v3 version of the model. Requires `WHISPERKITPRO_CLI_PATH` env var and depending on your permissions also `WHISPERKITPRO_API_KEY` env var.",
@@ -301,13 +332,30 @@ def register_pipeline_aliases() -> None:
301332
"repo_id": "argmaxinc/parakeetkit-pro",
302333
"model_variant": "nvidia_parakeet-v3_494MB",
303334
"cli_path": os.getenv("WHISPERKITPRO_CLI_PATH"),
304-
"orchestration_strategy": "segment",
305-
"clusterer_version_string": "pyannote4",
335+
"orchestration_strategy": "subsegment",
336+
"engine": "pyannote",
306337
"use_exclusive_reconciliation": True,
307338
},
308339
description="WhisperKitPro orchestration pipeline using the parakeet-v3 version of the model compressed to 494MB. Requires `WHISPERKITPRO_CLI_PATH` env var and depending on your permissions also `WHISPERKITPRO_API_KEY` env var.",
309340
)
310341

342+
PipelineRegistry.register_alias(
343+
"whisperkitpro-orchestration-parakeet-v3-compressed-sortformer-compressed",
344+
WhisperKitProOrchestrationPipeline,
345+
default_config={
346+
"repo_id": "argmaxinc/parakeetkit-pro",
347+
"model_variant": "nvidia_parakeet-v3_494MB",
348+
"cli_path": os.getenv("WHISPERKITPRO_CLI_PATH"),
349+
"orchestration_strategy": "subsegment",
350+
"engine": "sortformer",
351+
"diarization_mode": "prerecorded",
352+
},
353+
description=(
354+
"WhisperKitPro orchestration pipeline using the parakeet-v3 version of the model compressed to 494MB and using Sortformer for diarization. "
355+
"Requires `WHISPERKITPRO_CLI_PATH` env var and depending on your permissions also `WHISPERKITPRO_API_KEY` env var."
356+
),
357+
)
358+
311359
PipelineRegistry.register_alias(
312360
"openai-orchestration",
313361
OpenAIOrchestrationPipeline,

0 commit comments

Comments
 (0)