Add ability to check different model save

This commit is contained in:
Piotr Czapla
2019-05-13 17:47:57 +02:00
parent b55715e4f6
commit 29b600b8b5
2 changed files with 4 additions and 4 deletions
+3 -3
View File
@@ -141,7 +141,7 @@ class ULMFiT:
def eval(self, glob="data/mldoc/*-1/models/sp30k/lstm_nl4.m", dataset_template='${ds_name}', name=None,
num_lm_epochs=0, train=True, to_csv=None, return_df=False, label_smoothing_eps=0.0,
lmseed=None, ftseed=None, clsweightseed=None, clstrainseed=None,
lmseed=None, ftseed=None, clsweightseed=None, clstrainseed=None, save_name="cls_best",
skip_on_error=True, **trn_params):
results = []
model_args = {}
@@ -169,8 +169,8 @@ 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_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_tst = params.validate_cls(save_name=save_name, label_smoothing_eps=label_smoothing_eps, use_cache=True, mode="test")
d_val = params.validate_cls(save_name=save_name, label_smoothing_eps=label_smoothing_eps, use_cache=True, mode="valid")
d={}
d.update(d_val)
d.update(d_tst)
+1 -1
View File
@@ -155,7 +155,7 @@ class CLSHyperParams(LMHyperParams):
def validate_cls(self, save_name='cls_best', bs=40, data_tst=None, learn=None,
dump_preds=None, mode="test", label_smoothing_eps=None, use_cache=False):
cache_file = (self.model_dir / f'results_{mode}.json')
cache_file = (self.model_dir / f'results_{mode+("" if save_name == "cls_best" else save_name)}.json')
if use_cache and cache_file.exists():
with cache_file.open("r") as fp:
return json.load(fp)