Spaces:
Build error
Build error
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- 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 <
|
| 1077 |
-
return "Duration too short ({:.3f} <
|
| 1078 |
|
| 1079 |
-
if
|
| 1080 |
-
|
|
|
|
| 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
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 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
|
|
|
|
| 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 |
-
|
|
|
|
|
|
|
| 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,
|
| 1166 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 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 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1191 |
for segment in segments:
|
| 1192 |
-
file = filename.replace(".wav", f"_{pad(segment['id'], 4)}.wav") if
|
| 1193 |
path = f'{indir}/audio/{file}'
|
|
|
|
| 1194 |
if not os.path.exists(path):
|
| 1195 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1196 |
continue
|
| 1197 |
|
| 1198 |
text = segment['text'].strip()
|
|
|
|
|
|
|
| 1199 |
if len(text) > 200:
|
| 1200 |
-
|
|
|
|
|
|
|
| 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 |
-
|
|
|
|
|
|
|
| 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'])
|