Browse Source

[MNT] resolve some (more) comments

pull/3/head
troyyyyy 2 years ago
parent
commit
b43959ac1f
2 changed files with 8 additions and 9 deletions
  1. +5
    -5
      abl/reasoning/kb.py
  2. +3
    -4
      abl/reasoning/reasoner.py

+ 5
- 5
abl/reasoning/kb.py View File

@@ -5,7 +5,7 @@ import numpy as np
from collections import defaultdict
from itertools import product, combinations

from ..utils.utils import flatten, reform_idx, hamming_dist, check_equal, to_hashable, hashable_to_list
from abl.utils.utils import flatten, reform_idx, hamming_dist, check_equal, to_hashable, hashable_to_list

from multiprocessing import Pool

@@ -86,14 +86,14 @@ class KBBase(ABC):
for idx in range(key_idx - 1, 0, -1):
k = key_list[idx]
if abs(k - y) <= self.max_err:
all_candidates += potential_candidates[k]
all_candidates.extend(potential_candidates[k])
else:
break
for idx in range(key_idx, len(key_list)):
k = key_list[idx]
if abs(k - y) <= self.max_err:
all_candidates += potential_candidates[k]
all_candidates.extend(potential_candidates[k])
else:
break
return all_candidates
@@ -139,7 +139,7 @@ class KBBase(ABC):
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 += self._revision(revision_num, pred_res, y)
candidates.extend(self._revision(revision_num, pred_res, y))
if len(candidates) > 0:
min_revision_num = revision_num
break
@@ -149,7 +149,7 @@ class KBBase(ABC):
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 += self._revision(revision_num, pred_res, y)
candidates.extend(self._revision(revision_num, pred_res, y))
return candidates
@lru_cache(maxsize=None)


+ 3
- 4
abl/reasoning/reasoner.py View File

@@ -1,8 +1,7 @@
import abc
import numpy as np
from multiprocessing import Pool
from zoopt import Dimension, Objective, Parameter, Opt
from ..utils.utils import (
from abl.utils.utils import (
confidence_dist,
flatten,
reform_idx,
@@ -11,7 +10,7 @@ from ..utils.utils import (
)


class ReasonerBase(abc.ABC):
class ReasonerBase():
def __init__(self, kb, dist_func="hamming", mapping=None, zoopt=False):
if not (dist_func == "hamming" or dist_func == "confidence"):
raise NotImplementedError
@@ -50,7 +49,7 @@ class ReasonerBase(abc.ABC):
return hamming_dist(pseudo_label, candidates)

elif self.dist_func == "confidence":
candidates = [list(map(lambda x: self.remapping[x], c)) for c in candidates]
candidates = [[self.remapping[x] for x in c] for c in candidates]
return confidence_dist(pred_res_prob, candidates)

def _get_one_candidate(self, pseudo_label, pred_res_prob, candidates):


Loading…
Cancel
Save