diff --git a/.github/workflows/tests.yaml b/.github/workflows/tests.yaml new file mode 100644 index 0000000..5711e59 --- /dev/null +++ b/.github/workflows/tests.yaml @@ -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 \ No newline at end of file diff --git a/Makefile b/Makefile index f339e88..1660fbb 100644 --- a/Makefile +++ b/Makefile @@ -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..." diff --git a/pyproject.toml b/pyproject.toml index 936c0ab..4719300 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -68,6 +68,7 @@ build-backend = "uv_build" [tool.ruff] line-length = 119 unsafe-fixes = false +exclude = ["*.ipynb"] [tool.ruff.lint] diff --git a/src/openbench/engine/whisperkitpro_engine.py b/src/openbench/engine/whisperkitpro_engine.py index d178a80..bd56fcd 100644 --- a/src/openbench/engine/whisperkitpro_engine.py +++ b/src/openbench/engine/whisperkitpro_engine.py @@ -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) @@ -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), @@ -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: @@ -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. @@ -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) @@ -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: @@ -232,17 +224,12 @@ 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): @@ -250,10 +237,7 @@ class WhisperKitProInput(BaseModel): 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): @@ -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: diff --git a/src/openbench/metric/word_error_metrics/__init__.py b/src/openbench/metric/word_error_metrics/__init__.py index 159b603..305aace 100644 --- a/src/openbench/metric/word_error_metrics/__init__.py +++ b/src/openbench/metric/word_error_metrics/__init__.py @@ -4,7 +4,7 @@ from .english_abbreviations import ABBR from .text_normalizer import EnglishTextNormalizer from .word_error_metrics import ( - WordErrorRate, - WordDiarizationErrorRate, ConcatenatedMinimumPermutationWER, + WordDiarizationErrorRate, + WordErrorRate, ) diff --git a/src/openbench/pipeline/transcription/__init__.py b/src/openbench/pipeline/transcription/__init__.py index 7a5e72b..75f3dc8 100644 --- a/src/openbench/pipeline/transcription/__init__.py +++ b/src/openbench/pipeline/transcription/__init__.py @@ -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__ = [ diff --git a/src/openbench/pipeline/transcription/transcription_oss_whisper.py b/src/openbench/pipeline/transcription/transcription_oss_whisper.py index 7f8522f..8bda785 100644 --- a/src/openbench/pipeline/transcription/transcription_oss_whisper.py +++ b/src/openbench/pipeline/transcription/transcription_oss_whisper.py @@ -4,6 +4,7 @@ from pathlib import Path from typing import Callable +import whisper from argmaxtools.utils import get_logger from pydantic import Field @@ -12,7 +13,6 @@ from ...pipeline_prediction import Transcript from ...types import PipelineType from .common import TranscriptionConfig, TranscriptionOutput -import whisper logger = get_logger(__name__) @@ -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. @@ -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, @@ -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 = [] @@ -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). "), ) @@ -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 @@ -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) diff --git a/src/openbench/pipeline/transcription/transcription_whisperkitpro.py b/src/openbench/pipeline/transcription/transcription_whisperkitpro.py index bb2abfc..f57213e 100644 --- a/src/openbench/pipeline/transcription/transcription_whisperkitpro.py +++ b/src/openbench/pipeline/transcription/transcription_whisperkitpro.py @@ -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", diff --git a/tests/metric/word_error_metric/wder/Dockerfile b/tests/metric/word_error_metric/wder/Dockerfile deleted file mode 100644 index 6eb60e1..0000000 --- a/tests/metric/word_error_metric/wder/Dockerfile +++ /dev/null @@ -1,33 +0,0 @@ -FROM python:3.9-slim as builder - -# Install build dependencies and clean up in one layer -RUN apt-get update && apt-get install -y \ - git \ - build-essential \ - && rm -rf /var/lib/apt/lists/* - -# Create and activate virtual environment -RUN python -m venv /opt/venv -ENV PATH="/opt/venv/bin:$PATH" - -# Install Python dependencies -RUN pip install --no-cache-dir git+https://github.com/google/speaker-id.git#subdirectory=DiarizationLM - -# Start a new stage with minimal image -FROM python:3.9-slim - -# Copy virtual environment from builder -COPY --from=builder /opt/venv /opt/venv -ENV PATH="/opt/venv/bin:$PATH" - -# Set up working directory -WORKDIR /app - -# Copy only the necessary script -COPY wder_reference.py . - -# Make the script executable -RUN chmod +x wder_reference.py - -# Set the entrypoint to run the script -ENTRYPOINT ["python", "/app/wder_reference.py"] diff --git a/tests/metric/word_error_metric/wder/test_wder.py b/tests/metric/word_error_metric/wder/test_wder.py index 0e41079..ab8aee2 100644 --- a/tests/metric/word_error_metric/wder/test_wder.py +++ b/tests/metric/word_error_metric/wder/test_wder.py @@ -1,75 +1,16 @@ # For licensing see accompanying LICENSE.md file. # Copyright (C) 2025 Argmax, Inc. All Rights Reserved. -import os -import random -import subprocess +import json import unittest -from argmaxtools.utils import get_logger +from huggingface_hub import hf_hub_download from openbench.metric import WordDiarizationErrorRate from openbench.pipeline_prediction import Transcript -logger = get_logger(__name__) - -RANDOM_SEED = 69 - - -class ReferenceWDER: - """Wrapper class for running the reference WDER implementation through Docker.""" - - def __init__(self) -> None: - self.docker_image = "reference-wder:latest" - self._build_docker_image() - - def _build_docker_image(self) -> None: - """Build the Docker image if it doesn't exist.""" - script_dir = os.path.dirname(os.path.abspath(__file__)) - logger.info(f"Building Docker image in {script_dir}") - subprocess.run( - ["docker", "build", "-t", self.docker_image, "-f", "Dockerfile", "."], - cwd=script_dir, - check=True, - ) - logger.info("Docker image built successfully") - - def _format_input(self, transcript: Transcript) -> tuple[str, str]: - """Format a list of Words into text and speaker strings.""" - text = transcript.get_transcript_string() - speakers = transcript.get_speakers_string() - return text, speakers - - def __call__(self, reference: Transcript, hypothesis: Transcript) -> float: - """Compute WDER between reference and hypothesis using the reference implementation.""" - ref_text, ref_spk = self._format_input(reference) - hyp_text, hyp_spk = self._format_input(hypothesis) - - # Run the Docker container with the inputs - logger.info(f"Running Docker container with inputs: {ref_text}, {ref_spk}, {hyp_text}, {hyp_spk}") - result = subprocess.run( - [ - "docker", - "run", - "--rm", - self.docker_image, - "--reference-text", - ref_text, - "--reference-speaker", - ref_spk, - "--hypothesis-text", - hyp_text, - "--hypothesis-speaker", - hyp_spk, - ], - capture_output=True, - text=True, - check=True, - ) - - # Parse the output (which should be just the WDER value) - return float(result.stdout.strip()) +TEST_FIXTURES_REPO = "argmaxinc/test-fixtures" def create_transcript(text: str, speakers: str) -> Transcript: @@ -85,11 +26,9 @@ def create_transcript(text: str, speakers: str) -> Transcript: class TestWDER(unittest.TestCase): def setUp(self) -> None: self.wder = WordDiarizationErrorRate() - self.reference_wder = ReferenceWDER() def tearDown(self) -> None: self.wder = None - self.reference_wder = None def _compute_and_compare( self, @@ -97,16 +36,17 @@ def _compute_and_compare( reference_speakers: str, hypothesis_text: str, hypothesis_speakers: str, + expected_wder: float, ) -> None: reference = create_transcript(reference_text, reference_speakers) hypothesis = create_transcript(hypothesis_text, hypothesis_speakers) - main_result = self.wder(reference=reference, hypothesis=hypothesis) - ref_result = self.reference_wder(reference=reference, hypothesis=hypothesis) - self.assertEqual( - main_result, - ref_result, - f"Main result: {main_result}, Reference result: {ref_result}", + result = self.wder(reference=reference, hypothesis=hypothesis) + self.assertAlmostEqual( + result, + expected_wder, + places=4, + msg=f"WDER result: {result}, Expected: {expected_wder}", ) def test_wder_perfect_match(self) -> None: @@ -122,6 +62,7 @@ def test_wder_perfect_match(self) -> None: reference_speakers=reference_speakers, hypothesis_text=hypothesis_text, hypothesis_speakers=hypothesis_speakers, + expected_wder=0.0, ) def test_wder_completely_wrong_speaker(self) -> None: @@ -137,6 +78,7 @@ def test_wder_completely_wrong_speaker(self) -> None: reference_speakers=reference_speakers, hypothesis_text=hypothesis_text, hypothesis_speakers=hypothesis_speakers, + expected_wder=0.44444, ) def test_wder_completely_wrong_text(self) -> None: @@ -152,6 +94,7 @@ def test_wder_completely_wrong_text(self) -> None: reference_speakers=reference_speakers, hypothesis_text=hypothesis_text, hypothesis_speakers=hypothesis_speakers, + expected_wder=0.0, ) def test_wder_partially_wrong_text(self) -> None: @@ -167,6 +110,7 @@ def test_wder_partially_wrong_text(self) -> None: reference_speakers=reference_speakers, hypothesis_text=hypothesis_text, hypothesis_speakers=hypothesis_speakers, + expected_wder=0.0, ) def test_wder_partially_wrong_speaker(self) -> None: @@ -182,6 +126,7 @@ def test_wder_partially_wrong_speaker(self) -> None: reference_speakers=reference_speakers, hypothesis_text=hypothesis_text, hypothesis_speakers=hypothesis_speakers, + expected_wder=0.0, ) def test_wder_text_deletion(self) -> None: @@ -197,6 +142,7 @@ def test_wder_text_deletion(self) -> None: reference_speakers=reference_speakers, hypothesis_text=hypothesis_text, hypothesis_speakers=hypothesis_speakers, + expected_wder=0.0, ) def test_wder_text_insertion(self) -> None: @@ -212,6 +158,7 @@ def test_wder_text_insertion(self) -> None: reference_speakers=reference_speakers, hypothesis_text=hypothesis_text, hypothesis_speakers=hypothesis_speakers, + expected_wder=0.0, ) def test_wder_less_speakers(self) -> None: @@ -227,6 +174,7 @@ def test_wder_less_speakers(self) -> None: reference_speakers=reference_speakers, hypothesis_text=hypothesis_text, hypothesis_speakers=hypothesis_speakers, + expected_wder=0.33333, ) def test_wder_more_speakers(self) -> None: @@ -242,40 +190,53 @@ def test_wder_more_speakers(self) -> None: reference_speakers=reference_speakers, hypothesis_text=hypothesis_text, hypothesis_speakers=hypothesis_speakers, + expected_wder=0.0, ) - def test_wder_random_modifications(self) -> None: - # Reference text-speakers: - reference_text = "a b c d e f g h i j" - reference_speakers = "1 1 1 1 2 2 2 2 2" + def test_wder_real_more_hyp_speakers(self) -> None: + """Test WDER with a real example with more hypothesis speakers than reference.""" + json_path = hf_hub_download( + repo_id=TEST_FIXTURES_REPO, + repo_type="dataset", + filename="more_hyp_speakers.json", + subfolder="tests/wder", + ) + with open(json_path, "r") as f: + data = json.load(f) + + ref_text = data["reference"]["text"] + ref_speakers = data["reference"]["speakers"] + hyp_text = data["hypothesis"]["text"] + hyp_speakers = data["hypothesis"]["speakers"] + + self._compute_and_compare( + reference_text=ref_text, + reference_speakers=ref_speakers, + hypothesis_text=hyp_text, + hypothesis_speakers=hyp_speakers, + expected_wder=data["expected_results"]["wder"], + ) - # Set random seed - random.seed(RANDOM_SEED) - - # Run multiple iterations - for iteration in range(10): # Run 10 different random modifications - # Create hypothesis with random modifications - words = reference_text.split() - speakers = reference_speakers.split() - hypothesis_words = words.copy() - hypothesis_speakers = speakers.copy() - - # For each position, randomly choose whether to modify it and how - for pos in range(min(len(hypothesis_words), len(hypothesis_speakers))): - if random.random() < 0.5: # 50% chance to modify each position - modification_type = random.choice(["text", "speaker", "both"]) - - if modification_type in ["text", "both"]: - hypothesis_words[pos] = "x" - if modification_type in ["speaker", "both"]: - hypothesis_speakers[pos] = str(random.randint(1, 5)) - - hypothesis_text = " ".join(hypothesis_words) - hypothesis_speakers = " ".join(hypothesis_speakers) - - self._compute_and_compare( - reference_text=reference_text, - reference_speakers=reference_speakers, - hypothesis_text=hypothesis_text, - hypothesis_speakers=hypothesis_speakers, - ) + def test_wder_real_fewer_hyp_speakers(self) -> None: + """Test WDER with a real example with more hypothesis speakers than reference.""" + json_path = hf_hub_download( + repo_id=TEST_FIXTURES_REPO, + repo_type="dataset", + filename="fewer_hyp_speakers.json", + subfolder="tests/wder", + ) + with open(json_path, "r") as f: + data = json.load(f) + + ref_text = data["reference"]["text"] + ref_speakers = data["reference"]["speakers"] + hyp_text = data["hypothesis"]["text"] + hyp_speakers = data["hypothesis"]["speakers"] + + self._compute_and_compare( + reference_text=ref_text, + reference_speakers=ref_speakers, + hypothesis_text=hyp_text, + hypothesis_speakers=hyp_speakers, + expected_wder=data["expected_results"]["wder"], + ) diff --git a/tests/metric/word_error_metric/wder/wder_reference.py b/tests/metric/word_error_metric/wder/wder_reference.py deleted file mode 100644 index 265d34a..0000000 --- a/tests/metric/word_error_metric/wder/wder_reference.py +++ /dev/null @@ -1,48 +0,0 @@ -# For licensing see accompanying LICENSE.md file. -# Copyright (C) 2025 Argmax, Inc. All Rights Reserved. - -# This script is a helper to calculate reference WDER from diarizationlm library -# See https://github.com/google/speaker-id/blob/master/DiarizationLM/diarizationlm/metrics.py -# Doing it like this as word-levenhstein dependency of diarizationlm was causing troubles to install -import argparse - -import diarizationlm - - -def calculate_wder(hypothesis_text, hypothesis_speaker, reference_text, reference_speaker): - # Prepare the input in the format expected by diarizationlm - json_dict = { - "utterances": [ - { - "utterance_id": "utt1", - "hyp_text": hypothesis_text, - "hyp_spk": hypothesis_speaker, - "ref_text": reference_text, - "ref_spk": reference_speaker, - } - ] - } - - # Compute metrics - result = diarizationlm.compute_metrics_on_json_dict(json_dict) - - # Return just the WDER value - return result["WDER"] - - -if __name__ == "__main__": - parser = argparse.ArgumentParser(description="Calculate WDER using diarizationlm") - parser.add_argument("--hypothesis-text", required=True, help="The hypothesized transcription text") - parser.add_argument("--hypothesis-speaker", required=True, help="The hypothesized speaker labels") - parser.add_argument("--reference-text", required=True, help="The reference transcription text") - parser.add_argument("--reference-speaker", required=True, help="The reference speaker labels") - - args = parser.parse_args() - - wder = calculate_wder( - args.hypothesis_text, - args.hypothesis_speaker, - args.reference_text, - args.reference_speaker, - ) - print(f"{wder}")