1- from typing import Callable , Literal , Tuple , Union
1+ from collections .abc import Callable
2+ from typing import Literal
23import numpy as np
34import scipy .sparse as sp
45from numpy .typing import NDArray
89
910# Type aliases for clarity
1011Mode = Literal ["distance" , "similarity" ]
11- ConversionMethod = Union [
12- Literal ["reciprocal" , "negative" , "exp" , "gaussian" ],
13- Callable [[np .ndarray ], np .ndarray ]
14- ]
12+ ConversionMethod = Literal ["reciprocal" , "negative" , "exp" , "gaussian" ] | Callable [[np .ndarray ], np .ndarray ]
1513
1614
1715def _validate_square_matrix (M : np .ndarray ) -> None :
@@ -40,7 +38,7 @@ def _drop_diagonal(A: sp.csr_matrix) -> sp.csr_matrix:
4038 return sp .csr_matrix ((coo .data [mask ], (coo .row [mask ], coo .col [mask ])), shape = A .shape )
4139
4240
43- def _coerce_knn_inputs (indices , distances ) -> Tuple [np .ndarray , np .ndarray ]:
41+ def _coerce_knn_inputs (indices , distances ) -> tuple [np .ndarray , np .ndarray ]:
4442 ind = _to_numpy (indices )
4543 dist = _to_numpy (distances )
4644 if ind .shape != dist .shape :
@@ -60,7 +58,7 @@ def _csr_from_edges(n: int, rows: np.ndarray, cols: np.ndarray, weights: np.ndar
6058 return csr_matrix ((weights , (rows , cols )), shape = (n , n ))
6159
6260
63- def _as_csr_square (M : NDArray | spmatrix ) -> Tuple [sp .csr_matrix , int ]:
61+ def _as_csr_square (M : NDArray | spmatrix ) -> tuple [sp .csr_matrix , int ]:
6462 """Return (CSR, n) for a square matrix without densifying.
6563
6664 If `M` is dense, convert to CSR. If `M` is sparse, convert format to CSR
@@ -78,7 +76,7 @@ def _as_csr_square(M: NDArray | spmatrix) -> Tuple[sp.csr_matrix, int]:
7876 return sp .csr_matrix (arr ), arr .shape [0 ]
7977
8078
81- def _topk_per_row_sparse (csr : sp .csr_matrix , k : int , * , largest : bool ) -> Tuple [np .ndarray , np .ndarray ]:
79+ def _topk_per_row_sparse (csr : sp .csr_matrix , k : int , * , largest : bool ) -> tuple [np .ndarray , np .ndarray ]:
8280 """Return (indices, values) of top-k entries per row from CSR matrix.
8381
8482 This operates strictly on the row's nonzeros without densifying.
@@ -120,7 +118,7 @@ def _topk_per_row_sparse(csr: sp.csr_matrix, k: int, *, largest: bool) -> Tuple[
120118 return ind , vals
121119
122120
123- def _knn_from_matrix (M : NDArray | spmatrix , k : int , * , mode : MatrixMode ) -> Tuple [np .ndarray , np .ndarray ]:
121+ def _knn_from_matrix (M : NDArray | spmatrix , k : int , * , mode : MatrixMode ) -> tuple [np .ndarray , np .ndarray ]:
124122 """Compute kNN (indices, values) from a square distance/similarity matrix.
125123
126124 Supports dense and sparse inputs without densifying sparse matrices.
0 commit comments