Repository navigation
Expand file tree
/
Copy pathaugmentation.py
More file actions
253 lines (220 loc) · 9.57 KB
/
Copy pathaugmentation.py
File metadata and controls
253 lines (220 loc) · 9.57 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
"""Classes for implementing data augmentation pipelines.
Authors
* Mirco Ravanelli 2020
"""
import torch
import random
from speechbrain.utils.callchains import lengths_arg_exists
class Augmenter(torch.nn.Module):
"""Applies pipelines of data augmentation.
Arguments
---------
**augmentations: dict
The inputs are treated as a dictionary containing the name assigned to
the augmentation and the corresponding objects.
The augmentations are applied in sequence (or parallel).
parallel_augment: bool
If False, the augmentations are applied sequentially with
the order specified in the pipeline argument (one orignal input, one
augmented output).
When True, all the N augmentations are concatenated in the output
on the batch axis (one orignal input, N augmented output)
parallel_augment_fixed_bs: bool
If False, each augmenter (performed in parallel) generates a number of
augmented examples equal to the batch size. Thus, overall, with this option N*batch size artificial data are
generated, where N is the number of augmenters.
When True, the number of total augmented examples is kept fixed at
the batch size, thus, for each augmenter, fixed at batch size // N examples.
This option is useful to keep controlled the number of synthetic examples
with respect to the original data distribution, as it keep always
50% of original data, and 50% of augmented data.
concat_original: bool
if True, the original input is concatenated with the
augmented outputs (on the batch axis).
min_augmentations: int
The number of augmentations applied to the input signal is randomly
sampled between min_augmentations and max_augmentations. For instance,
if the augmentation dict contains N=6 augmentations and we set
select min_augmentations=1 and max_augmentations=4 we apply up to
M=4 augmentations. The selected augmentations are applied in the order
specified in the augmentations dict. If shuffle_augmentations = True,
a random set of M augmentations is selected.
max_augmentations: int
Maximum number of augmentations to apply. See min_augmentations for
more details.
shuffle_augmentations: bool
If True, it shuffles the entries of the augmentations dictionary.
The effect is to randomply select the order of the augmentations
to apply.
repeat_augment: int
Applies the augmentation algorithm N times. This can be used to
perform more data augmentation.
Example
-------
>>> from speechbrain.processing.speech_augmentation import DropFreq, DropChunk
>>> freq_dropper = DropFreq()
>>> chunk_dropper = DropChunk(drop_start=100, drop_end=16000)
>>> augment = Augmenter(parallel_augment=False, concat_original=False, freq_dropper=freq_dropper, chunk_dropper= chunk_dropper)
>>> signal = torch.rand([4, 16000])
>>> output_signal, lenghts = augment(signal, lengths=torch.tensor([0.2,0.5,0.7,1.0]))
"""
def __init__(
self,
parallel_augment=False,
parallel_augment_fixed_bs=False,
concat_original=False,
min_augmentations=None,
max_augmentations=None,
shuffle_augmentations=False,
repeat_augment=1,
**augmentations,
):
super().__init__()
self.parallel_augment = parallel_augment
self.parallel_augment_fixed_bs = parallel_augment_fixed_bs
self.concat_original = concat_original
self.augmentations = augmentations
self.min_augmentations = min_augmentations
self.max_augmentations = max_augmentations
self.shuffle_augmentations = shuffle_augmentations
self.repeat_augment = repeat_augment
# Check min and max augmentations
self.check_min_max_augmentations()
# Check repeat augment arguments
if not isinstance(self.repeat_augment, int):
raise ValueError("repeat_augment must be an integer.")
if self.repeat_augment < 0:
raise ValueError("repeat_augment must be greater than 0.")
# Check if augmentation modules need the length argument
self.require_lengths = {}
for aug_key, aug_fun in self.augmentations.items():
self.require_lengths[aug_key] = lengths_arg_exists(aug_fun.forward)
def augment(self, x, lengths, selected_augmentations):
"""Applies data augmentation on the seleted augmentations.
Arguments
---------
x : torch.Tensor (batch, time, channel)
input to augment.
lengths : torch.Tensor
The length of each sequence in the batch.
selected_augmentations: dict
Dictionary containg the selected augmentation to apply.
"""
next_input = x
next_lengths = lengths
output = []
output_lengths = []
out_lengths = lengths
for k, augment_name in enumerate(selected_augmentations):
augment_fun = self.augmentations[augment_name]
idx = torch.arange(x.shape[0])
if self.parallel_augment and self.parallel_augment_fixed_bs:
idx_startstop = torch.linspace(
0, x.shape[0], len(selected_augmentations) + 1
).to(torch.int)
idx_start = idx_startstop[k]
idx_stop = idx_startstop[k + 1]
idx = idx[idx_start:idx_stop]
# Check input arguments
if self.require_lengths[augment_name]:
out = augment_fun(
next_input[idx, ...], lengths=next_lengths[idx]
)
else:
out = augment_fun(next_input[idx, ...])
# Check output arguments
if isinstance(out, tuple):
if len(out) == 2:
out, out_lengths = out
else:
raise ValueError(
"The function must return max two arguments (Tensor, Length[optional])"
)
# Manage sequential or parallel augmentation
if not self.parallel_augment:
next_input = out
next_lengths = out_lengths[idx]
else:
output.append(out)
output_lengths.append(out_lengths[idx])
if self.parallel_augment:
# Concatenate all the augmented data
output = torch.cat(output, dim=0)
output_lengths = torch.cat(output_lengths, dim=0)
else:
# Take the last agumented signal of the pipeline
output = out
output_lengths = out_lengths
return output, output_lengths
def forward(self, x, lengths):
"""Applies data augmentation.
Arguments
---------
x : torch.Tensor (batch, time, channel)
input to augment.
lengths : torch.Tensor
The length of each sequence in the batch.
"""
# Select the number of augmentations to apply
N_augment = torch.randint(
low=self.min_augmentations,
high=self.max_augmentations + 1,
size=(1,),
)
# Get augmentations list
augmentations_lst = list(self.augmentations.keys())
# No augmentation
if (
self.repeat_augment == 0
or N_augment == 0
or len(self.augmentations) == 0
):
return x, lengths
# Shuffle augmentation
if self.shuffle_augmentations:
random.shuffle(augmentations_lst)
# Select the augmentations to apply
selected_augmentations = augmentations_lst[0:N_augment]
# # Select the augmentations to apply
# selected_augmentations = list(self.augmentations.keys())[0:N_augment]
#
# # Shuffle augmentation
# if self.shuffle_augmentations:
# random.shuffle(selected_augmentations)
# Lists to collect the outputs
output_lst = []
output_len_lst = []
# Concatenate the original signal if required
if self.concat_original:
output_lst.append(x)
output_len_lst.append(lengths)
# Perform augmentations
for i in range(self.repeat_augment):
output, output_lengths = self.augment(
x, lengths, selected_augmentations
)
output_lst.append(output)
output_len_lst.append(output_lengths)
# Concatenate the final outputs
output = torch.cat(output_lst, dim=0)
output_lengths = torch.cat(output_len_lst, dim=0)
return output, output_lengths
def check_min_max_augmentations(self):
"""Checks the min_augmentations and max_augmentations arguments.
"""
if self.min_augmentations is None:
self.min_augmentations = 1
if self.max_augmentations is None:
self.max_augmentations = len(self.augmentations)
if self.max_augmentations > len(self.augmentations):
self.max_augmentations = len(self.augmentations)
if self.min_augmentations > len(self.augmentations):
self.min_augmentations = len(self.augmentations)
if self.max_augmentations < self.min_augmentations:
raise ValueError(
"max_augmentations cannot be smaller than min_augmentations "
)
if self.min_augmentations < 0:
raise ValueError("min_augmentations cannot be smaller than 0.")
if self.max_augmentations < 0:
raise ValueError("max_augmentations cannot be smaller than 0.")