U
    ½mœdÎY  ã                   @   sv  d Z ddlZddlZddlZddlmZ ddlmZ ddlmZ ddlm	Z	 ddlm
Z
 ddlmZmZ dd	lmZ dd
lmZ ddlmZ ddlmZmZ ddlmZ ddlmZ ddlmZ ddlmZ ddlmZ ddlmZ ddlm Z m!Z! ddl"m#Z#m$Z$ ddl%m&Z& ddl'm(Z( ddl)m*Z* ej+ ,d¡Z-ddgddgddgddgddgddggZ.ddddddgZ/ddddddgZ0ddgddgddggZ1dddgZ2dddgZ3e* 4¡ Z5e- 6e5j7j8¡Z9e&e5j:e5j7e-d�\e5_:e5_7e* ;¡ Z<e&e<j:e<j7e-d�\e<_:e<_7dd „ Z=d!d"„ Z>ej? @d#d$d%g¡d&d'„ ƒZAd(d)„ ZBd*d+„ ZCej? @d,d-d.d/g¡d0d1„ ƒZDej? @d#d$d%g¡d2d3„ ƒZEd4d5„ ZFd6d7„ ZGd8d9„ ZHd:d;„ ZId<d=„ ZJd>d?„ ZKd@dA„ ZLdBdC„ ZMdDdE„ ZNdFdG„ ZOej? @d#d$d%g¡dHdI„ ƒZPdJdK„ ZQej? @d#d$d%g¡dLdM„ ƒZRej? @dNeƒ e5j:e5j7feƒ e<j:e<j7fg¡dOdP„ ƒZSdQdR„ ZTej? @dSee#fee$fg¡dTdU„ ƒZUej? @dVeeg¡dWdX„ ƒZVdYdZ„ ZWdS )[z6Testing for the boost module (sklearn.ensemble.boost).é    N)Ú
csc_matrix)Ú
csr_matrix)Ú
coo_matrix)Ú
dok_matrix)Ú
lil_matrix)Úassert_array_equalÚassert_array_less)Úassert_array_almost_equal)ÚBaseEstimator)Úclone)ÚDummyClassifierÚDummyRegressor)ÚLinearRegression)Útrain_test_split)ÚGridSearchCV)ÚAdaBoostClassifier)ÚAdaBoostRegressor)Ú_samme_proba)ÚSVCÚSVR)ÚDecisionTreeClassifierÚDecisionTreeRegressor)Úshuffle)ÚNoSampleWeightWrapper)Údatasetséþÿÿÿéÿÿÿÿé   é   Úfooé   ©Úrandom_statec                     sÔ   t  dddgdddgddd	gddd
gg¡‰ ˆ t  ˆ jdd�¡d d …t jf  ‰ G ‡ fdd„dƒ} | ƒ }t|dt  ˆ ¡ƒ}t|jˆ jƒ t  	|¡ 
¡ s˜t‚tt j|dd�ddddgƒ tt j|dd�ddddgƒ d S )Nr   g�íµ ÷Æ°>r   gR¸…ëQÈ?g333333ã?çš™™™™™É?iüÿÿgR¸…ëQà?g      à?g•Ö&è.>©Zaxisc                       s   e Zd Z‡ fdd„ZdS )z'test_samme_proba.<locals>.MockEstimatorc                    s   t |jˆ jƒ ˆ S ©N)r   Úshape©ÚselfÚX©Zprobs© úd/home/sam/Atlas/atlas_env/lib/python3.8/site-packages/sklearn/ensemble/tests/test_weight_boosting.pyÚpredict_probaC   s    z5test_samme_proba.<locals>.MockEstimator.predict_probaN)Ú__name__Ú
__module__Ú__qualname__r-   r+   r*   r+   r,   ÚMockEstimatorB   s   r1   r    r   )ÚnpÚarrayÚabsÚsumÚnewaxisr   Ú	ones_liker   r&   ÚisfiniteÚallÚAssertionErrorZargminÚargmax)r1   ZmockZsamme_probar+   r*   r,   Útest_samme_proba7   s    "ÿ$r<   c                  C   s>   t  ttƒ¡} tƒ  t| ¡}t| t¡t  ttƒdf¡ƒ d S )Nr   )r2   ZonesÚlenr)   r   Úfitr	   r-   )Zy_tÚclfr+   r+   r,   Útest_oneclass_adaboost_probaT   s    r@   Ú	algorithmÚSAMMEúSAMME.Rc                 C   sz   t | dd�}| tt¡ t| t¡tƒ tt 	t 
t¡¡|jƒ | t¡jttƒdfks\t‚| t¡jttƒfksvt‚d S )Nr   ©rA   r"   r   )r   r>   r)   Úy_classr   ÚpredictÚTÚ	y_t_classr2   ÚuniqueÚasarrayÚclasses_r-   r&   r=   r:   Údecision_function)rA   r?   r+   r+   r,   Útest_classification_toy]   s    rM   c                  C   s*   t dd�} |  tt¡ t|  t¡tƒ d S )Nr   r!   )r   r>   r)   Úy_regrr   rF   rG   Úy_t_regr©r?   r+   r+   r,   Útest_regression_toyh   s    
rQ   c                  C   s  t  tj¡} d  }}dD ]Ú}t|d�}| tjtj¡ t| |jƒ | 	tj¡}|dkr^|}|}|j
d t| ƒkstt‚| tj¡j
d t| ƒks’t‚| tjtj¡}|dksºtd||f ƒ‚t|jƒdksÌt‚ttdd„ |jD ƒƒƒt|jƒkst‚qd	|_td
t  | 	tj¡| ¡ƒ d S )N©rB   rC   ©rA   rB   r   gÍÌÌÌÌÌì?z'Failed with algorithm %s and score = %fc                 s   s   | ]}|j V  qd S r%   r!   ©Ú.0Zestr+   r+   r,   Ú	<genexpr>†   s     ztest_iris.<locals>.<genexpr>rC   r   )r2   rI   ÚirisÚtargetr   r>   Údatar   rK   r-   r&   r=   r:   rL   ÚscoreÚestimators_ÚsetrA   r   r4   )ÚclassesZ	clf_sammeZ
prob_sammeÚalgr?   ÚprobarZ   r+   r+   r,   Ú	test_iriso   s(    
ÿr`   ÚlossZlinearZsquareZexponentialc                 C   st   t | dd�}| tjtj¡ | tjtj¡}|dks8t‚t|jƒdksJt‚tt	dd„ |jD ƒƒƒt|jƒkspt‚d S )Nr   )ra   r"   gš™™™™™á?r   c                 s   s   | ]}|j V  qd S r%   r!   rT   r+   r+   r,   rV   œ   s     z test_diabetes.<locals>.<genexpr>)
r   r>   ÚdiabetesrY   rX   rZ   r:   r=   r[   r\   )ra   ÚregrZ   r+   r+   r,   Útest_diabetes‘   s    rd   c                 C   sÚ  t j d¡}|jdtjjd�}|jdtjjd�}t| dd�}|j	tj
tj|d� | tj
¡}dd„ | tj
¡D ƒ}| tj
¡}dd„ | tj
¡D ƒ}|jtj
tj|d�}	d	d„ |jtj
tj|d�D ƒ}
t|ƒdksÖt‚t||d
 ƒ t|ƒdksôt‚t||d
 ƒ t|
ƒdk�st‚t|	|
d
 ƒ tddd�}|j	tj
tj|d� | tj
¡}dd„ | tj
¡D ƒ}|jtj
tj|d�}	dd„ |jtj
tj|d�D ƒ}
t|ƒdk�s¨t‚t||d
 ƒ t|
ƒdk�sÈt‚t|	|
d
 ƒ d S )Nr   é
   ©Úsize)rA   Ún_estimators©Úsample_weightc                 S   s   g | ]}|‘qS r+   r+   ©rU   Úpr+   r+   r,   Ú
<listcomp>ª   s     z'test_staged_predict.<locals>.<listcomp>c                 S   s   g | ]}|‘qS r+   r+   rk   r+   r+   r,   rm   ¬   s     c                 S   s   g | ]}|‘qS r+   r+   ©rU   Úsr+   r+   r,   rm   ®   s    r   )rh   r"   c                 S   s   g | ]}|‘qS r+   r+   rk   r+   r+   r,   rm   ¾   s     c                 S   s   g | ]}|‘qS r+   r+   rn   r+   r+   r,   rm   À   s   ÿ)r2   ÚrandomÚRandomStateÚrandintrW   rX   r&   rb   r   r>   rY   rF   Ústaged_predictr-   Ústaged_predict_probarZ   Ústaged_scorer=   r:   r	   r   )rA   ÚrngZiris_weightsZdiabetes_weightsr?   ZpredictionsZstaged_predictionsr_   Zstaged_probasrZ   Zstaged_scoresr+   r+   r,   Útest_staged_predictŸ   sF    ÿ  ÿþrw   c                  C   sh   t tƒ d�} ddddœ}t| |ƒ}| tjtj¡ ttƒ dd�} dddœ}t| |ƒ}| t	jt	j¡ d S )N)Ú	estimator)r   r   rR   )rh   Úestimator__max_depthrA   r   ©rx   r"   )rh   ry   )
r   r   r   r>   rW   rY   rX   r   r   rb   )ÚboostÚ
parametersr?   r+   r+   r,   Útest_gridsearchÍ   s    ý


r}   c                  C   sî   dd l } dD ]p}t|d�}| tjtj¡ | tjtj¡}|  |¡}|  |¡}t	|ƒ|j
ks`t‚| tjtj¡}||kst‚qtdd�}| tjtj¡ | tjtj¡}|  |¡}|  |¡}t	|ƒ|j
ksÎt‚| tjtj¡}||ksêt‚d S )Nr   rR   rS   r!   )Úpickler   r>   rW   rY   rX   rZ   ÚdumpsÚloadsÚtypeÚ	__class__r:   r   rb   )r~   r^   ÚobjrZ   ro   Úobj2Zscore2r+   r+   r,   Útest_pickleà   s$    





r…   c               	   C   s~   t jdddddddd�\} }dD ]X}t|d	�}| | |¡ |j}|jd dksRt‚|d d…tjf |dd … k 	¡ s t‚q d S )
NiÐ  re   r    r   Fr   )Ú	n_samplesÚ
n_featuresZn_informativeZn_redundantZ
n_repeatedr   r"   rR   rS   )
r   Zmake_classificationr   r>   Úfeature_importances_r&   r:   r2   r6   r9   )r)   Úyr^   r?   Zimportancesr+   r+   r,   Útest_importancesü   s    ù


rŠ   c               	   C   sF   t ƒ } t d¡}tjt|d�� | jttt	 
dg¡d� W 5 Q R X d S )Nz*sample_weight.shape == (1,), expected (6,)©Úmatchr   ri   )r   ÚreÚescapeÚpytestÚraisesÚ
ValueErrorr>   r)   rE   r2   rJ   )r?   Úmsgr+   r+   r,   Ú,test_adaboost_classifier_sample_weight_error  s    
r“   c               	   C   sÜ   ddl m}  t| ƒ ƒ}| tt¡ ttƒ dd�}| tt¡ ddl m} t	|ƒ dd�}| tt¡ t	t
ƒ dd�}| tt¡ ddgddgddgddgg}dd	dd
g}ttƒ dd�}tjtdd�� | ||¡ W 5 Q R X d S )Nr   )ÚRandomForestClassifierrB   rS   )ÚRandomForestRegressorr!   r   r   Úbarr   zworse than randomr‹   )Úsklearn.ensembler”   r   r>   r)   rN   r   rE   r•   r   r   r�   r�   r‘   )r”   r?   r•   ZX_failZy_failr+   r+   r,   Útest_estimator  s    
r˜   c               	   C   s@   d} t dddd�}tjt| d�� | tjtj¡ W 5 Q R X d S )Nz+Sample weights have reached infinite valuesé   g      7@rB   )rh   Zlearning_raterA   r‹   )r   r�   ÚwarnsÚUserWarningr>   rW   rY   rX   )r’   r?   r+   r+   r,   Útest_sample_weights_infinite6  s    rœ   c                  C   s<  G dd„ dt ƒ} tjddddd�\}}t |¡}t||dd	�\}}}}tttt	t
fD �]à}||ƒ}||ƒ}	t| d
d�ddd� ||¡}
t| d
d�ddd� ||¡}|
 |	¡}| |¡}t||ƒ |
 |	¡}| |¡}t||ƒ |
 |	¡}| |¡}t||ƒ |
 |	¡}| |¡}t||ƒ |
 |	|¡}| ||¡}t||ƒ |
 |	¡}| |¡}t||ƒD ]\}}t||ƒ �qZ|
 |	¡}| |¡}t||ƒD ]\}}t||ƒ �qŽ|
 |	¡}| |¡}t||ƒD ]\}}t||ƒ �qÂ|
 |	|¡}| ||¡}t||ƒD ]\}}t||ƒ �qúdd„ |
jD ƒ}tdd„ |D ƒƒsTt‚qTd S )Nc                       s"   e Zd ZdZd‡ fdd„	Z‡  ZS )z-test_sparse_classification.<locals>.CustomSVCz8SVC variant that records the nature of the training set.Nc                    s    t ƒ j|||d� t|ƒ| _| S ©z<Modification on fit caries data type for later verification.ri   ©Úsuperr>   r�   Ú
data_type_©r(   r)   r‰   rj   ©r‚   r+   r,   r>   C  s    
z1test_sparse_classification.<locals>.CustomSVC.fit)N©r.   r/   r0   Ú__doc__r>   Ú__classcell__r+   r+   r¢   r,   Ú	CustomSVC@  s   r¦   r   é   é   é*   )Z	n_classesr†   r‡   r"   r   r!   T)ZprobabilityrB   )rx   r"   rA   c                 S   s   g | ]
}|j ‘qS r+   ©r    ©rU   Úir+   r+   r,   rm   •  s     z.test_sparse_classification.<locals>.<listcomp>c                 S   s   g | ]}|t kp|tk‘qS r+   ©r   r   ©rU   Útr+   r+   r,   rm   —  s     )r   r   Zmake_multilabel_classificationr2   Zravelr   r   r   r   r   r   r   r>   rF   r   rL   r	   Zpredict_log_probar-   rZ   Zstaged_decision_functionÚziprs   rt   ru   r[   r9   r:   )r¦   r)   r‰   ÚX_trainÚX_testÚy_trainÚy_testÚsparse_formatÚX_train_sparseÚX_test_sparseÚsparse_classifierÚdense_classifierÚsparse_resultsÚdense_resultsÚ
sprase_resÚ	dense_resÚtypesr+   r+   r,   Útest_sparse_classification=  sz    	   ÿ

ý üý ü


















r¿   c                  C   s
  G dd„ dt ƒ} tjddddd�\}}t||dd	�\}}}}tttttfD ]º}||ƒ}||ƒ}	t	| ƒ dd
� 
||¡}
t	| ƒ dd
� 
||¡ }}|
 |	¡}| |¡}t||ƒ |
 |	¡}| |¡}t||ƒD ]\}}t||ƒ qÊdd„ |
jD ƒ}tdd„ |D ƒƒsJt‚qJd S )Nc                       s"   e Zd ZdZd‡ fdd„	Z‡  ZS )z)test_sparse_regression.<locals>.CustomSVRz8SVR variant that records the nature of the training set.Nc                    s    t ƒ j|||d� t|ƒ| _| S r�   rž   r¡   r¢   r+   r,   r>      s    
z-test_sparse_regression.<locals>.CustomSVR.fit)Nr£   r+   r+   r¢   r,   Ú	CustomSVR�  s   rÀ   r§   é2   r   r©   )r†   r‡   Ú	n_targetsr"   r   r!   rz   c                 S   s   g | ]
}|j ‘qS r+   rª   r«   r+   r+   r,   rm   Å  s     z*test_sparse_regression.<locals>.<listcomp>c                 S   s   g | ]}|t kp|tk‘qS r+   r­   r®   r+   r+   r,   rm   Ç  s     )r   r   Zmake_regressionr   r   r   r   r   r   r   r>   rF   r	   rs   r°   r[   r9   r:   )rÀ   r)   r‰   r±   r²   r³   r´   rµ   r¶   r·   r¸   r¹   r»   rº   r¼   r½   r¾   r+   r+   r,   Útest_sparse_regressionš  sD    	   ÿ
 ÿ þ ÿ þ




rÃ   c                  C   sF   G dd„ dt ƒ} t| ƒ dd�}| tt¡ t|jƒt|jƒksBt‚dS )z·
    AdaBoostRegressor should work without sample_weights in the base estimator
    The random weighted sampling is done internally in the _boost method in
    AdaBoostRegressor.
    c                   @   s   e Zd Zdd„ Zdd„ ZdS )z=test_sample_weight_adaboost_regressor.<locals>.DummyEstimatorc                 S   s   d S r%   r+   )r(   r)   r‰   r+   r+   r,   r>   Ò  s    zAtest_sample_weight_adaboost_regressor.<locals>.DummyEstimator.fitc                 S   s   t  |jd ¡S )Nr   )r2   Zzerosr&   r'   r+   r+   r,   rF   Õ  s    zEtest_sample_weight_adaboost_regressor.<locals>.DummyEstimator.predictN)r.   r/   r0   r>   rF   r+   r+   r+   r,   ÚDummyEstimatorÑ  s   rÄ   r    )rh   N)	r
   r   r>   r)   rN   r=   Zestimator_weights_Zestimator_errors_r:   )rÄ   r{   r+   r+   r,   Ú%test_sample_weight_adaboost_regressorÊ  s    rÅ   c                  C   s†   t j d¡} |  ddd¡}|  ddgd¡}|  d¡}ttdd�ƒ}| ||¡ | |¡ | 	|¡ t
tƒ ƒ}| ||¡ | |¡ dS )zX
    Check that the AdaBoost estimators can work with n-dimensional
    data matrix
    r   rÁ   r    r   Zmost_frequent)ZstrategyN)r2   rp   rq   ZrandnÚchoicer   r   r>   rF   r-   r   r   )rv   r)   ZycÚyrr{   r+   r+   r,   Útest_multidimensional_XÝ  s    



rÈ   c              	   C   s\   t jt j }}ttƒ ƒ}t|| d�}d |jj¡}t	j
t|d�� | ||¡ W 5 Q R X d S )N)rx   rA   z {} doesn't support sample_weightr‹   )rW   rY   rX   r   r   r   Úformatr‚   r.   r�   r�   r‘   r>   )rA   r)   r‰   rx   r?   Úerr_msgr+   r+   r,   Ú-test_adaboostclassifier_without_sample_weightò  s    
rË   c            
      C   sR  t j d¡} t jdddd�}d| d |  |jd ¡d  }| d	d
¡}|d	  d9  < d|d	< ttƒ d
dd�}t	|ƒ}t	|ƒ}| 
||¡ | 
|d d	… |d d	… ¡ t  |¡}d|d	< |j
|||d� | |d d	… |d d	… ¡}| |d d	… |d d	… ¡}| |d d	… |d d	… ¡}	||k �s,t‚||	k �s:t‚|t |	¡k�sNt‚d S )Nr©   r   éd   éè  )Únumgš™™™™™é?r#   g-Cëâ6?r   r   re   i'  ©rx   rh   r"   ri   )r2   rp   rq   ZlinspaceZrandr&   Zreshaper   r   r   r>   r7   rZ   r:   r�   Zapprox)
rv   r)   r‰   Zregr_no_outlierZregr_with_weightZregr_with_outlierrj   Zscore_with_outlierZscore_no_outlierZscore_with_weightr+   r+   r,   Ú$test_adaboostregressor_sample_weightü  s0       ÿ
rÐ   c                 C   sZ   t tjdd�ddiŽ\}}}}t| dd�}| ||¡ ttj| |¡dd�| 	|¡ƒ d S )NT)Z
return_X_yr"   r©   rD   r   r$   )
r   r   Zload_digitsr   r>   r   r2   r;   r-   rF   )rA   r±   r²   r³   r´   Úmodelr+   r+   r,   Ú test_adaboost_consistent_predict"  s    
ÿÿ ÿrÒ   zmodel, X, yc              	   C   sD   t  |¡}d|d< d}tjt|d�� | j|||d� W 5 Q R X d S )Niöÿÿÿr   z1Negative values in data passed to `sample_weight`r‹   ri   )r2   r7   r�   r�   r‘   r>   )rÑ   r)   r‰   rj   rÊ   r+   r+   r,   Ú#test_adaboost_negative_weight_error2  s
    
rÓ   c                  C   s~   t j d¡} | jdd�}| jddgdd�}t  |¡d }tdd	d
�}t|dd	d�}|j|||d� t  	|j
¡ ¡ dkszt‚dS )z¸Check that we don't create NaN feature importance with numerically
    instable inputs.

    Non-regression test for:
    https://github.com/scikit-learn/scikit-learn/issues/20320
    r©   )rÍ   re   rf   r   r   rÍ   gtDíS 'T	re   é   )Ú	max_depthr"   é   rÏ   ri   N)r2   rp   rq   ÚnormalrÆ   r7   r   r   r>   Úisnanrˆ   r5   r:   )rv   r)   r‰   rj   ÚtreeZ	ada_modelr+   r+   r,   ÚFtest_adaboost_numerically_stable_feature_importance_with_small_weightsB  s    rÚ   zAdaBoost, Estimatorc              	   C   s^   t  ddgddgg¡}t  ddg¡}| |ƒ d�}d}tjt|d�� | ||¡ W 5 Q R X d S )	Nr   r   r    é   r   )Zbase_estimatorzV`base_estimator` was renamed to `estimator` in version 1.2 and will be removed in 1.4.r‹   )r2   r3   r�   rš   ÚFutureWarningr>   )ÚAdaBoostZ	Estimatorr)   r‰   rÑ   Úwarn_msgr+   r+   r,   Ú'test_base_estimator_argument_deprecatedT  s    ÿrß   rÝ   c              	   C   s^   t  ddgddgg¡}t  ddg¡}| ƒ }| ||¡ d}tjt|d�� |j W 5 Q R X d S )Nr   r   r    rÛ   r   zoAttribute `base_estimator_` was deprecated in version 1.2 and will be removed in 1.4. Use `estimator_` instead.r‹   )r2   r3   r>   r�   rš   rÜ   Zbase_estimator_)rÝ   r)   r‰   rÑ   rÞ   r+   r+   r,   Ú'test_base_estimator_property_deprecatedi  s    ÿrà   c               	   C   s4   t tƒ ƒ} tjtdd�� | jdd� W 5 Q R X dS )zåCheck that setting base_estimator parameters works.

    During the deprecation cycle setting "base_estimator__*" params should
    work.

    Non-regression test for https://github.com/scikit-learn/scikit-learn/issues/25470
    zParameter 'base_estimator' ofr‹   r   )Zbase_estimator__max_depthN)r   r   r�   rš   rÜ   Z
set_paramsrP   r+   r+   r,   Ú4test_deprecated_base_estimator_parameters_can_be_set|  s    
rá   )Xr¤   Únumpyr2   r�   r�   Zscipy.sparser   r   r   r   r   Zsklearn.utils._testingr   r   r	   Zsklearn.baser
   r   Zsklearn.dummyr   r   Zsklearn.linear_modelr   Zsklearn.model_selectionr   r   r—   r   r   Z!sklearn.ensemble._weight_boostingr   Zsklearn.svmr   r   Zsklearn.treer   r   Zsklearn.utilsr   Zsklearn.utils._mockingr   Zsklearnr   rp   rq   rv   r)   rE   rN   rG   rH   rO   Z	load_irisrW   ZpermutationrX   rg   ÚpermrY   Zload_diabetesrb   r<   r@   ÚmarkZparametrizerM   rQ   r`   rd   rw   r}   r…   rŠ   r“   r˜   rœ   r¿   rÃ   rÅ   rÈ   rË   rÐ   rÒ   rÓ   rÚ   rß   rà   rá   r+   r+   r+   r,   Ú<module>   s¬   (

  ÿ	

"

-]0
	&
þþ
	þþ
þ
