Skip to content
51 changes: 51 additions & 0 deletions nemo/collections/tts/parts/utils/tts_dataset_utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -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:
Expand Down
30 changes: 10 additions & 20 deletions scripts/ssl_tts/make_supdata.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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)
Expand Down
32 changes: 9 additions & 23 deletions scripts/ssl_tts/ssl_tts_vc.py
Original file line number Diff line number Diff line change
Expand Up @@ -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):
Expand Down Expand Up @@ -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
)
Expand Down
29 changes: 29 additions & 0 deletions tests/collections/tts/parts/utils/test_tts_dataset_utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down Expand Up @@ -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."""
Expand Down
Loading