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
50 changes: 50 additions & 0 deletions .github/workflows/tests.yaml
Original file line number Diff line number Diff line change
@@ -0,0 +1,50 @@
name: Tests and Formatting

on:
pull_request:
branches:
- main

jobs:
check-format:
name: Check Format
runs-on: ubuntu-latest

steps:
- name: Checkout code
uses: actions/checkout@v4

- name: Run Ruff linter
uses: astral-sh/ruff-action@v3
with:
args: "check --output-format github src/"

- name: Run Ruff formatter
run: ruff format --check --diff src/

tests:
name: Tests
runs-on: ubuntu-latest

steps:
- name: Checkout code
uses: actions/checkout@v4

- name: Setup uv
uses: astral-sh/setup-uv@v6
with:
version: "0.8.3"
enable-cache: true

- name: Login to Hugging Face
env:
HF_TOKEN: ${{ secrets.HF_TOKEN }}
shell: bash -e {0}
run: |
uv run huggingface-cli login --token "$HF_TOKEN" > /dev/null 2>&1

- name: Run tests
shell: bash -e {0}
run: |
make test
continue-on-error: false
2 changes: 1 addition & 1 deletion Makefile
Original file line number Diff line number Diff line change
Expand Up @@ -5,7 +5,7 @@ setup:
@make install-pre-commit

test:
@uv run pytest tests/ -v
@uv run pytest tests/ -v -p no:warnings

install-pre-commit:
@echo "Installing pre-commit..."
Expand Down
1 change: 1 addition & 0 deletions pyproject.toml
Original file line number Diff line number Diff line change
Expand Up @@ -68,6 +68,7 @@ build-backend = "uv_build"
[tool.ruff]
line-length = 119
unsafe-fixes = false
exclude = ["*.ipynb"]


[tool.ruff.lint]
Expand Down
76 changes: 26 additions & 50 deletions src/openbench/engine/whisperkitpro_engine.py
Original file line number Diff line number Diff line change
Expand Up @@ -25,7 +25,7 @@
# the CLI just the ones that are most commonly used
class WhisperKitProConfig(BaseModel):
"""Configuration for transcription operations.

Supports two modes:
1. Legacy: model_version, model_prefix, model_repo_name
2. New: repo_id, model_variant (downloads models locally)
Expand Down Expand Up @@ -128,9 +128,7 @@ def generate_cli_args(self, model_path: Path | None = None) -> list[str]:
# Use either --model-path (new) or legacy model args
if self.use_model_path:
if model_path is None:
raise ValueError(
"model_path required when using repo_id/model_variant"
)
raise ValueError("model_path required when using repo_id/model_variant")
args = [
"--model-path",
str(model_path),
Expand All @@ -147,19 +145,21 @@ def generate_cli_args(self, model_path: Path | None = None) -> list[str]:
]

# Common args
args.extend([
"--report", # Always generate the report files
"--report-path", # Report path should always be provided
self.report_path,
"--chunking-strategy",
self.chunking_strategy,
"--audio-encoder-compute-units",
COMPUTE_UNITS_MAPPER[self.audio_encoder_compute_units],
"--text-decoder-compute-units",
COMPUTE_UNITS_MAPPER[self.text_decoder_compute_units],
"--fast-load",
str(self.fast_load).lower(),
])
args.extend(
[
"--report", # Always generate the report files
"--report-path", # Report path should always be provided
self.report_path,
"--chunking-strategy",
self.chunking_strategy,
"--audio-encoder-compute-units",
COMPUTE_UNITS_MAPPER[self.audio_encoder_compute_units],
"--text-decoder-compute-units",
COMPUTE_UNITS_MAPPER[self.text_decoder_compute_units],
"--fast-load",
str(self.fast_load).lower(),
]
)

# Add optional args
if self.word_timestamps:
Expand All @@ -186,9 +186,7 @@ def generate_cli_args(self, model_path: Path | None = None) -> list[str]:
@property
def use_model_path(self) -> bool:
"""Check if we should use --model-path vs legacy args."""
return (
self.repo_id is not None and self.model_variant is not None
)
return self.repo_id is not None and self.model_variant is not None

def download_and_prepare_model(self) -> Path:
"""Download model from HuggingFace and prepare folder.
Expand All @@ -197,10 +195,7 @@ def download_and_prepare_model(self) -> Path:
Path to model directory for --model-path
"""
if not self.use_model_path:
raise ValueError(
"download_and_prepare_model requires "
"repo_id and model_variant"
)
raise ValueError("download_and_prepare_model requires repo_id and model_variant")

cache_dir = Path(self.models_cache_dir or "./models_cache")
cache_dir.mkdir(parents=True, exist_ok=True)
Expand All @@ -214,10 +209,7 @@ def download_and_prepare_model(self) -> Path:
logger.info(f"Model already exists at: {model_path}")
return model_path

logger.info(
f"Downloading model from {self.repo_id}, "
f"variant: {self.model_variant}"
)
logger.info(f"Downloading model from {self.repo_id}, variant: {self.model_variant}")

# Download specific model variant folder from HuggingFace
try:
Expand All @@ -232,28 +224,20 @@ def download_and_prepare_model(self) -> Path:
logger.info(f"Model path for CLI: {model_path}")

if not model_path.exists():
raise RuntimeError(
f"Model download succeeded but path doesn't exist: "
f"{model_path}"
)
raise RuntimeError(f"Model download succeeded but path doesn't exist: {model_path}")

return model_path

except Exception as e:
raise RuntimeError(
f"Failed to download model from {self.repo_id}: {e}"
) from e
raise RuntimeError(f"Failed to download model from {self.repo_id}: {e}") from e


class WhisperKitProInput(BaseModel):
"""Input for transcription CLI."""

audio_path: Path
keep_audio: bool = False
custom_vocabulary_path: str | None = Field(
None,
description="Optional path to custom vocabulary file"
)
custom_vocabulary_path: str | None = Field(None, description="Optional path to custom vocabulary file")


class WhisperKitProOutput(BaseModel):
Expand Down Expand Up @@ -287,21 +271,13 @@ def __init__(
# Download and prepare model if using new model management
self.model_path = None
if self.transcription_config.use_model_path:
logger.info(
"Using new model management with repo_id/model_variant"
)
self.model_path = (
self.transcription_config.download_and_prepare_model()
)
logger.info("Using new model management with repo_id/model_variant")
self.model_path = self.transcription_config.download_and_prepare_model()
else:
logger.info("Using legacy model management")

# Generate CLI args (with model_path if available)
self.transcription_args = (
self.transcription_config.generate_cli_args(
model_path=self.model_path
)
)
self.transcription_args = self.transcription_config.generate_cli_args(model_path=self.model_path)
self.transcription_config.create_report_path()

def __call__(self, input: WhisperKitProInput) -> WhisperKitProOutput:
Expand Down
4 changes: 2 additions & 2 deletions src/openbench/metric/word_error_metrics/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -4,7 +4,7 @@
from .english_abbreviations import ABBR
from .text_normalizer import EnglishTextNormalizer
from .word_error_metrics import (
WordErrorRate,
WordDiarizationErrorRate,
ConcatenatedMinimumPermutationWER,
WordDiarizationErrorRate,
WordErrorRate,
)
2 changes: 1 addition & 1 deletion src/openbench/pipeline/transcription/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -8,9 +8,9 @@
from .transcription_groq import GroqTranscriptionConfig, GroqTranscriptionPipeline
from .transcription_nemo import NeMoTranscriptionPipeline, NeMoTranscriptionPipelineConfig
from .transcription_openai import OpenAITranscriptionPipeline, OpenAITranscriptionPipelineConfig
from .transcription_oss_whisper import WhisperOSSTranscriptionPipeline, WhisperOSSTranscriptionPipelineConfig
from .transcription_whisperkitpro import WhisperKitProTranscriptionConfig, WhisperKitProTranscriptionPipeline
from .whisperkit import WhisperKitTranscriptionConfig, WhisperKitTranscriptionPipeline
from .transcription_oss_whisper import WhisperOSSTranscriptionPipeline, WhisperOSSTranscriptionPipelineConfig


__all__ = [
Expand Down
36 changes: 9 additions & 27 deletions src/openbench/pipeline/transcription/transcription_oss_whisper.py
Original file line number Diff line number Diff line change
Expand Up @@ -4,6 +4,7 @@
from pathlib import Path
from typing import Callable

import whisper
from argmaxtools.utils import get_logger
from pydantic import Field

Expand All @@ -12,7 +13,6 @@
from ...pipeline_prediction import Transcript
from ...types import PipelineType
from .common import TranscriptionConfig, TranscriptionOutput
import whisper


logger = get_logger(__name__)
Expand All @@ -21,9 +21,7 @@


class WhisperOSSApi:
def __init__(
self, model_version: str = "base", device: str | None = None
):
def __init__(self, model_version: str = "base", device: str | None = None):
"""
Initialize OpenAI Whisper (open-source) engine.

Expand All @@ -42,15 +40,10 @@ def __init__(
else:
self.device = device

logger.info(
f"Loading Whisper model '{model_version}' "
f"on device '{self.device}'"
)
logger.info(f"Loading Whisper model '{model_version}' on device '{self.device}'")

# Load the model with download_root to suppress progress bars
self.model = whisper.load_model(
model_version, device=self.device, download_root=None
)
self.model = whisper.load_model(model_version, device=self.device, download_root=None)

def transcribe(
self,
Expand Down Expand Up @@ -86,9 +79,7 @@ def transcribe(
logger.debug(f"Using language: {language}")

# Transcribe
result = self.model.transcribe(
str(audio_path), **transcribe_options
)
result = self.model.transcribe(str(audio_path), **transcribe_options)

# Extract words and timestamps
words = []
Expand Down Expand Up @@ -121,16 +112,11 @@ def transcribe(
class WhisperOSSTranscriptionPipelineConfig(TranscriptionConfig):
model_version: str = Field(
default="base",
description=(
"The version of the Whisper model to use "
"(tiny, base, small, medium, large, turbo)"
),
description=("The version of the Whisper model to use (tiny, base, small, medium, large, turbo)"),
)
device: str | None = Field(
default=None,
description=(
"Device to use for inference (cuda, cpu, mps). "
),
description=("Device to use for inference (cuda, cpu, mps). "),
)


Expand All @@ -140,9 +126,7 @@ class WhisperOSSTranscriptionPipeline(Pipeline):
pipeline_type = PipelineType.TRANSCRIPTION

def build_pipeline(self) -> Callable[[Path], dict]:
whisper_api = WhisperOSSApi(
model_version=self.config.model_version, device=self.config.device
)
whisper_api = WhisperOSSApi(model_version=self.config.model_version, device=self.config.device)

def transcribe(audio_path: Path) -> dict:
language = None
Expand Down Expand Up @@ -173,9 +157,7 @@ def parse_input(self, input_sample: TranscriptionSample) -> Path:
# Extract language if force_language is enabled
self.current_language = None
if self.config.force_language:
self.current_language = input_sample.extra_info.get(
"language", None
)
self.current_language = input_sample.extra_info.get("language", None)

return input_sample.save_audio(TEMP_AUDIO_DIR)

Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -93,12 +93,8 @@ def build_pipeline(self) -> WhisperKitPro:
repo_id=self.config.repo_id,
model_variant=self.config.model_variant,
models_cache_dir=self.config.models_cache_dir,
audio_encoder_compute_units=(
self.config.audio_encoder_compute_units
),
text_decoder_compute_units=(
self.config.text_decoder_compute_units
),
audio_encoder_compute_units=(self.config.audio_encoder_compute_units),
text_decoder_compute_units=(self.config.text_decoder_compute_units),
report_path="whisperkitpro_transcription_reports",
word_timestamps=True,
chunking_strategy="vad",
Expand Down
33 changes: 0 additions & 33 deletions tests/metric/word_error_metric/wder/Dockerfile

This file was deleted.

Loading
Loading