From 5524006d81375f942fba721a97d873232c122eff Mon Sep 17 00:00:00 2001 From: Piotr Czapla Date: Tue, 4 Dec 2018 16:41:22 +0100 Subject: [PATCH] Respect max_vocab --- ulmfit/pretrain_lm.py | 11 +++++++---- 1 file changed, 7 insertions(+), 4 deletions(-) diff --git a/ulmfit/pretrain_lm.py b/ulmfit/pretrain_lm.py index 05af939..d40ad8a 100644 --- a/ulmfit/pretrain_lm.py +++ b/ulmfit/pretrain_lm.py @@ -62,6 +62,7 @@ class LMHyperParams: qrnn: bool = True max_vocab: int = 60000 tokenizer: Tokenizers = Tokenizers.MOSES + pretrained_model: str = None emb_sz:int = 400 nh: int = None @@ -132,7 +133,7 @@ class LMHyperParams: if num_epochs > 0: if self.pretrained_fnames : - learn.fit_one_cycle(1, 1e-2, moms=(0.8, 0.7)) + learn.fit_one_cycle(1, 1e-2, moms=(0.8, 0.7)) # TODO Fix the learning rates learn.unfreeze() if num_epochs > 0: learn.fit_one_cycle(num_epochs, 1e-3, moms=(0.8, 0.7)) else: @@ -158,7 +159,8 @@ class LMHyperParams: learn = lm_learner(data_lm, bptt=self.bptt, emb_sz=self.emb_sz, nh=self.nh, nl=self.nl, pad_token=PAD_TOKEN_ID, drop_mult=self.drop_mult, tie_weights=True, model_dir= self.model_dir.relative_to(data_lm.path), - bias=True, qrnn=self.qrnn, clip=self.clip, pretrained_fnames=self.pretrained_fnames) + bias=True, qrnn=self.qrnn, clip=self.clip, pretrained_fnames=self.pretrained_fnames, + pretrained_model=self.pretrained_model) # compared to standard Adam, we set beta_1 to 0.8 learn.opt_fn = partial(optim.Adam, betas=(0.8, 0.99)) learn.true_wd = False @@ -223,7 +225,8 @@ class LMHyperParams: pretokenized = Tokenizer(tok_func=BaseTokenizer, lang='en', pre_rules=None, post_rules=None) data_lm = TextLMDataBunch.from_df(path=self.cache_dir, train_df=read_file(trn_path), valid_df=read_file(val_path), tokenizer=pretokenized, - test_df=read_file(tst_path), classes=None, lm_type=self.lm_type) + test_df=read_file(tst_path), classes=None, lm_type=self.lm_type, + max_vocab=self.max_vocab) data_lm.save('.') elif self.tokenizer is Tokenizers.FASTAI: try: @@ -232,7 +235,7 @@ class LMHyperParams: except FileNotFoundError: print("Running tokenization") data_lm = TextLMDataBunch.from_df(path=self.cache_dir, train_df=read_file(trn_path), valid_df=read_file(val_path), - test_df=read_file(tst_path), classes=None, lm_type=self.lm_type) + test_df=read_file(tst_path), classes=None, lm_type=self.lm_type, max_vocab=self.max_vocab,) data_lm.save('.') else: raise ValueError(f"self.tokenizer has wrong value {self.tokenizer}, Allowed values are taken from {Tokenizers}")