diff --git a/nemo/collections/tts/parts/utils/tts_dataset_utils.py b/nemo/collections/tts/parts/utils/tts_dataset_utils.py index a8c04ef86bdf..6ebf4a93eb7e 100644 --- a/nemo/collections/tts/parts/utils/tts_dataset_utils.py +++ b/nemo/collections/tts/parts/utils/tts_dataset_utils.py @@ -478,6 +478,57 @@ def filter_dataset_by_duration(entries: List[Dict[str, Any]], min_duration: floa return filtered_entries, total_hours, filtered_hours +def segment_wav( + wav, segment_length: int = 44100, segment_hop_size: int = 44100, min_segment_length: int = 22050 +) -> List[torch.Tensor]: + """ + Splits a waveform into fixed-length, zero-padded segments using a sliding window. + + If input waveform is shorter than ``segment_length``, it's zero-padded and + returned as a single segment. + Otherwise, waveform is split into overlapping or non-overlapping segments + of ``segment_length`` using a hop of ``segment_hop_size``. + + Sliding window stops once fewer than ``min_segment_length`` samples remain. + Final segment, if any, is zero-padded to ``segment_length``. + + Args: + wav: 1D waveform tensor to segment. + + segment_length: Length of each output segment, in samples. Defaults to 44100 — 1 second at 44.1kHz. + + segment_hop_size: Number of samples to advance the sliding window between + segments. Defaults to 44100 — no overlap at the default segment_length. + + min_segment_length: Minimum number of remaining samples required to extract + another segment. Once fewer samples than this remain, the loop stops. + Defaults to 22050 — 0.5 seconds at 44.1kHz. + + Returns: + A list of 1D tensors, each of length ``segment_length``. + """ + + if len(wav) < segment_length: + pad = torch.zeros(segment_length - len(wav)) + segment = torch.cat([wav, pad]) + return [segment] + + segment_start_idx = 0 + segments = [] + + while segment_start_idx < len(wav) - min_segment_length: + segment = wav[segment_start_idx : segment_start_idx + segment_length] + + if len(segment) < segment_length: + pad = torch.zeros(segment_length - len(segment)) + segment = torch.cat([segment, pad]) + + segments.append(segment) + segment_start_idx += segment_hop_size + + return segments + + def get_weighted_sampler( sample_weights: List[float], batch_size: int, world_size: int, num_steps: int ) -> torch.utils.data.WeightedRandomSampler: diff --git a/scripts/ssl_tts/make_supdata.py b/scripts/ssl_tts/make_supdata.py index f178057c1afe..65b89ebbb1f6 100644 --- a/scripts/ssl_tts/make_supdata.py +++ b/scripts/ssl_tts/make_supdata.py @@ -29,7 +29,7 @@ from nemo.collections.asr.parts.preprocessing.segment import AudioSegment from nemo.collections.tts.models import ssl_tts -from nemo.collections.tts.parts.utils.tts_dataset_utils import get_base_dir +from nemo.collections.tts.parts.utils.tts_dataset_utils import get_base_dir, segment_wav from nemo.core.classes import Dataset from nemo.core.classes.common import safe_instantiate from nemo.utils import logging @@ -126,33 +126,23 @@ def __getitem__(self, index): } -def segment_wav(wav, segment_length, segment_hop_size, min_segment_length): - if len(wav) < segment_length: - pad = torch.zeros(segment_length - len(wav)) - segment = torch.cat([wav, pad]) - return [segment] - else: - si = 0 - segments = [] - while si < len(wav) - min_segment_length: - segment = wav[si : si + segment_length] - if len(segment) < segment_length: - pad = torch.zeros(segment_length - len(segment)) - segment = torch.cat([segment, pad]) - segments.append(segment) - si += segment_hop_size - return segments - - def segment_batch(batch, segment_length=44100, segment_hop_size=22050, min_segment_length=22050): all_segments = [] segment_indices = [] si = 0 + for bidx in range(len(batch['audio'])): audio = batch['audio'][bidx] audio_length = batch['audio_len'][bidx] audio_actual = audio[:audio_length] - audio_segments = segment_wav(audio_actual, segment_length, segment_hop_size, min_segment_length) + + audio_segments = segment_wav( + wav=audio_actual, + segment_length=segment_length, + segment_hop_size=segment_hop_size, + min_segment_length=min_segment_length, + ) + all_segments += audio_segments segment_indices.append((si, si + len(audio_segments) - 1)) si += len(audio_segments) diff --git a/scripts/ssl_tts/ssl_tts_vc.py b/scripts/ssl_tts/ssl_tts_vc.py index 66e551304130..fcf84586a009 100644 --- a/scripts/ssl_tts/ssl_tts_vc.py +++ b/scripts/ssl_tts/ssl_tts_vc.py @@ -26,7 +26,7 @@ from nemo.collections.asr.parts.preprocessing.features import WaveformFeaturizer from nemo.collections.tts.models import fastpitch_ssl, hifigan, ssl_tts -from nemo.collections.tts.parts.utils.tts_dataset_utils import get_base_dir +from nemo.collections.tts.parts.utils.tts_dataset_utils import get_base_dir, segment_wav def load_wav(wav_path, wav_featurizer, pad_multiple=1024): @@ -63,43 +63,29 @@ def get_pitch_contour(wav, pitch_mean=None, pitch_std=None, compute_mean_std=Fal return pitch_contour -def segment_wav(wav, segment_length=44100, hop_size=44100, min_segment_size=22050): - if len(wav) < segment_length: - pad = torch.zeros(segment_length - len(wav)) - segment = torch.cat([wav, pad]) - return [segment] - else: - si = 0 - segments = [] - while si < len(wav) - min_segment_size: - segment = wav[si : si + segment_length] - if len(segment) < segment_length: - pad = torch.zeros(segment_length - len(segment)) - segment = torch.cat([segment, pad]) - - segments.append(segment) - si += hop_size - return segments - - def get_speaker_embedding(ssl_model, wav_featurizer, audio_paths, duration=None, device="cpu"): all_segments = [] all_wavs = [] + for audio_path in audio_paths: wav = load_wav(audio_path, wav_featurizer) - segments = segment_wav(wav) + segments = segment_wav(wav=wav) + all_segments += segments all_wavs.append(wav) + if duration is not None and len(all_segments) >= duration: - # each segment is 2 seconds with one second overlap. - # so 10 segments would mean 0 to 2, 1 to 3.. 9 to 11 (11 seconds.) + # Each segment is 2 seconds with one second overlap. + # So 10 segments would mean 0 to 2, 1 to 3...9 to 11 (11 seconds). all_segments = all_segments[: int(duration)] break signal_batch = torch.stack(all_segments) signal_length_batch = torch.stack([torch.tensor(signal_batch.shape[1]) for _ in range(len(all_segments))]) + signal_batch = signal_batch.to(device) signal_length_batch = signal_length_batch.to(device) + _, speaker_embeddings, _, _, _ = ssl_model.forward_for_export( input_signal=signal_batch, input_signal_length=signal_length_batch, normalize_content=True ) diff --git a/tests/collections/tts/parts/utils/test_tts_dataset_utils.py b/tests/collections/tts/parts/utils/test_tts_dataset_utils.py index 5e502221f771..01feb665b2c1 100644 --- a/tests/collections/tts/parts/utils/test_tts_dataset_utils.py +++ b/tests/collections/tts/parts/utils/test_tts_dataset_utils.py @@ -30,6 +30,7 @@ chunk_and_tokenize_text_by_sentence, chunk_text_for_inference, filter_dataset_by_duration, + segment_wav, get_abs_rel_paths, get_audio_filepaths, get_tokenizer_for_language, @@ -264,6 +265,34 @@ def test_filter_dataset_by_duration(self): assert total_hours == (135.6 / 3600.0) assert filtered_hours == (15.0 / 3600.0) + @pytest.mark.run_only_on('CPU') + @pytest.mark.unit + def test_segment_wav_pads_short_input(self): + wav = torch.arange(4, dtype=torch.float) + + segments = segment_wav(wav, segment_length=10, segment_hop_size=5, min_segment_length=3) + + assert len(segments) == 1 + assert len(segments[0]) == 10 + assert torch.equal(segments[0][:4], wav) + assert torch.equal(segments[0][4:], torch.zeros(6)) + + @pytest.mark.run_only_on('CPU') + @pytest.mark.unit + def test_segment_wav_sliding_window(self): + wav = torch.arange(20, dtype=torch.float) + + segments = segment_wav(wav, segment_length=10, segment_hop_size=5, min_segment_length=3) + + assert len(segments) == 4 + assert torch.equal(segments[0], wav[0:10]) + assert torch.equal(segments[1], wav[5:15]) + assert torch.equal(segments[2], wav[10:20]) + + assert len(segments[3]) == 10 + assert torch.equal(segments[3][:5], wav[15:20]) + assert torch.equal(segments[3][5:], torch.zeros(5)) + class TestLanguageThresholds: """Test cases for LanguageThresholds dataclass."""