From 6c6a4e08ca19d9eaa33298752cad29d46b99b06f Mon Sep 17 00:00:00 2001 From: Tilman Kamp <5991088+tilmankamp@users.noreply.github.com> Date: Mon, 18 Mar 2019 18:53:04 +0100 Subject: [PATCH] Fix #1962 --- bin/import_cv.py | 33 +++++++++++++++++++++++++-------- bin/import_cv2.py | 13 ++++++++++--- util/text.py | 1 + 3 files changed, 36 insertions(+), 11 deletions(-) diff --git a/bin/import_cv.py b/bin/import_cv.py index f00fd815..f07ebeb3 100755 --- a/bin/import_cv.py +++ b/bin/import_cv.py @@ -8,13 +8,13 @@ import sys sys.path.insert(1, os.path.join(sys.path[0], '..')) import csv +import sox import tarfile import subprocess import progressbar from glob import glob from os import path -from sox import Transformer from threading import RLock from multiprocessing.dummy import Pool from multiprocessing import cpu_count @@ -66,7 +66,7 @@ def _maybe_convert_set(extracted_dir, source_csv, target_csv): samples.append((row['filename'], row['text'])) # Mutable counters for the concurrent embedded routine - counter = { 'all': 0, 'too_short': 0, 'too_long': 0 } + counter = { 'all': 0, 'failed': 0, 'invalid_label': 0, 'too_short': 0, 'too_long': 0 } lock = RLock() num_samples = len(samples) rows = [] @@ -78,9 +78,19 @@ def _maybe_convert_set(extracted_dir, source_csv, target_csv): wav_filename = path.splitext(mp3_filename)[0] + ".wav" _maybe_convert_wav(mp3_filename, wav_filename) frames = int(subprocess.check_output(['soxi', '-s', wav_filename], stderr=subprocess.STDOUT)) - file_size = path.getsize(wav_filename) + file_size = -1 + if path.exists(wav_filename): + file_size = path.getsize(wav_filename) + frames = int(subprocess.check_output(['soxi', '-s', wav_filename], stderr=subprocess.STDOUT)) + label = validate_label(sample[1]) with lock: - if int(frames/SAMPLE_RATE*1000/10/2) < len(str(sample[1])): + if file_size == -1: + # Excluding samples that failed upon conversion + counter['failed'] += 1 + elif label is None: + # Excluding samples that failed on label validation + counter['invalid_label'] += 1 + elif int(frames/SAMPLE_RATE*1000/10/2) < len(str(label)): # Excluding samples that are too short to fit the transcript counter['too_short'] += 1 elif frames/SAMPLE_RATE > MAX_SECS: @@ -88,7 +98,7 @@ def _maybe_convert_set(extracted_dir, source_csv, target_csv): counter['too_long'] += 1 else: # This one is good - keep it for the target CSV - rows.append((wav_filename, file_size, sample[1])) + rows.append((wav_filename, file_size, label)) counter['all'] += 1 print('Importing mp3 files...') @@ -108,7 +118,11 @@ def _maybe_convert_set(extracted_dir, source_csv, target_csv): for filename, file_size, transcript in bar(rows): writer.writerow({ 'wav_filename': filename, 'wav_filesize': file_size, 'transcript': transcript }) - print('Imported %d samples.' % (counter['all'] - counter['too_short'] - counter['too_long'])) + print('Imported %d samples.' % (counter['all'] - counter['failed'] - counter['too_short'] - counter['too_long'])) + if counter['failed'] > 0: + print('Skipped %d samples that failed upon conversion.' % counter['failed']) + if counter['invalid_label'] > 0: + print('Skipped %d samples that failed on transcript validation.' % counter['invalid_label']) if counter['too_short'] > 0: print('Skipped %d samples that were too short to match the transcript.' % counter['too_short']) if counter['too_long'] > 0: @@ -116,9 +130,12 @@ def _maybe_convert_set(extracted_dir, source_csv, target_csv): def _maybe_convert_wav(mp3_filename, wav_filename): if not path.exists(wav_filename): - transformer = Transformer() + transformer = sox.Transformer() transformer.convert(samplerate=SAMPLE_RATE) - transformer.build(mp3_filename, wav_filename) + try: + transformer.build(mp3_filename, wav_filename) + except sox.core.SoxError: + pass if __name__ == "__main__": _download_and_preprocess_data(sys.argv[1]) diff --git a/bin/import_cv2.py b/bin/import_cv2.py index 50c21521..00c3a541 100755 --- a/bin/import_cv2.py +++ b/bin/import_cv2.py @@ -17,6 +17,7 @@ from threading import RLock from multiprocessing.dummy import Pool from multiprocessing import cpu_count from util.downloader import SIMPLE_BAR +from util.text import validate_label ''' Broadly speaking, this script takes the audio downloaded from Common Voice @@ -62,7 +63,7 @@ def _maybe_convert_set(audio_dir, input_tsv): samples.append((row['path'], row['sentence'])) # Keep track of how many samples are good vs. problematic - counter = { 'all': 0, 'failed': 0, 'too_short': 0, 'too_long': 0 } + counter = { 'all': 0, 'failed': 0, 'invalid_label': 0, 'too_short': 0, 'too_long': 0 } lock = RLock() num_samples = len(samples) rows = [] @@ -79,11 +80,15 @@ def _maybe_convert_set(audio_dir, input_tsv): if path.exists(wav_filename): file_size = path.getsize(wav_filename) frames = int(subprocess.check_output(['soxi', '-s', wav_filename], stderr=subprocess.STDOUT)) + label = validate_label(sample[1]) with lock: if file_size == -1: # Excluding samples that failed upon conversion counter['failed'] += 1 - elif int(frames/SAMPLE_RATE*1000/10/2) < len(str(sample[1])): + elif label is None: + # Excluding samples that failed on label validation + counter['invalid_label'] += 1 + elif int(frames/SAMPLE_RATE*1000/10/2) < len(str(label)): # Excluding samples that are too short to fit the transcript counter['too_short'] += 1 elif frames/SAMPLE_RATE > MAX_SECS: @@ -91,7 +96,7 @@ def _maybe_convert_set(audio_dir, input_tsv): counter['too_long'] += 1 else: # This one is good - keep it for the target CSV - rows.append((wav_filename, file_size, sample[1])) + rows.append((wav_filename, file_size, label)) counter['all'] += 1 print("Importing mp3 files...") @@ -114,6 +119,8 @@ def _maybe_convert_set(audio_dir, input_tsv): print('Imported %d samples.' % (counter['all'] - counter['failed'] - counter['too_short'] - counter['too_long'])) if counter['failed'] > 0: print('Skipped %d samples that failed upon conversion.' % counter['failed']) + if counter['invalid_label'] > 0: + print('Skipped %d samples that failed on transcript validation.' % counter['invalid_label']) if counter['too_short'] > 0: print('Skipped %d samples that were too short to match the transcript.' % counter['too_short']) if counter['too_long'] > 0: diff --git a/util/text.py b/util/text.py index 904ad18f..ec2be8e0 100644 --- a/util/text.py +++ b/util/text.py @@ -133,6 +133,7 @@ def validate_label(label): label = label.replace(".", "") label = label.replace(",", "") label = label.replace("?", "") + label = label.replace("\"", "") label = label.strip() return label.lower()