mirror of
https://github.com/mozilla/DeepSpeech.git
synced 2025-10-26 11:19:39 +00:00
updates from transfer-learning
This commit is contained in:
parent
f2f35ac24e
commit
a4caebceb4
4
.compute
4
.compute
@ -35,9 +35,9 @@
|
||||
source ../tmp/venv/bin/activate
|
||||
|
||||
# if num_layers_dropped=6, then we train from scratch (you dropped all the layers!)
|
||||
num_layers_dropped=5
|
||||
num_layers_dropped=6
|
||||
|
||||
target_lang='de'
|
||||
target_lang='fr'
|
||||
source_model='en-v030'
|
||||
|
||||
|
||||
|
||||
2
.install
2
.install
@ -1,6 +1,6 @@
|
||||
#!/bin/bash
|
||||
|
||||
transfer=1
|
||||
transfer=0
|
||||
|
||||
## SETUP VIRETUALENV ##
|
||||
virtualenv -p python3 ../tmp/venv
|
||||
|
||||
@ -117,7 +117,8 @@ def evaluate(test_data, inference_graph, alphabet):
|
||||
sparse_labels = tf.cast(ctc_label_dense_to_sparse(labels_ph, label_lengths_ph, FLAGS.test_batch_size), tf.int32)
|
||||
loss = tf.nn.ctc_loss(labels=sparse_labels,
|
||||
inputs=layers['raw_logits'],
|
||||
sequence_length=inputs['input_lengths'])
|
||||
sequence_length=inputs['input_lengths'],
|
||||
ignore_longer_outputs_than_inputs=True)
|
||||
|
||||
# Create a saver using variables from the above newly created graph
|
||||
mapping = {v.op.name: v for v in tf.global_variables() if not v.op.name.startswith('previous_state_')}
|
||||
|
||||
@ -2,6 +2,7 @@ import pandas
|
||||
import sys
|
||||
import os
|
||||
import ntpath
|
||||
import glob
|
||||
|
||||
'''
|
||||
USAGE: $ python3 filter_cv1_dev_test.py DATA_DIR LOCALE CLIPS.TSV SAVE_TO_DIR
|
||||
@ -57,7 +58,6 @@ output_folder = sys.argv[4]
|
||||
|
||||
print("Looking for clips.tsv here: ", clips_tsv)
|
||||
clips = pandas.read_csv(clips_tsv, sep='\t')
|
||||
|
||||
print("### Total valid+invalid clips per lang ###")
|
||||
print(clips['locale'].value_counts())
|
||||
print("############################")
|
||||
@ -74,16 +74,6 @@ speaker_counts.columns= ['ID', 'counts']
|
||||
speaker_counts['cum_sum'] = speaker_counts.counts.cumsum()
|
||||
speaker_counts['cum_perc'] = 100 * speaker_counts.cum_sum / speaker_counts.counts.sum()
|
||||
|
||||
### BIG LANGS ###
|
||||
# num_spks = len(speaker_counts.index)
|
||||
# train = ['train']*8
|
||||
# dev = ['dev']*1
|
||||
# test = ['test']*1
|
||||
# splits = train + dev + test
|
||||
# speaker_counts['new_bucket'] = pandas.Series((splits*num_spks)[:num_spks]).values
|
||||
##^ BIG LANGS ^##
|
||||
|
||||
### LIL LANGS ###
|
||||
def bucket(row):
|
||||
if (row.cum_perc > 0.0) and( row.cum_perc < 80.0) :
|
||||
return "train"
|
||||
@ -91,12 +81,11 @@ def bucket(row):
|
||||
return "dev"
|
||||
elif (row.cum_perc > 90.0) and( row.cum_perc <= 100.0) :
|
||||
return "test"
|
||||
|
||||
speaker_counts.loc[:, 'new_bucket'] = speaker_counts.apply(bucket, axis = 1)
|
||||
##^ LIL LANGS ^##
|
||||
# speaker_counts.to_csv("speaker_count.csv", sep="\t")
|
||||
|
||||
|
||||
locale['new_bucket']=locale['ID'].map(speaker_counts.set_index('ID')['new_bucket'])
|
||||
### REASSIGN TEXT / DEV / TRAIN ###
|
||||
print(locale.head())
|
||||
|
||||
|
||||
locale['path'] = locale['path'].str.replace('/', '___')
|
||||
@ -112,13 +101,22 @@ train_paths = locale[locale['new_bucket'] == 'train'].loc[:, ['path']]
|
||||
#### ####
|
||||
|
||||
# cv_LANG_valid.csv == wav_filename,wav_filesize,transcript
|
||||
### don't use splits on cluster quite yet ###
|
||||
all_files = glob.glob(os.path.join(data_dir, LOCALE, "*_valid_*.csv"))
|
||||
print((str(i) for i in all_files))
|
||||
df_from_each_file = (pandas.read_csv(f) for f in all_files)
|
||||
validated_clips = pandas.concat(df_from_each_file, ignore_index=True)
|
||||
# validated_clips = pandas.read_csv('{}/{}/cv_{}_valid.csv'.format(data_dir, LOCALE, LOCALE))
|
||||
### don't use splits on cluster quite yet ###
|
||||
print(validated_clips.head())
|
||||
|
||||
validated_clips = pandas.read_csv('{}/{}/cv_{}_valid.csv'.format(data_dir, LOCALE, LOCALE))
|
||||
# validated_clips = pandas.read_csv('{}/{}/cv_{}_valid.csv'.format(data_dir, LOCALE, LOCALE))
|
||||
validated_clips['path'] = validated_clips['wav_filename'].apply(ntpath.basename)
|
||||
validated_clips['transcript'] = validated_clips['transcript'].str.replace(u'\xa0', ' ') # kyrgyz
|
||||
validated_clips['transcript'] = validated_clips['transcript'].str.replace(u'\xad', ' ') # catalan
|
||||
validated_clips['transcript'] = validated_clips['transcript'].str.replace(u'\\', ' ') # welsh
|
||||
|
||||
print(validated_clips.head())
|
||||
|
||||
|
||||
#### ####
|
||||
@ -129,10 +127,11 @@ validated_clips['transcript'] = validated_clips['transcript'].str.replace(u'\\'
|
||||
dev_indices = validated_clips['path'].isin(dev_paths['path'])
|
||||
test_indices = validated_clips['path'].isin(test_paths['path'])
|
||||
train_indices = validated_clips['path'].isin(train_paths['path'])
|
||||
validated_clips['wav_filename'] = data_dir + "/" + LOCALE + "/valid/" + validated_clips['path'].astype(str)
|
||||
validated_clips = validated_clips.drop(columns=['path'])
|
||||
validated_clips['wav_filename'] = data_dir + "/" + LOCALE + "/" + validated_clips['wav_filename'].astype(str)
|
||||
|
||||
|
||||
print(validated_clips.head())
|
||||
|
||||
|
||||
#### ####
|
||||
|
||||
Loading…
Reference in New Issue
Block a user