mrq commited on
Commit
ee1b048
·
1 Parent(s): 0cf9db5

when creating the train/validatio datasets, use segments if the main audio's duration is too long, and slice to make the segments if they don't exist

Browse files
Files changed (1) hide show
  1. src/utils.py +89 -30
src/utils.py CHANGED
@@ -51,6 +51,9 @@ LEARNING_RATE_SCHEDULE = [ 9, 18, 25, 33 ]
51
 
52
  RESAMPLERS = {}
53
 
 
 
 
54
  args = None
55
  tts = None
56
  tts_loading = False
@@ -62,6 +65,9 @@ training_state = None
62
  current_voice = None
63
 
64
  def resample( waveform, input_rate, output_rate=44100 ):
 
 
 
65
  if input_rate == output_rate:
66
  return waveform, output_rate
67
 
@@ -1066,18 +1072,19 @@ def whisper_transcribe( file, language=None ):
1066
  result['segments'].append(reparsed)
1067
  return result
1068
 
1069
- def validate_waveform( waveform, sample_rate ):
1070
  if not torch.any(waveform < 0):
1071
  return "Waveform is empty"
1072
 
1073
  num_channels, num_frames = waveform.shape
1074
  duration = num_channels * num_frames / sample_rate
1075
 
1076
- if duration < 0.6:
1077
- return "Duration too short ({:.3f} < 0.6s)".format(duration)
1078
 
1079
- if duration > 11:
1080
- return "Duration too long (11s < {:.3f})".format(duration)
 
1081
 
1082
  return
1083
 
@@ -1122,7 +1129,24 @@ def transcribe_dataset( voice, language=None, skip_existings=False, progress=Non
1122
 
1123
  return f"Processed dataset to: {indir}"
1124
 
1125
- def slice_dataset( voice, trim_silence=True, start_offset=0, end_offset=0 ):
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1126
  indir = f'./training/{voice}/'
1127
  infile = f'{indir}/whisper.json'
1128
  messages = []
@@ -1130,7 +1154,8 @@ def slice_dataset( voice, trim_silence=True, start_offset=0, end_offset=0 ):
1130
  if not os.path.exists(infile):
1131
  raise Exception(f"Missing dataset: {infile}")
1132
 
1133
- results = json.load(open(infile, 'r', encoding="utf-8"))
 
1134
 
1135
  files = 0
1136
  segments = 0
@@ -1140,37 +1165,35 @@ def slice_dataset( voice, trim_silence=True, start_offset=0, end_offset=0 ):
1140
  path = f'./training/{voice}/{filename}'
1141
 
1142
  if not os.path.exists(path):
1143
- messages.append(f"Missing source audio: {filename}")
 
 
1144
  continue
1145
 
1146
  files += 1
1147
  result = results[filename]
1148
  waveform, sample_rate = torchaudio.load(path)
 
 
1149
 
1150
  for segment in result['segments']:
1151
- start = int((segment['start'] + start_offset) * sample_rate)
1152
- end = int((segment['end'] + end_offset) * sample_rate)
1153
-
1154
- if start < 0:
1155
- start = 0
1156
- if end >= waveform.shape[-1]:
1157
- end = waveform.shape[-1] - 1
1158
-
1159
- sliced = waveform[:, start:end]
1160
  file = filename.replace(".wav", f"_{pad(segment['id'], 4)}.wav")
1161
-
1162
- if trim_silence:
1163
- sliced = torchaudio.functional.vad( sliced, sample_rate )
1164
 
1165
- sliced, sample_rate = resample( sliced, sample_rate, 22050 )
1166
- torchaudio.save(f"{indir}/audio/{file}", sliced, sample_rate)
 
 
 
 
 
 
1167
 
1168
  segments +=1
1169
 
1170
  messages.append(f"Sliced segments: {files} => {segments}.")
1171
  return "\n".join(messages)
1172
 
1173
- def prepare_dataset( voice, use_segments, text_length, audio_length ):
1174
  indir = f'./training/{voice}/'
1175
  infile = f'{indir}/whisper.json'
1176
  messages = []
@@ -1187,31 +1210,67 @@ def prepare_dataset( voice, use_segments, text_length, audio_length ):
1187
 
1188
  for filename in results:
1189
  result = results[filename]
1190
- segments = result['segments'] if use_segments else [{'text': result['text']}]
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1191
  for segment in segments:
1192
- file = filename.replace(".wav", f"_{pad(segment['id'], 4)}.wav") if use_segments else filename
1193
  path = f'{indir}/audio/{file}'
 
1194
  if not os.path.exists(path):
1195
- messages.append(f"Missing source audio: {file}")
 
 
 
 
 
 
 
 
 
 
1196
  continue
1197
 
1198
  text = segment['text'].strip()
 
 
1199
  if len(text) > 200:
1200
- messages.append(f"[{file}] Text length too long (200 < {len(text)}), skipping...")
 
 
1201
 
1202
  waveform, sample_rate = torchaudio.load(path)
1203
- num_channels, num_frames = waveform.shape
1204
- duration = num_channels * num_frames / sample_rate
1205
 
1206
  error = validate_waveform( waveform, sample_rate )
1207
  if error:
1208
- messages.append(f"[{file}]: {error}, skipping...")
 
 
1209
  continue
1210
 
1211
  culled = len(text) < text_length
1212
  if not culled and audio_length > 0:
 
 
1213
  culled = duration < audio_length
1214
 
 
1215
  lines['training' if not culled else 'validation'].append(f'audio/{file}|{text}')
1216
 
1217
  training_joined = "\n".join(lines['training'])
 
51
 
52
  RESAMPLERS = {}
53
 
54
+ MIN_TRAINING_DURATION = 0.6
55
+ MAX_TRAINING_DURATION = 11.6097505669
56
+
57
  args = None
58
  tts = None
59
  tts_loading = False
 
65
  current_voice = None
66
 
67
  def resample( waveform, input_rate, output_rate=44100 ):
68
+ # mono-ize
69
+ waveform = torch.mean(waveform, dim=0, keepdim=True)
70
+
71
  if input_rate == output_rate:
72
  return waveform, output_rate
73
 
 
1072
  result['segments'].append(reparsed)
1073
  return result
1074
 
1075
+ def validate_waveform( waveform, sample_rate, min_only=False ):
1076
  if not torch.any(waveform < 0):
1077
  return "Waveform is empty"
1078
 
1079
  num_channels, num_frames = waveform.shape
1080
  duration = num_channels * num_frames / sample_rate
1081
 
1082
+ if duration < MIN_TRAINING_DURATION:
1083
+ return "Duration too short ({:.3f}s < {:.3f}s)".format(duration, MIN_TRAINING_DURATION)
1084
 
1085
+ if not min_only:
1086
+ if duration > MAX_TRAINING_DURATION:
1087
+ return "Duration too long ({:.3f}s < {:.3f}s)".format(MAX_TRAINING_DURATION, duration)
1088
 
1089
  return
1090
 
 
1129
 
1130
  return f"Processed dataset to: {indir}"
1131
 
1132
+ def slice_waveform( waveform, sample_rate, start, end, trim ):
1133
+ start = int(start * sample_rate)
1134
+ end = int(end * sample_rate)
1135
+
1136
+ if start < 0:
1137
+ start = 0
1138
+ if end >= waveform.shape[-1]:
1139
+ end = waveform.shape[-1] - 1
1140
+
1141
+ sliced = waveform[:, start:end]
1142
+
1143
+ error = validate_waveform( sliced, sample_rate, min_only=True )
1144
+ if trim and not error:
1145
+ sliced = torchaudio.functional.vad( sliced, sample_rate )
1146
+
1147
+ return sliced, error
1148
+
1149
+ def slice_dataset( voice, trim_silence=True, start_offset=0, end_offset=0, results=None ):
1150
  indir = f'./training/{voice}/'
1151
  infile = f'{indir}/whisper.json'
1152
  messages = []
 
1154
  if not os.path.exists(infile):
1155
  raise Exception(f"Missing dataset: {infile}")
1156
 
1157
+ if results is None:
1158
+ results = json.load(open(infile, 'r', encoding="utf-8"))
1159
 
1160
  files = 0
1161
  segments = 0
 
1165
  path = f'./training/{voice}/{filename}'
1166
 
1167
  if not os.path.exists(path):
1168
+ message = f"Missing source audio: {filename}"
1169
+ print(message)
1170
+ messages.append(message)
1171
  continue
1172
 
1173
  files += 1
1174
  result = results[filename]
1175
  waveform, sample_rate = torchaudio.load(path)
1176
+ num_channels, num_frames = waveform.shape
1177
+ duration = num_channels * num_frames / sample_rate
1178
 
1179
  for segment in result['segments']:
 
 
 
 
 
 
 
 
 
1180
  file = filename.replace(".wav", f"_{pad(segment['id'], 4)}.wav")
 
 
 
1181
 
1182
+ sliced, error = slice_waveform( waveform, sample_rate, segment['start'] + start_offset, segment['end'] + end_offset, trim_silence )
1183
+ if error:
1184
+ message = f"{error}, skipping... {file}"
1185
+ print(message)
1186
+ messages.append(message)
1187
+ continue
1188
+ sliced, _ = resample( sliced, sample_rate, 22050 )
1189
+ torchaudio.save(f"{indir}/audio/{file}", sliced, 22050)
1190
 
1191
  segments +=1
1192
 
1193
  messages.append(f"Sliced segments: {files} => {segments}.")
1194
  return "\n".join(messages)
1195
 
1196
+ def prepare_dataset( voice, use_segments, text_length, audio_length, normalize=True ):
1197
  indir = f'./training/{voice}/'
1198
  infile = f'{indir}/whisper.json'
1199
  messages = []
 
1210
 
1211
  for filename in results:
1212
  result = results[filename]
1213
+ use_segment = use_segments
1214
+
1215
+ # check if unsegmented audio exceeds 11.6s
1216
+ if not use_segment:
1217
+ path = f'{indir}/audio/{filename}'
1218
+ if not os.path.exists(path):
1219
+ messages.append(f"Missing source audio: {filename}")
1220
+ continue
1221
+
1222
+ metadata = torchaudio.info(path)
1223
+ duration = metadata.num_channels * metadata.num_frames / metadata.sample_rate
1224
+ if duration >= MAX_TRAINING_DURATION:
1225
+ message = f"Audio too large, using segments: {filename}"
1226
+ print(message)
1227
+ messages.append(message)
1228
+ use_segment = True
1229
+
1230
+ segments = result['segments'] if use_segment else [{'text': result['text']}]
1231
+
1232
  for segment in segments:
1233
+ file = filename.replace(".wav", f"_{pad(segment['id'], 4)}.wav") if use_segment else filename
1234
  path = f'{indir}/audio/{file}'
1235
+ # segment when needed
1236
  if not os.path.exists(path):
1237
+ tmp_results = {}
1238
+ tmp_results[filename] = result
1239
+ print(f"Audio not segmented, segmenting: {filename}")
1240
+ message = slice_dataset( voice, results=tmp_results )
1241
+ print(message)
1242
+ messages = messages + message.split("\n")
1243
+
1244
+ if not os.path.exists(path):
1245
+ message = f"Missing source audio: {file}"
1246
+ print(message)
1247
+ messages.append(message)
1248
  continue
1249
 
1250
  text = segment['text'].strip()
1251
+ normalized_text = text
1252
+
1253
  if len(text) > 200:
1254
+ message = f"Text length too long (200 < {len(text)}), skipping... {file}"
1255
+ print(message)
1256
+ messages.append(message)
1257
 
1258
  waveform, sample_rate = torchaudio.load(path)
 
 
1259
 
1260
  error = validate_waveform( waveform, sample_rate )
1261
  if error:
1262
+ message = f"{error}, skipping... {file}"
1263
+ print(message)
1264
+ messages.append(message)
1265
  continue
1266
 
1267
  culled = len(text) < text_length
1268
  if not culled and audio_length > 0:
1269
+ num_channels, num_frames = waveform.shape
1270
+ duration = num_channels * num_frames / sample_rate
1271
  culled = duration < audio_length
1272
 
1273
+ # lines['training' if not culled else 'validation'].append(f'audio/{file}|{text}|{normalized_text}')
1274
  lines['training' if not culled else 'validation'].append(f'audio/{file}|{text}')
1275
 
1276
  training_joined = "\n".join(lines['training'])