personal memory agent
0

Configure Feed

Select the types of activity you want to include in your feed.

solstone / tests / test_transcribe_overlap.py
8.0 kB 280 lines
1# SPDX-License-Identifier: AGPL-3.0-only 2# Copyright (c) 2026 sol pbc 3 4"""Tests for pyannote overlap-fraction inference.""" 5 6from __future__ import annotations 7 8import sys 9import types 10 11import numpy as np 12import pytest 13 14from solstone.observe.utils import SAMPLE_RATE 15 16 17class _Input: 18 def __init__(self, name: str): 19 self.name = name 20 21 22class _StubSession: 23 def __init__(self, log_probs: np.ndarray | list[np.ndarray]): 24 if isinstance(log_probs, list): 25 self._log_probs = [item.astype(np.float32) for item in log_probs] 26 self._repeat = False 27 else: 28 self._log_probs = [log_probs.astype(np.float32)] 29 self._repeat = True 30 self._idx = 0 31 32 def get_inputs(self): 33 return [_Input("input_values")] 34 35 def run(self, _outputs, _inputs): 36 if self._idx >= len(self._log_probs): 37 if not self._repeat: 38 raise AssertionError("unexpected pyannote run") 39 idx = len(self._log_probs) - 1 40 else: 41 idx = self._idx 42 self._idx += 1 43 return [self._log_probs[idx][None, :, :]] 44 45 46def _dominant_log_probs(classes: np.ndarray) -> np.ndarray: 47 log_probs = np.full((classes.shape[0], 7), -10.0, dtype=np.float32) 48 log_probs[np.arange(classes.shape[0]), classes] = 0.0 49 return log_probs 50 51 52def test_compute_overlap_fraction_silent_audio_returns_zero(monkeypatch): 53 from solstone.observe.transcribe import overlap 54 55 monkeypatch.setattr( 56 overlap, 57 "_get_overlap_session", 58 lambda: _StubSession(_dominant_log_probs(np.zeros(589, dtype=np.int64))), 59 ) 60 61 result = overlap.compute_overlap_fraction( 62 np.zeros(12 * SAMPLE_RATE, dtype=np.float32) 63 ) 64 65 assert result == 0.0 66 67 68def test_compute_overlap_fraction_short_audio_padded(monkeypatch): 69 from solstone.observe.transcribe import overlap 70 71 monkeypatch.setattr( 72 overlap, 73 "_get_overlap_session", 74 lambda: _StubSession(_dominant_log_probs(np.zeros(589, dtype=np.int64))), 75 ) 76 77 result = overlap.compute_overlap_fraction( 78 np.zeros(3 * SAMPLE_RATE, dtype=np.float32) 79 ) 80 81 assert isinstance(result, float) 82 assert result == 0.0 83 84 85def test_compute_overlap_fraction_non_aligned_length(monkeypatch): 86 from solstone.observe.transcribe import overlap 87 88 monkeypatch.setattr( 89 overlap, 90 "_get_overlap_session", 91 lambda: _StubSession(_dominant_log_probs(np.zeros(589, dtype=np.int64))), 92 ) 93 94 audio = np.zeros(int(13.7 * SAMPLE_RATE), dtype=np.float32) 95 result = overlap.compute_overlap_fraction(audio) 96 97 assert result == 0.0 98 99 100def test_compute_overlap_fraction_rejects_wrong_sample_rate(): 101 from solstone.observe.transcribe.overlap import compute_overlap_fraction 102 103 with pytest.raises(ValueError, match="requires 16000 Hz audio"): 104 compute_overlap_fraction(np.zeros(16000, dtype=np.float32), sample_rate=8000) 105 106 107def test_get_overlap_session_loads_and_caches(monkeypatch, tmp_path): 108 from solstone.observe import model_assets 109 from solstone.observe.transcribe import overlap 110 111 model = tmp_path / "seg.onnx" 112 model.write_bytes(b"stub") 113 constructions = [] 114 115 class _CountingSession: 116 def __init__(self, *args, **kwargs): 117 constructions.append((args, kwargs)) 118 119 def get_providers(self): 120 return ["CPUExecutionProvider"] 121 122 fake_ort = types.ModuleType("onnxruntime") 123 fake_ort.InferenceSession = _CountingSession 124 monkeypatch.setattr(overlap, "_overlap_session", None) 125 monkeypatch.setattr( 126 model_assets, 127 "resolve_pyannote_segmentation_model", 128 lambda: model, 129 ) 130 monkeypatch.setitem(sys.modules, "onnxruntime", fake_ort) 131 132 first = overlap._get_overlap_session() 133 second = overlap._get_overlap_session() 134 135 assert len(constructions) == 1 136 assert first is second 137 138 139def test_compute_overlap_fraction_uses_conditioned_formula(monkeypatch): 140 from solstone.observe.transcribe import overlap 141 142 classes = np.concatenate( 143 [ 144 np.full(300, 1, dtype=np.int64), 145 np.full(100, 4, dtype=np.int64), 146 np.zeros(189, dtype=np.int64), 147 ] 148 ) 149 monkeypatch.setattr( 150 overlap, 151 "_get_overlap_session", 152 lambda: _StubSession(_dominant_log_probs(classes)), 153 ) 154 155 result = overlap.compute_overlap_fraction( 156 np.zeros(10 * SAMPLE_RATE, dtype=np.float32) 157 ) 158 159 assert result == pytest.approx(100 / 400) 160 161 162def test_compute_overlap_and_logprobs_returns_fraction_and_logprobs(monkeypatch): 163 from solstone.observe.transcribe import overlap 164 165 classes = np.concatenate( 166 [ 167 np.full(300, 1, dtype=np.int64), 168 np.full(100, 4, dtype=np.int64), 169 np.zeros(189, dtype=np.int64), 170 ] 171 ) 172 monkeypatch.setattr( 173 overlap, 174 "_get_overlap_session", 175 lambda: _StubSession(_dominant_log_probs(classes)), 176 ) 177 178 result = overlap.compute_overlap_and_logprobs( 179 np.zeros(10 * SAMPLE_RATE, dtype=np.float32) 180 ) 181 182 assert result.overlap_fraction == pytest.approx(100 / 400) 183 assert result.avg_log_probs.shape == (589, 7) 184 assert result.avg_log_probs.dtype == np.float32 185 assert result.window_stats == (overlap.SpeakerWindowStats(400, 1, 100),) 186 187 188def test_decide_speaker_evidence_solo_one_slot_returns_single(): 189 from solstone.observe.transcribe import overlap 190 191 decision = overlap.decide_speaker_evidence( 192 0.0, 193 (overlap.SpeakerWindowStats(100, 1, 0),), 194 ) 195 196 assert decision.speaker_evidence == "single" 197 assert decision.multi_window_fraction == 0.0 198 199 200def test_decide_speaker_evidence_slot_permuted_windows_return_single(monkeypatch): 201 from solstone.observe.transcribe import overlap 202 203 monkeypatch.setattr( 204 overlap, 205 "_get_overlap_session", 206 lambda: _StubSession( 207 [ 208 _dominant_log_probs(np.full(589, 1, dtype=np.int64)), 209 _dominant_log_probs(np.full(589, 2, dtype=np.int64)), 210 ] 211 ), 212 ) 213 214 result = overlap.compute_overlap_and_logprobs( 215 np.zeros(12 * SAMPLE_RATE, dtype=np.float32) 216 ) 217 decision = overlap.decide_speaker_evidence( 218 result.overlap_fraction, 219 result.window_stats, 220 ) 221 222 assert len(result.window_stats) == 2 223 assert {row.active_slot_count for row in result.window_stats} == {1} 224 assert decision.speaker_evidence == "single" 225 226 227def test_decide_speaker_evidence_turn_taking_returns_multi(): 228 from solstone.observe.transcribe import overlap 229 230 decision = overlap.decide_speaker_evidence( 231 0.0, 232 (overlap.SpeakerWindowStats(100, 2, 0),), 233 ) 234 235 assert decision.speaker_evidence == "multi" 236 237 238def test_decide_speaker_evidence_overlap_heavy_returns_multi(): 239 from solstone.observe.transcribe import overlap 240 241 decision = overlap.decide_speaker_evidence( 242 0.8, 243 (overlap.SpeakerWindowStats(100, 1, 80),), 244 ) 245 246 assert decision.speaker_evidence == "multi" 247 248 249def test_decide_speaker_evidence_all_silence_returns_none(): 250 from solstone.observe.transcribe import overlap 251 252 decision = overlap.decide_speaker_evidence( 253 0.0, 254 (overlap.SpeakerWindowStats(0, 0, 0),), 255 ) 256 257 assert decision.speaker_evidence == "none" 258 assert decision.multi_window_fraction == 0.0 259 260 261def test_decide_speaker_evidence_overlap_fraction_term_engages_multi(): 262 from solstone.observe.transcribe import overlap 263 264 decision = overlap.decide_speaker_evidence( 265 0.5, 266 (overlap.SpeakerWindowStats(100, 1, 0),), 267 ) 268 269 assert decision.speaker_evidence == "multi" 270 271 272def test_decide_speaker_evidence_branch_four_overlap_ambiguity_returns_multi(): 273 from solstone.observe.transcribe import overlap 274 275 decision = overlap.decide_speaker_evidence( 276 0.0, 277 (overlap.SpeakerWindowStats(100, 1, 5),), 278 ) 279 280 assert decision.speaker_evidence == "multi"