Read speaker channel count from the model instead of assuming 10
Browse files
python/ls_eend_sdk/postprocess.py
CHANGED
|
@@ -11,19 +11,21 @@ from .feature import FRAME_SEC
|
|
| 11 |
|
| 12 |
SILENCE_CHANNEL = 0
|
| 13 |
FIRST_SPEAKER_CHANNEL = 1
|
| 14 |
-
|
|
|
|
|
|
|
| 15 |
|
| 16 |
|
| 17 |
def to_activity(logits, max_speakers=None, threshold=0.5, median=11):
|
| 18 |
-
"""logits (T,
|
| 19 |
|
| 20 |
-
Channel 0 is silence and
|
| 21 |
-
channels 1..
|
| 22 |
"""
|
| 23 |
from scipy.signal import medfilt
|
| 24 |
|
| 25 |
probs = 1.0 / (1.0 + np.exp(-np.asarray(logits, dtype=np.float32)))
|
| 26 |
-
last =
|
| 27 |
if max_speakers is not None:
|
| 28 |
last = min(last, FIRST_SPEAKER_CHANNEL + max_speakers - 1)
|
| 29 |
active = (probs[:, FIRST_SPEAKER_CHANNEL:last + 1] > threshold).astype(int)
|
|
|
|
| 11 |
|
| 12 |
SILENCE_CHANNEL = 0
|
| 13 |
FIRST_SPEAKER_CHANNEL = 1
|
| 14 |
+
# The last channel is the non-speaker slot, so speakers are channels
|
| 15 |
+
# 1 .. n_channels-2. n_channels depends on the checkpoint's max_speakers
|
| 16 |
+
# (simu 8 -> 10 channels, AMI 4 -> 6, CALLHOME 7 -> 9, DIHARD 10 -> 12).
|
| 17 |
|
| 18 |
|
| 19 |
def to_activity(logits, max_speakers=None, threshold=0.5, median=11):
|
| 20 |
+
"""logits (T, C) -> binary activity grid (T, n_speakers).
|
| 21 |
|
| 22 |
+
Channel 0 is silence and the last channel is the non-speaker slot, so the
|
| 23 |
+
speakers are channels 1..C-2. ``max_speakers`` keeps the first N of them.
|
| 24 |
"""
|
| 25 |
from scipy.signal import medfilt
|
| 26 |
|
| 27 |
probs = 1.0 / (1.0 + np.exp(-np.asarray(logits, dtype=np.float32)))
|
| 28 |
+
last = probs.shape[1] - 2
|
| 29 |
if max_speakers is not None:
|
| 30 |
last = min(last, FIRST_SPEAKER_CHANNEL + max_speakers - 1)
|
| 31 |
active = (probs[:, FIRST_SPEAKER_CHANNEL:last + 1] > threshold).astype(int)
|