personal memory agent
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"