HY-2012 commited on
Commit
550c1bd
·
verified ·
1 Parent(s): fad0776

Read speaker channel count from the model instead of assuming 10

Browse files
Files changed (1) hide show
  1. python/ls_eend_sdk/postprocess.py +7 -5
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
- LAST_SPEAKER_CHANNEL = 8 # ch9 is the non-speaker slot
 
 
15
 
16
 
17
  def to_activity(logits, max_speakers=None, threshold=0.5, median=11):
18
- """logits (T,10) -> binary activity grid (T, n_speakers).
19
 
20
- Channel 0 is silence and channel 9 is the non-speaker slot, so only
21
- channels 1..8 are speakers. ``max_speakers`` keeps the first N of them.
22
  """
23
  from scipy.signal import medfilt
24
 
25
  probs = 1.0 / (1.0 + np.exp(-np.asarray(logits, dtype=np.float32)))
26
- last = LAST_SPEAKER_CHANNEL
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)