| @@ -76,7 +76,6 @@ class AbducerBase(abc.ABC): | |||
| return solution | |||
| def address_by_idx(self, pred_res, key, address_idx): | |||
| # print(pred_res, address_idx) | |||
| return self.kb.address_by_idx(pred_res, key, address_idx) | |||
| def abduce(self, data, max_address=-1, require_more_address=0): | |||
| @@ -102,9 +101,8 @@ class AbducerBase(abc.ABC): | |||
| candidate = self._get_one_candidate(pred_res, pred_res_prob, candidates) | |||
| return candidate | |||
| def batch_abduce(self, data, max_address=-1, require_more_address=0): | |||
| Z1, Z2, Y = data | |||
| return [self.abduce((z, prob, y), max_address, require_more_address) for z, prob, y in zip(Z1, Z2, Y)] | |||
| def batch_abduce(self, Z, Y, max_address=-1, require_more_address=0): | |||
| return [self.abduce((z, prob, y), max_address, require_more_address) for z, prob, y in zip(Z['cls'], Z['prob'], Y)] | |||
| def __call__(self, Z, Y, max_address=-1, require_more_address=0): | |||
| return self.batch_abduce(Z, Y, max_address, require_more_address) | |||
| @@ -169,60 +167,60 @@ if __name__ == '__main__': | |||
| print('add_KB with GKB:') | |||
| kb = add_KB(GKB_flag=True) | |||
| abd = AbducerBase(kb, 'confidence') | |||
| res = abd.batch_abduce(([[1, 1]], prob1, [8]), max_address=2, require_more_address=0) | |||
| res = abd.batch_abduce({'cls':[[1, 1]], 'prob':prob1}, [8], max_address=2, require_more_address=0) | |||
| print(res) | |||
| res = abd.batch_abduce(([[1, 1]], prob2, [8]), max_address=2, require_more_address=0) | |||
| res = abd.batch_abduce({'cls':[[1, 1]], 'prob':prob2}, [8], max_address=2, require_more_address=0) | |||
| print(res) | |||
| res = abd.batch_abduce(([[1, 1]], prob1, [17]), max_address=2, require_more_address=0) | |||
| res = abd.batch_abduce({'cls':[[1, 1]], 'prob':prob1}, [17], max_address=2, require_more_address=0) | |||
| print(res) | |||
| res = abd.batch_abduce(([[1, 1]], prob1, [17]), max_address=1, require_more_address=0) | |||
| res = abd.batch_abduce({'cls':[[1, 1]], 'prob':prob1}, [17], max_address=1, require_more_address=0) | |||
| print(res) | |||
| res = abd.batch_abduce(([[1, 1]], prob1, [20]), max_address=2, require_more_address=0) | |||
| res = abd.batch_abduce({'cls':[[1, 1]], 'prob':prob1}, [20], max_address=2, require_more_address=0) | |||
| print(res) | |||
| print() | |||
| print('add_KB without GKB:') | |||
| kb = add_KB() | |||
| abd = AbducerBase(kb, 'confidence') | |||
| res = abd.batch_abduce(([[1, 1]], prob1, [8]), max_address=2, require_more_address=0) | |||
| res = abd.batch_abduce({'cls':[[1, 1]], 'prob':prob1}, [8], max_address=2, require_more_address=0) | |||
| print(res) | |||
| res = abd.batch_abduce(([[1, 1]], prob2, [8]), max_address=2, require_more_address=0) | |||
| res = abd.batch_abduce({'cls':[[1, 1]], 'prob':prob2}, [8], max_address=2, require_more_address=0) | |||
| print(res) | |||
| res = abd.batch_abduce(([[1, 1]], prob1, [17]), max_address=2, require_more_address=0) | |||
| res = abd.batch_abduce({'cls':[[1, 1]], 'prob':prob1}, [17], max_address=2, require_more_address=0) | |||
| print(res) | |||
| res = abd.batch_abduce(([[1, 1]], prob1, [17]), max_address=1, require_more_address=0) | |||
| res = abd.batch_abduce({'cls':[[1, 1]], 'prob':prob1}, [17], max_address=1, require_more_address=0) | |||
| print(res) | |||
| res = abd.batch_abduce(([[1, 1]], prob1, [20]), max_address=2, require_more_address=0) | |||
| res = abd.batch_abduce({'cls':[[1, 1]], 'prob':prob1}, [20], max_address=2, require_more_address=0) | |||
| print(res) | |||
| print() | |||
| print('prolog_KB with add.pl:') | |||
| kb = prolog_KB(pseudo_label_list=list(range(10)), pl_file='../examples/datasets/mnist_add/add.pl') | |||
| abd = AbducerBase(kb, 'confidence') | |||
| res = abd.batch_abduce(([[1, 1]], prob1, [8]), max_address=2, require_more_address=0) | |||
| res = abd.batch_abduce({'cls':[[1, 1]], 'prob':prob1}, [8], max_address=2, require_more_address=0) | |||
| print(res) | |||
| res = abd.batch_abduce(([[1, 1]], prob2, [8]), max_address=2, require_more_address=0) | |||
| res = abd.batch_abduce({'cls':[[1, 1]], 'prob':prob2}, [8], max_address=2, require_more_address=0) | |||
| print(res) | |||
| res = abd.batch_abduce(([[1, 1]], prob1, [17]), max_address=2, require_more_address=0) | |||
| res = abd.batch_abduce({'cls':[[1, 1]], 'prob':prob1}, [17], max_address=2, require_more_address=0) | |||
| print(res) | |||
| res = abd.batch_abduce(([[1, 1]], prob1, [17]), max_address=1, require_more_address=0) | |||
| res = abd.batch_abduce({'cls':[[1, 1]], 'prob':prob1}, [17], max_address=1, require_more_address=0) | |||
| print(res) | |||
| res = abd.batch_abduce(([[1, 1]], prob1, [20]), max_address=2, require_more_address=0) | |||
| res = abd.batch_abduce({'cls':[[1, 1]], 'prob':prob1}, [20], max_address=2, require_more_address=0) | |||
| print(res) | |||
| print() | |||
| print('prolog_KB with add.pl using zoopt:') | |||
| kb = prolog_KB(pseudo_label_list=list(range(10)), pl_file='../examples/datasets/mnist_add/add.pl') | |||
| abd = AbducerBase(kb, 'confidence', zoopt=True) | |||
| res = abd.batch_abduce(([[1, 1]], prob1, [8]), max_address=2, require_more_address=0) | |||
| res = abd.batch_abduce({'cls':[[1, 1]], 'prob':prob1}, [8], max_address=2, require_more_address=0) | |||
| print(res) | |||
| res = abd.batch_abduce(([[1, 1]], prob2, [8]), max_address=2, require_more_address=0) | |||
| res = abd.batch_abduce({'cls':[[1, 1]], 'prob':prob2}, [8], max_address=2, require_more_address=0) | |||
| print(res) | |||
| res = abd.batch_abduce(([[1, 1]], prob1, [17]), max_address=2, require_more_address=0) | |||
| res = abd.batch_abduce({'cls':[[1, 1]], 'prob':prob1}, [17], max_address=2, require_more_address=0) | |||
| print(res) | |||
| res = abd.batch_abduce(([[1, 1]], prob1, [17]), max_address=1, require_more_address=0) | |||
| res = abd.batch_abduce({'cls':[[1, 1]], 'prob':prob1}, [17], max_address=1, require_more_address=0) | |||
| print(res) | |||
| res = abd.batch_abduce(([[1, 1]], prob1, [20]), max_address=2, require_more_address=0) | |||
| res = abd.batch_abduce({'cls':[[1, 1]], 'prob':prob1}, [20], max_address=2, require_more_address=0) | |||
| print(res) | |||
| print() | |||
| @@ -232,70 +230,70 @@ if __name__ == '__main__': | |||
| kb = add_KB() | |||
| abd = AbducerBase(kb, 'confidence') | |||
| res = abd.batch_abduce(([[1, 1], [1, 2]], multiple_prob, [4, 8]), max_address=4, require_more_address=0) | |||
| res = abd.batch_abduce({'cls':[[1, 1], [1, 2]], 'prob':multiple_prob}, [4, 8], max_address=4, require_more_address=0) | |||
| print(res) | |||
| res = abd.batch_abduce(([[1, 1], [1, 2]], multiple_prob, [4, 8]), max_address=4, require_more_address=1) | |||
| res = abd.batch_abduce({'cls':[[1, 1], [1, 2]], 'prob':multiple_prob}, [4, 8], max_address=4, require_more_address=1) | |||
| print(res) | |||
| print() | |||
| print('HWF_KB with GKB, max_err=0.1') | |||
| kb = HWF_KB(len_list=[1, 3, 5], GKB_flag=True, max_err = 0.1) | |||
| abd = AbducerBase(kb, 'hamming') | |||
| res = abd.batch_abduce(([['5', '+', '2']], [None], [3]), max_address=2, require_more_address=0) | |||
| res = abd.batch_abduce({'cls':[['5', '+', '2']], 'prob':[None]}, [3], max_address=2, require_more_address=0) | |||
| print(res) | |||
| res = abd.batch_abduce(([['5', '+', '9']], [None], [65]), max_address=3, require_more_address=0) | |||
| res = abd.batch_abduce({'cls':[['5', '+', '9']], 'prob':[None]}, [65], max_address=3, require_more_address=0) | |||
| print(res) | |||
| res = abd.batch_abduce(([['5', '8', '8', '8', '8']], [None], [3.17]), max_address=5, require_more_address=3) | |||
| res = abd.batch_abduce({'cls':[['5', '8', '8', '8', '8']], 'prob':[None]}, [3.17], max_address=5, require_more_address=3) | |||
| print(res) | |||
| print() | |||
| print('HWF_KB without GKB, max_err=0.1') | |||
| kb = HWF_KB(len_list=[1, 3, 5], max_err = 0.1) | |||
| abd = AbducerBase(kb, 'hamming') | |||
| res = abd.batch_abduce(([['5', '+', '2']], [None], [3]), max_address=2, require_more_address=0) | |||
| res = abd.batch_abduce({'cls':[['5', '+', '2']], 'prob':[None]}, [3], max_address=2, require_more_address=0) | |||
| print(res) | |||
| res = abd.batch_abduce(([['5', '+', '9']], [None], [65]), max_address=3, require_more_address=0) | |||
| res = abd.batch_abduce({'cls':[['5', '+', '9']], 'prob':[None]}, [65], max_address=3, require_more_address=0) | |||
| print(res) | |||
| res = abd.batch_abduce(([['5', '8', '8', '8', '8']], [None], [3.17]), max_address=5, require_more_address=3) | |||
| res = abd.batch_abduce({'cls':[['5', '8', '8', '8', '8']], 'prob':[None]}, [3.17], max_address=5, require_more_address=3) | |||
| print(res) | |||
| print() | |||
| print('HWF_KB with GKB, max_err=1') | |||
| kb = HWF_KB(len_list=[1, 3, 5], GKB_flag=True, max_err = 1) | |||
| abd = AbducerBase(kb, 'hamming') | |||
| res = abd.batch_abduce(([['5', '+', '9']], [None], [65]), max_address=3, require_more_address=0) | |||
| res = abd.batch_abduce({'cls':[['5', '+', '9']], 'prob':[None]}, [65], max_address=3, require_more_address=0) | |||
| print(res) | |||
| res = abd.batch_abduce(([['5', '+', '2']], [None], [1.67]), max_address=3, require_more_address=0) | |||
| res = abd.batch_abduce({'cls':[['5', '+', '2']], 'prob':[None]}, [1.67], max_address=3, require_more_address=0) | |||
| print(res) | |||
| res = abd.batch_abduce(([['5', '8', '8', '8', '8']], [None], [3.17]), max_address=5, require_more_address=3) | |||
| res = abd.batch_abduce({'cls':[['5', '8', '8', '8', '8']], 'prob':[None]}, [3.17], max_address=5, require_more_address=3) | |||
| print(res) | |||
| print() | |||
| print('HWF_KB without GKB, max_err=1') | |||
| kb = HWF_KB(len_list=[1, 3, 5], max_err = 1) | |||
| abd = AbducerBase(kb, 'hamming') | |||
| res = abd.batch_abduce(([['5', '+', '9']], [None], [65]), max_address=3, require_more_address=0) | |||
| res = abd.batch_abduce({'cls':[['5', '+', '9']], 'prob':[None]}, [65], max_address=3, require_more_address=0) | |||
| print(res) | |||
| res = abd.batch_abduce(([['5', '+', '2']], [None], [1.67]), max_address=3, require_more_address=0) | |||
| res = abd.batch_abduce({'cls':[['5', '+', '2']], 'prob':[None]}, [1.67], max_address=3, require_more_address=0) | |||
| print(res) | |||
| res = abd.batch_abduce(([['5', '8', '8', '8', '8']], [None], [3.17]), max_address=5, require_more_address=3) | |||
| res = abd.batch_abduce({'cls':[['5', '8', '8', '8', '8']], 'prob':[None]}, [3.17], max_address=5, require_more_address=3) | |||
| print(res) | |||
| print() | |||
| print('HWF_KB with multiple inputs at once:') | |||
| kb = HWF_KB(len_list=[1, 3, 5], max_err = 0.1) | |||
| abd = AbducerBase(kb, 'hamming') | |||
| res = abd.batch_abduce(([['5', '+', '2'], ['5', '+', '9']], [None, None], [3, 64]), max_address=1, require_more_address=0) | |||
| res = abd.batch_abduce({'cls':[['5', '+', '2'], ['5', '+', '9']], 'prob':[None, None]}, [3, 64], max_address=1, require_more_address=0) | |||
| print(res) | |||
| res = abd.batch_abduce(([['5', '+', '2'], ['5', '+', '9']], [None, None], [3, 64]), max_address=3, require_more_address=0) | |||
| res = abd.batch_abduce({'cls':[['5', '+', '2'], ['5', '+', '9']], 'prob':[None, None]}, [3, 64], max_address=3, require_more_address=0) | |||
| print(res) | |||
| res = abd.batch_abduce(([['5', '+', '2'], ['5', '+', '9']], [None, None], [3, 65]), max_address=3, require_more_address=0) | |||
| res = abd.batch_abduce({'cls':[['5', '+', '2'], ['5', '+', '9']], 'prob':[None, None]}, [3, 65], max_address=3, require_more_address=0) | |||
| print(res) | |||
| print() | |||
| print('max_address is float') | |||
| res = abd.batch_abduce(([['5', '+', '2'], ['5', '+', '9']], [None, None], [3, 64]), max_address=0.5, require_more_address=0) | |||
| res = abd.batch_abduce({'cls':[['5', '+', '2'], ['5', '+', '9']], 'prob':[None, None]}, [3, 64], max_address=0.5, require_more_address=0) | |||
| print(res) | |||
| res = abd.batch_abduce(([['5', '+', '2'], ['5', '+', '9']], [None, None], [3, 64]), max_address=0.9, require_more_address=0) | |||
| res = abd.batch_abduce({'cls':[['5', '+', '2'], ['5', '+', '9']], 'prob':[None, None]}, [3, 64], max_address=0.9, require_more_address=0) | |||
| print(res) | |||
| print() | |||