taejinp's picture
Update Nemotron streaming multispeaker ASR Space
2dd1cd7 verified
Raw History Blame Contribute Delete
22.9 kB
import time
import math
import os
import re
from copy import deepcopy
import librosa
import numpy as np
from omegaconf import ListConfig, OmegaConf, open_dict
import nemo.collections.asr as nemo_asr
from nemo.collections.asr.models import SortformerEncLabelModel, ASRModel
from nemo.collections.asr.parts.submodules.subsampling import FeatureStacking
from nemo.collections.asr.parts.utils.streaming_utils import CacheAwareStreamingAudioBuffer
from nemo.collections.asr.parts.utils.rnnt_utils import Hypothesis
import torch
from nemo.collections.asr.parts.utils.multispk_transcribe_utils import (
SpeakerTaggedASR,
configure_diar_streaming,
validate_feature_frame_strides,
)
from typing import List, Optional, Union
from dataclasses import dataclass, field
ASR_LOCAL_MODEL_EXTENSIONS = (".nemo",)
DIAR_LOCAL_MODEL_EXTENSIONS = (".nemo", ".ckpt")
LANG_TAG_PATTERN = re.compile(r"^(?:<\|[^|]+?\|>\s*)+")
def resolve_model_reference(model_reference: str, model_label: str, local_extensions) -> str:
"""Resolve a local checkpoint path or preserve a NeMo pretrained model identifier."""
if model_reference is None or not model_reference.strip():
raise ValueError(f"{model_label} model reference must not be empty")
model_reference = model_reference.strip()
expanded_reference = os.path.expanduser(model_reference)
if os.path.isfile(expanded_reference):
extension = os.path.splitext(expanded_reference)[1].lower()
if extension not in local_extensions:
raise ValueError(
f"Unsupported local {model_label.lower()} model file: {model_reference}. "
f"Expected one of: {', '.join(local_extensions)}"
)
return expanded_reference
if os.path.exists(expanded_reference):
raise ValueError(f"{model_label} model reference must point to a checkpoint file: {model_reference}")
extension = os.path.splitext(expanded_reference)[1].lower()
looks_like_local_path = (
extension in local_extensions
or model_reference.startswith(("/", "./", "../", "~"))
)
if looks_like_local_path:
raise FileNotFoundError(f"{model_label} model not found at: {model_reference}")
# NeMo resolves identifiers containing "/" from Hugging Face Hub and bare
# identifiers from its registered pretrained-model catalog.
return model_reference
def format_time(seconds):
minutes = math.floor(seconds / 60)
sec = seconds % 60
return f"{minutes}:{sec:05.2f}"
def extract_transcriptions(hyps):
"""
The transcribed_texts returned by CTC and RNNT models are different.
This method would extract and return the text section of the hypothesis.
"""
if isinstance(hyps[0], Hypothesis):
transcriptions = []
for hyp in hyps:
transcriptions.append(hyp.text)
else:
transcriptions = hyps
return transcriptions
def configure_asr_for_multitalker_streaming(config, asr_model):
"""Apply ASR-model-specific settings for the shared multitalker streaming path."""
if config.parallel_speaker_strategy and not config.masked_asr and not hasattr(asr_model, "set_speaker_targets"):
raise ValueError(
"parallel_speaker_strategy=True with masked_asr=False requires an ASR model that supports "
"speaker-target injection. Use masked_asr=True for conventional streaming ASR models."
)
if hasattr(asr_model, "set_inference_prompt"):
target_lang = config.target_lang or "auto"
asr_model.set_inference_prompt(target_lang)
if config.strip_lang_tags and hasattr(asr_model, "cfg"):
with open_dict(asr_model.cfg.decoding):
asr_model.cfg.decoding.strip_lang_tags = True
asr_model.cfg.decoding.lang_tag_pattern = config.lang_tag_pattern or LANG_TAG_PATTERN.pattern
@dataclass
class MultitalkerTranscriptionConfig:
"""
Configuration for Multi-talker transcription with an ASR model and a diarization model.
"""
# Required configs
diar_model: Optional[str] = None # Local checkpoint or pretrained model identifier
diar_pretrained_name: Optional[str] = None # Name of a pretrained model
max_num_of_spks: Optional[int] = 4 # maximum number of speakers
parallel_speaker_strategy: bool = True # whether to use parallel speaker strategy
masked_asr: bool = False # whether to use masked ASR
mask_preencode: bool = False # whether to mask preencode or mask features
cache_gating: bool = True # whether to use cache gating
cache_gating_buffer_size: int = 2 # buffer size for cache gating
single_speaker_mode: bool = False # whether to use single speaker mode
displayed_transcript_lines: int = 10 # number of transcript lines to display
# General configs
session_len_sec: float = -1 # End-to-end diarization session length in seconds
num_workers: int = 8
random_seed: Optional[int] = None # seed number going to be used in seed_everything()
log: bool = False # If True,log will be printed
precision: str = "bf16"
# Streaming diarization configs
streaming_mode: bool = True # If True, streaming diarization will be used.
spkcache_len: Optional[int] = None
spkcache_update_period: int = 222
fifo_len: int = 264
diar_right_context: int = 0
# If `cuda` is a negative number, inference will be on CPU only.
cuda: Optional[int] = None
allow_mps: bool = False # allow to select MPS device (Apple Silicon M-series GPU)
matmul_precision: str = "highest" # Literal["highest", "high", "medium"]
# ASR Configs
asr_model: Optional[str] = None # Local checkpoint or pretrained model identifier
device: str = 'cuda'
audio_file: Optional[str] = None
manifest_file: Optional[str] = None
use_amp: bool = True
debug_mode: bool = False
batch_size: int = 32
online_normalization: bool = False
output_path: Optional[str] = None
pad_and_drop_preencoded: bool = False
set_decoder: Optional[str] = None # ["ctc", "rnnt"]
att_context_size: Optional[List[int]] = field(default_factory=lambda: [56, 13])
target_lang: Optional[str] = "auto"
strip_lang_tags: bool = True
lang_tag_pattern: Optional[str] = LANG_TAG_PATTERN.pattern
generate_realtime_scripts: bool = True # This should be True for huggingface space
word_window: int = 50
sent_break_sec: float = 1.0
fix_prev_words_count: int = 5
update_prev_words_sentence: int = 5
left_frame_shift: int = -1
right_frame_shift: int = 0
min_sigmoid_val: float = 1e-2
discarded_frames: int = 8
print_time: bool = True
print_sample_indices: List[int] = field(default_factory=lambda: [0])
colored_text: bool = True
real_time_mode: bool = False
print_path: Optional[str] = None
ignored_initial_frame_steps: int = 5
verbose: bool = False
feat_len_sec: float = 0.01
finetune_realtime_ratio: float = 0.01
spk_supervision: str = "diar" # ["diar", "rttm"]
binary_diar_preds: bool = False
# Microphone streaming configs
sample_rate: int = 16000
microphone_streaming: bool = False
microphone_chunk_size: int = 0 # Number of samples per chunk
microphone_hop_length: int = 0 # Hop length for sliding window
# Mix config
deploy_mode: bool = True
# Audio pre-loading configs
audio_pre_loading: bool = True
class StreamingSession:
def __init__(self):
self.config: Optional[MultitalkerTranscriptionConfig] = None
self.asr_model: Optional[ASRModel] = None
self.diar_model: Optional[SortformerEncLabelModel] = None
self.streaming_buffer = None
self.multispk_asr_streamer: Optional[SpeakerTaggedASR] = None
self.transcribed_speaker_texts = ""
self.step_num = 0
self.feat_frame_count = 0
self._max_num_of_spks = 4
self.chunk_audio_buffer = np.array([])
self.all_chunk_audio = np.array([])
self.session_start_time = None
self.silence_level = 1e-4
self.silence_second = 0.5
self._is_result_buffer_called = False
self._audio_pre_loading = True
def load_asr_model(
self,
asr_model: str,
device: Union[str, torch.device] = 'cuda',
reload: bool = False,
):
asr_model = resolve_model_reference(
asr_model,
model_label="ASR",
local_extensions=ASR_LOCAL_MODEL_EXTENSIONS,
)
if not reload:
if os.path.isfile(asr_model):
self.asr_model = nemo_asr.models.ASRModel.restore_from(restore_path=asr_model)
else:
self.asr_model = nemo_asr.models.ASRModel.from_pretrained(
model_name=asr_model,
map_location=device,
)
self.asr_model.eval()
self.asr_model = self.asr_model.to(device)
# configure the decoding config
decoding_cfg = self.asr_model.cfg.decoding
with open_dict(decoding_cfg):
decoding_cfg.strategy = "greedy_batch"
decoding_cfg.preserve_alignments = False
decoding_cfg.strip_lang_tags = self.config.strip_lang_tags
decoding_cfg.lang_tag_pattern = self.config.lang_tag_pattern or LANG_TAG_PATTERN.pattern
if hasattr(self.asr_model, 'joint'): # if an RNNT model
decoding_cfg.greedy.max_symbols = 10
decoding_cfg.fused_batch_size = -1
self.asr_model.change_decoding_strategy(decoding_cfg)
if self.config.att_context_size:
self.asr_model.encoder.set_default_att_context_size(self.config.att_context_size)
configure_asr_for_multitalker_streaming(self.config, self.asr_model)
def load_diar_model(
self,
diar_model: str,
device: Union[str, torch.device],
reload: bool = False,
):
diar_model = resolve_model_reference(
diar_model,
model_label="Diarization",
local_extensions=DIAR_LOCAL_MODEL_EXTENSIONS,
)
if not reload:
if os.path.isfile(diar_model) and diar_model.endswith(".ckpt"):
self.diar_model = SortformerEncLabelModel.load_from_checkpoint(checkpoint_path=diar_model, map_location=device, strict=False)
elif os.path.isfile(diar_model):
self.diar_model = SortformerEncLabelModel.restore_from(restore_path=diar_model, map_location=device)
else:
self.diar_model = SortformerEncLabelModel.from_pretrained(
model_name=diar_model,
map_location=device,
)
self.diar_model.eval()
self.diar_model = self.diar_model.to(device)
if (
self.config.precision.startswith("bf16")
and self.diar_model.device.type == "cuda"
and torch.cuda.is_bf16_supported()
):
self.diar_model = self.diar_model.to(dtype=torch.bfloat16)
validate_feature_frame_strides(asr_model=self.asr_model, diar_model=self.diar_model)
diar_chunk_len = (
self.asr_model.encoder.streaming_cfg.valid_out_len
+ self.asr_model.encoder.streaming_cfg.cache_drop_size
)
configure_diar_streaming(
diar_model=self.diar_model,
cfg=self.config,
output_subsampling_factor=self.asr_model.encoder.subsampling_factor,
diar_chunk_len=diar_chunk_len,
)
self.config.spkcache_len = int(self.diar_model.sortformer_modules.spkcache_len)
def add_file_to_streaming_buffer(self, audio_file_path: str):
# Load audio file at 16000 Hz using librosa (handles resampling automatically)
audio, SR = librosa.load(audio_file_path, sr=self.config.sample_rate, mono=True)
# Add very short silence (0.5 second) at the beginning of the audio for better streaming
silence = np.random.randn(int(SR * self.silence_second)) * self.silence_level
audio = np.concatenate([silence, audio])
# Preprocess and append to buffer
self.streaming_buffer.append_audio(audio, stream_id=-1)
# Create iterator for streaming
self.streaming_buffer_iter = iter(self.streaming_buffer)
def _setup_streaming_buffer(self):
streaming_buffer = CacheAwareStreamingAudioBuffer(
model=self.asr_model,
online_normalization=False,
pad_and_drop_preencoded=self.config.pad_and_drop_preencoded,
)
return streaming_buffer
def reset_buffer(self):
# if self.streaming_buffer.buffer_idx != 0:
print(f"[DEBUG] reset_buffer() Called at buffer_idx {self.streaming_buffer.buffer_idx}")
self.streaming_buffer.reset_buffer()
self.step_num = 0
self.feat_frame_count = 0
self._reset_microphone_audio_buffer()
self.all_chunk_audio = np.array([])
self.session_start_time = None
self._is_result_buffer_called = True
self.transcribed_speaker_texts = ""
def setup_streaming_session(self, config: MultitalkerTranscriptionConfig, reload: bool = False, reset_buffer: bool = False):
"""Setup streaming parameters based on model configuration"""
self.transcribed_speaker_texts = ""
self.config = OmegaConf.structured(config)
self.load_asr_model(self.config.asr_model, self.config.device, reload=reload)
self._configure_microphone_geometry()
config.microphone_chunk_size = self.config.microphone_chunk_size
config.microphone_hop_length = self.config.microphone_hop_length
self.load_diar_model(self.config.diar_model, self.config.device, reload=reload)
config.spkcache_len = self.config.spkcache_len
self.config.max_num_of_spks = min(
self.config.max_num_of_spks,
int(self.diar_model._cfg.max_num_of_spks),
)
config.max_num_of_spks = self.config.max_num_of_spks
if isinstance(self.diar_model.encoder.pre_encode, FeatureStacking):
self.config.pad_and_drop_preencoded = True
# Always create new SpeakerTaggedASR instance to ensure clean state
self.multispk_asr_streamer = SpeakerTaggedASR(self.config, self.asr_model, self.diar_model)
print(f"[DEBUG] setup_streaming_session() Called, >>>>>> creating new SpeakerTaggedASR instance...")
self.streaming_buffer = self._setup_streaming_buffer()
self._max_num_of_spks = self.config.max_num_of_spks
self._audio_pre_loading = self.config.audio_pre_loading
self._displayed_transcript_lines = self.config.displayed_transcript_lines
self._reset_microphone_audio_buffer()
if reset_buffer:
self.reset_buffer()
def _configure_microphone_geometry(self):
if not self.config.microphone_streaming:
return
streaming_cfg = self.asr_model.encoder.streaming_cfg
feature_stride = float(self.asr_model.cfg.preprocessor.window_stride)
hop_feature_frames = streaming_cfg.valid_out_len * self.asr_model.encoder.subsampling_factor
hop_samples = round(hop_feature_frames * feature_stride * self.config.sample_rate)
pre_encode_cache_size = streaming_cfg.pre_encode_cache_size
if isinstance(pre_encode_cache_size, (list, tuple, ListConfig)):
pre_encode_cache_size = pre_encode_cache_size[-1]
cache_samples = round(pre_encode_cache_size * feature_stride * self.config.sample_rate)
self.config.microphone_hop_length = int(hop_samples)
self.config.microphone_chunk_size = int(hop_samples + cache_samples)
def _reset_microphone_audio_buffer(self):
if self.config is None or not self.config.microphone_streaming:
self.chunk_audio_buffer = np.array([])
return
cache_samples = self.config.microphone_chunk_size - self.config.microphone_hop_length
self.chunk_audio_buffer = np.zeros(cache_samples, dtype=np.float32)
def _get_left_right_offset(self, step_num):
if step_num == 0:
self.session_start_time = time.time()
left_offset = 0
right_offset = 0
else:
left_offset = 8
right_offset = 0
return left_offset, right_offset
def _get_transcribed_texts(self):
"""
Extract transcribed texts from the multispk_asr_streamer's internal state.
Returns a formatted string with speaker-tagged transcriptions.
"""
if self.multispk_asr_streamer is None:
return ""
# Get the word and timestamp sequences from the streamer
word_and_ts_seq = self.multispk_asr_streamer._word_and_ts_seq
transcriptions = []
for _uniq_id, data in word_and_ts_seq.items():
if data.get('sentences') is not None:
for sentence in data['sentences']:
speaker = sentence.get('speaker', 'Unknown')
text = sentence.get('words', '').strip()
if text:
transcriptions.append(f"Speaker {speaker}: {text}")
return "\n".join(transcriptions) if transcriptions else ""
def _get_diar_visualization_state(self):
"""Collect retained cache predictions and fresh FIFO/chunk predictions for visualization."""
if (
self.multispk_asr_streamer is None
or self.multispk_asr_streamer.instance_manager is None
or self.multispk_asr_streamer.instance_manager.diar_states is None
):
return None
diar_state = self.multispk_asr_streamer.instance_manager.diar_states
streaming_state = diar_state.streaming_state
if streaming_state is None:
return None
def get_valid_length(embeddings, lengths):
if lengths is not None:
return int(lengths[0].item())
return 0 if embeddings is None else embeddings.shape[1]
def trim_predictions(predictions, valid_length):
if predictions is None:
return None
return predictions[:, :valid_length].detach()
spkcache_preds = streaming_state.spkcache_preds
fifo_preds = streaming_state.fifo_preds
chunk_capacity = self.multispk_asr_streamer._nframes_per_chunk
chunk_preds = (
None
if diar_state.diar_pred_out_stream is None
else diar_state.diar_pred_out_stream[:, -chunk_capacity:]
)
retained_chunk_frames = 0
if fifo_preds is not None and chunk_preds is not None:
retained_chunk_frames = min(fifo_preds.shape[1], chunk_preds.shape[1])
fifo_preds = fifo_preds[:, : fifo_preds.shape[1] - retained_chunk_frames]
spkcache_valid = get_valid_length(streaming_state.spkcache, streaming_state.spkcache_lengths)
fifo_valid = 0 if fifo_preds is None else fifo_preds.shape[1]
chunk_valid = 0 if chunk_preds is None else chunk_preds.shape[1]
return {
"spkcache": trim_predictions(spkcache_preds, spkcache_valid),
"fifo": trim_predictions(fifo_preds, fifo_valid),
"chunk": None if chunk_preds is None else chunk_preds.detach(),
"valid_frames": {
"spkcache": spkcache_valid,
"fifo": fifo_valid,
"chunk": chunk_valid,
},
"capacities": {
"spkcache": self.config.spkcache_len,
"fifo": self.config.fifo_len,
"chunk": chunk_capacity,
},
}
@torch.inference_mode()
def process_audio_chunk(self, chunk_audio, chunk_lengths, is_buffer_empty=None):
"""
Process an audio chunk using the SpeakerTaggedASR streamer.
The streamer handles all internal state management through its instance_manager.
Args:
chunk_audio: torch.Tensor of shape (batch, channels, time)
chunk_lengths: torch.Tensor of shape (batch,) with lengths
"""
# Determine drop_extra_pre_encoded based on offsets
# drop_extra_pre_encoded = left_offset
drop_extra_pre_encoded = (
0
if self.step_num == 0 and not self.config.pad_and_drop_preencoded
else self.asr_model.encoder.streaming_cfg.drop_extra_pre_encoded
)
if is_buffer_empty is None:
is_buffer_empty = self.streaming_buffer.buffer is not None and self.streaming_buffer.is_buffer_empty()
use_bf16 = (
self.config.precision.startswith("bf16")
and self.asr_model.device.type == "cuda"
and torch.cuda.is_bf16_supported()
)
# Call the appropriate streaming method based on config
with torch.amp.autocast(
device_type=self.asr_model.device.type,
dtype=torch.bfloat16 if use_bf16 else self.asr_model.dtype,
enabled=self.config.use_amp and use_bf16,
):
if self.config.parallel_speaker_strategy:
# Parallel streaming: multiple ASR instances for multiple speakers
transcribed_speaker_texts = self.multispk_asr_streamer.perform_parallel_streaming_stt_spk(
step_num=self.step_num,
chunk_audio=chunk_audio,
chunk_lengths=chunk_lengths,
is_buffer_empty=is_buffer_empty,
drop_extra_pre_encoded=drop_extra_pre_encoded,
)
else:
# Serial streaming: single ASR instance for all speakers
transcribed_speaker_texts = self.multispk_asr_streamer.perform_serial_streaming_stt_spk(
step_num=self.step_num,
chunk_audio=chunk_audio,
chunk_lengths=chunk_lengths,
is_buffer_empty=is_buffer_empty,
drop_extra_pre_encoded=drop_extra_pre_encoded,
)
# Extract transcriptions from the instance manager's state
if transcribed_speaker_texts is not None:
self.transcribed_speaker_texts = transcribed_speaker_texts[0]
self.feat_frame_count += (chunk_audio.shape[-1] - self.config.discarded_frames)
if self.config.log:
print(f'[Deploy Mode: {self.config.deploy_mode}] transcribed_speaker_texts:', self.transcribed_speaker_texts)
self.step_num += 1
return self._get_diar_visualization_state()
def flush_microphone_buffer(self):
"""Process microphone audio that has not yet filled a complete streaming hop."""
if not self.config.microphone_streaming:
return None
cache_samples = self.config.microphone_chunk_size - self.config.microphone_hop_length
if len(self.chunk_audio_buffer) <= cache_samples:
return None
frame = self.chunk_audio_buffer.copy()
self.chunk_audio_buffer = (
frame[-cache_samples:].copy() if cache_samples > 0 else np.array([], dtype=np.float32)
)
chunk_audio, chunk_lengths = self.streaming_buffer.preprocess_audio(frame)
chunk_audio = chunk_audio[:, :, :chunk_lengths[0]]
return self.process_audio_chunk(chunk_audio, chunk_lengths, is_buffer_empty=True)
def process_microphone_chunk(self, chunk_audio, sr):
# Convert to float32 and normalize
if chunk_audio.dtype == np.int16:
chunk_audio = chunk_audio.astype(np.float32) / 32768.0
elif chunk_audio.dtype == np.int32:
chunk_audio = chunk_audio.astype(np.float32) / 2147483648.0
if sr != self.config.sample_rate:
chunk_audio = librosa.resample(chunk_audio, orig_sr=sr, target_sr=self.config.sample_rate)
self.all_chunk_audio = np.concatenate([self.all_chunk_audio, chunk_audio])
self.chunk_audio_buffer = np.concatenate([self.chunk_audio_buffer, chunk_audio])
processed_flag = False
diar_pred_out = None
while len(self.chunk_audio_buffer) >= self.config.microphone_chunk_size:
frame = self.chunk_audio_buffer[:self.config.microphone_chunk_size]
print(f"[DEBUG] process_microphone_chunk() L0: frame.shape: {frame.shape}")
self.chunk_audio_buffer = self.chunk_audio_buffer[self.config.microphone_hop_length:]
chunk_audio, chunk_lengths = self.streaming_buffer.preprocess_audio(frame)
print(f"[DEBUG] process_microphone_chunk() L1: original chunk_audio.shape: {chunk_audio.shape}, chunk_lengths: {chunk_lengths} frame.shape: {frame.shape}")
chunk_audio = chunk_audio[:, :, :chunk_lengths[0]]
print(f"[DEBUG] process_microphone_chunk() L2: self.config.microphone_hop_length {self.config.microphone_hop_length} self.config.microphone_chunk_size {self.config.microphone_chunk_size}")
print(f"[DEBUG] process_microphone_chunk() L3: Truncated chunk_audio.shape: {chunk_audio.shape}")
print(f"[DEBUG] self._max_num_of_spks: {self._max_num_of_spks} and self.multispk_asr_streamer._max_num_of_spks: {self.multispk_asr_streamer._max_num_of_spks}")
if self.step_num == 0:
self.session_start_time = time.time()
diar_pred_out = self.process_audio_chunk(chunk_audio, chunk_lengths, is_buffer_empty=False)
processed_flag = True
if processed_flag:
print(f">>>> Finished processing microphone chunk: {chunk_audio.shape} {sr}")
else:
print(f">>>> Skipping processing microphone chunk: {chunk_audio.shape} {sr}")
return diar_pred_out