Skip to content

Commit eee0114

Browse files
committed
Fix probabilities depend on other samples (#159)
1 parent e6f6187 commit eee0114

1 file changed

Lines changed: 7 additions & 14 deletions

File tree

hiclass/LocalClassifierPerParentNode.py

Lines changed: 7 additions & 14 deletions
Original file line numberDiff line numberDiff line change
@@ -217,9 +217,6 @@ def predict_proba(self, X):
217217

218218
self.logger_.info("Predicting Probability")
219219

220-
# Initialize array that holds predictions
221-
y = np.empty((X.shape[0], self.max_levels_), dtype=self.dtype_)
222-
223220
# Predict first level
224221
classifier = self.hierarchy_.nodes[self.root_]["classifier"]
225222
# use classifier as a fallback if no calibrator is available
@@ -229,8 +226,7 @@ def predict_proba(self, X):
229226
)
230227
proba = calibrator.predict_proba(X)
231228

232-
y[:, 0] = calibrator.classes_[np.argmax(proba, axis=1)]
233-
level_probability_list = [proba] + self._predict_proba_remaining_levels(X, y)
229+
level_probability_list = [proba] + self._predict_proba_remaining_levels(X)
234230

235231
level_probability_list = self._combine_and_reorder(level_probability_list)
236232

@@ -252,18 +248,16 @@ def predict_proba(self, X):
252248
else level_probability_list[-1]
253249
)
254250

255-
def _predict_proba_remaining_levels(self, X, y):
251+
def _predict_proba_remaining_levels(self, X):
256252
level_probability_list = []
257-
for level in range(1, y.shape[1]):
258-
predecessors = set(y[:, level - 1])
253+
for level in range(1, len(self.global_classes_)):
254+
predecessors = set(self.global_classes_[level - 1])
259255
predecessors.discard("")
260256
level_dimension = self.max_level_dimensions_[level]
261257
cur_level_probabilities = np.zeros((X.shape[0], level_dimension))
262258

263259
for predecessor in predecessors:
264-
mask = np.isin(y[:, level - 1], self.global_classes_[level - 1])
265-
predecessor_x = X[mask]
266-
if predecessor_x.shape[0] > 0:
260+
if X.shape[0] > 0:
267261
successors = list(self.hierarchy_.successors(predecessor))
268262
if len(successors) > 0:
269263
classifier = self.hierarchy_.nodes[predecessor]["classifier"]
@@ -275,8 +269,7 @@ def _predict_proba_remaining_levels(self, X, y):
275269
or classifier
276270
)
277271

278-
proba = calibrator.predict_proba(predecessor_x)
279-
y[mask, level] = calibrator.classes_[np.argmax(proba, axis=1)]
272+
proba = calibrator.predict_proba(X)
280273

281274
for successor in successors:
282275
class_index = self.global_class_to_index_mapping_[level][
@@ -286,7 +279,7 @@ def _predict_proba_remaining_levels(self, X, y):
286279
proba_index = np.where(calibrator.classes_ == successor)[0][
287280
0
288281
]
289-
cur_level_probabilities[mask, class_index] = proba[
282+
cur_level_probabilities[:, class_index] = proba[
290283
:, proba_index
291284
]
292285

0 commit comments

Comments
 (0)