Add no-test + ability to start training from the classifcation data set without wiki

This commit is contained in:
Piotr Czapla
2019-05-18 13:10:29 +02:00
parent 75b934bb83
commit 7a33ea5d1f
3 changed files with 19 additions and 9 deletions
+9 -4
View File
@@ -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
+3 -3
View File
@@ -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
+7 -2
View File
@@ -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)