Browse Source

[ENH] change parameters passing in reasoning

pull/1/head
troyyyyy 2 years ago
parent
commit
869b1ba695
1 changed files with 38 additions and 52 deletions
  1. +38
    -52
      abl/reasoning/kb.py

+ 38
- 52
abl/reasoning/kb.py View File

@@ -64,18 +64,18 @@ class KBBase(ABC):

Parameters
----------
pred_pseudo_label : List[Any]
pseudo_label : List[Any]
Pseudo label.
"""
pass

def abduce_candidates(self, pred_pseudo_label, y, max_revision_num, require_more_revision=0):
def abduce_candidates(self, pseudo_label, y, max_revision_num, require_more_revision=0):
"""
Perform abductive reasoning to get a candidate consistent with the knowledge base.

Parameters
----------
pred_pseudo_label : List[Any]
pseudo_label : List[Any]
Predicted pseudo label.
y : any
Ground truth for the logical result.
@@ -91,12 +91,7 @@ class KBBase(ABC):
A list of candidates, i.e. revised pseudo labels that are consistent with the
knowledge base.
"""
# if self.use_cache:
# return self._abduce_by_search_cache(to_hashable(pred_pseudo_label),
# to_hashable(y),
# max_revision_num, require_more_revision)
# else:
return self._abduce_by_search(pred_pseudo_label, y, max_revision_num, require_more_revision)
return self._abduce_by_search(pseudo_label, y, max_revision_num, require_more_revision)

def _check_equal(self, logic_result, y):
"""
@@ -111,13 +106,13 @@ class KBBase(ABC):
else:
return logic_result == y

def revise_at_idx(self, pred_pseudo_label, y, revision_idx):
def revise_at_idx(self, pseudo_label, y, revision_idx):
"""
Revise the predicted pseudo label at specified index positions.

Parameters
----------
pred_pseudo_label : List[Any]
pseudo_label : List[Any]
Predicted pseudo label.
y : Any
Ground truth for the logical result.
@@ -127,28 +122,28 @@ class KBBase(ABC):
candidates = []
abduce_c = product(self.pseudo_label_list, repeat=len(revision_idx))
for c in abduce_c:
candidate = pred_pseudo_label.copy()
candidate = pseudo_label.copy()
for i, idx in enumerate(revision_idx):
candidate[idx] = c[i]
if self._check_equal(self.logic_forward(candidate), y):
candidates.append(candidate)
return candidates

def _revision(self, revision_num, pred_pseudo_label, y):
def _revision(self, revision_num, pseudo_label, y):
"""
For a specified number of pseudo label to revise, iterate through all possible
indices to find any candidates that are consistent with the knowledge base.
"""
new_candidates = []
revision_idx_list = combinations(range(len(pred_pseudo_label)), revision_num)
revision_idx_list = combinations(range(len(pseudo_label)), revision_num)

for revision_idx in revision_idx_list:
candidates = self.revise_at_idx(pred_pseudo_label, y, revision_idx)
candidates = self.revise_at_idx(pseudo_label, y, revision_idx)
new_candidates.extend(candidates)
return new_candidates

@abl_cache()
def _abduce_by_search(self, pred_pseudo_label, y, max_revision_num, require_more_revision):
def _abduce_by_search(self, pseudo_label, y, max_revision_num, require_more_revision):
"""
Perform abductive reasoning by exhastive search. Specifically, begin with 0 and
continuously increase the number of pseudo labels to revise, until candidates
@@ -156,7 +151,7 @@ class KBBase(ABC):

Parameters
----------
pred_pseudo_label : List[Any]
pseudo_label : List[Any]
Predicted pseudo label.
y : Any
Ground truth for the logical result.
@@ -174,11 +169,11 @@ class KBBase(ABC):
knowledge base.
"""
candidates = []
for revision_num in range(len(pred_pseudo_label) + 1):
if revision_num == 0 and self._check_equal(self.logic_forward(pred_pseudo_label), y):
candidates.append(pred_pseudo_label)
for revision_num in range(len(pseudo_label) + 1):
if revision_num == 0 and self._check_equal(self.logic_forward(pseudo_label), y):
candidates.append(pseudo_label)
elif revision_num > 0:
candidates.extend(self._revision(revision_num, pred_pseudo_label, y))
candidates.extend(self._revision(revision_num, pseudo_label, y))
if len(candidates) > 0:
min_revision_num = revision_num
break
@@ -190,18 +185,9 @@ class KBBase(ABC):
):
if revision_num > max_revision_num:
return candidates
candidates.extend(self._revision(revision_num, pred_pseudo_label, y))
candidates.extend(self._revision(revision_num, pseudo_label, y))
return candidates

# @abl_cache(max_size=4096)
# def _abduce_by_search_cache(self, pred_pseudo_label, y, max_revision_num, require_more_revision):
# """
# `_abduce_by_search` with cache.
# """
# pred_pseudo_label = restore_from_hashable(pred_pseudo_label)
# y = restore_from_hashable(y)
# return self._abduce_by_search(pred_pseudo_label, y, max_revision_num, require_more_revision)

def __repr__(self):
return (
f"{self.__class__.__name__} is a KB with "
@@ -277,7 +263,7 @@ class GroundKB(KBBase):
X, Y = zip(*sorted(zip(X, Y), key=lambda pair: pair[1]))
return X, Y

def abduce_candidates(self, pred_pseudo_label, y, max_revision_num, require_more_revision=0):
def abduce_candidates(self, pseudo_label, y, max_revision_num, require_more_revision=0):
"""
Perform abductive reasoning by directly retrieving consistent candidates from
the prebuilt GKB. In this way, the time-consuming exhaustive search can be
@@ -285,28 +271,28 @@ class GroundKB(KBBase):
This is an overridden function. For more information about the parameters and
returns, refer to the function of the same name in class `KBBase`.
"""
if self.GKB == {} or len(pred_pseudo_label) not in self.GKB_len_list:
if self.GKB == {} or len(pseudo_label) not in self.GKB_len_list:
return []

all_candidates = self._find_candidate_GKB(pred_pseudo_label, y)
all_candidates = self._find_candidate_GKB(pseudo_label, y)
if len(all_candidates) == 0:
return []

cost_list = hamming_dist(pred_pseudo_label, all_candidates)
cost_list = hamming_dist(pseudo_label, all_candidates)
min_revision_num = np.min(cost_list)
revision_num = min(max_revision_num, min_revision_num + require_more_revision)
idxs = np.where(cost_list <= revision_num)[0]
candidates = [all_candidates[idx] for idx in idxs]
return candidates

def _find_candidate_GKB(self, pred_pseudo_label, y):
def _find_candidate_GKB(self, pseudo_label, y):
"""
Retrieve consistent candidates from the prebuilt GKB. For numerical logical results,
return all candidates whose logical results fall within the
[y - max_err, y + max_err] range.
"""
if isinstance(y, (int, float)):
potential_candidates = self.GKB[len(pred_pseudo_label)]
potential_candidates = self.GKB[len(pseudo_label)]
key_list = list(potential_candidates.keys())

low_key = bisect.bisect_left(key_list, y - self.max_err)
@@ -320,7 +306,7 @@ class GroundKB(KBBase):
return all_candidates

else:
return self.GKB[len(pred_pseudo_label)][y]
return self.GKB[len(pseudo_label)][y]

def __repr__(self):
return (
@@ -380,20 +366,20 @@ class PrologKB(KBBase):
return False
return result

def _revision_pred_pseudo_label(self, pred_pseudo_label, revision_idx):
def _revision_pseudo_label(self, pseudo_label, revision_idx):
import re

revision_pred_pseudo_label = pred_pseudo_label.copy()
revision_pred_pseudo_label = flatten(revision_pred_pseudo_label)
revision_pseudo_label = pseudo_label.copy()
revision_pseudo_label = flatten(revision_pseudo_label)

for idx in revision_idx:
revision_pred_pseudo_label[idx] = "P" + str(idx)
revision_pred_pseudo_label = reform_list(revision_pred_pseudo_label, pred_pseudo_label)
revision_pseudo_label[idx] = "P" + str(idx)
revision_pseudo_label = reform_list(revision_pseudo_label, pseudo_label)

regex = r"'P\d+'"
return re.sub(regex, lambda x: x.group().replace("'", ""), str(revision_pred_pseudo_label))
return re.sub(regex, lambda x: x.group().replace("'", ""), str(revision_pseudo_label))

def get_query_string(self, pred_pseudo_label, y, revision_idx):
def get_query_string(self, pseudo_label, y, revision_idx):
"""
Consult prolog with `logic_forward([kept_labels, Revise_labels], Res).`, and set
the returned `Revise_labels` together with the kept labels as the candidates. This is
@@ -401,27 +387,27 @@ class PrologKB(KBBase):
Prolog file.
"""
query_string = "logic_forward("
query_string += self._revision_pred_pseudo_label(pred_pseudo_label, revision_idx)
query_string += self._revision_pseudo_label(pseudo_label, revision_idx)
key_is_none_flag = y is None or (type(y) == list and y[0] is None)
query_string += ",%s)." % y if not key_is_none_flag else ")."
return query_string

def revise_at_idx(self, pred_pseudo_label, y, revision_idx):
def revise_at_idx(self, pseudo_label, y, revision_idx):
"""
Revise the predicted pseudo label at specified index positions by querying Prolog.
This is an overridden function. For more information about the parameters, refer to
the function of the same name in class `KBBase`.
"""
candidates = []
query_string = self.get_query_string(pred_pseudo_label, y, revision_idx)
save_pred_pseudo_label = pred_pseudo_label
pred_pseudo_label = flatten(pred_pseudo_label)
query_string = self.get_query_string(pseudo_label, y, revision_idx)
save_pseudo_label = pseudo_label
pseudo_label = flatten(pseudo_label)
abduce_c = [list(z.values()) for z in self.prolog.query(query_string)]
for c in abduce_c:
candidate = pred_pseudo_label.copy()
candidate = pseudo_label.copy()
for i, idx in enumerate(revision_idx):
candidate[idx] = c[i]
candidate = reform_list(candidate, save_pred_pseudo_label)
candidate = reform_list(candidate, save_pseudo_label)
candidates.append(candidate)
return candidates



Loading…
Cancel
Save