diff --git a/ulmfit/__main__.py b/ulmfit/__main__.py index 7a692b4..9affd82 100644 --- a/ulmfit/__main__.py +++ b/ulmfit/__main__.py @@ -46,10 +46,14 @@ class ULMFiT: lm2 = LMHyperParams @wraps(CLSHyperParams) - def cls(self, dataset_path, base_lm_path, **changes): - params = CLSHyperParams.from_lm(dataset_path, base_lm_path, **changes) + def cls(self, dataset_path, base_lm_path=None, **changes): + if base_lm_path is not None: + params = CLSHyperParams.from_lm(dataset_path, base_lm_path, **changes) + else: + params = CLSHyperParams(dataset_path=dataset_path, **changes) return FireView(train=params.train_cls, validate_cls=params.validate_cls) + @wraps(CLSHyperParams) def load_cls(self, model_path, **changes): params = CLSHyperParams.from_json(model_path, **changes) @@ -86,10 +90,11 @@ class ULMFiT: tar.add(f, dest) - def poleval19_full(self, base, num_lm_epochs=6, lmtype=None, **kwargs): + def poleval19_full(self, base, num_lm_epochs=6, lmtype=None, skip_train_seed=False, **kwargs): clsbase = self.poleval19_init(base, num_lm_epochs=num_lm_epochs, lmtype=lmtype, **kwargs) self.poleval19_seeds(clsbase, seed_name='clsweightseed', **kwargs) - self.poleval19_seeds(clsbase, seed_name='clstrainseed', **kwargs) + if skip_train_seed: + self.poleval19_seeds(clsbase, seed_name='clstrainseed', **kwargs) def poleval19_init(self, base, name=None, lmseed=None, lmtype=None, **kwargs): clstrainseed = clsweightseed = ftseed = 0 diff --git a/ulmfit/pretrain_lm.py b/ulmfit/pretrain_lm.py index 03bf4d2..12ed153 100644 --- a/ulmfit/pretrain_lm.py +++ b/ulmfit/pretrain_lm.py @@ -191,10 +191,10 @@ class LMHyperParams: print("Saving info", self.model_dir / 'info.json') def train_lm(self, num_epochs=20, data_lm=None, bs=70, true_wd=False, drop_mult=0.0, lr=5e-3, label_smoothing_eps=0.0): - if not hasattr(self, 'ftseed'): - self.set_seed(self.lmseed, "LM") - else: + if self.pretrained_fnames or self.pretrained_model: self.set_seed(self.ftseed, "fine-tune") + else: + self.set_seed(self.lmseed, "LM") self.model_dir.mkdir(exist_ok=True, parents=True) data_lm = self.load_wiki_data(bs=bs) if data_lm is None else data_lm diff --git a/ulmfit/train_clas.py b/ulmfit/train_clas.py index 33aecc1..2fddf01 100644 --- a/ulmfit/train_clas.py +++ b/ulmfit/train_clas.py @@ -24,6 +24,8 @@ class CLSHyperParams(LMHyperParams): bicls_head:str = 'BiPoolingLinearClassifier' + use_tst_for_lm:bool = True + def __post_init__(self, *args, **kwargs): super().__post_init__(*args, **kwargs) self.dataset_dir=self.dataset_path @@ -238,7 +240,7 @@ class CLSHyperParams(LMHyperParams): self.model_dir.mkdir(exist_ok=True, parents=True) add_trn_to_lm = True lang = self.lang - use_moses = False #True + 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: @@ -247,6 +249,8 @@ class CLSHyperParams(LMHyperParams): if 'mldoc' in str(self.dataset_dir): add_trn_to_lm = False # False as trn_df is contained in unsup already lang = self.lang + if 'hate' in str(self.dataset_dir): + use_moses = False data = self.load_data(lang=lang, add_trn_to_lm=add_trn_to_lm, @@ -287,7 +291,7 @@ class CLSHyperParams(LMHyperParams): return trn_df 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 [])) + lm_trn_df = pd.concat([unsup_df, val_df] + ([tst_df] if self.use_tst_for_lm else []) + ([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] @@ -308,6 +312,7 @@ class CLSHyperParams(LMHyperParams): args['text_cols'] = list(trn_df.columns.values)[1:] args['mark_fields'] = True lm_suffix = self.bptt if self.bptt != 70 else "" + lm_suffix = self.use_tst_for_lm if "" else "-notst" data_lm = self.lm_databunch(f'lm{lm_suffix}', train_df=lm_trn_df, valid_df=lm_val_df, bs=bs, force=force, bptt=self.bptt, **args) args['vocab'] = data_lm.vocab data_cls = self.cls_databunch(cls_name, train_df=trn_df, valid_df=val_df, bs=bs, force=force, **args)