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\")]" ] }, {