diff --git a/.github/workflows/test-moss-adapter.yml b/.github/workflows/test-moss-adapter.yml index c13a3a37d..48acbe57c 100644 --- a/.github/workflows/test-moss-adapter.yml +++ b/.github/workflows/test-moss-adapter.yml @@ -4,6 +4,8 @@ on: pull_request: paths: - "funasr/models/moss_transcribe_diarize/**" + - "funasr/models/campplus/utils.py" + - "tests/test_campplus_utils.py" - "funasr/auto/auto_model.py" - "funasr/bin/_server_app.py" - "funasr/bin/server.py" @@ -21,6 +23,8 @@ on: - main paths: - "funasr/models/moss_transcribe_diarize/**" + - "funasr/models/campplus/utils.py" + - "tests/test_campplus_utils.py" - "funasr/auto/auto_model.py" - "funasr/bin/_server_app.py" - "funasr/bin/server.py" @@ -58,6 +62,7 @@ jobs: - name: Check formatting and syntax run: | python -m black --check \ + tests/test_campplus_utils.py \ funasr/models/moss_transcribe_diarize/model.py \ tests/test_moss_transcribe_diarize_model.py \ tests/test_moss_transcribe_diarize_docs.py \ @@ -71,6 +76,7 @@ jobs: - name: Run adapter and documentation contracts run: | python -m pytest -q \ + tests/test_campplus_utils.py \ tests/test_moss_transcribe_diarize_model.py \ tests/test_moss_transcribe_diarize_docs.py \ tests/test_server_app_openai_segments.py \ diff --git a/funasr/models/campplus/utils.py b/funasr/models/campplus/utils.py index dc2ff0300..9ea12a7ed 100644 --- a/funasr/models/campplus/utils.py +++ b/funasr/models/campplus/utils.py @@ -254,26 +254,21 @@ def smooth(res, mindur=0.7): def distribute_spk(sentence_list, sd_time_list): - """Distribute spk. - - Args: - sentence_list: TODO. - sd_time_list: TODO. - """ + """Assign each sentence to the speaker with the greatest total overlap. + + Sentence times are milliseconds; diarization times are seconds. Equal + totals retain the first overlapping speaker, and no overlap defaults to 0. + """ sd_time_list = [(spk_st * 1000, spk_ed * 1000, spk) for spk_st, spk_ed, spk in sd_time_list] for d in sentence_list: sentence_start = d['start'] sentence_end = d['end'] - sentence_spk = 0 - max_overlap = 0 + speaker_overlaps = {} for spk_st, spk_ed, spk in sd_time_list: overlap = max(min(sentence_end, spk_ed) - max(sentence_start, spk_st), 0) - if overlap > max_overlap: - max_overlap = overlap - sentence_spk = spk - if overlap > 0 and sentence_spk == spk: - max_overlap += overlap - d['spk'] = int(sentence_spk) + if overlap > 0: + speaker_overlaps[spk] = speaker_overlaps.get(spk, 0) + overlap + d['spk'] = int(max(speaker_overlaps, key=speaker_overlaps.get, default=0)) return sentence_list diff --git a/tests/test_campplus_utils.py b/tests/test_campplus_utils.py new file mode 100644 index 000000000..94c60b47c --- /dev/null +++ b/tests/test_campplus_utils.py @@ -0,0 +1,71 @@ +"""Speaker assignment uses the actual overlap with each sentence.""" + +from itertools import permutations + +import numpy as np +import pytest + +from funasr.models.campplus.utils import distribute_spk + + +def test_distribute_spk_does_not_double_count_the_first_overlap(): + sentences = [{"start": 0, "end": 10000, "text": "first and second"}] + + result = distribute_spk(sentences, [(0, 4, 0), (4, 10, 1)]) + + assert result == [{"start": 0, "end": 10000, "text": "first and second", "spk": 1}] + + +@pytest.mark.parametrize( + "timeline", + list(permutations([(0, 2, 7), (2, 5, 9), (5, 7, 7), (7, 10, 9)])), +) +def test_distribute_spk_totals_nonconsecutive_turns_independent_of_order(timeline): + sentences = [{"start": 0, "end": 10000}] + + assert distribute_spk(sentences, timeline) == [{"start": 0, "end": 10000, "spk": 9}] + + +@pytest.mark.parametrize( + ("timeline", "expected"), + [ + ([(0, 2, 7), (2, 7, 9), (7, 10, 7)], 7), + ([(2, 7, 9), (0, 2, 7), (7, 10, 7)], 9), + ([(-2, 0, 9), (0, 5, 7), (5, 10, 9)], 7), + ], +) +def test_distribute_spk_breaks_total_ties_by_first_positive_overlap(timeline, expected): + assert distribute_spk([{"start": 0, "end": 10000}], timeline) == [ + {"start": 0, "end": 10000, "spk": expected} + ] + + +def test_distribute_spk_clips_to_each_sentence_and_preserves_objects(): + first = {"start": 4000, "end": 6000, "text": "clipped", "spk": 99} + second = {"start": 6000, "end": 7000, "timestamp": [[6000, 7000]]} + sentences = [first, second] + timeline = [(0, 4, 8), (4, 5.5, np.int64(3)), (5.5, 100, np.int64(9))] + + result = distribute_spk(sentences, timeline) + + assert result is sentences + assert result[0] is first + assert result[1] is second + assert result == [ + {"start": 4000, "end": 6000, "text": "clipped", "spk": 3}, + {"start": 6000, "end": 7000, "timestamp": [[6000, 7000]], "spk": 9}, + ] + assert type(first["spk"]) is int + assert type(second["spk"]) is int + + +@pytest.mark.parametrize("timeline", [[], [(0, 1, 8), (2, 3, 9)]]) +def test_distribute_spk_defaults_to_zero_without_positive_overlap(timeline): + assert distribute_spk([{"start": 1000, "end": 2000}], timeline) == [ + {"start": 1000, "end": 2000, "spk": 0} + ] + + +def test_distribute_spk_accepts_empty_sentences(): + sentences = [] + assert distribute_spk(sentences, [(0, 1, 7)]) is sentences