Apply Code Formatter Change
This commit is contained in:
committed by
github-actions[bot]
parent
d82ba3457a
commit
eef212fd00
153
data_utils.py
153
data_utils.py
@@ -16,9 +16,9 @@ from text import cleaned_text_to_sequence, get_bert
|
||||
|
||||
class TextAudioSpeakerLoader(torch.utils.data.Dataset):
|
||||
"""
|
||||
1) loads audio, speaker_id, text pairs
|
||||
2) normalizes text and converts them to sequences of integers
|
||||
3) computes spectrograms from audio files.
|
||||
1) loads audio, speaker_id, text pairs
|
||||
2) normalizes text and converts them to sequences of integers
|
||||
3) computes spectrograms from audio files.
|
||||
"""
|
||||
|
||||
def __init__(self, audiopaths_sid_text, hparams):
|
||||
@@ -32,7 +32,9 @@ class TextAudioSpeakerLoader(torch.utils.data.Dataset):
|
||||
self.spk_map = hparams.spk2id
|
||||
self.hparams = hparams
|
||||
|
||||
self.use_mel_spec_posterior = getattr(hparams, "use_mel_posterior_encoder", False)
|
||||
self.use_mel_spec_posterior = getattr(
|
||||
hparams, "use_mel_posterior_encoder", False
|
||||
)
|
||||
if self.use_mel_spec_posterior:
|
||||
self.n_mel_channels = getattr(hparams, "n_mel_channels", 80)
|
||||
|
||||
@@ -58,17 +60,26 @@ class TextAudioSpeakerLoader(torch.utils.data.Dataset):
|
||||
lengths = []
|
||||
skipped = 0
|
||||
logger.info("Init dataset...")
|
||||
for _id, spk, language, text, phones, tone, word2ph in tqdm(self.audiopaths_sid_text):
|
||||
audiopath = f'{_id}'
|
||||
for _id, spk, language, text, phones, tone, word2ph in tqdm(
|
||||
self.audiopaths_sid_text
|
||||
):
|
||||
audiopath = f"{_id}"
|
||||
if self.min_text_len <= len(phones) and len(phones) <= self.max_text_len:
|
||||
phones = phones.split(" ")
|
||||
tone = [int(i) for i in tone.split(" ")]
|
||||
word2ph = [int(i) for i in word2ph.split(" ")]
|
||||
audiopaths_sid_text_new.append([audiopath, spk, language, text, phones, tone, word2ph])
|
||||
audiopaths_sid_text_new.append(
|
||||
[audiopath, spk, language, text, phones, tone, word2ph]
|
||||
)
|
||||
lengths.append(os.path.getsize(audiopath) // (2 * self.hop_length))
|
||||
else:
|
||||
skipped += 1
|
||||
logger.info("skipped: " + str(skipped) + ", total: " + str(len(self.audiopaths_sid_text)))
|
||||
logger.info(
|
||||
"skipped: "
|
||||
+ str(skipped)
|
||||
+ ", total: "
|
||||
+ str(len(self.audiopaths_sid_text))
|
||||
)
|
||||
self.audiopaths_sid_text = audiopaths_sid_text_new
|
||||
self.lengths = lengths
|
||||
|
||||
@@ -76,7 +87,9 @@ class TextAudioSpeakerLoader(torch.utils.data.Dataset):
|
||||
# separate filename, speaker_id and text
|
||||
audiopath, sid, language, text, phones, tone, word2ph = audiopath_sid_text
|
||||
|
||||
bert, ja_bert, phones, tone, language = self.get_text(text, word2ph, phones, tone, language, audiopath)
|
||||
bert, ja_bert, phones, tone, language = self.get_text(
|
||||
text, word2ph, phones, tone, language, audiopath
|
||||
)
|
||||
|
||||
spec, wav = self.get_audio(audiopath)
|
||||
sid = torch.LongTensor([int(self.spk_map[sid])])
|
||||
@@ -85,8 +98,11 @@ class TextAudioSpeakerLoader(torch.utils.data.Dataset):
|
||||
def get_audio(self, filename):
|
||||
audio, sampling_rate = load_wav_to_torch(filename)
|
||||
if sampling_rate != self.sampling_rate:
|
||||
raise ValueError("{} {} SR doesn't match target {} SR".format(
|
||||
sampling_rate, self.sampling_rate))
|
||||
raise ValueError(
|
||||
"{} {} SR doesn't match target {} SR".format(
|
||||
sampling_rate, self.sampling_rate
|
||||
)
|
||||
)
|
||||
audio_norm = audio / self.max_wav_value
|
||||
audio_norm = audio_norm.unsqueeze(0)
|
||||
spec_filename = filename.replace(".wav", ".spec.pt")
|
||||
@@ -96,13 +112,26 @@ class TextAudioSpeakerLoader(torch.utils.data.Dataset):
|
||||
spec = torch.load(spec_filename)
|
||||
except:
|
||||
if self.use_mel_spec_posterior:
|
||||
spec = mel_spectrogram_torch(audio_norm, self.filter_length,
|
||||
self.n_mel_channels, self.sampling_rate, self.hop_length,
|
||||
self.win_length, self.hparams.mel_fmin, self.hparams.mel_fmax, center=False)
|
||||
spec = mel_spectrogram_torch(
|
||||
audio_norm,
|
||||
self.filter_length,
|
||||
self.n_mel_channels,
|
||||
self.sampling_rate,
|
||||
self.hop_length,
|
||||
self.win_length,
|
||||
self.hparams.mel_fmin,
|
||||
self.hparams.mel_fmax,
|
||||
center=False,
|
||||
)
|
||||
else:
|
||||
spec = spectrogram_torch(audio_norm, self.filter_length,
|
||||
self.sampling_rate, self.hop_length, self.win_length,
|
||||
center=False)
|
||||
spec = spectrogram_torch(
|
||||
audio_norm,
|
||||
self.filter_length,
|
||||
self.sampling_rate,
|
||||
self.hop_length,
|
||||
self.win_length,
|
||||
center=False,
|
||||
)
|
||||
spec = torch.squeeze(spec, 0)
|
||||
torch.save(spec, spec_filename)
|
||||
return spec, audio_norm
|
||||
@@ -125,17 +154,29 @@ class TextAudioSpeakerLoader(torch.utils.data.Dataset):
|
||||
torch.save(bert, bert_path)
|
||||
assert bert.shape[-1] == len(phone), phone
|
||||
|
||||
if language_str=='ZH':
|
||||
if language_str == "ZH":
|
||||
bert = bert
|
||||
ja_bert = torch.zeros(768, len(phone))
|
||||
elif language_str=="JA":
|
||||
elif language_str == "JA":
|
||||
ja_bert = bert
|
||||
bert = torch.zeros(1024, len(phone))
|
||||
else:
|
||||
bert = torch.zeros(1024, len(phone))
|
||||
ja_bert = torch.zeros(768, len(phone))
|
||||
assert bert.shape[-1] == len(phone), (
|
||||
bert.shape, len(phone), sum(word2ph), p1, p2, t1, t2, pold, pold2, word2ph, text, w2pho)
|
||||
bert.shape,
|
||||
len(phone),
|
||||
sum(word2ph),
|
||||
p1,
|
||||
p2,
|
||||
t1,
|
||||
t2,
|
||||
pold,
|
||||
pold2,
|
||||
word2ph,
|
||||
text,
|
||||
w2pho,
|
||||
)
|
||||
phone = torch.LongTensor(phone)
|
||||
tone = torch.LongTensor(tone)
|
||||
language = torch.LongTensor(language)
|
||||
@@ -152,9 +193,8 @@ class TextAudioSpeakerLoader(torch.utils.data.Dataset):
|
||||
return len(self.audiopaths_sid_text)
|
||||
|
||||
|
||||
class TextAudioSpeakerCollate():
|
||||
""" Zero-pads model inputs and targets
|
||||
"""
|
||||
class TextAudioSpeakerCollate:
|
||||
"""Zero-pads model inputs and targets"""
|
||||
|
||||
def __init__(self, return_ids=False):
|
||||
self.return_ids = return_ids
|
||||
@@ -167,8 +207,8 @@ class TextAudioSpeakerCollate():
|
||||
"""
|
||||
# Right zero-pad all one-hot text sequences to max input length
|
||||
_, ids_sorted_decreasing = torch.sort(
|
||||
torch.LongTensor([x[1].size(1) for x in batch]),
|
||||
dim=0, descending=True)
|
||||
torch.LongTensor([x[1].size(1) for x in batch]), dim=0, descending=True
|
||||
)
|
||||
|
||||
max_text_len = max([len(x[0]) for x in batch])
|
||||
max_spec_len = max([x[1].size(1) for x in batch])
|
||||
@@ -198,32 +238,44 @@ class TextAudioSpeakerCollate():
|
||||
row = batch[ids_sorted_decreasing[i]]
|
||||
|
||||
text = row[0]
|
||||
text_padded[i, :text.size(0)] = text
|
||||
text_padded[i, : text.size(0)] = text
|
||||
text_lengths[i] = text.size(0)
|
||||
|
||||
spec = row[1]
|
||||
spec_padded[i, :, :spec.size(1)] = spec
|
||||
spec_padded[i, :, : spec.size(1)] = spec
|
||||
spec_lengths[i] = spec.size(1)
|
||||
|
||||
wav = row[2]
|
||||
wav_padded[i, :, :wav.size(1)] = wav
|
||||
wav_padded[i, :, : wav.size(1)] = wav
|
||||
wav_lengths[i] = wav.size(1)
|
||||
|
||||
sid[i] = row[3]
|
||||
|
||||
tone = row[4]
|
||||
tone_padded[i, :tone.size(0)] = tone
|
||||
tone_padded[i, : tone.size(0)] = tone
|
||||
|
||||
language = row[5]
|
||||
language_padded[i, :language.size(0)] = language
|
||||
language_padded[i, : language.size(0)] = language
|
||||
|
||||
bert = row[6]
|
||||
bert_padded[i, :, :bert.size(1)] = bert
|
||||
bert_padded[i, :, : bert.size(1)] = bert
|
||||
|
||||
ja_bert = row[7]
|
||||
ja_bert_padded[i, :, :ja_bert.size(1)] = ja_bert
|
||||
ja_bert_padded[i, :, : ja_bert.size(1)] = ja_bert
|
||||
|
||||
return text_padded, text_lengths, spec_padded, spec_lengths, wav_padded, wav_lengths, sid, tone_padded, language_padded, bert_padded, ja_bert_padded
|
||||
return (
|
||||
text_padded,
|
||||
text_lengths,
|
||||
spec_padded,
|
||||
spec_lengths,
|
||||
wav_padded,
|
||||
wav_lengths,
|
||||
sid,
|
||||
tone_padded,
|
||||
language_padded,
|
||||
bert_padded,
|
||||
ja_bert_padded,
|
||||
)
|
||||
|
||||
|
||||
class DistributedBucketSampler(torch.utils.data.distributed.DistributedSampler):
|
||||
@@ -236,7 +288,15 @@ class DistributedBucketSampler(torch.utils.data.distributed.DistributedSampler):
|
||||
Ex) boundaries = [b1, b2, b3] -> any x s.t. length(x) <= b1 or length(x) > b3 are discarded.
|
||||
"""
|
||||
|
||||
def __init__(self, dataset, batch_size, boundaries, num_replicas=None, rank=None, shuffle=True):
|
||||
def __init__(
|
||||
self,
|
||||
dataset,
|
||||
batch_size,
|
||||
boundaries,
|
||||
num_replicas=None,
|
||||
rank=None,
|
||||
shuffle=True,
|
||||
):
|
||||
super().__init__(dataset, num_replicas=num_replicas, rank=rank, shuffle=shuffle)
|
||||
self.lengths = dataset.lengths
|
||||
self.batch_size = batch_size
|
||||
@@ -254,7 +314,7 @@ class DistributedBucketSampler(torch.utils.data.distributed.DistributedSampler):
|
||||
if idx_bucket != -1:
|
||||
buckets[idx_bucket].append(i)
|
||||
|
||||
try:
|
||||
try:
|
||||
for i in range(len(buckets) - 1, 0, -1):
|
||||
if len(buckets[i]) == 0:
|
||||
buckets.pop(i)
|
||||
@@ -262,7 +322,7 @@ class DistributedBucketSampler(torch.utils.data.distributed.DistributedSampler):
|
||||
assert all(len(bucket) > 0 for bucket in buckets)
|
||||
# When one bucket is not traversed
|
||||
except Exception as e:
|
||||
print('Bucket warning ', e)
|
||||
print("Bucket warning ", e)
|
||||
for i in range(len(buckets) - 1, -1, -1):
|
||||
if len(buckets[i]) == 0:
|
||||
buckets.pop(i)
|
||||
@@ -272,7 +332,9 @@ class DistributedBucketSampler(torch.utils.data.distributed.DistributedSampler):
|
||||
for i in range(len(buckets)):
|
||||
len_bucket = len(buckets[i])
|
||||
total_batch_size = self.num_replicas * self.batch_size
|
||||
rem = (total_batch_size - (len_bucket % total_batch_size)) % total_batch_size
|
||||
rem = (
|
||||
total_batch_size - (len_bucket % total_batch_size)
|
||||
) % total_batch_size
|
||||
num_samples_per_bucket.append(len_bucket + rem)
|
||||
return buckets, num_samples_per_bucket
|
||||
|
||||
@@ -293,21 +355,30 @@ class DistributedBucketSampler(torch.utils.data.distributed.DistributedSampler):
|
||||
for i in range(len(self.buckets)):
|
||||
bucket = self.buckets[i]
|
||||
len_bucket = len(bucket)
|
||||
if (len_bucket == 0):
|
||||
if len_bucket == 0:
|
||||
continue
|
||||
ids_bucket = indices[i]
|
||||
num_samples_bucket = self.num_samples_per_bucket[i]
|
||||
|
||||
# add extra samples to make it evenly divisible
|
||||
rem = num_samples_bucket - len_bucket
|
||||
ids_bucket = ids_bucket + ids_bucket * (rem // len_bucket) + ids_bucket[:(rem % len_bucket)]
|
||||
ids_bucket = (
|
||||
ids_bucket
|
||||
+ ids_bucket * (rem // len_bucket)
|
||||
+ ids_bucket[: (rem % len_bucket)]
|
||||
)
|
||||
|
||||
# subsample
|
||||
ids_bucket = ids_bucket[self.rank::self.num_replicas]
|
||||
ids_bucket = ids_bucket[self.rank :: self.num_replicas]
|
||||
|
||||
# batching
|
||||
for j in range(len(ids_bucket) // self.batch_size):
|
||||
batch = [bucket[idx] for idx in ids_bucket[j * self.batch_size:(j + 1) * self.batch_size]]
|
||||
batch = [
|
||||
bucket[idx]
|
||||
for idx in ids_bucket[
|
||||
j * self.batch_size : (j + 1) * self.batch_size
|
||||
]
|
||||
]
|
||||
batches.append(batch)
|
||||
|
||||
if self.shuffle:
|
||||
|
||||
Reference in New Issue
Block a user