aboutsummaryrefslogtreecommitdiff
path: root/working/meeting-transcription-service/tests/test_diarize.py
blob: ef009b800ea0d18b26798765f0a9ac87d5a5e949 (plain)
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
"""Tests for diarize's pure parts. The pyannote pipeline itself is not loaded here."""

import sys
from collections import namedtuple
from pathlib import Path

import pytest

sys.path.insert(0, str(Path(__file__).resolve().parent.parent / "src"))

import diarize  # noqa: E402

Segment = namedtuple("Segment", "start end")


class TestTurnsFromTracks:
    def test_diarize_turns_from_tracks_converts_and_sorts(self):
        """Normal: (segment, track, label) triples become sorted turn dicts."""
        tracks = [(Segment(5.0, 9.25), "_", "SPEAKER_01"), (Segment(0.5, 4.0), "_", "SPEAKER_00")]
        assert diarize.turns_from_tracks(tracks) == [
            {"start": 0.5, "end": 4.0, "speaker": "SPEAKER_00"},
            {"start": 5.0, "end": 9.25, "speaker": "SPEAKER_01"},
        ]

    def test_diarize_turns_from_tracks_rounds_to_milliseconds(self):
        """Boundary: float noise from the model is rounded away."""
        tracks = [(Segment(0.03096875, 1.9980000000000002), "_", "SPEAKER_00")]
        assert diarize.turns_from_tracks(tracks) == [
            {"start": 0.031, "end": 1.998, "speaker": "SPEAKER_00"}
        ]

    def test_diarize_turns_from_tracks_drops_empty_segments(self):
        """Boundary: zero or negative length segments carry no speech."""
        tracks = [(Segment(2.0, 2.0), "_", "SPEAKER_00"), (Segment(3.0, 2.5), "_", "SPEAKER_00")]
        assert diarize.turns_from_tracks(tracks) == []

    def test_diarize_turns_from_tracks_empty_input_is_empty_list(self):
        """Boundary: no tracks."""
        assert diarize.turns_from_tracks([]) == []


class TestParseArgs:
    def test_diarize_parse_args_speakers_sets_exact_count(self):
        """Normal: --speakers pins the count."""
        args = diarize.parse_args(["a.wav", "out.json", "--speakers", "3"])
        assert (args.audio, args.out, args.speakers) == ("a.wav", "out.json", 3)

    def test_diarize_parse_args_defaults_let_the_model_estimate(self):
        """Normal: no count given."""
        args = diarize.parse_args(["a.wav", "out.json"])
        assert args.speakers is None and args.min_speakers is None and args.max_speakers is None

    @pytest.mark.parametrize("argv", [
        ["a.wav", "out.json", "--speakers", "0"],
        ["a.wav", "out.json", "--speakers", "-2"],
        ["a.wav", "out.json", "--speakers", "three"],
        ["a.wav", "out.json", "--speakers", "3", "--max-speakers", "5"],
        ["a.wav", "out.json", "--min-speakers", "4", "--max-speakers", "2"],
        ["a.wav"],
    ])
    def test_diarize_parse_args_rejects_bad_counts(self, argv):
        """Error: non-positive, non-numeric, contradictory or missing arguments."""
        with pytest.raises(SystemExit):
            diarize.parse_args(argv)


class TestPipelineKwargs:
    def test_diarize_pipeline_kwargs_only_passes_what_was_given(self):
        """Normal: unset options are not forwarded to the model."""
        args = diarize.parse_args(["a.wav", "o.json", "--min-speakers", "2", "--max-speakers", "4"])
        assert diarize.pipeline_kwargs(args) == {"min_speakers": 2, "max_speakers": 4}
        args = diarize.parse_args(["a.wav", "o.json", "--speakers", "3"])
        assert diarize.pipeline_kwargs(args) == {"num_speakers": 3}
        assert diarize.pipeline_kwargs(diarize.parse_args(["a.wav", "o.json"])) == {}