o
    Ú­jî  ã                   @   s  d Z ddlmZmZmZ ddlZddlZddlm	Z	 ddl
mZ ddlmZmZmZ dd	lmZmZ dd
lmZ ddlmZmZ dededejdejdee ddfdd„Zdedededejdejdee ddfdd„Zdededdfdd„Zdeddfdd„Zdeddfdd„ZdS )z$Tests for compatiblity with sklearn.é    )ÚCallableÚOptionalÚTypeN)Úsoftmaxé   )ÚDMatrix)ÚXGBClassifierÚXGBRegressorÚXGBRFRegressoré   )Úget_california_housingÚmake_batches)Úmake_recoded)ÚDeviceÚassert_allcloseÚtree_methodÚdeviceÚXÚyÚas_frameÚreturnc                 C   sª   t ddd| |d�}|j||d� |j|dd�}|dur||ƒ}t ddd| |d�}|j|||d	� |j||d
�}t ddd| |d�}	|	j||d� |	 |¡}
tj ||
¡ dS )z•
    Parameters
    ----------

    as_frame: A callable function to convert margin into DataFrame, useful for different
    df implementations.
    ç333333Ó?r   é   ©Úlearning_rateÚrandom_stateÚn_estimatorsr   r   ©r   r   T©Úoutput_marginN©r   r   Úbase_margin©r!   é   )r   ÚfitÚpredictÚnpÚtestingr   )r   r   r   r   r   Úmodel_0ÚmarginÚmodel_1Úpredictions_1Úcls_2Úpredictions_2© r.   úU/var/www/html/CropPilot/venv/lib/python3.10/site-packages/xgboost/testing/with_skl.pyÚ run_boost_from_prediction_binary   s<   ûûû
r0   Ú	estimatorc                 C   sê   | ddd||d�}|j ||d� | ¡ j|dd�}|dur!||ƒ}| ddd||d�}|j |||d	� | ¡ jt||d
�dd�}	| ddd||d�}
|
j ||d� |
 ¡ j|dd�}t|	dƒra|	 ¡ }	t|dƒrj| ¡ }tjj	|	|dd� dS )z.Boosting from prediction with multi-class clf.r   r   r   r   r   r)   )Úpredict_typeNr    r"   Tr   r#   Úgetç�íµ ÷Æ°>©Úatol)
r$   Úget_boosterÚinplace_predictr%   r   Úhasattrr3   r&   r'   r   )r1   r   r   r   r   r   r(   r)   r*   r+   Úmodel_2r-   r.   r.   r/   Ú&run_boost_from_prediction_multi_clasasB   sH   
ûûÿû

r;   c                 C   sê   ddl m} ddlm} tƒ \}}tj d¡}|dd|d�}| ||¡D ]'\}}	t	d| |d	� 
|| || ¡}
|
 ||	 ¡}||	 }|||ƒd
k sKJ ‚q$t	|d�}t t¡� |jdd� | 
||¡ W d  ƒ dS 1 snw   Y  dS )z"Testwith the cali housing dataset.r   )Úmean_squared_error)ÚKFoldéÊ  r   T)Ún_splitsÚshuffler   é*   )r   r   r   é#   ©r   é
   )Úearly_stopping_roundsN)Úsklearn.metricsr<   Úsklearn.model_selectionr=   r   r&   ÚrandomÚRandomStateÚsplitr
   r$   r%   ÚpytestÚraisesÚNotImplementedErrorÚ
set_params)r   r   r<   r=   r   r   ÚrngÚkfÚtrain_indexÚ
test_indexÚ	xgb_modelÚpredsÚlabelsÚrfregr.   r.   r/   Úrun_housing_rf_regressionu   s&   
ÿþ
"þrW   c           
      C   s>  t | dd�\}}}}}tdd| d�}|j||||fgd� | ¡ }| ¡ }| ¡  ¡ r-J ‚tdd| d�}|j|||||fgd� | ¡ }| ¡ }| ¡ dksPJ ‚| ¡  ¡ rXJ ‚tdd| d�}|j||||fgd� | ¡ }	tj	 
|	d	 d
 |d	 d
 |d	 d
  ¡ tj	 
| |¡| |¡¡ tj	 
| |¡| |¡¡ dS )z)Test re-coding for training continuation.é   )Ú
n_featuresTr   )Úenable_categoricalr   r   )Úeval_set)rS   r[   r   Úvalidation_0ÚrmseN)r   r	   r$   Úevals_resultr7   Úget_categoriesÚemptyÚnum_boosted_roundsr&   r'   r   r%   Úapply)
r   ÚencÚreencr   Ú_ÚregÚ	results_0ÚboosterÚ	results_1Ú	results_2r.   r.   r/   Úrun_recoding‹   s*   
þrk   c                 C   s,  ddl m}m} dd„ tddddd	�D ƒ\}}}t| d
�}|j|||d� |j}|jtj	ks0J ‚|d dk s8J ‚td| d�}|j|||d� |j}t
|tjƒsQJ ‚|jtj	ksYJ ‚|d dk saJ ‚d}|ddd|ddd�\}}tdd| d�}	|	 ||¡ |	j}t
|tjƒs‡J ‚t|ƒdks�J ‚t|ƒdk ¡ s™J ‚tjjt|ƒddd� tj tt|ƒƒd¡ tj|tj	d�| }
| dkrÆddl}| |
¡}
td|
d�}	|	 ||¡ t| |
|	jƒ |ddd|d�\}}tj|tj	d�d  }
| dkrúddl}| |
¡}
t|
d!�}	|	 ||¡ t| |
|	jƒ |	jd"k�sJ ‚dS )#zTests for the intercept.r   )Úmake_classificationÚmake_multilabel_classificationc                 S   s   g | ]}|d  ‘qS )r   r.   )Ú.0Úvr.   r.   r/   Ú
<listcomp>®   s    z!run_intercept.<locals>.<listcomp>é   é   r   F)Úuse_cupyrC   )Úsample_weightg      à?Úgblinear)rh   r   r   r>   é€   rX   )r   Ú	n_samplesrY   Ú	n_classesÚn_informativeÚn_redundantÚgbtreezmulti:softprob)rh   Ú	objectiver   g        r4   r5   g      ð?)ÚshapeÚdtypeÚcudaN)r|   Ú
base_score)r   rw   rY   rx   r   )r€   zbinary:logistic)Úsklearn.datasetsrl   rm   r   r	   r$   Ú
intercept_r~   r&   Úfloat32Ú
isinstanceÚndarrayr   Úlenr   Úallr'   r   ÚsumÚonesÚcupyÚarrayr|   )r   rl   rm   r   r   Úwrf   Úresultrx   ÚclfÚ	interceptÚcpr.   r.   r/   Úrun_interceptª   s`    

ú	

ÿ

r‘   )Ú__doc__Útypingr   r   r   Únumpyr&   rK   Úscipy.specialr   Úcorer   Úsklearnr   r	   r
   Údatar   r   Úordinalr   Úutilsr   r   Ústrr…   r0   r;   rW   rk   r‘   r.   r.   r.   r/   Ú<module>   sR   ÿþýüû
ú1ÿþýüûú
ù3