From c87d3772f5f2c9241bd716d6293688235eef2874 Mon Sep 17 00:00:00 2001 From: Gao Enhao Date: Wed, 15 Nov 2023 22:41:12 +0800 Subject: [PATCH] [FIX] fix bug in SemanticMetric --- abl/evaluation/semantics_metric.py | 6 ++++-- examples/mnist_add/mnist_add_example.ipynb | 4 ++-- 2 files changed, 6 insertions(+), 4 deletions(-) diff --git a/abl/evaluation/semantics_metric.py b/abl/evaluation/semantics_metric.py index 21ecabf..271ca1c 100644 --- a/abl/evaluation/semantics_metric.py +++ b/abl/evaluation/semantics_metric.py @@ -10,8 +10,10 @@ class SemanticsMetric(BaseMetric): self.kb = kb def process(self, data_samples: Sequence[dict]) -> None: - for data_sample in data_samples: - if self.kb.check_equal(data_sample, data_sample.Y[0]): + pred_psedudo_label_list = data_samples.pred_pseudo_label + y_list = data_samples.Y + for pred_psedudo_label, y in zip(pred_psedudo_label_list, y_list): + if self.kb._check_equal(self.kb.logic_forward(pred_psedudo_label), y): self.results.append(1) else: self.results.append(0) diff --git a/examples/mnist_add/mnist_add_example.ipynb b/examples/mnist_add/mnist_add_example.ipynb index 295a3d2..0927cb5 100644 --- a/examples/mnist_add/mnist_add_example.ipynb +++ b/examples/mnist_add/mnist_add_example.ipynb @@ -15,7 +15,7 @@ "\n", "from abl.learning import BasicNN, ABLModel\n", "from abl.bridge import SimpleBridge\n", - "from abl.evaluation import SymbolMetric\n", + "from abl.evaluation import SymbolMetric, SemanticsMetric\n", "from abl.utils import ABLLogger, print_log\n", "\n", "from examples.models.nn import LeNet5\n", @@ -135,7 +135,7 @@ "outputs": [], "source": [ "# Add metric\n", - "metric = [SymbolMetric(prefix=\"mnist_add\")]" + "metric = [SymbolMetric(prefix=\"mnist_add\"), SemanticsMetric(kb=kb, prefix=\"mnist_add\")]" ] }, {