Browse Source

Update batch_abduce

pull/3/head
troyyyyy 3 years ago
parent
commit
5884584dfe
1 changed files with 41 additions and 43 deletions
  1. +41
    -43
      abl/abducer/abducer_base.py

+ 41
- 43
abl/abducer/abducer_base.py View File

@@ -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()



Loading…
Cancel
Save