Browse Source

[FIX] fix bug in SemanticMetric

pull/4/head
Gao Enhao 2 years ago
parent
commit
c87d3772f5
2 changed files with 6 additions and 4 deletions
  1. +4
    -2
      abl/evaluation/semantics_metric.py
  2. +2
    -2
      examples/mnist_add/mnist_add_example.ipynb

+ 4
- 2
abl/evaluation/semantics_metric.py View File

@@ -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)


+ 2
- 2
examples/mnist_add/mnist_add_example.ipynb View File

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


Loading…
Cancel
Save