diff --git a/ulmfit/pretrain_lm.py b/ulmfit/pretrain_lm.py index c9d3b47..08530f2 100644 --- a/ulmfit/pretrain_lm.py +++ b/ulmfit/pretrain_lm.py @@ -113,17 +113,17 @@ class LMHyperParams: def lm_type(self): return contrib_data.LanguageModelType.BiLM if self.bidir else contrib_data.LanguageModelType.FwdLM - def tokenzier_to_fastai_args(self, trn_data_loading_func, add_moses): - tok_func = MosesTokenizerFunc if add_moses else BaseTokenizer + def tokenzier_to_fastai_args(self, sp_data_func, use_moses): + tok_func = MosesTokenizerFunc if use_moses else BaseTokenizer if self.tokenizer is Tokenizers.SUBWORD: - if self.base_lm_path: # ensure we are using the same sentence piece model + if self.base_lm_path and not(self.cache_dir/"spm.model").exists(): # ensure we are using the same sentence piece model shutil.copy(self.base_lm_path / '..' / 'itos.pkl', self.cache_dir) shutil.copy(self.base_lm_path / '..' / 'spm.model', self.cache_dir) shutil.copy(self.base_lm_path / '..' / 'spm.vocab', self.cache_dir) args = get_sentencepiece(self.cache_dir, - trn_data_loading_func, + sp_data_func, vocab_size=self.max_vocab, - use_moses=add_moses, + use_moses=use_moses, lang=self.lang) elif self.tokenizer is Tokenizers.MOSES: @@ -207,7 +207,7 @@ class LMHyperParams: for path_ in [trn_path, val_path, tst_path]: assert path_.exists(), f'Error: {path_} does not exist.' - args = self.tokenzier_to_fastai_args(trn_data_loading_func=self.load_train_text, add_moses=False) + 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) print("Tokenized data loaded") diff --git a/ulmfit/train_clas.py b/ulmfit/train_clas.py index d3a280b..3de435d 100644 --- a/ulmfit/train_clas.py +++ b/ulmfit/train_clas.py @@ -38,10 +38,10 @@ class CLSHyperParams(LMHyperParams): def need_fine_tune_lm(self): return not (self.model_dir/f"enc_best.pth").exists() def train_cls(self, num_lm_epochs, unfreeze=True, bs=40, true_wd=True, drop_mul_lm=0.3, drop_mul_cls=0.5, - use_test_for_validation=False): + use_test_for_validation=False, num_cls_epochs=2, limit=None, noise=0.0): assert use_test_for_validation == False, "use_test_for_validation=True is not supported" - data_clas, data_lm, data_tst = self.load_cls_data(bs) + data_clas, data_lm, data_tst = self.load_cls_data(bs, limit=limit, noise=noise) if self.need_fine_tune_lm: self.train_lm(num_lm_epochs, data_lm=data_lm, true_wd=true_wd, drop_mult=drop_mul_lm) learn = self.create_cls_learner(data_clas, drop_mult=drop_mul_cls) @@ -61,7 +61,7 @@ class CLSHyperParams(LMHyperParams): learn.freeze_to(-3) learn.fit_one_cycle(1, slice(5e-3 / (2.6 ** 4), 5e-3), moms=(0.8, 0.7)) learn.unfreeze() - learn.fit_one_cycle(2, slice(1e-3 / (2.6 ** 4), 1e-3), moms=(0.8, 0.7)) + learn.fit_one_cycle(num_cls_epochs, slice(1e-3 / (2.6 ** 4), 1e-3), moms=(0.8, 0.7)) else: learn.true_wd = False print("Starting classifier training") @@ -72,16 +72,21 @@ class CLSHyperParams(LMHyperParams): learn.freeze_to(-3) learn.fit_one_cycle(1, slice(5e-4 / (2.6 ** 4), 5e-4), moms=(0.8, 0.7), wd=1e-7) learn.unfreeze() - learn.fit_one_cycle(2, slice(1e-2 / (2.6 ** 4), 1e-2), moms=(0.8, 0.7), wd=1e-7) + 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) - self.validate_cls(learn, data_tst, 'cls_last', bs=bs) - self.validate_cls(learn, data_tst, 'cls_best', bs=bs) - return learn - def validate_cls(self, learn, data_tst, save_name='cls_last', bs=40): + self.validate_cls('cls_best', bs=bs, limit=limit, data_tst=data_tst, learn=learn) + return None + + def validate_cls(self, save_name='cls_last', limit=None, bs=40, data_tst=None, learn=None): + if data_tst is None: + _, _, data_tst = self.load_cls_data(bs, limit=limit) + if learn is None: + learn = self.create_cls_learner(data_tst, drop_mult=0.3) + learn.unfreeze() learn.load(save_name) - print(f"Loss and accuracy using ({save_name}):", learn.validate(data_tst.test_dl)) + 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 @@ -101,6 +106,7 @@ class CLSHyperParams(LMHyperParams): def load_cls_data(self, bs, **kwargs): add_trn_to_lm = True lang = self.lang + use_moses = True if 'xnli' in str(self.dataset_dir): NotImplementedError("Support for Xnli is not implemented yet") if 'imdb' in self.dataset_dir.name: @@ -109,7 +115,10 @@ class CLSHyperParams(LMHyperParams): add_trn_to_lm = False # False as trn_df is contained in unsup already lang = self.lang - data = self.load_data(lang=lang, add_trn_to_lm=add_trn_to_lm,**kwargs) + data = self.load_data(lang=lang, + add_trn_to_lm=add_trn_to_lm, + use_moses=use_moses, + **kwargs) return self.databunches(bs, **data) def load_data(self, lang='', **kwargs): @@ -134,13 +143,13 @@ class CLSHyperParams(LMHyperParams): kwargs.update(dict(trn_df=trn_df, val_df=val_df, tst_df=tst_df, unsup_df=unsup_df)) return kwargs - def databunches(self, bs, trn_df, val_df, tst_df, unsup_df, add_trn_to_lm=True, force=False): + def databunches(self, bs, trn_df, val_df, tst_df, unsup_df, add_trn_to_lm=True, use_moses=False, force=False, limit=None, noise=0.0): lm_trn_df = pd.concat([unsup_df, val_df, tst_df] + ([trn_df] if add_trn_to_lm else [])) val_len = max(int(len(lm_trn_df) * 0.1), 2) lm_trn_df = lm_trn_df[val_len:] lm_val_df = lm_trn_df[:val_len] - args = self.tokenzier_to_fastai_args(trn_data_loading_func=lambda: trn_df[1], add_moses=True) + 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) @@ -152,9 +161,24 @@ class CLSHyperParams(LMHyperParams): print(f"Saving tokenized: cls.trn {len(data_lm.train_ds)}, cls.val {len(data_lm.valid_ds)}") data_lm.save('lm') + cls_name="." + if limit is not None: + print("Limiting data set to:", limit) + trn_df = trn_df[:limit] + val_df = val_df[:limit] + cls_name=f'{cls_name}limit{limit}' + if noise > 0.0: + count = len(trn_df) + labels = trn_df[0].unique() + assert np.issubdtype(labels.dtype, np.integer), "noise only works on numerical numbers" + modulo = labels.max()+1 + idx_to_distrub = np.random.permutation(count)[:int(count * noise)] + trn_df.loc[idx_to_distrub, [0]] = (trn_df.loc[idx_to_distrub, [0]] + 1) % modulo + print(f"Added noise to {len(idx_to_distrub)} examples, only {(count-len(idx_to_distrub))/count} have correct labels") + cls_name = f'{cls_name}noise{noise}' try: if force: raise FileNotFoundError("Forcing reloading of caches") - data_cls = TextClasDataBunch.load(self.cache_dir, '.', bs=bs) + data_cls = TextClasDataBunch.load(self.cache_dir, cls_name, bs=bs) print(f"Tokenized data loaded, cls.trn {len(data_cls.train_ds)}, cls.val {len(data_cls.valid_ds)}") except FileNotFoundError: print(f"Running tokenization...") @@ -162,7 +186,7 @@ class CLSHyperParams(LMHyperParams): data_cls = TextClasDataBunch.from_df(path=self.cache_dir, train_df=trn_df, valid_df=val_df, max_vocab=self.max_vocab, bs=bs, **args) print(f"Saving tokenized: cls.trn {len(data_cls.train_ds)}, cls.val {len(data_cls.valid_ds)}") - data_cls.save('.') + data_cls.save(cls_name) # Hack to load test dataset with labels try: