
    MpjO                        S SK r S SKJr  S SKJr  S SKrS SKJr  S SK	J
r
JrJrJrJrJr  S SKJrJr  S SKJr  S SKJrJr  S S	KJr  S S
KJrJrJr  S SKJr  S5S jr \!" 5       4S jr"\" \ SS9r#S6S jr$S r% " S S\&5      r' " S S\\
5      r( " S S\\
5      r) " S S\\
5      r* " S S\\
5      r+ " S S\+5      r, " S S\+5      r- " S  S!\+5      r. " S" S#\\
5      r/ " S$ S%\
5      r0S6S& jr1 " S' S(\5      r2 " S) S*\\5      r3 " S+ S,\3\5      r4 " S- S.\\\
5      r5 " S/ S0\\\
5      r6 " S1 S2\\\
5      r7 " S3 S4\\\
5      r8g)7    N)defaultdict)partial)assert_array_equal)BaseEstimatorClassifierMixinMetaEstimatorMixinRegressorMixinTransformerMixinclone)_Scorermean_squared_error)BaseCrossValidator)
GroupKFoldGroupsConsumerMixin)SIMPLE_METHODS)MetadataRouterMethodMappingprocess_routing)_check_partial_fit_first_callc                    [         R                  " 5       nUS   R                  nUS   R                  n[        U S5      (       d  [	        S 5      U l        U(       dA  UR                  5        VVs0 s H$  u  pg[        U[        5      (       a  US:w  d  M"  Xg_M&     nnnU R
                  U   U   R                  U5        gs  snnf )zUtility function to store passed metadata to a method of obj.

If record_default is False, kwargs whose values are "default" are skipped.
This is so that checks on keyword arguments whose default was not changed
are skipped.

      _recordsc                       [        [        5      $ N)r   list     a/var/www/html/pdf-tiff/venv/lib/python3.13/site-packages/sklearn/tests/metadata_routing_common.py<lambda>!record_metadata.<locals>.<lambda>*   s	    ;t+<r   defaultN)
inspectstackfunctionhasattrr   r   items
isinstancestrappend)objrecord_defaultkwargsr$   calleecallerkeyvals           r   record_metadatar2      s     MMOE1XF1XF3
##"#<= #LLN
*c3''C9,< CH* 	 

 LL ''/
s   0!C C c           	      ~   [        U S[        5       5      R                  U[        5       5      R                  U[        5       5      nU H  n[	        UR                  5       5      [	        UR                  5       5      :X  d)   SUR                  5        SUR                  5        35       eUR                  5        H~  u  pxXg   n	Xs;   a0  U	b-  [        R                  " X5      R                  5       (       d   eM>  [        U	[        R                  5      (       a  [        X5        Mj  XL a  Mp   SU	 SU SU 35       e   M     g)an  Check whether the expected metadata is passed to the object's method.

Parameters
----------
obj : estimator object
    sub-estimator to check routed params for
method : str
    sub-estimator's method where metadata is routed to, or otherwise in
    the context of metadata routing referred to as 'callee'
parent : str
    the parent method which should have called `method`, or otherwise in
    the context of metadata routing referred to as 'caller'
split_params : tuple, default=empty
    specifies any parameters which are to be checked as being a subset
    of the original values
**kwargs : dict
    passed metadata
r   z	Expected z vs Nz
. Method: )getattrdictgetr   setkeysr'   npisinallr(   ndarrayr   )
r+   methodparentsplit_paramsr-   all_recordsrecordr0   valuerecorded_values
             r   check_recorded_metadatarD   4   s   ( 	Z(,,VTV<@@P   6;;=!S%77 	
d6;;=/:	
7 !,,.JC#[N "~'Aww~599;;;;nbjj99&~=)2 #N#34wjQ2 ) r   F)r,   c           	         [        U [        5      (       a/  U  H(  u  p#Ub
  X!;   a  X   nOSn[        UR                  US9  M*     gUc  / OUn[         Hf  nXQ;   a  M
  [        X5      nUR                  R                  5        VVs/ s H!  u  px[        U[        5      (       d  Uc  M  UPM#     n	nnU	(       d  Mf   e   gs  snnf )zCheck if a metadata request dict is empty.

One can exclude a method or a list of methods from the check using the
``exclude`` parameter. If metadata_request is a MetadataRouter, then
``exclude`` can be of the form ``{"object" : [method, ...]}``.
N)exclude)	r(   r   assert_request_is_emptyrouterr   r4   requestsr'   r)   )
metadata_requestrF   nameroute_mapping_excluder=   mmrpropaliaspropss
             r   rG   rG   b   s     "N33#3D"t"=#M$8$8(K $4 	ObG &/  #||113
3%%% 3 	 

 5y !
s   B=&B=c                    UR                  5        H"  u  p#[        X5      nUR                  U:X  a  M"   e   [         Vs/ s H  o"U;  d  M
  UPM     nnU H(  n[	        [        X5      R                  5      (       d  M(   e   g s  snf r   )r'   r4   rI   r   len)request
dictionaryr=   rI   rN   empty_methodss         r   assert_request_equalrW      s{    &,,.g&||x''' / +9U.*<TV.MUww/889999   Vs   	BBc                        \ rS rSrS rS rSrg)	_Registry   c                     U $ r   r   )selfmemos     r   __deepcopy___Registry.__deepcopy__       r   c                     U $ r   r   r\   s    r   __copy___Registry.__copy__   r`   r   r   N)__name__
__module____qualname____firstlineno__r^   rc   __static_attributes__r   r   r   rY   rY      s    r   rY   c                   J    \ rS rSrSrS
S jrSS jrSS jrSS jrSS jr	S	r
g)ConsumingRegressor   aC  A regressor consuming metadata.

Parameters
----------
registry : list, default=None
    If a list, the estimator will append itself to the list in order to have
    a reference to the estimator later on. Since that reference is not
    required in all tests, registration can be skipped by leaving this value
    as None.
Nc                     Xl         g r   registryr\   ro   s     r   __init__ConsumingRegressor.__init__        r   c                 j    U R                   b  U R                   R                  U 5        [        XUS9  U $ Nsample_weightmetadataro   r*   record_metadata_not_defaultr\   Xyrw   rx   s        r   partial_fitConsumingRegressor.partial_fit   2    ==$MM  &#	
 r   c                 j    U R                   b  U R                   R                  U 5        [        XUS9  U $ ru   ry   r{   s        r   fitConsumingRegressor.fit   r   r   c                 R    [        XUS9  [        R                  " [        U5      4S9$ )Nrv   shape)rz   r9   zerosrS   r{   s        r   predictConsumingRegressor.predict   s&    #	
 xxs1vi((r   c                     [        XUS9  gNrv   r   rz   r{   s        r   scoreConsumingRegressor.score       #	
 r   rn   r   r"   r"   Nr"   r"   )re   rf   rg   rh   __doc__rq   r~   r   r   r   ri   r   r   r   rk   rk      s     	!)r   rk   c                   J    \ rS rSrSrSS jrS rSS jrS rS r	S	 r
S
 rSrg)NonConsumingClassifier   5A classifier which accepts no metadata on any method.c                     Xl         g r   )alpha)r\   r   s     r   rq   NonConsumingClassifier.__init__   s    
r   c                 r    [         R                  " U5      U l        [         R                  " U5      U l        U $ r   )r9   uniqueclasses_	ones_likecoef_r\   r|   r}   s      r   r   NonConsumingClassifier.fit   s%    		!\\!_
r   Nc                     U $ r   r   )r\   r|   r}   classess       r   r~   "NonConsumingClassifier.partial_fit   r`   r   c                 $    U R                  U5      $ r   )r   r\   r|   s     r   decision_function(NonConsumingClassifier.decision_function   s    ||Ar   c                     [         R                  " [        U5      4S9nSUS [        U5      S-  & SU[        U5      S-  S & U$ )Nr   r   r   r   )r9   emptyrS   )r\   r|   y_preds      r   r   NonConsumingClassifier.predict   sC    Q	* !}Q1 !s1v{}r   c                 *   [         R                  " [        U5      [        U R                  5      4[         R                  S9n[         R
                  R                  [         R                  " [        U R                  5      5      [        U5      S9US S & U$ )Nr   dtyper   size)r9   r   rS   r   float32random	dirichletones)r\   r|   y_probas      r   predict_proba$NonConsumingClassifier.predict_proba   sc    ((#a&#dmm*<!=RZZPYY((rwws4==7I/JQTUVQW(X
r   c                 $    U R                  U5      $ r   )r   r   s     r   predict_log_proba(NonConsumingClassifier.predict_log_proba   s    !!!$$r   )r   r   r   )        r   )re   rf   rg   rh   r   rq   r   r~   r   r   r   r   ri   r   r   r   r   r      s(    ?
%r   r   c                   *    \ rS rSrSrS rS rS rSrg)NonConsumingRegressor   r   c                     U $ r   r   r   s      r   r   NonConsumingRegressor.fit   r`   r   c                     U $ r   r   r   s      r   r~   !NonConsumingRegressor.partial_fit   r`   r   c                 @    [         R                  " [        U5      5      $ r   )r9   r   rS   r   s     r   r   NonConsumingRegressor.predict   s    wws1vr   r   N)	re   rf   rg   rh   r   r   r~   r   ri   r   r   r   r   r      s    ?r   r   c                   j    \ rS rSrSrSS jr SS jrSS jrSS jrSS jr	SS	 jr
SS
 jrSS jrSrg)ConsumingClassifier   a  A classifier consuming metadata.

Parameters
----------
registry : list, default=None
    If a list, the estimator will append itself to the list in order to have
    a reference to the estimator later on. Since that reference is not
    required in all tests, registration can be skipped by leaving this value
    as None.

alpha : float, default=0
    This parameter is only used to test the ``*SearchCV`` objects, and
    doesn't do anything.
Nc                     X l         Xl        g r   )r   ro   )r\   ro   r   s      r   rq   ConsumingClassifier.__init__  s    
 r   c                     U R                   b  U R                   R                  U 5        [        XUS9  [        X5        U $ ru   )ro   r*   rz   r   )r\   r|   r}   r   rw   rx   s         r   r~   ConsumingClassifier.partial_fit  s<     ==$MM  &#	
 	&d4r   c                     U R                   b  U R                   R                  U 5        [        XUS9  [        R                  " U5      U l        [        R                  " U5      U l        U $ ru   )ro   r*   rz   r9   r   r   r   r   r{   s        r   r   ConsumingClassifier.fit  sP    ==$MM  &#	
 		!\\!_
r   c                     [        XUS9  [        R                  " [        U5      4SS9nSU[        U5      S-  S & SUS [        U5      S-  & U$ )Nrv   int8r   r   r   r   rz   r9   r   rS   r\   r|   rw   rx   y_scores        r   r   ConsumingClassifier.predict   sT    #	
 ((#a&&9!"A!!"#a&A+r   c                 >   [        XUS9  [        R                  " [        U5      [        U R                  5      4[        R
                  S9n[        R                  R                  [        R                  " [        U R                  5      5      [        U5      S9US S & U$ )Nrv   r   r   )	rz   r9   r   rS   r   r   r   r   r   )r\   r|   rw   rx   r   s        r   r   !ConsumingClassifier.predict_proba)  sr    #	
 ((#a&#dmm*<!=RZZPYY((rwws4==7I/JQTUVQW(X
r   c                 8    [        XUS9  U R                  U5      $ ru   )rz   r   r\   r|   rw   rx   s       r   r   %ConsumingClassifier.predict_log_proba2  s"    #	
 !!!$$r   c                     [        XUS9  [        R                  " [        U5      4S9nSU[        U5      S-  S & SUS [        U5      S-  & U$ )Nrv   r   r   r   r   r   r   s        r   r   %ConsumingClassifier.decision_function8  sR    #	
 ((#a&+!"A!!"#a&A+r   c                     [        XUS9  gr   r   r{   s        r   r   ConsumingClassifier.scoreA  r   r   )r   r   r   ro   )Nr   r   r   )re   rf   rg   rh   r   rq   r~   r   r   r   r   r   r   ri   r   r   r   r   r      s6    !
 EN

%r   r   c                   (    \ rS rSrSr\S 5       rSrg)&ConsumingClassifierWithoutPredictProbaiH  zConsumingClassifier without a predict_proba method, but with predict_log_proba.

Used to mimic dynamic method selection such as in the `_parallel_predict_proba()`
function called by `BaggingClassifier`.
c                     [        S5      eNz-This estimator does not support predict_probaAttributeErrorrb   s    r   r   4ConsumingClassifierWithoutPredictProba.predict_probaO      LMMr   r   N)re   rf   rg   rh   r   propertyr   ri   r   r   r   r   r   H  s     N Nr   r   c                   (    \ rS rSrSr\S 5       rSrg))ConsumingClassifierWithoutPredictLogProbaiT  zConsumingClassifier without a predict_log_proba method, but with predict_proba.

Used to mimic dynamic method selection such as in
`BaggingClassifier.predict_log_proba()`.
c                     [        S5      eNz1This estimator does not support predict_log_probar   rb   s    r   r   ;ConsumingClassifierWithoutPredictLogProba.predict_log_proba[      PQQr   r   N)re   rf   rg   rh   r   r   r   ri   r   r   r   r   r   T  s     R Rr   r   c                   8    \ rS rSrSr\S 5       r\S 5       rSrg)"ConsumingClassifierWithOnlyPredicti`  zConsumingClassifier with only a predict method.

Used to mimic dynamic method selection such as in
`BaggingClassifier.predict_log_proba()`.
c                     [        S5      er   r   rb   s    r   r   0ConsumingClassifierWithOnlyPredict.predict_probag  r   r   c                     [        S5      er   r   rb   s    r   r   4ConsumingClassifierWithOnlyPredict.predict_log_probak  r   r   r   N)	re   rf   rg   rh   r   r   r   r   ri   r   r   r   r   r   `  s3     N N R Rr   r   c                   J    \ rS rSrSrS
S jrSS jrSS jrSS jrSS jr	S	r
g)ConsumingTransformerip  a^  A transformer which accepts metadata on fit and transform.

Parameters
----------
registry : list, default=None
    If a list, the estimator will append itself to the list in order to have
    a reference to the estimator later on. Since that reference is not
    required in all tests, registration can be skipped by leaving this value
    as None.
Nc                     Xl         g r   rn   rp   s     r   rq   ConsumingTransformer.__init__|  rs   r   c                 x    U R                   b  U R                   R                  U 5        [        XUS9  SU l        U $ )Nrv   T)ro   r*   rz   fitted_r{   s        r   r   ConsumingTransformer.fit  s9    ==$MM  &#	
 r   c                      [        XUS9  US-   $ r   r   r   s       r   	transformConsumingTransformer.transform      #	
 1ur   c                 R    [        XUS9  U R                  XX4S9R                  XUS9$ ru   )rz   r   r   r{   s        r   fit_transform"ConsumingTransformer.fit_transform  s>    
 	$	
 xxMxMWWX X 
 	
r   c                      [        XUS9  US-
  $ r   r   r   s       r   inverse_transform&ConsumingTransformer.inverse_transform  r   r   )r   ro   r   r   r   NN)re   rf   rg   rh   r   rq   r   r   r   r  ri   r   r   r   r   r   p  s     	!

r   r   c                   6    \ rS rSrSrSS jrS	S jrS
S jrSrg)"ConsumingNoFitTransformTransformeri  zA metadata consuming transformer that doesn't inherit from
TransformerMixin, and thus doesn't implement `fit_transform`. Note that
TransformerMixin's `fit_transform` doesn't route metadata to `transform`.Nc                     Xl         g r   rn   rp   s     r   rq   +ConsumingNoFitTransformTransformer.__init__  rs   r   c                 j    U R                   b  U R                   R                  U 5        [        XUS9  U $ ru   )ro   r*   r2   r{   s        r   r   &ConsumingNoFitTransformTransformer.fit  s-    ==$MM  &HMr   c                     [        XUS9  U$ ru   )r2   r   s       r   r   ,ConsumingNoFitTransformTransformer.transform  s    HMr   rn   r   NNNr  )	re   rf   rg   rh   r   rq   r   r   ri   r   r   r   r  r    s    Q!r   r  c                     Ub  UR                  [        5        [        [        40 UD6  UR                  SS 5      n[	        XUS9$ )Nrw   rw   )r*   consuming_metricrz   r6   r   )r   y_truero   r-   rw   s        r   r  r    s@    () 0;F;JJ5MfMJJr   c                   ,   ^  \ rS rSrSU 4S jjrSrU =r$ )ConsumingScoreri  c                 N   > [        [        US9n[        TU ]  US0 SS9  Xl        g )Nrn   r   r   )
score_funcsignr-   response_method)r   r  superrq   ro   )r\   ro   r  	__class__s      r   rq   ConsumingScorer.__init__  s2    -A
!"i 	 	
 !r   rn   r   )re   rf   rg   rh   rq   ri   __classcell__)r  s   @r   r  r    s    ! !r   r  c                   <    \ rS rSrSS jrS	S jrS
S jrSS jrSrg)ConsumingSplitteri  Nc                     Xl         g r   rn   rp   s     r   rq   ConsumingSplitter.__init__  rs   r   c              #     #    U R                   b  U R                   R                  U 5        [        XUS9  [        U5      S-  n[	        [        SU5      5      n[	        [        U[        U5      5      5      nXv4v   Xg4v   g 7f)N)groupsrx   r   r   )ro   r*   rz   rS   r   range)r\   r|   r}   r   rx   split_indextrain_indicestest_indicess           r   splitConsumingSplitter.split  sp     ==$MM  &#D(K!fkU1k23E+s1v67))))s   A?Bc                     g)Nr   r   )r\   r|   r}   r   rx   s        r   get_n_splitsConsumingSplitter.get_n_splits  s    r   c              #      #    [        U5      S-  n[        [        SU5      5      n[        [        U[        U5      5      5      nUv   Uv   g 7f)Nr   r   )rS   r   r!  )r\   r|   r}   r   r"  r#  r$  s          r   _iter_test_indices$ConsumingSplitter._iter_test_indices  sD     !fkU1k23E+s1v67s   AArn   r   r   )NNNNr  )	re   rf   rg   rh   rq   r%  r(  r+  ri   r   r   r   r  r    s    !
*r   r  c                       \ rS rSrSrSrg))ConsumingSplitterInheritingFromGroupKFoldi  zXHelper class that can be used to test TargetEncoder, that only takes specific
splitters.r   N)re   rf   rg   rh   r   ri   r   r   r   r.  r.    s    r   r.  c                   *    \ rS rSrSrS rS rS rSrg)MetaRegressori  z(A meta-regressor which is only a router.c                     Xl         g r   )	estimator)r\   r2  s     r   rq   MetaRegressor.__init__  s    "r   c                     [        U S40 UD6n[        U R                  5      R                  " X40 UR                  R                  D6U l        g Nr   )r   r   r2  r   
estimator_r\   r|   r}   
fit_paramsparamss        r   r   MetaRegressor.fit  s?     u;
;/33AQF<L<L<P<PQr   c                 t    [        U S9R                  U R                  [        5       R                  SSS9S9nU$ Nownerr   r/   r.   r2  method_mapping)r   addr2  r   r\   rH   s     r   get_metadata_routing"MetaRegressor.get_metadata_routing  s?    d+//nn(?..eE.J 0 
 r   )r2  r6  N	re   rf   rg   rh   r   rq   r   rD  ri   r   r   r   r0  r0    s    2#Rr   r0  c                   8    \ rS rSrSrS	S jrS	S jrS rS rSr	g)
WeightedMetaRegressori  z*A meta-regressor which is also a consumer.Nc                     Xl         X l        g r   r2  ro   r\   r2  ro   s      r   rq   WeightedMetaRegressor.__init__      " r   c                    U R                   b  U R                   R                  U 5        [        XS9  [        U S4SU0UD6n[	        U R
                  5      R                  " X40 UR
                  R                  D6U l        U $ Nr  r   rw   ro   r*   r2   r   r   r2  r   r6  )r\   r|   r}   rw   r8  r9  s         r   r   WeightedMetaRegressor.fit  sm    ==$MM  &: uXMXZX/33AQF<L<L<P<PQr   c                 ~    [        U S40 UD6nU R                  R                  " U40 UR                  R                  D6$ )Nr   )r   r6  r   r2  )r\   r|   predict_paramsr9  s       r   r   WeightedMetaRegressor.predict
  s9     yCNC&&qEF,<,<,D,DEEr   c                     [        U S9R                  U 5      R                  U R                  [	        5       R                  SSS9R                  SSS9S9nU$ )Nr=  r   r?  r   r@  r   add_self_requestrB  r2  r   rC  s     r   rD  *WeightedMetaRegressor.get_metadata_routing  sX    &d#S..,E%0Ii8	   	 r   r2  r6  ro   r   )
re   rf   rg   rh   r   rq   r   r   rD  ri   r   r   r   rH  rH    s    4!Fr   rH  c                   2    \ rS rSrSrSS jrSS jrS rSrg)	WeightedMetaClassifieri  zEA meta-estimator which also consumes sample_weight itself in ``fit``.Nc                     Xl         X l        g r   rJ  rK  s      r   rq   WeightedMetaClassifier.__init__  rM  r   c                    U R                   b  U R                   R                  U 5        [        XS9  [        U S4SU0UD6n[	        U R
                  5      R                  " X40 UR
                  R                  D6U l        U $ rO  rP  )r\   r|   r}   rw   r-   r9  s         r   r   WeightedMetaClassifier.fit#  sm    ==$MM  &: uTMTVT/33AQF<L<L<P<PQr   c                     [        U S9R                  U 5      R                  U R                  [	        5       R                  SSS9S9nU$ r<  rV  rC  s     r   rD  +WeightedMetaClassifier.get_metadata_routing,  sL    &d#S..,22%2N   	 r   rY  r   rF  r   r   r   r[  r[    s    O!	r   r[  c                   8    \ rS rSrSrS rS	S jrS	S jrS rSr	g)
MetaTransformeri8  zA simple meta-transformer.c                     Xl         g r   )transformer)r\   re  s     r   rq   MetaTransformer.__init__;  s    &r   Nc                     [        U S40 UD6n[        U R                  5      R                  " X40 UR                  R                  D6U l        U $ r5  )r   r   re  r   transformer_r7  s        r   r   MetaTransformer.fit>  sG     u;
;!$"2"2377W@R@R@V@VWr   c                 ~    [        U S40 UD6nU R                  R                  " U40 UR                  R                  D6$ )Nr   )r   rh  r   re  )r\   r|   r}   transform_paramsr9  s        r   r   MetaTransformer.transformC  s<     {G6FG  **1M0B0B0L0LMMr   c                     [        U S9R                  U R                  [        5       R                  SSS9R                  SSS9S9$ )Nr=  r   r?  r   )re  rA  )r   rB  re  r   rb   s    r   rD  $MetaTransformer.get_metadata_routingG  sI    D)--(((?SeS,SKS8	 . 
 	
r   )re  rh  r   )
re   rf   rg   rh   r   rq   r   r   rD  ri   r   r   r   rc  rc  8  s    $'
N
r   rc  )Tr   )9r#   collectionsr   	functoolsr   numpyr9   numpy.testingr   sklearn.baser   r   r   r	   r
   r   sklearn.metrics._scorerr   r   sklearn.model_selectionr   sklearn.model_selection._splitr   r    sklearn.utils._metadata_requestsr   sklearn.utils.metadata_routingr   r   r   sklearn.utils.multiclassr   r2   tuplerD   rz   rG   rW   r   rY   rk   r   r   r   r   r   r   r   r  r  r  r  r.  r0  rH  r[  rc  r   r   r   <module>r{     st    #   ,  @ 6 J 
 C0, ?Dg (V &oeL ::
 
+ +\ %_m  %F
NM 
Q/= Qh	N-@ 	N	R0C 	RR)< R /+] /d *K!g !+-? 60A: 
& $. D/- 8
(*:M 
r   