Spaces:
Running on Zero
Running on Zero
Download init_nemo_models.py from nvidia/NeMo-Speech-Streaming-Multispeaker-ASR: direct link, hf CLI and curl.
- Browser
- Download file 22.9 kB
-
https://huggingface.co/spaces/nvidia/NeMo-Speech-Streaming-Multispeaker-ASR/resolve/main/init_nemo_models.py
- Command line
-
hf download hf://spaces/nvidia/NeMo-Speech-Streaming-Multispeaker-ASR/init_nemo_models.py
-
curl -L -o init_nemo_models.py https://huggingface.co/spaces/nvidia/NeMo-Speech-Streaming-Multispeaker-ASR/resolve/main/init_nemo_models.py
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 | |
| 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, | |
| }, | |
| } | |
| 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 | |