diff --git a/ulmfit/pretrain_lm.py b/ulmfit/pretrain_lm.py index d0bb6dd..cb133bd 100644 --- a/ulmfit/pretrain_lm.py +++ b/ulmfit/pretrain_lm.py @@ -228,7 +228,8 @@ class LMHyperParams: learn.opt_fn = partial(optim.Adam, betas=(0.8, 0.99)) learn.metrics = [accuracy_fwd, accuracy_bwd] if self.bidir else [accuracy] learn.callback_fns += [partial(CSVLogger, filename=f"{learn.model_dir}/lm-history"), - partial(SaveModelCallback, every='improvement', name='lm')] + # partial(SaveModelCallback, every='improvement', name='lm') disabled due to Memory issues + ] return learn def load_train_text(self): diff --git a/ulmfit/train_clas.py b/ulmfit/train_clas.py index d47e8ab..46f4ab5 100644 --- a/ulmfit/train_clas.py +++ b/ulmfit/train_clas.py @@ -65,6 +65,7 @@ class CLSHyperParams(LMHyperParams): learn.fit_one_cycle(num_cls_epochs, slice(1e-2 / (2.6 ** 4), 1e-2), moms=(0.8, 0.7), wd=1e-7) print(f"Saving models at {learn.path / learn.model_dir}") learn.save('cls_last', with_opt=False) + learn.save('cls_best', with_opt=False) # we don't use early stopping for the time being return self.validate_cls('cls_best', bs=bs, data_tst=data_tst, learn=learn) @@ -96,7 +97,8 @@ class CLSHyperParams(LMHyperParams): learn.freeze() learn.callback_fns += [partial(CSVLogger, filename=f"{learn.model_dir}/cls-history"), - partial(SaveModelCallback, every='improvement', name='cls_best')] + #partial(SaveModelCallback, every='improvement', name='cls_best') disabled due to memory issues + ] return learn def load_cls_data(self, bs, **kwargs):