44from typing import TYPE_CHECKING , cast
55
66import numpy as np
7- from scipy .sparse import csr_matrix
87from sklearn .base import BaseEstimator , ClusterMixin
98from sklearn .neighbors import KNeighborsTransformer
10- from sklearn .utils import Tags
119from sklearn .utils .validation import validate_data
1210
1311from ..utils import get_sparse_row
1412
1513if TYPE_CHECKING :
16- from collections .abc import Iterator
17- from typing import Literal , Self
14+ from collections .abc import Callable , Iterator
15+ from typing import Any , Literal , Self
1816
1917 from numpy .typing import NDArray
18+ from scipy .sparse import csr_matrix
2019 from sklearn .pipeline import Pipeline
20+ from sklearn .utils import Tags
2121
2222
2323UNCLASSIFIED = - 2
2424NOISE = - 1
2525
2626
27- def join (it1 , it2 ):
27+ def join (
28+ it1 : Iterator [tuple [int , float ]], it2 : Iterator [tuple [int , float ]]
29+ ) -> Iterator [tuple [int , float ]]:
2830 cur_it1 = next (it1 , None )
2931 cur_it2 = next (it2 , None )
3032 while 1 :
31- if cur_it1 is None and cur_it2 is None :
32- break
33- elif cur_it1 is None :
33+ if cur_it1 is None :
34+ if cur_it2 is None :
35+ break
3436 yield cur_it2
3537 cur_it2 = next (it2 , None )
3638 elif cur_it2 is None :
@@ -66,7 +68,10 @@ def neighborhood(
6668
6769
6870def rnn_dbscan_inner (
69- is_core : NDArray [np .bool_ ], knns : csr_matrix , rev_knns : csr_matrix , labels
71+ is_core : NDArray [np .bool_ ],
72+ knns : csr_matrix ,
73+ rev_knns : csr_matrix ,
74+ labels : NDArray [np .int32 ],
7075) -> list [float ]:
7176 cluster = 0
7277 cur_dens = 0.0
@@ -78,7 +83,7 @@ def rnn_dbscan_inner(
7883 labels [x_idx ] = cluster
7984 # TODO: Make this inner bit faster - can just assume
8085 # sorted an keep sorted
81- seeds = deque ()
86+ seeds : deque [ int ] = deque ()
8287 for neighbor_idx , dist in neighborhood (is_core , knns , rev_knns , x_idx ):
8388 labels [neighbor_idx ] = cluster
8489 if dist > cur_dens :
@@ -166,37 +171,39 @@ def __init__(
166171 self .keep_knns = keep_knns
167172
168173 def fit (self , X : NDArray [np .float64 ] | csr_matrix , y : None = None ) -> Self :
169- X = cast ( csr_matrix , validate_data (self , X , accept_sparse = "csr" ) )
174+ X = validate_data (self , X , accept_sparse = "csr" )
170175 if self .input_guarantee == "none" :
171176 algorithm = KNeighborsTransformer (n_neighbors = self .n_neighbors )
172- X = algorithm .fit_transform (X )
177+ knns = cast ( "csr_matrix" , algorithm .fit_transform (X ) )
173178 elif self .input_guarantee == "kneighbors" :
174- pass
179+ knns = cast ( "csr_matrix" , X )
175180 else :
176181 raise ValueError (
177182 "Expected input_guarantee to be one of 'none', 'kneighbors'"
178183 )
179184
180- XT = cast ( csr_matrix , X .transpose ().tocsr (copy = True ) )
185+ rev_knns = knns .transpose ().tocsr (copy = True )
181186 if self .keep_knns :
182- self .knns_ = X
183- self .rev_knns_ = XT
187+ self .knns_ = knns
188+ self .rev_knns_ = rev_knns
184189
185190 # Initially, all samples are unclassified.
186- labels = np .full (X .shape [0 ], UNCLASSIFIED , dtype = np .int32 )
191+ labels = np .full (knns .shape [0 ], UNCLASSIFIED , dtype = np .int32 )
187192
188193 # A list of all core samples found. -1 is to account for diagonal.
189- core_samples = XT .getnnz (1 ) - 1 >= self .n_neighbors
194+ core_samples = rev_knns .getnnz (1 ) - 1 >= self .n_neighbors
190195
191- dens = rnn_dbscan_inner (core_samples , X , XT , labels )
196+ dens = rnn_dbscan_inner (core_samples , knns , rev_knns , labels )
192197
193198 self .core_sample_indices_ = core_samples .nonzero ()
194199 self .labels_ = labels
195200 self .dens_ = dens
196201
197202 return self
198203
199- def fit_predict (self , X , y = None ) -> NDArray [np .int32 ]:
204+ def fit_predict ( # type: ignore[override]
205+ self , X : NDArray [np .float64 ] | csr_matrix , y : None = None
206+ ) -> NDArray [np .int32 ]:
200207 self .fit (X , y = y )
201208 return self .labels_
202209
@@ -205,13 +212,13 @@ def drop_knns(self) -> None:
205212 del self .rev_knns_
206213
207214 def __sklearn_tags__ (self ) -> Tags :
208- tags = cast (Tags , super ().__sklearn_tags__ ())
215+ tags = cast (" Tags" , super ().__sklearn_tags__ ()) # type: ignore[no-untyped-call]
209216 tags .input_tags .sparse = True
210217 return tags
211218
212219
213220def simple_rnn_dbscan_pipeline (
214- neighbor_transformer : object ,
221+ neighbor_transformer : Callable [..., Any ] ,
215222 n_neighbors : int ,
216223 * ,
217224 n_jobs : int | None = None ,
@@ -236,12 +243,17 @@ class implementing KNeighborsTransformer interface
236243 """
237244 from sklearn .pipeline import make_pipeline
238245
239- return make_pipeline (
240- neighbor_transformer (n_neighbors = n_neighbors , n_jobs = n_jobs , ** kwargs ),
241- RnnDBSCAN (
242- n_neighbors = n_neighbors ,
243- input_guarantee = "kneighbors" ,
244- n_jobs = n_jobs ,
245- keep_knns = keep_knns ,
246+ return cast (
247+ "Pipeline" ,
248+ make_pipeline (
249+ neighbor_transformer (
250+ n_neighbors = n_neighbors , n_jobs = n_jobs , input_guarantee = input_guarantee
251+ ),
252+ RnnDBSCAN (
253+ n_neighbors = n_neighbors ,
254+ input_guarantee = "kneighbors" ,
255+ n_jobs = n_jobs ,
256+ keep_knns = keep_knns ,
257+ ),
246258 ),
247259 )
0 commit comments