Skip to content

Commit 4d85b21

Browse files
committed
allow larger amount of workers
1 parent dd4293d commit 4d85b21

2 files changed

Lines changed: 39 additions & 28 deletions

File tree

delft/sequenceLabelling/trainer.py

Lines changed: 5 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -30,7 +30,8 @@ def __init__(self,
3030
save_path='',
3131
preprocessor: Preprocessor=None,
3232
transformer_preprocessor=None,
33-
enable_wandb = False
33+
enable_wandb = False,
34+
nb_workers=6
3435
):
3536

3637
# for single model training
@@ -39,6 +40,8 @@ def __init__(self,
3940
# for n-folds training
4041
self.models = models
4142

43+
self.nb_workers = nb_workers
44+
4245
self.embeddings = embeddings
4346
self.model_config = model_config
4447
self.training_config = training_config
@@ -192,7 +195,7 @@ def train_model(self, local_model, x_train, y_train, f_train=None,
192195
model=local_model,
193196
external_callbacks=callbacks
194197
)
195-
nb_workers = 6
198+
nb_workers = self.nb_workers
196199
multiprocessing = self.training_config.multiprocessing
197200

198201
# multiple workers should work with transformer layers, but not with ELMo due to GPU memory limit (with GTX 1080Ti 11GB)

delft/sequenceLabelling/wrapper.py

Lines changed: 34 additions & 26 deletions
Original file line numberDiff line numberDiff line change
@@ -1,3 +1,4 @@
1+
import multiprocessing
12
import os
23

34
from packaging import version
@@ -60,7 +61,7 @@
6061
class 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

Comments
 (0)