mirror of
https://github.com/wassname/multifit.git
synced 2026-09-09 11:27:26 +08:00
Move LM and classifier parameters to configs
This commit is contained in:
Regular → Executable
+1
-1
@@ -23,4 +23,4 @@ class ULMFiT:
|
||||
return FireView(train=params.train_cls, validate_cls=params.validate_cls)
|
||||
|
||||
if __name__ == '__main__':
|
||||
fire.Fire(ULMFiT())
|
||||
fire.Fire(ULMFiT())
|
||||
|
||||
+24
-13
@@ -68,7 +68,7 @@ class LMHyperParams:
|
||||
|
||||
# these hyperparameters are for training on ~100M tokens (e.g. WikiText-103)
|
||||
# for training on smaller datasets, more dropout is necessary
|
||||
dps = (0.25, 0.1, 0.2, 0.02, 0.15) # consider removing dps & clip from the default hyperparams and put them to train
|
||||
dps = dict(output_p=0.25, hidden_p=0.1, input_p=0.2, embed_p=0.02, weight_p=0.15) # consider removing dps & clip from the default hyperparams and put them to train
|
||||
clip: float = 0.12
|
||||
bptt: int = 70
|
||||
|
||||
@@ -93,7 +93,6 @@ class LMHyperParams:
|
||||
print('Max vocab:', self.max_vocab)
|
||||
print('Cache dir:', self.cache_dir)
|
||||
print('Model dir:', self.model_dir)
|
||||
self.dps = np.array(self.dps)
|
||||
if self.nh is None: self.nh = 1550 if self.qrnn else 1150
|
||||
if self.name is None: self.name = self.lang
|
||||
|
||||
@@ -175,19 +174,29 @@ class LMHyperParams:
|
||||
print(learn.path)
|
||||
|
||||
self.save_info()
|
||||
return learn
|
||||
# do we need to return `learn'? it adds noise to Fire output
|
||||
#return learn
|
||||
|
||||
def create_lm_learner(self, data_lm, dps=None, **kwargs):
|
||||
fastai.text.learner.default_dropout['language'] = dps or self.dps
|
||||
lm_learner = bilm_learner if self.bidir else language_model_learner
|
||||
|
||||
trn_args = dict(tie_weights=True, clip=self.clip, bptt=self.bptt,
|
||||
pretrained_fnames=self.pretrained_fnames,
|
||||
pretrained_model=self.pretrained_model)
|
||||
assert self.bidir == False, "bidirectional model is not yet supported"
|
||||
config = dict(emb_sz=self.emb_sz, n_hid=self.nh, n_layers=self.nl, pad_token=PAD_TOKEN_ID, qrnn=self.qrnn, bidir=self.bidir,
|
||||
tie_weights=True, out_bias=True)
|
||||
config.update(dps or self.dps)
|
||||
trn_args = dict(clip=self.clip)
|
||||
trn_args.update(kwargs)
|
||||
print ("Training args: ", trn_args, "dps: ", dps or self.dps)
|
||||
learn = lm_learner(data_lm, emb_sz=self.emb_sz, nh=self.nh, nl=self.nl, pad_token=PAD_TOKEN_ID,
|
||||
bias=True, qrnn=self.qrnn, model_dir=self.model_dir.relative_to(data_lm.path), **trn_args)
|
||||
learn = language_model_learner(data_lm, AWD_LSTM, config=config, model_dir=self.model_dir.relative_to(data_lm.path), pretrained=False, **trn_args)
|
||||
if self.pretrained_model is not None:
|
||||
print("Loading pretrained model")
|
||||
model_path = untar_data(self.pretrained_model, data=False)
|
||||
fnames = [list(model_path.glob(f'*.{ext}'))[0] for ext in ['pth', 'pkl']]
|
||||
learn.load_pretrained(*fnames)
|
||||
learn.freeze()
|
||||
if self.pretrained_fnames is not None:
|
||||
print("Loading pretrained model")
|
||||
fnames = [learn.path/learn.model_dir/f'{fn}.{ext}' for fn,ext in zip(self.pretrained_fnames, ['pth', 'pkl'])]
|
||||
learn.load_pretrained(*fnames)
|
||||
learn.freeze()
|
||||
# compared to standard Adam, we set beta_1 to 0.8
|
||||
learn.opt_fn = partial(optim.Adam, betas=(0.8, 0.99))
|
||||
learn.metrics = [accuracy_fwd, accuracy_bwd] if self.bidir else [accuracy]
|
||||
@@ -209,13 +218,15 @@ class LMHyperParams:
|
||||
|
||||
args = self.tokenzier_to_fastai_args(sp_data_func=self.load_train_text, use_moses=False)
|
||||
try:
|
||||
data_lm = TextLMDataBunch.load(self.cache_dir, '.', lm_type=self.lm_type, bs=bs)
|
||||
data_lm = TextLMDataBunch.load(self.cache_dir, '.',
|
||||
bs=bs)
|
||||
print("Tokenized data loaded")
|
||||
except FileNotFoundError:
|
||||
print("Running tokenization")
|
||||
data_lm = TextLMDataBunch.from_df(path=self.cache_dir, train_df=read_wiki_articles(trn_path),
|
||||
valid_df=read_wiki_articles(val_path),
|
||||
classes=None, lm_type=self.lm_type, max_vocab=self.max_vocab,
|
||||
classes=None,
|
||||
max_vocab=self.max_vocab,
|
||||
bs=bs, text_cols='texts', **args)
|
||||
data_lm.save('.')
|
||||
|
||||
|
||||
+16
-11
@@ -89,16 +89,21 @@ class CLSHyperParams(LMHyperParams):
|
||||
print(f"Loss and accuracy using ({save_name}):", learn.validate(data_tst.valid_dl))
|
||||
|
||||
def create_cls_learner(self, data_clas, dps=None, **kwargs):
|
||||
fastai.text.learner.default_dropout['language'] = dps or self.dps
|
||||
trn_args=dict(bptt=self.bptt, clip=self.clip,)
|
||||
assert self.bidir == False, "bidirectional model is not yet supported"
|
||||
config = dict(emb_sz=self.emb_sz, n_hid=self.nh, n_layers=self.nl, pad_token=PAD_TOKEN_ID, qrnn=self.qrnn, bidir=self.bidir)
|
||||
config.update(dps or self.dps)
|
||||
trn_args=dict(bptt=self.bptt, clip=self.clip)
|
||||
trn_args.update(kwargs)
|
||||
classifier_learner = text_classifier_learner
|
||||
if self.bidir:
|
||||
classifier_learner = bilm_text_classifier_learner
|
||||
trn_args['bicls_head'] = self.bicls_head
|
||||
learn = classifier_learner(data_clas, pad_token=PAD_TOKEN_ID,
|
||||
path=self.model_dir.parent, model_dir=self.model_dir.name,
|
||||
qrnn=self.qrnn, emb_sz=self.emb_sz, nh=self.nh, nl=self.nl, **trn_args)
|
||||
learn = text_classifier_learner(data_clas, AWD_LSTM, config=config,
|
||||
pretrained=False, path=self.model_dir.parent, model_dir=self.model_dir.name, **trn_args)
|
||||
|
||||
if self.pretrained_model is not None:
|
||||
print("Loading pretrained model")
|
||||
model_path = untar_data(self.pretrained_model, data=False)
|
||||
fnames = [list(model_path.glob(f'*.{ext}'))[0] for ext in ['pth', 'pkl']]
|
||||
learn.load_pretrained(*fnames, strict=False)
|
||||
learn.freeze()
|
||||
|
||||
learn.callback_fns += [partial(CSVLogger, filename=f"{learn.model_dir}/cls-history"),
|
||||
partial(SaveModelCallback, every='improvement', name='cls_best')]
|
||||
return learn
|
||||
@@ -152,12 +157,12 @@ class CLSHyperParams(LMHyperParams):
|
||||
args = self.tokenzier_to_fastai_args(sp_data_func=lambda: trn_df[1], use_moses=use_moses)
|
||||
try:
|
||||
if force: raise FileNotFoundError("Forcing reloading of caches")
|
||||
data_lm = TextLMDataBunch.load(self.cache_dir, 'lm', lm_type=self.lm_type, bs=bs)
|
||||
data_lm = TextLMDataBunch.load(self.cache_dir, 'lm', bs=bs)
|
||||
print(f"Tokenized data loaded, lm.trn {len(data_lm.train_ds)}, lm.val {len(data_lm.valid_ds)}")
|
||||
except FileNotFoundError:
|
||||
print(f"Running tokenization...")
|
||||
data_lm = TextLMDataBunch.from_df(path=self.cache_dir, train_df=lm_trn_df, valid_df=lm_val_df,
|
||||
max_vocab=self.max_vocab, bs=bs, lm_type=self.lm_type, **args)
|
||||
max_vocab=self.max_vocab, bs=bs, **args)
|
||||
print(f"Saving tokenized: cls.trn {len(data_lm.train_ds)}, cls.val {len(data_lm.valid_ds)}")
|
||||
data_lm.save('lm')
|
||||
|
||||
|
||||
Reference in New Issue
Block a user