Browse Source

Add zoopt option

pull/3/head
troyyyyy GitHub 3 years ago
parent
commit
3b8f77e4ff
No known key found for this signature in database GPG Key ID: 4AEE18F83AFDEB23
1 changed files with 69 additions and 34 deletions
  1. +69
    -34
      abducer/kb.py

+ 69
- 34
abducer/kb.py View File

@@ -20,6 +20,7 @@ from itertools import product, combinations

import pyswip


class KBBase(ABC):
def __init__(self):
pass
@@ -31,10 +32,20 @@ class KBBase(ABC):
@abstractmethod
def abduce_candidates(self):
pass


def address(self, address_num, pred_res, key):
new_candidates = []
address_idx_list = list(combinations(list(range(len(pred_res))), address_num))
for address_idx in address_idx_list:
candidates = self.address_by_idx(pred_res, key, address_idx)
new_candidates += candidates
return new_candidates
def abduction(self, pred_res, key, max_address_num, require_more_address):
candidates = []
for address_num in range(len(pred_res) + 1):
if address_num == 0:
if abs(self.logic_forward(pred_res) - key) <= 1e-3:
@@ -106,6 +117,11 @@ class ClsKB(KBBase):



def hamming_dist(self, A, B):
B = np.array(B)
A = np.expand_dims(A, axis = 0).repeat(axis=0, repeats=(len(B)))
return np.sum(A != B, axis = 1)

def abduce_from_GKB(self, pred_res, key, max_address_num, require_more_address):
if self.base == {} or len(pred_res) not in self.len_list:
return []
@@ -124,28 +140,18 @@ class ClsKB(KBBase):
candidates = [all_candidates[idx] for idx in idxs]
return candidates, min_address_num, address_num
def hamming_dist(self, A, B):
B = np.array(B)
A = np.expand_dims(A, axis = 0).repeat(axis=0, repeats=(len(B)))
return np.sum(A != B, axis = 1)

def address(self, address_num, pred_res, key):
new_candidates = []
all_address_candidate = self.all_address_candidate_dict[address_num]
address_idx_list = list(combinations(list(range(len(pred_res))), address_num))
for address_idx in address_idx_list:
for c in all_address_candidate:
candidate = pred_res.copy()
for i, idx in enumerate(address_idx):
candidate[idx] = c[i]
if self.logic_forward(candidate) == key:
if sum(pred_res[idx] != candidate[idx] for idx in range(len(pred_res))) == address_num:
new_candidates.append(candidate)
return new_candidates
def address_by_idx(self, pred_res, key, address_idx):
candidates = []
abduce_c = self.all_address_candidate_dict[len(address_idx)]
for c in abduce_c:
candidate = pred_res.copy()
for i, idx in enumerate(address_idx):
candidate[idx] = c[i]
if self.logic_forward(candidate) == key:
candidates.append(candidate)
return candidates

def _dict_len(self, dic):
@@ -208,19 +214,32 @@ class prolog_KB(KBBase):
def abduce_candidates(self, pred_res, key, max_address_num, require_more_address):
return self.abduction(pred_res, key, max_address_num, require_more_address)
def address(self, address_num, pred_res, key):
new_candidates = []
address_idx_list = list(combinations(list(range(len(pred_res))), address_num))
for address_idx in address_idx_list:
query_string = self.get_query_string(pred_res, address_idx)
abduce_c = [list(z.values()) for z in list(self.prolog.query(query_string % key))]
for c in abduce_c:
candidate = pred_res.copy()
for i, idx in enumerate(address_idx):
candidate[idx] = c[i]
if sum(pred_res[idx] != candidate[idx] for idx in range(len(pred_res))) == address_num:
new_candidates.append(candidate)
return new_candidates
def address_by_idx(self, pred_res, key, address_idx, verbose=True):
candidates = []
query_string = self.get_query_string(pred_res, address_idx)
abduce_c = [list(z.values()) for z in list(self.prolog.query(query_string % key))]
for c in abduce_c:
candidate = pred_res.copy()
for i, idx in enumerate(address_idx):
candidate[idx] = c[i]
candidates.append(candidate)
return candidates
def address_by_idx2(self, pred_res, key, address_idx):
candidates = []
query_string = self.get_query_string(pred_res, address_idx)
abduce_c = [list(z.values()) for z in list(self.prolog.query(query_string % key))]
# print('abduce_c:', abduce_c)
for c in abduce_c:
candidate = pred_res.copy()
for i, idx in enumerate(address_idx):
candidate[idx] = c[i]
candidates.append(candidate)
return candidates


class add_prolog_KB(prolog_KB):
@@ -240,6 +259,22 @@ class add_prolog_KB(prolog_KB):
return query_string

class hed_prolog_KB(prolog_KB):
def __init__(self, pseudo_label_list = list(range(10))):
super().__init__(pseudo_label_list)
self.prolog.assertz("addition(Z1, Z2, Res) :- pseudo_label(Z1), pseudo_label(Z2), Res is Z1+Z2")
def logic_forward(self, nums):
return list(self.prolog.query("addition(%s, %s, Res)." %(nums[0], nums[1])))[0]['Res']
def get_query_string(self, pred_res, address_idx):
query_string = "addition("
for idx, i in enumerate(pred_res):
tmp = 'Z' + str(idx) + ',' if idx in address_idx else str(i) + ','
query_string += tmp
query_string += "%s)."
return query_string


class RegKB(KBBase):
def __init__(self, GKB_flag = False, X = None, Y = None):


Loading…
Cancel
Save