diff --git a/src/anonymeter/evaluators/inference_evaluator.py b/src/anonymeter/evaluators/inference_evaluator.py index 51682ea..af716e2 100644 --- a/src/anonymeter/evaluators/inference_evaluator.py +++ b/src/anonymeter/evaluators/inference_evaluator.py @@ -324,7 +324,7 @@ def risk_for_groups(self, confidence_level: float = 0.95) -> dict[str, Evaluatio # Count the number of success control attacks n_control = evaluate_inference_guesses( guesses=self._guesses_control.loc[common_indices], - secrets=data_control[self._secret], + secrets=data_control[self._secret].loc[common_indices], regression=self._regression, ).sum() else: diff --git a/src/anonymeter/neighbors/mixed_types_kneighbors.py b/src/anonymeter/neighbors/mixed_types_kneighbors.py index 318628a..be7c402 100644 --- a/src/anonymeter/neighbors/mixed_types_kneighbors.py +++ b/src/anonymeter/neighbors/mixed_types_kneighbors.py @@ -277,4 +277,7 @@ def predict(self, x: pd.DataFrame) -> pd.Series: guesses_idx = self._nn.kneighbors(queries=x[self._columns]) if isinstance(guesses_idx, tuple): raise RuntimeError("guesses_idx cannot be a tuple") - return self._target_series.iloc[guesses_idx.flatten()] + + guesses = self._target_series.iloc[guesses_idx.flatten()] + guesses.index = x.index + return guesses diff --git a/tests/test_mixed_types_kneigbors.py b/tests/test_mixed_types_kneigbors.py index 21e0b26..a7315b7 100644 --- a/tests/test_mixed_types_kneigbors.py +++ b/tests/test_mixed_types_kneigbors.py @@ -5,7 +5,7 @@ import pandas as pd import pytest -from anonymeter.neighbors.mixed_types_kneighbors import MixedTypeKNeighbors, gower_distance +from anonymeter.neighbors.mixed_types_kneighbors import KNNInferencePredictor, MixedTypeKNeighbors, gower_distance from tests.fixtures import get_adult @@ -79,3 +79,14 @@ def test_gower_distance_numerical(): r0, r1 = rng.random(size=10), rng.random(size=10) dist = gower_distance(r0=r0, r1=r1, cat_cols_index=10) np.testing.assert_almost_equal(dist, np.sum(np.abs(r0 - r1))) + + +def test_knn_inference_predictor_prediction_index_alignment(): + df = get_adult("ori", n_samples=10) + aux_cols = ["age", "education", "sex"] + predictor = KNNInferencePredictor(data=df, columns=aux_cols, target_col="income", n_jobs=1) + queries = df[aux_cols] + + guesses = predictor.predict(queries) + + pd.testing.assert_index_equal(guesses.index, queries.index)