HY-2012's picture
Upload folder using huggingface_hub
7c268e9 verified
Raw
History Blame Contribute Delete
6.73 kB
"""Numpy port of NeMo's Sortformer/VAD post-processing (preds -> RTTM).
Mirrors ``predlist_to_timestamps`` / ``binarization_vectorized`` / ``filtering``
from ``nemo/collections/asr/parts/utils/vad_utils.py`` and
``generate_diarization_output_lines`` from ``speaker_utils.py`` (NeMo Speech 3.0).
"""
from dataclasses import dataclass
import numpy as np
FRAME_LENGTH_IN_SEC = 0.08
@dataclass
class PostProcessingParams:
onset: float = 0.5
offset: float = 0.5
pad_onset: float = 0.0
pad_offset: float = 0.0
min_duration_on: float = 0.0
min_duration_off: float = 0.0
filter_speech_first: float = 1.0
def merge_overlap_segment(segments: np.ndarray) -> np.ndarray:
if segments.shape[0] <= 1:
return segments
segments = segments[np.argsort(segments[:, 0])]
merge_boundary = segments[:-1, 1] >= segments[1:, 0]
head_padded = np.concatenate([[False], merge_boundary])
tail_padded = np.concatenate([merge_boundary, [False]])
head = segments[~head_padded, 0]
tail = segments[~tail_padded, 1]
return np.stack([head, tail], axis=1)
def filter_short_segments(segments: np.ndarray, threshold: float) -> np.ndarray:
return segments[segments[:, 1] - segments[:, 0] >= threshold]
def get_gap_segments(segments: np.ndarray) -> np.ndarray:
segments = segments[np.argsort(segments[:, 0])]
return np.column_stack((segments[:-1, 1], segments[1:, 0]))
def remove_segments(original_segments: np.ndarray, to_be_removed: np.ndarray) -> np.ndarray:
keep = np.ones(original_segments.shape[0], dtype=bool)
for segment in to_be_removed:
keep &= ~(original_segments == segment).all(axis=1)
return original_segments[keep]
def filtering(segments: np.ndarray, params: PostProcessingParams) -> np.ndarray:
if segments.shape[0] == 0:
return segments
def filter_speech():
if params.min_duration_on > 0:
return filter_short_segments(segments, params.min_duration_on)
return segments
def restore_short_gaps(current: np.ndarray) -> np.ndarray:
if params.min_duration_off <= 0 or current.shape[0] == 0:
return current
non_speech = get_gap_segments(current)
short_gaps = remove_segments(non_speech, filter_short_segments(non_speech, params.min_duration_off))
if short_gaps.shape[0] == 0:
return current
return merge_overlap_segment(np.concatenate([current, short_gaps], axis=0))
if params.filter_speech_first == 1.0:
segments = filter_speech()
segments = restore_short_gaps(segments)
else:
segments = restore_short_gaps(segments)
segments = filter_speech()
return segments
def binarization_vectorized(sequence: np.ndarray, params: PostProcessingParams) -> np.ndarray:
empty = np.empty((0, 2), dtype=np.float32)
num_frames = sequence.shape[0]
if num_frames == 0:
return empty
positions = np.arange(1, num_frames + 1)
onset, offset = params.onset, params.offset
if onset >= offset:
force_on = sequence > onset
force_off = sequence < offset
has_event = force_on | force_off
event_positions = np.where(has_event, positions, 0)
last_event_positions = np.maximum.accumulate(event_positions)
event_states = np.concatenate([[False], force_on])
above = event_states[last_event_positions]
else:
force_on = sequence >= offset
force_off = sequence <= onset
toggle = (sequence > onset) & (sequence < offset)
has_reset = force_on | force_off
reset_positions = np.where(has_reset, positions, 0)
last_reset_positions = np.maximum.accumulate(reset_positions)
reset_states = np.concatenate([[0], force_on.astype(np.int64)])
base_state = reset_states[last_reset_positions]
toggle_prefix = np.concatenate([[0], np.cumsum(toggle.astype(np.int64))])
toggles_since_reset = toggle_prefix[positions] - toggle_prefix[last_reset_positions]
above = np.logical_xor(base_state.astype(bool), (toggles_since_reset % 2).astype(bool))
padded = np.pad(above.astype(np.float32), (1, 1))
diff = padded[1:] - padded[:-1]
starts = np.where(diff > 0.5)[0]
ends = np.where(diff < -0.5)[0]
if starts.shape[0] == 0:
return empty
start_times = np.clip(starts.astype(np.float32) * FRAME_LENGTH_IN_SEC - params.pad_onset, 0.0, None)
end_times = ends.astype(np.float32) * FRAME_LENGTH_IN_SEC + params.pad_offset
valid = end_times > start_times
if not valid.any():
return empty
segments = np.stack([start_times[valid], end_times[valid]], axis=1)
if params.pad_onset > 0 or params.pad_offset > 0:
segments = merge_overlap_segment(segments)
return segments
def predlist_to_timestamps(
preds: np.ndarray,
offset: float = 0.0,
params: PostProcessingParams = None,
bypass_postprocessing: bool = False,
precision: int = 2,
):
"""Convert (num_frames, num_speakers) probabilities to per-speaker timestamps."""
if params is None:
params = PostProcessingParams()
if bypass_postprocessing:
params = PostProcessingParams(onset=0.5, offset=0.5)
timestamps = []
for spk in range(preds.shape[1]):
segments = binarization_vectorized(preds[:, spk], params)
if not bypass_postprocessing:
segments = filtering(segments, params)
if segments.shape[0] == 0:
timestamps.append([])
continue
segments = segments + offset
timestamps.append([[round(float(start), precision), round(float(end), precision)] for start, end in segments])
return timestamps
def generate_diarization_output_lines(timestamps, model_spk_num: int):
lines = []
for spk_idx in range(model_spk_num):
if not timestamps[spk_idx]:
continue
intervals = np.asarray(timestamps[spk_idx], dtype=np.float32).reshape(-1, 2)
for start, end in merge_overlap_segment(intervals):
lines.append(f"{start:.3f} {end:.3f} speaker_{int(spk_idx)}")
return lines
def timestamps_to_rttm_lines(timestamps, uri: str, model_spk_num: int):
lines = []
for spk_idx in range(model_spk_num):
intervals = timestamps[spk_idx]
if not intervals:
continue
merged = merge_overlap_segment(np.asarray(intervals, dtype=np.float32).reshape(-1, 2))
for start, end in merged:
duration = float(end) - float(start)
if duration > 0:
lines.append(
f"SPEAKER {uri} 1 {float(start):.3f} {duration:.3f} <NA> <NA> speaker_{int(spk_idx)} <NA>"
)
return lines