From 50b73510de5a88359de045b453b34b687c583fbf Mon Sep 17 00:00:00 2001 From: Marcin Date: Sun, 12 May 2019 21:43:54 +0200 Subject: [PATCH] Add workaround for metrics init --- ulmfit/train_clas.py | 7 ++++--- 1 file changed, 4 insertions(+), 3 deletions(-) diff --git a/ulmfit/train_clas.py b/ulmfit/train_clas.py index c20a1bb..044f856 100644 --- a/ulmfit/train_clas.py +++ b/ulmfit/train_clas.py @@ -79,14 +79,15 @@ class CLSHyperParams(LMHyperParams): if num_cls_epochs > 5: learn.fit_one_cycle(num_cls_epochs-4, slice(1e-2 / (2.6 ** 4), 1e-2), moms=(0.8, 0.7), wd=1e-7) - def get_metrics(self): + def get_metrics(self, init=False): f1_score = FBeta(beta=1.0) precision = Precision() recall = Recall() metrics = [f1_score, precision, recall] # TODO: fix this in fast.ai - for metric in metrics: metric.on_train_begin() + if init: + for metric in metrics: metric.on_train_begin() metrics.append(accuracy) return metrics @@ -151,7 +152,7 @@ class CLSHyperParams(LMHyperParams): if data_tst is None: data_clas , _, data_tst = self.load_cls_data(bs) if learn is None: - learn = self.create_cls_learner(data_tst, drop_mult=0.3, metrics=self.get_metrics()) + learn = self.create_cls_learner(data_tst, drop_mult=0.3, metrics=self.get_metrics(True)) learn.unfreeze() learn.load(save_name) probs, targets = learn.get_preds(ordered=True)