From 6ef75e4dd7a919f67b29abefc704bcab54efdba2 Mon Sep 17 00:00:00 2001 From: Piotr Czapla Date: Mon, 13 May 2019 09:17:59 +0200 Subject: [PATCH] Fix issue in poleval_eval + add valid metrics --- ulmfit/__main__.py | 19 +++++++++++-------- ulmfit/train_clas.py | 10 +++++++++- 2 files changed, 20 insertions(+), 9 deletions(-) diff --git a/ulmfit/__main__.py b/ulmfit/__main__.py index 18ccb16..25b7bf8 100644 --- a/ulmfit/__main__.py +++ b/ulmfit/__main__.py @@ -109,6 +109,9 @@ class ULMFiT: print("Setting lmseed ", lmseed) dataset_template=f"../hate/pl-10-{lmtype}" + if name is None: + name = f"ft{kwargs.get('num_lm_epochs',6)}_cl{kwargs.get('num_cls_epochs',6)}" + print("Setting name to ", name) return self.poleval19_eval(glob=base, name=name, @@ -127,17 +130,13 @@ class ULMFiT: print("Seed: ", seed_name, seed) self.poleval19_eval(glob=base, name=name, num_lm_epochs=0, **kwargs) - def poleval19_eval(self, glob, name=None, num_lm_epochs=6, num_cls_epochs=8, bs=160, **kwargs): - if name is None: - name = f"ft{num_lm_epochs}_cl{num_cls_epochs}" - print("Setting name to ", name) - + def poleval19_eval(self, glob, name=None, num_lm_epochs=6, num_cls_epochs=8, bs=160, lr_sched="1cycle", **kwargs): return self.eval(glob=glob, name=name, num_lm_epochs=num_lm_epochs, num_cls_epochs=num_cls_epochs, bs=bs, - lr_sched="1cycle", + lr_sched=lr_sched, **kwargs) def eval(self, glob="data/mldoc/*-1/models/sp30k/lstm_nl4.m", dataset_template='${ds_name}', name=None, @@ -170,7 +169,11 @@ class ULMFiT: last_model_dir = params.model_dir.relative_to(data_dir.parent) if (params.model_dir/"cls_best.pth").exists(): print("Evaluating previously trained model") - d = params.validate_cls(label_smoothing_eps=label_smoothing_eps, use_cache=True) + d_tst = params.validate_cls(label_smoothing_eps=label_smoothing_eps, use_cache=True, mode="test") + d_val = params.validate_cls(label_smoothing_eps=label_smoothing_eps, use_cache=True, mode="valid") + d={} + d.update(d_val) + d.update(d_tst) elif train: print("Training") d = params.train_cls(num_lm_epochs=num_lm_epochs, label_smoothing_eps=label_smoothing_eps, **trn_params) @@ -195,7 +198,7 @@ class ULMFiT: df.to_csv(to_csv) if return_df: return last_model_dir, df - return last_model_dir + return str(last_model_dir) def remove_lm_saves(self): for lm_save in Path("data").glob("**/lm_*.pth"): diff --git a/ulmfit/train_clas.py b/ulmfit/train_clas.py index 581a924..7dfed85 100644 --- a/ulmfit/train_clas.py +++ b/ulmfit/train_clas.py @@ -171,7 +171,15 @@ class CLSHyperParams(LMHyperParams): if dump_preds: with open(dump_preds, 'w') as f: f.write('\n'.join([str(x) for x in preds])) - results = learn.validate(data_tst.valid_dl if mode == "test" else data_clas.valid_dl) + if mode == "test": + ds = data_tst.valid_dl + elif mode == "valid" or mode == "dev": + ds = data_clas.valid_dl + elif mode == "train": + ds = data_clas.valid_dl + else: + raise AttributeError(f"Unrecognized mode {mode}, valid options: test, valid, train optionally dev==valid") + results = learn.validate(ds) print(f"Model: {self.name}") print(f"Validation on: {mode}") labeled_results = self.output_metrics(results, mode=mode)