From b43959ac1f7b1261f032bbccf430e7c110c991df Mon Sep 17 00:00:00 2001 From: troyyyyy Date: Sun, 8 Oct 2023 16:36:07 +0800 Subject: [PATCH] [MNT] resolve some (more) comments --- abl/reasoning/kb.py | 10 +++++----- abl/reasoning/reasoner.py | 7 +++---- 2 files changed, 8 insertions(+), 9 deletions(-) diff --git a/abl/reasoning/kb.py b/abl/reasoning/kb.py index 1a44a37..33b74a9 100644 --- a/abl/reasoning/kb.py +++ b/abl/reasoning/kb.py @@ -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) diff --git a/abl/reasoning/reasoner.py b/abl/reasoning/reasoner.py index 8964dbd..79ab960 100644 --- a/abl/reasoning/reasoner.py +++ b/abl/reasoning/reasoner.py @@ -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):