1+ import multiprocessing
12import os
23
34from packaging import version
6061class Sequence (object ):
6162
6263 # number of parallel worker for the data generator
63- nb_workers = 6
64+ nb_workers = multiprocessing . cpu_count () - 1
6465
6566 def __init__ (
6667 self ,
@@ -89,7 +90,8 @@ def __init__(
8990 multiprocessing = True ,
9091 features_indices = None ,
9192 transformer_name : str = None ,
92- report_to_wandb = False
93+ report_to_wandb = False ,
94+ nb_workers = 6
9395 ):
9496
9597 if model_name is None :
@@ -107,6 +109,7 @@ def __init__(
107109 self .embeddings_name = embeddings_name
108110
109111 word_emb_size = 0
112+ self .nb_workers = nb_workers
110113 self .embeddings = None
111114 self .model_local_path = None
112115
@@ -127,22 +130,23 @@ def __init__(
127130 else :
128131 learning_rate = 2e-5
129132
130- self .model_config = ModelConfig (model_name = model_name ,
131- architecture = architecture ,
132- embeddings_name = embeddings_name ,
133- word_embedding_size = word_emb_size ,
134- char_emb_size = char_emb_size ,
135- char_lstm_units = char_lstm_units ,
136- max_char_length = max_char_length ,
137- word_lstm_units = word_lstm_units ,
138- max_sequence_length = max_sequence_length ,
139- dropout = dropout ,
140- recurrent_dropout = recurrent_dropout ,
141- fold_number = fold_number ,
142- batch_size = batch_size ,
143- use_ELMo = use_ELMo ,
144- features_indices = features_indices ,
145- transformer_name = transformer_name )
133+ self .model_config = ModelConfig (
134+ model_name = model_name ,
135+ architecture = architecture ,
136+ embeddings_name = embeddings_name ,
137+ word_embedding_size = word_emb_size ,
138+ char_emb_size = char_emb_size ,
139+ char_lstm_units = char_lstm_units ,
140+ max_char_length = max_char_length ,
141+ word_lstm_units = word_lstm_units ,
142+ max_sequence_length = max_sequence_length ,
143+ dropout = dropout ,
144+ recurrent_dropout = recurrent_dropout ,
145+ fold_number = fold_number ,
146+ batch_size = batch_size ,
147+ use_ELMo = use_ELMo ,
148+ features_indices = features_indices ,
149+ transformer_name = transformer_name )
146150
147151 self .training_config = TrainingConfig (learning_rate , batch_size , optimizer ,
148152 lr_decay , clip_gradients , max_epoch ,
@@ -259,7 +263,8 @@ def train_(self, x_train, y_train, f_train=None, x_valid=None, y_valid=None, f_v
259263 checkpoint_path = self .log_dir ,
260264 preprocessor = self .p ,
261265 transformer_preprocessor = self .model .transformer_preprocessor ,
262- enable_wandb = self .report_to_wandb
266+ enable_wandb = self .report_to_wandb ,
267+ nb_workers = self .nb_workers
263268 )
264269 trainer .train (x_train , y_train , x_valid , y_valid , features_train = f_train , features_valid = f_valid , callbacks = callbacks )
265270 if self .embeddings and self .embeddings .use_ELMo :
@@ -298,13 +303,16 @@ def train_nfold_(self, x_train, y_train, x_valid=None, y_valid=None, f_train=Non
298303 self .model_config .case_vocab_size = len (self .p .vocab_case )
299304 self .models = []
300305
301- trainer = Trainer (self .model ,
302- self .models ,
303- self .embeddings ,
304- self .model_config ,
305- self .training_config ,
306- checkpoint_path = self .log_dir ,
307- preprocessor = self .p )
306+ trainer = Trainer (
307+ self .model ,
308+ self .models ,
309+ self .embeddings ,
310+ self .model_config ,
311+ self .training_config ,
312+ checkpoint_path = self .log_dir ,
313+ preprocessor = self .p ,
314+ nb_workers = self .nb_workers
315+ )
308316
309317 trainer .train_nfold (x_train , y_train , x_valid , y_valid , f_train = f_train , f_valid = f_valid , callbacks = callbacks )
310318 if self .embeddings and self .embeddings .use_ELMo :
0 commit comments