From 13f23756de1751fa8d01b3d1f357b0325f8fe47f Mon Sep 17 00:00:00 2001 From: troyyyyy Date: Sun, 29 Oct 2023 10:46:37 +0800 Subject: [PATCH] [MNT] add ground_KB as subclass of KBBase --- abl/reasoning/kb.py | 162 ++++++++++++++++++------------------ abl/reasoning/reasoner.py | 167 ++++++++++++++++++-------------------- 2 files changed, 160 insertions(+), 169 deletions(-) diff --git a/abl/reasoning/kb.py b/abl/reasoning/kb.py index 2d65357..b700aa1 100644 --- a/abl/reasoning/kb.py +++ b/abl/reasoning/kb.py @@ -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) diff --git a/abl/reasoning/reasoner.py b/abl/reasoning/reasoner.py index dbfb968..968b19a 100644 --- a/abl/reasoning/reasoner.py +++ b/abl/reasoning/reasoner.py @@ -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()