diff --git a/sknetwork/classification/knn.py b/sknetwork/classification/knn.py index f3c3aefd..a5a036b9 100644 --- a/sknetwork/classification/knn.py +++ b/sknetwork/classification/knn.py @@ -65,13 +65,14 @@ def _instantiate_vars(labels: np.ndarray): def _fit_core(self, embedding, labels, index_train, index_test): n_neighbors = check_n_neighbors(self.n_neighbors, len(index_train)) - norms_train = get_norms(embedding[index_train], p=2) + embedding_train = embedding[index_train] + norms_train = get_norms(embedding_train, p=2) neighbors = [] for i in index_test: vector = embedding[i] if sparse.issparse(vector): vector = vector.toarray().ravel() - distances = norms_train**2 - 2 * embedding[index_train].dot(vector) + np.sum(vector**2) + distances = norms_train**2 - 2 * embedding_train.dot(vector) + np.sum(vector**2) neighbors += list(index_train[np.argpartition(distances, n_neighbors)[:n_neighbors]]) labels_neighbor = labels[neighbors] diff --git a/sknetwork/linkpred/nn.py b/sknetwork/linkpred/nn.py index ca6296d5..94a17586 100644 --- a/sknetwork/linkpred/nn.py +++ b/sknetwork/linkpred/nn.py @@ -67,6 +67,7 @@ def _fit_core(self, embedding, mask): index_col = np.arange(n) n_col = n n_neighbors = check_n_neighbors(self.n_neighbors, n_col) + embedding_col = embedding[index_col] row = [] col = [] @@ -76,7 +77,7 @@ def _fit_core(self, embedding, mask): vector = embedding[i] if sparse.issparse(vector): vector = vector.toarray().ravel() - similarities = embedding[index_col].dot(vector) + similarities = embedding_col.dot(vector) nn = np.argpartition(-similarities, n_neighbors)[:n_neighbors] mask_nn = np.zeros(n_col, dtype=bool) mask_nn[nn] = 1