Browse Source

[MNT] add ground_KB as subclass of KBBase

pull/3/head
troyyyyy 2 years ago
parent
commit
13f23756de
2 changed files with 160 additions and 169 deletions
  1. +82
    -80
      abl/reasoning/kb.py
  2. +78
    -89
      abl/reasoning/reasoner.py

+ 82
- 80
abl/reasoning/kb.py View File

@@ -13,23 +13,92 @@ from functools import lru_cache
import pyswip

class KBBase(ABC):
def __init__(self, pseudo_label_list, prebuild_GKB=False, GKB_len_list=None, max_err=0, use_cache=True):
def __init__(self, pseudo_label_list, max_err=0, use_cache=True):
# TODO:添加一下类型检查,比如
# if not isinstance(X, (np.ndarray, spmatrix)):
# raise TypeError("X should be numpy array or sparse matrix")

self.pseudo_label_list = pseudo_label_list
self.prebuild_GKB = prebuild_GKB
self.GKB_len_list = GKB_len_list
self.max_err = max_err
self.use_cache = use_cache
self.use_cache = use_cache

@abstractmethod
def logic_forward(self, pseudo_labels):
pass

if prebuild_GKB:
self.base = {}
X, Y = self._get_GKB()
for x, y in zip(X, Y):
self.base.setdefault(len(x), defaultdict(list))[y].append(x)
def abduce_candidates(self, pred_res, y, max_revision_num, require_more_revision=0):
if not self.use_cache:
return self._abduce_by_search(pred_res, y, max_revision_num, require_more_revision)
else:
return self._abduce_by_search_cache(to_hashable(pred_res), to_hashable(y), max_revision_num, require_more_revision)
def revise_by_idx(self, pred_res, y, revision_idx):
candidates = []
abduce_c = product(self.pseudo_label_list, repeat=len(revision_idx))
for c in abduce_c:
candidate = pred_res.copy()
for i, idx in enumerate(revision_idx):
candidate[idx] = c[i]
if check_equal(self.logic_forward(candidate), y, self.max_err):
candidates.append(candidate)
return candidates

def _revision(self, revision_num, pred_res, y):
new_candidates = []
revision_idx_list = combinations(range(len(pred_res)), revision_num)

for revision_idx in revision_idx_list:
candidates = self.revise_by_idx(pred_res, y, revision_idx)
new_candidates.extend(candidates)
return new_candidates

def _abduce_by_search(self, pred_res, y, max_revision_num, require_more_revision):
candidates = []
for revision_num in range(len(pred_res) + 1):
if revision_num == 0 and check_equal(self.logic_forward(pred_res), y, self.max_err):
candidates.append(pred_res)
elif revision_num > 0:
candidates.extend(self._revision(revision_num, pred_res, y))
if len(candidates) > 0:
min_revision_num = revision_num
break
if revision_num >= max_revision_num:
return []

for revision_num in range(min_revision_num + 1, min_revision_num + require_more_revision + 1):
if revision_num > max_revision_num:
return candidates
candidates.extend(self._revision(revision_num, pred_res, y))
return candidates
@lru_cache(maxsize=None)
def _abduce_by_search_cache(self, pred_res, y, max_revision_num, require_more_revision):
pred_res = hashable_to_list(pred_res)
y = hashable_to_list(y)
return self._abduce_by_search(pred_res, y, max_revision_num, require_more_revision)
def _dict_len(self, dic):
if not self.GKB_flag:
return 0
else:
return sum(len(c) for c in dic.values())

def __len__(self):
if not self.GKB_flag:
return 0
else:
return sum(self._dict_len(v) for v in self.base.values())
class ground_KB(KBBase):
def __init__(self, pseudo_label_list, GKB_len_list=None, max_err=0):
super().__init__(pseudo_label_list, max_err)
self.GKB_len_list = GKB_len_list
self.base = {}
X, Y = self._get_GKB()
for x, y in zip(X, Y):
self.base.setdefault(len(x), defaultdict(list))[y].append(x)
# For parallel version of _get_GKB
def _get_XY_list(self, args):
pre_x, post_x_it = args[0], args[1]
@@ -60,19 +129,9 @@ class KBBase(ABC):
if Y and isinstance(Y[0], (int, float)):
X, Y = zip(*sorted(zip(X, Y), key=lambda pair: pair[1]))
return X, Y

@abstractmethod
def logic_forward(self, pseudo_labels):
pass

def abduce_candidates(self, pred_res, y, max_revision_num, require_more_revision=0):
if self.prebuild_GKB:
return self._abduce_by_GKB(pred_res, y, max_revision_num, require_more_revision)
else:
if not self.use_cache:
return self._abduce_by_search(pred_res, y, max_revision_num, require_more_revision)
else:
return self._abduce_by_search_cache(to_hashable(pred_res), to_hashable(y), max_revision_num, require_more_revision)
return self._abduce_by_GKB(pred_res, y, max_revision_num, require_more_revision)
def _find_candidate_GKB(self, pred_res, y):
if self.max_err == 0:
@@ -113,67 +172,10 @@ class KBBase(ABC):
candidates = [all_candidates[idx] for idx in idxs]
return candidates

def revise_by_idx(self, pred_res, y, revision_idx):
candidates = []
abduce_c = product(self.pseudo_label_list, repeat=len(revision_idx))
for c in abduce_c:
candidate = pred_res.copy()
for i, idx in enumerate(revision_idx):
candidate[idx] = c[i]
if check_equal(self.logic_forward(candidate), y, self.max_err):
candidates.append(candidate)
return candidates

def _revision(self, revision_num, pred_res, y):
new_candidates = []
revision_idx_list = combinations(range(len(pred_res)), revision_num)

for revision_idx in revision_idx_list:
candidates = self.revise_by_idx(pred_res, y, revision_idx)
new_candidates.extend(candidates)
return new_candidates

def _abduce_by_search(self, pred_res, y, max_revision_num, require_more_revision):
candidates = []
for revision_num in range(len(pred_res) + 1):
if revision_num == 0 and check_equal(self.logic_forward(pred_res), y, self.max_err):
candidates.append(pred_res)
elif revision_num > 0:
candidates.extend(self._revision(revision_num, pred_res, y))
if len(candidates) > 0:
min_revision_num = revision_num
break
if revision_num >= max_revision_num:
return []

for revision_num in range(min_revision_num + 1, min_revision_num + require_more_revision + 1):
if revision_num > max_revision_num:
return candidates
candidates.extend(self._revision(revision_num, pred_res, y))
return candidates
@lru_cache(maxsize=None)
def _abduce_by_search_cache(self, pred_res, y, max_revision_num, require_more_revision):
pred_res = hashable_to_list(pred_res)
y = hashable_to_list(y)
return self._abduce_by_search(pred_res, y, max_revision_num, require_more_revision)
def _dict_len(self, dic):
if not self.GKB_flag:
return 0
else:
return sum(len(c) for c in dic.values())

def __len__(self):
if not self.GKB_flag:
return 0
else:
return sum(self._dict_len(v) for v in self.base.values())


class prolog_KB(KBBase):
def __init__(self, pseudo_label_list, pl_file):
super().__init__(pseudo_label_list)
def __init__(self, pseudo_label_list, pl_file, max_err=0):
super().__init__(pseudo_label_list, max_err)
self.prolog = pyswip.Prolog()
self.prolog.consult(pl_file)



+ 78
- 89
abl/reasoning/reasoner.py View File

@@ -12,7 +12,7 @@ from ..utils.utils import (
class ReasonerBase:
def __init__(self, kb, dist_func="hamming", mapping=None, use_zoopt=False):
"""
Root class for all reasoner in the ABL system.
Base class for all reasoner in the ABL system.

Parameters
----------
@@ -47,7 +47,7 @@ class ReasonerBase:

def _get_cost_list(self, pred_pseudo_label, pred_prob, candidates):
"""
Get the list of costs between pseudo label and each candidate.
Get the list of costs between each pseudo label and candidate.

Parameters
----------
@@ -284,33 +284,26 @@ class ReasonerBase:


if __name__ == "__main__":
from kb import KBBase, prolog_KB
from kb import KBBase, ground_KB, prolog_KB

prob1 = [
[
[0, 0.99, 0.01, 0, 0, 0, 0, 0, 0, 0],
[0.1, 0.1, 0.1, 0.1, 0.1, 0.1, 0.1, 0.1, 0.1, 0.1],
]
]
prob2 = [
[
[0, 0, 0.01, 0, 0, 0, 0, 0.99, 0, 0],
[0.1, 0.1, 0.1, 0.1, 0.1, 0.1, 0.1, 0.1, 0.1, 0.1],
]
]
prob1 = [[[0, 0.99, 0.01, 0, 0, 0, 0, 0, 0, 0],
[0.1, 0.1, 0.1, 0.1, 0.1, 0.1, 0.1, 0.1, 0.1, 0.1]]]
prob2 = [[[0, 0, 0.01, 0, 0, 0, 0, 0.99, 0, 0],
[0.1, 0.1, 0.1, 0.1, 0.1, 0.1, 0.1, 0.1, 0.1, 0.1]]]

class add_KB(KBBase):
def __init__(
self,
pseudo_label_list=list(range(10)),
prebuild_GKB=False,
GKB_len_list=[2],
max_err=0,
use_cache=True,
):
super().__init__(
pseudo_label_list, prebuild_GKB, GKB_len_list, max_err, use_cache
)
def __init__(self, pseudo_label_list=list(range(10)),
use_cache=True):
super().__init__(pseudo_label_list, use_cache=use_cache)

def logic_forward(self, nums):
return sum(nums)
class add_ground_KB(ground_KB):
def __init__(self, pseudo_label_list=list(range(10)),
GKB_len_list=[2]):
super().__init__(pseudo_label_list, GKB_len_list)

def logic_forward(self, nums):
return sum(nums)
@@ -329,7 +322,7 @@ if __name__ == "__main__":
print()

print("add_KB with GKB:")
kb = add_KB(prebuild_GKB=True)
kb = add_ground_KB()
reasoner = ReasonerBase(kb, "confidence")
test_add(reasoner)

@@ -338,16 +331,14 @@ if __name__ == "__main__":
reasoner = ReasonerBase(kb, "confidence")
test_add(reasoner)

print("add_KB without GKB:, no cache")
print("add_KB without GKB, no cache")
kb = add_KB(use_cache=False)
reasoner = ReasonerBase(kb, "confidence")
test_add(reasoner)

print("prolog_KB with add.pl:")
kb = prolog_KB(
pseudo_label_list=list(range(10)),
pl_file="examples/mnist_add/datasets/add.pl",
)
kb = prolog_KB(pseudo_label_list=list(range(10)),
pl_file="examples/mnist_add/datasets/add.pl")
reasoner = ReasonerBase(kb, "confidence")
test_add(reasoner)

@@ -360,16 +351,14 @@ if __name__ == "__main__":
test_add(reasoner)

print("add_KB with multiple inputs at once:")
multiple_prob = [
[
[0, 0.99, 0.01, 0, 0, 0, 0, 0, 0, 0],
[0.1, 0.1, 0.1, 0.1, 0.1, 0.1, 0.1, 0.1, 0.1, 0.1],
],
[
[0, 0, 0.01, 0, 0, 0, 0, 0.99, 0, 0],
[0.1, 0.1, 0.1, 0.1, 0.1, 0.1, 0.1, 0.1, 0.1, 0.1],
],
]
multiple_prob = [[
[0, 0.99, 0.01, 0, 0, 0, 0, 0, 0, 0],
[0.1, 0.1, 0.1, 0.1, 0.1, 0.1, 0.1, 0.1, 0.1, 0.1],
],
[
[0, 0, 0.01, 0, 0, 0, 0, 0.99, 0, 0],
[0.1, 0.1, 0.1, 0.1, 0.1, 0.1, 0.1, 0.1, 0.1, 0.1],
]]

kb = add_KB()
reasoner = ReasonerBase(kb, "confidence")
@@ -394,45 +383,45 @@ if __name__ == "__main__":
class HWF_KB(KBBase):
def __init__(
self,
pseudo_label_list=[
"1",
"2",
"3",
"4",
"5",
"6",
"7",
"8",
"9",
"+",
"-",
"times",
"div",
],
prebuild_GKB=False,
pseudo_label_list=["1", "2", "3", "4", "5", "6", "7", "8", "9",
"+", "-", "times", "div"],
max_err=1e-3,
):
super().__init__(pseudo_label_list, max_err)

def _valid_candidate(self, formula):
if len(formula) % 2 == 0:
return False
for i in range(len(formula)):
if i % 2 == 0 and formula[i] not in ["1", "2", "3", "4", "5", "6", "7", "8", "9"]:
return False
if i % 2 != 0 and formula[i] not in ["+", "-", "times", "div"]:
return False
return True

def logic_forward(self, formula):
if not self._valid_candidate(formula):
return np.inf
mapping = {str(i): str(i) for i in range(1, 10)}
mapping.update({"+": "+", "-": "-", "times": "*", "div": "/"})
formula = [mapping[f] for f in formula]
return eval("".join(formula))
class HWF_ground_KB(ground_KB):
def __init__(
self,
pseudo_label_list=["1", "2", "3", "4", "5", "6", "7", "8", "9",
"+", "-", "times", "div"],
GKB_len_list=[1, 3, 5, 7],
max_err=1e-3,
use_cache=True,
):
super().__init__(
pseudo_label_list, prebuild_GKB, GKB_len_list, max_err, use_cache
)
super().__init__(pseudo_label_list, GKB_len_list, max_err)

def _valid_candidate(self, formula):
if len(formula) % 2 == 0:
return False
for i in range(len(formula)):
if i % 2 == 0 and formula[i] not in [
"1",
"2",
"3",
"4",
"5",
"6",
"7",
"8",
"9",
]:
if i % 2 == 0 and formula[i] not in ["1", "2", "3", "4", "5", "6", "7", "8", "9"]:
return False
if i % 2 != 0 and formula[i] not in ["+", "-", "times", "div"]:
return False
@@ -473,12 +462,12 @@ if __name__ == "__main__":
print(res)
print()
def test_hwf_multiple(reasoner):
def test_hwf_multiple(reasoner, max_revisions):
res = reasoner.batch_abduce(
[None, None],
[["5", "+", "2"], ["5", "+", "9"]],
[3, 64],
max_revision=1,
max_revision=max_revisions[0],
require_more_revision=0,
)
print(res)
@@ -486,7 +475,7 @@ if __name__ == "__main__":
[None, None],
[["5", "+", "2"], ["5", "+", "9"]],
[3, 64],
max_revision=3,
max_revision=max_revisions[1],
require_more_revision=0,
)
print(res)
@@ -494,39 +483,39 @@ if __name__ == "__main__":
[None, None],
[["5", "+", "2"], ["5", "+", "9"]],
[3, 65],
max_revision=3,
max_revision=max_revisions[2],
require_more_revision=0,
)
print(res)
print()

print("HWF_KB with GKB, max_err=0.1")
kb = HWF_KB(prebuild_GKB=True, GKB_len_list=[1, 3, 5], max_err=0.1)
kb = HWF_ground_KB(GKB_len_list=[1, 3, 5], max_err=0.1)
reasoner = ReasonerBase(kb, "hamming")
test_hwf(reasoner)

print("HWF_KB without GKB, max_err=0.1")
kb = HWF_KB(GKB_len_list=[1, 3, 5], max_err=0.1)
kb = HWF_KB(max_err=0.1)
reasoner = ReasonerBase(kb, "hamming")
test_hwf(reasoner)

print("HWF_KB with GKB, max_err=1")
kb = HWF_KB(GKB_len_list=[1, 3, 5], prebuild_GKB=True, max_err=1)
kb = HWF_ground_KB(GKB_len_list=[1, 3, 5], max_err=1)
reasoner = ReasonerBase(kb, "hamming")
test_hwf(reasoner)

print("HWF_KB without GKB, max_err=1")
kb = HWF_KB(GKB_len_list=[1, 3, 5], max_err=1)
kb = HWF_KB(max_err=1)
reasoner = ReasonerBase(kb, "hamming")
test_hwf(reasoner)

print("HWF_KB with multiple inputs at once:")
kb = HWF_KB(GKB_len_list=[1, 3, 5], max_err=0.1)
kb = HWF_KB(max_err=0.1)
reasoner = ReasonerBase(kb, "hamming")
test_hwf_multiple(reasoner)
test_hwf_multiple(reasoner, max_revisions=[1,3,3])
print("max_revision is float")
test_hwf_multiple(reasoner)
test_hwf_multiple(reasoner, max_revisions=[0.5,0.9,0.9])

class HED_prolog_KB(prolog_KB):
def __init__(self, pseudo_label_list, pl_file):
@@ -548,7 +537,7 @@ if __name__ == "__main__":

class HED_Reasoner(ReasonerBase):
def __init__(self, kb, dist_func="hamming"):
super().__init__(kb, dist_func, zoopt=True)
super().__init__(kb, dist_func, use_zoopt=True)

def _revise_by_idxs(self, pred_res, y, all_revision_flag, idxs):
pred = []
@@ -562,7 +551,7 @@ if __name__ == "__main__":
candidate = self.revise_by_idx(pred, k, revision_idx)
return candidate

def zoopt_revision_score(self, pred_res, pred_prob, y, sol):
def zoopt_revision_score(self, symbol_num, pred_res, pred_prob, y, sol):
all_revision_flag = reform_idx(sol.get_x(), pred_res)
lefted_idxs = [i for i in range(len(pred_res))]
candidate_size = []
@@ -631,15 +620,15 @@ if __name__ == "__main__":

print("HED_Reasoner abduce")
res = reasoner.abduce(
(consist_exs, [[[None]]] * len(consist_exs), [None] * len(consist_exs))
[[[None]]] * len(consist_exs), consist_exs, [None] * len(consist_exs)
)
print(res)
res = reasoner.abduce(
(inconsist_exs1, [[[None]]] * len(inconsist_exs1), [None] * len(inconsist_exs1))
[[[None]]] * len(inconsist_exs1), inconsist_exs1, [None] * len(inconsist_exs1)
)
print(res)
res = reasoner.abduce(
(inconsist_exs2, [[[None]]] * len(inconsist_exs2), [None] * len(inconsist_exs2))
[[[None]]] * len(inconsist_exs2), inconsist_exs2, [None] * len(inconsist_exs2)
)
print(res)
print()


Loading…
Cancel
Save