o
    Ù­jª2 ã                   @   sf  U d Z ddl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mZmZmZmZmZmZmZmZmZmZmZmZmZmZmZm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/m0Z0m1Z1 ddl2m3Z3 ddl*m4Z5 ddl*m6Z7 ddl*m8Z9 ddl:m;Z;m<Z<m=Z= ddl>m?Z?m@Z@mAZAmBZBmCZCmDZDmEZEmFZF ddlGmHZHmIZImJZJmKZKmLZLmMZMmNZNmOZOmPZPmQZQmRZRmSZS ddlTmUZU ddlVmWZX ddlYmZZZm[Z[ ddl\m]Z]m^Z^m_Z_m`Z` ee!jae%jbe%jcf Zdeeed< ee!jae%jbf Zfeeed< ede?ed œƒZgg d!¢Zhe id"¡Zjd#ekd$eeeel eeelekf  f  d%eek d&e7fd'd(„Zmd#ekd)eel d*eeelekf  d%eek d&e7f
d+d,„ZnG d-d.„ d.e*joƒZod/ed0 d&d0fd1d2„ZpG d3d4„ d4ƒZqed5ƒZred6ƒZsd/ed0 d7eeserf d8ed9eel d&erf
d:d;„ZtG d<d=„ d=eqƒZu		d©d/d0d#ekd>eeelef  d?ee5 d&eeleelekf f f
d@dA„Zvd&eeelef  fdBdC„ZwdDeqdEeeeeqelf   d&eel fdFdG„Zxd9eel d/d0d&dfdHdI„Zyd/d0dJeelef d>eeelef  dKeelef dDeqdLekdEeeeeqelf   dMeeB dNeek dOeekezf dPee? dQeee3  dReeA d?ee5 d&eeg fdSdT„Z{eE	UdªdddddVddddWœd/d0dKeelef dDeqdLekdEeeeeqelf   dMeeB dNeek dPee? dOeekezf dQeee3  dReeA d?ee5 d&efdXdY„ƒZWdZezd[ed&ezfd\d]„Z|d^ed_ed`eek dZezd&ef
dadb„Z}dceddded^efdfeed d[eekdgf dheekelf d&edfdidj„Z~dde?dkekdZezdlezdmed&eeekdgf eekelf f fdndo„Zd/d0dpee?edef d&defdqdr„Z€d/d0dJeelef dpee?edef d^efdsezdte�duezdvezdwezdxezdyezdze1d{ezd&edfd|d}„Z‚eEd~ejƒd~d~d~d~dVdd~d€œ	d/ed0 dpeege?def d^eeqeff dsezdte�duezdvezdwezdxezdyezdze1d{ezd&efd�d‚„ƒZ„d/d0dJeelef dpee?edef d^efdze1dƒeldte�dyezdfeed d{ezd&edfd„d…„Z…eEdd†ejƒdVdd~d‡œd/ed0 dpeege?def d^efdze1dƒeldte�dyezdfeed d{ezd&efdˆd‰„ƒZ†d/ed0 dŠeel d‹eel dŒeek dmed&eeqeeeeqelf   f fd�dŽ„Z‡edpd�d/d0d&efd�d‘„ƒZˆG d’d�„ d�eJƒZ‰eSd“d”dpgƒG d•d–„ d–eMe‰ƒƒZŠeSd—d”dpgƒG d˜d™„ d™eIe‰ƒƒZ‹eSdšd”dpgd›dœd��G dždŸ„ dŸeLe‰ƒƒZŒeSd dpd¡gd¢d£�G d¤d¥„ d¥eŠƒƒZ�eSd¦dpd¡gd¢d£�G d§d¨„ d¨e‹ƒƒZŽdS )«aS  
Dask extensions for distributed training
----------------------------------------

See :doc:`Distributed XGBoost with Dask </tutorials/dask>` for simple tutorial.  Also
:doc:`/python/dask-examples/index` for some examples.

There are two sets of APIs in this module, one is the functional API including
``train`` and ``predict`` methods.  Another is stateful Scikit-Learner wrapper
inherited from single-node Scikit-Learn interface.

The implementation is heavily influenced by dask_xgboost:
https://github.com/dask/dask-xgboost

Optional dask configuration
===========================

- **coll_cfg**:
    Specify the scheduler address along with communicator configurations. This can be
    used as a replacement of the existing global Dask configuration
    `xgboost.scheduler_address` (see below). See :ref:`tracker-ip` for more info. The
    `tracker_host_ip` should specify the IP address of the Dask scheduler node.

  .. versionadded:: 3.0.0

  .. code-block:: python

    from xgboost import dask as dxgb
    from xgboost.collective import Config

    coll_cfg = Config(
        retry=1, timeout=20, tracker_host_ip="10.23.170.98", tracker_port=0
    )

    clf = dxgb.DaskXGBClassifier(coll_cfg=coll_cfg)
    # or
    dxgb.train(client, {}, Xy, num_boost_round=10, coll_cfg=coll_cfg)

- **xgboost.scheduler_address**: Specify the scheduler address

  .. versionadded:: 1.6.0

  .. deprecated:: 3.0.0

  .. code-block:: python

      dask.config.set({"xgboost.scheduler_address": "192.0.0.100"})
      # We can also specify the port.
      dask.config.set({"xgboost.scheduler_address": "192.0.0.100:12345"})

é    N)Údefaultdict)Úcontextmanager)ÚpartialÚupdate_wrapper)ÚThread)ÚAnyÚ	AwaitableÚCallableÚDictÚ	GeneratorÚIterableÚListÚOptionalÚ	ParamSpecÚSequenceÚSetÚTupleÚ	TypeAliasÚ	TypedDictÚ	TypeGuardÚTypeVarÚUnion)Úarray)Úbag)Ú	dataframe)ÚDelayed)ÚFutureé   )Ú
collectiveÚconfig)Ú
Categories)ÚFeatureNamesÚFeatureTypesÚIterationRange)ÚTrainingCallback)ÚConfig)Ú_Args)Ú_ArgVals)Ú_is_cudf_dfÚ_is_cudf_serÚ_is_cupy_alike)ÚBoosterÚDMatrixÚMetricÚPlainObjÚXGBoostErrorÚ_check_distributed_paramsÚ_deprecate_positional_argsÚ_expect)ÚXGBClassifierÚXGBClassifierBaseÚXGBModelÚ	XGBRankerÚXGBRankerMixInÚXGBRegressorBaseÚ_can_use_qdmÚ_check_rf_callbackÚ_cls_predict_probaÚ_objective_decoratorÚ_wrap_evaluation_matricesÚxgboost_model_doc)ÚRabitTracker)Útrainé   )Ú_get_dmatricesÚno_group_split)Ú_DASK_2024_12_1Ú_DASK_2025_3_0Úget_address_from_userÚget_n_threadsÚ_DaskCollectionÚ_DataTÚTrainReturnT©ÚboosterÚhistory)ÚCommunicatorContextÚDaskDMatrixÚDaskQuantileDMatrixÚDaskXGBRegressorÚDaskXGBClassifierÚDaskXGBRankerÚDaskXGBRFRegressorÚDaskXGBRFClassifierr@   ÚpredictÚinplace_predictz[xgboost.dask]Ú	n_workersÚaddrsÚtimeoutÚreturnc           
   
   C   s(  i }z[t |d tƒr&|d d }|d d }t| ||d|d u r!dn|d�}n|d }t |tƒs5|d u s5J ‚t| |d|d u r?dn|d�}| ¡  t|jd�}d|_| ¡  | | 	¡ ¡ W |S  t
y“ }	 z*t|ƒdk rl‚ t d	t|d ƒt|d ƒt|	ƒ¡ t| |dd … |ƒ}W Y d }	~	|S d }	~	ww )
Nr   rA   Útask)rX   Úhost_ipÚportÚsortbyrZ   )rX   r]   r_   rZ   )ÚtargetTr   zCFailed to bind address '%s', trying to use '%s' instead. Error:
 %s)Ú
isinstanceÚtupler?   ÚstrÚstartr   Úwait_forÚdaemonÚupdateÚworker_argsr/   ÚlenÚLOGGERÚwarningÚ_try_start_tracker)
rX   rY   rZ   Úenvr]   r^   Úrabit_trackerÚaddrÚthreadÚe© rr   úR/var/www/html/CropPilot/venv/lib/python3.10/site-packages/xgboost/dask/__init__.pyrl   ¯   sN   ûüõ

ü€õrl   Úaddr_from_daskÚaddr_from_userc                 C   s   t | ||g|ƒ}|S )z8Start Rabit tracker, recurse to try different addresses.)rl   )rX   rt   ru   rZ   rm   rr   rr   rs   Ú_start_trackerß   s   rv   c                       s*   e Zd ZdZdeddf‡ fdd„Z‡  ZS )rN   zNA context controlling collective communicator initialization and finalization.Úargsr[   Nc                    s8   t ƒ jdi |¤Ž t ¡ }d|j› d|j› �| jd< d S )Nz[xgboost.dask-z]:ÚDMLC_TASK_IDrr   )ÚsuperÚ__init__ÚdistributedÚ
get_workerÚnameÚaddressrw   )Úselfrw   Úworker©Ú	__class__rr   rs   rz   í   s   zCommunicatorContext.__init__)Ú__name__Ú
__module__Ú__qualname__Ú__doc__ÚCollArgsValsrz   Ú__classcell__rr   rr   r�   rs   rN   ê   s    rN   Úclientúdistributed.Clientc                 C   sX   t | tt ¡ ƒtdƒfƒstttt ¡ ƒtdƒgt| ƒƒƒ‚| du r(t ¡ }|S | }|S )z#Simple wrapper around testing None.N)ra   Útyper{   Ú
get_clientÚ	TypeErrorr2   )r‰   Úretrr   rr   rs   Ú_get_clientø   s   ÿÿr�   c                #   @   sH  e Zd ZdZe	d$dddddddddddddœded dedee d	ee d
ee dee de	dee
 dee dee dee dee dee dee de	ddf dd„ƒZded fdd„Zddddddddœdddedee dee d
ee dee dee dee dee dd fdd„Zdedeeef fd d!„Zdefd"d#„ZdS )%rO   aq  DMatrix holding on references to Dask DataFrame or Dask Array.  Constructing a
    `DaskDMatrix` forces all lazy computation to be carried out.  Wait for the input
    data explicitly if you want to see actual computation of constructing `DaskDMatrix`.

    See doc for :py:obj:`xgboost.DMatrix` constructor for other parameters.  DaskDMatrix
    accepts only dask collection.

    .. note::

        `DaskDMatrix` does not repartition or move data between workers.  It's the
        caller's responsibility to balance the data.

    .. note::

        For aligning partitions with ranking query groups, use the
        :py:class:`DaskXGBRanker` and its ``allow_group_split`` option.

    .. versionadded:: 1.0.0

    Parameters
    ----------
    client :
        Specify the dask client used for training.  Use default client returned from
        dask if it's set to None.

    NF)ÚweightÚbase_marginÚmissingÚsilentÚfeature_namesÚfeature_typesÚgroupÚqidÚlabel_lower_boundÚlabel_upper_boundÚfeature_weightsÚenable_categoricalr‰   rŠ   ÚdataÚlabelr�   r‘   r’   r“   r”   r•   r–   r—   r˜   r™   rš   r›   r[   c                C   s>  t |ƒ}|| _|	| _t|	tƒrtdƒ‚|d ur|ntj| _|| _	|d ur,|d ur,t
dƒ‚|
d ur4t
dƒ‚t|jƒdkrCtd|j› �ƒ‚t|tjtjfƒsYtttjtjft|ƒƒƒ‚t|tjtjtjtd ƒfƒsvtttjtjtjft|ƒƒƒ‚|jd | _t| jtƒs„J ‚ttƒ| _d| _|j| j|||||||||d�
| _d S )	NzFThe Dask interface can handle categories from DataFrame automatically.z$per-group weight is not implemented.z4group structure is not implemented, use qid instead.r   z$Expecting 2 dimensional input, got: rA   F)	r‰   rœ   r�   Úweightsr‘   r—   rš   r˜   r™   )r�   r”   r•   ra   r    r�   ÚnumpyÚnanr’   r›   ÚNotImplementedErrorri   ÚshapeÚ
ValueErrorÚddÚ	DataFrameÚdaÚArrayr2   r‹   ÚSeriesÚ_n_colsÚintr   ÚlistÚ
worker_mapÚis_quantileÚsyncÚ_map_local_dataÚ_init)r   r‰   rœ   r�   r�   r‘   r’   r“   r”   r•   r–   r—   r˜   r™   rš   r›   rr   rr   rs   rz   &  sJ   
ÿÿ
özDaskDMatrix.__init__)NNrO   c                 C   s
   | j  ¡ S ©N)r°   Ú	__await__©r   rr   rr   rs   r²   f  s   
zDaskDMatrix.__await__)r�   rž   r‘   r—   rš   r˜   r™   rž   c       	      
   ƒ   sH  �dt t dtdt t dtdtf
dd„‰dtdt t f‡fd	d
„‰dtt dtt t  f‡fdd„}
ˆ|ƒ‰ |
|ƒ}|
|ƒ}|
|ƒ}|
|ƒ}|
|ƒ}|
|	ƒ}dˆ i‰dtt t  dtddf‡ ‡‡fdd„}||dƒ ||dƒ ||dƒ ||dƒ ||dƒ ||dƒ g }ttˆ ƒƒD ]}i }ˆ ¡ D ]
\}}|| ||< q“| 	|¡ q‹t
ttj|ƒƒ}ˆ |¡}t |¡I dH  |D ]}|jdksÇJ |jƒ‚q»i | _t|ƒD ]
\}}|| j|j< qÏdd„ |D ƒ}ˆjjdd„ |D ƒd�I dH }tt
ƒ}| ¡ D ]\}}|tt|ƒƒ  	|| ¡ qø|| _|du �rd| _| S ˆ |¡ ¡ I dH | _| S ) z Obtain references to local data.ÚleftÚ	left_nameÚrightÚ
right_namer[   c              	   S   s*   d|› d|› dt | ƒ› dt |ƒ› d�	}|S )NzPartitions between z and z are not consistent: z != z/.  Please try to repartition/rechunk your data.©ri   )r´   rµ   r¶   r·   Úmsgrr   rr   rs   Úinconsistentx  s   ÿÿÿz1DaskDMatrix._map_local_data.<locals>.inconsistentÚdc                    sH   ˆ   | ¡} t| jdƒrt| jjƒdkr| jjd dkrtdƒ‚ˆ  | ¡S )zBreaking data into partitions.r¢   rA   z�Data should be partitioned by row. To avoid this specify the number of columns for your dask Array explicitly. e.g. chunks=(partition_size, -1]))ÚpersistÚhasattrÚ
partitionsri   r¢   r£   Ú
futures_of)r»   ©r‰   rr   rs   Ú
to_futures‚  s   

ÿÿ
z/DaskDMatrix._map_local_data.<locals>.to_futuresÚmetac                    s   | d ur
ˆ | ƒ}|S d S r±   rr   )rÂ   Ú
meta_parts)rÁ   rr   rs   Úflatten_meta’  s   z1DaskDMatrix._map_local_data.<locals>.flatten_metarœ   Úm_partsr}   Nc                    s:   | d urt ˆ ƒt | ƒksJ ˆˆ d| |ƒƒ‚| ˆ|< d S d S )NÚXr¸   )rÅ   r}   )ÚX_partsrº   Úpartsrr   rs   Úappend_meta¢  s   ÿüz0DaskDMatrix._map_local_data.<locals>.append_metar�   r�   r‘   r—   r˜   r™   Úfinishedc                 S   s   i | ]}|j |“qS rr   ©Úkey©Ú.0Úpartrr   rr   rs   Ú
<dictcomp>Ð  ó    z/DaskDMatrix._map_local_data.<locals>.<dictcomp>c                 S   s   g | ]}|j ‘qS rr   rË   rÍ   rr   rr   rs   Ú
<listcomp>Ò  s    z/DaskDMatrix._map_local_data.<locals>.<listcomp>)Úkeys)r   r   rc   rH   r   r   Úrangeri   ÚitemsÚappendr«   ÚmapÚdaskÚdelayedÚcomputer{   ÚwaitÚstatusÚpartition_orderÚ	enumeraterÌ   Ú	schedulerÚwho_hasr   ÚnextÚiterr¬   rš   Úresult)r   r‰   rœ   r�   rž   r‘   r—   rš   r˜   r™   rÄ   Úy_partsÚw_partsÚmargin_partsÚ	qid_partsÚll_partsÚlu_partsrÉ   Úpacked_partsÚiÚ	part_dictrÌ   ÚvalueÚdelayed_partsÚ	fut_partsrÏ   Úkey_to_partitionrà   r¬   Úworkersrr   )rÇ   r‰   rº   rÈ   rÁ   rs   r¯   i  sp   €ÿÿÿÿ
þ
"&






ÿ
þzDaskDMatrix._map_local_dataÚworker_addrc              	   C   s*   | j | j| j| j| j| j |d¡| jdœS )z\Create a dictionary of objects that can be pickled for function
        arguments.

        N)r”   r•   rš   r’   r›   rÈ   r­   )r”   r•   rš   r’   r›   r¬   Úgetr­   )r   rò   rr   rr   rs   Ú_create_fn_argsã  s   ùzDaskDMatrix._create_fn_argsc                 C   s   | j S )zxGet the number of columns (features) in the DMatrix.

        Returns
        -------
        number of columns
        )r©   r³   rr   rr   rs   Únum_colò  s   zDaskDMatrix.num_colr±   )rƒ   r„   r…   r†   r1   r   rI   rH   ÚfloatÚboolr!   r"   rz   r   r²   r¯   rc   r
   r   rô   rª   rõ   rr   rr   rr   rs   rO   	  s¢    üïþýüúùø	÷
öõôóòñðïî?õýüûúùø	÷
öõ
ôzrO   Ú_MapRetTÚ_PÚfuncÚrefsrñ   c             
   ‡   sø   �t | ƒ} g }|D ]I}g }|D ]}t|tƒr| | |¡¡ q| |¡ qdtdtjdtjdt	t
 f‡ fdd„}| jtt||ƒ|ƒg|¢R d|gdd	œŽ}	| |	¡ q	d
ttt
  dtt
 fdd„}
t |¡}| |
|
¡I dH }	|  |	¡ ¡ I dH }|S )z.Map a function onto partitions of each worker.Ú_addressrw   Úkwargsr[   c                    s:   t  ¡ }|j| krtd|j› d| › d�ƒ‚ˆ |i |¤ŽgS )NzInvalid worker address: z, expecting zz. This is likely caused by one of the workers died and Dask re-scheduled a different one. Resilience is not yet supported.)r{   r|   r~   r£   )rü   rw   rý   r€   ©rú   rr   rs   Úfn  s   
ÿz!map_worker_partitions.<locals>.fnFT)Úpurerñ   Úallow_other_workersÚresultsc                 S   s   | D ]
}|d ur|  S qd S r±   rr   )r  Úvrr   rr   rs   Úfirst_valid:  s
   ÿz*map_worker_partitions.<locals>.first_validN)r�   ra   rO   rÖ   rô   rc   rù   rw   rý   r   rø   Úsubmitr   r   r   r   ÚdbÚfrom_delayedÚ	reductionrÚ   rã   )r‰   rú   rñ   rû   Úfuturesro   rw   Úrefrÿ   Úfutr  r   rã   rr   rþ   rs   Úmap_worker_partitions   s2   €	
&ÿþû
r  c                )       sþ   e Zd ZdZe	dddddddddddddddddœded dedee d	ee d
ee dee de	dee
 deeeee f  dee dee dee dee dee dee dee de	dee ddf&‡ fdd„ƒZdedeeef f‡ fdd„Z‡  ZS )rP   zmA dask version of :py:class:`QuantileDMatrix`. See :py:class:`DaskDMatrix` for
    parameter documents.

    NF)r�   r‘   r’   r“   r”   r•   Úmax_binr
  r–   r—   r˜   r™   rš   r›   Úmax_quantile_batchesr‰   rŠ   rœ   r�   r�   r‘   r’   r“   r”   r•   r  r
  r–   r—   r˜   r™   rš   r›   r  r[   c                   s\   t ƒ j||||||||||||||	|d� |
| _|| _d| _|d ur)t|ƒ| _d S d | _d S )N)r‰   rœ   r�   r�   r‘   r–   r—   r˜   r™   r’   r“   rš   r”   r•   r›   T)ry   rz   r  r  r­   ÚidÚ_ref)r   r‰   rœ   r�   r�   r‘   r’   r“   r”   r•   r  r
  r–   r—   r˜   r™   rš   r›   r  r�   rr   rs   rz   M  s*   ñ zDaskQuantileDMatrix.__init__rò   c                    s8   t ƒ  |¡}| j|d< | j|d< | jd ur| j|d< |S )Nr  r  r
  )ry   rô   r  r  r  )r   rò   rw   r�   rr   rs   rô   z  s   



z#DaskQuantileDMatrix._create_fn_argsr±   )rƒ   r„   r…   r†   r1   r   rI   rH   rö   r÷   r!   r   r   r   rª   rO   rz   rc   r
   rô   rˆ   rr   rr   r�   rs   rP   G  sx    üìþýüúùø	÷
öõôóòñðïîíìë&,rP   ÚdconfigÚcoll_cfgc           	      Ã   sª   �|du rt ƒ n|}d}d}t||ƒ\}}|dur||f}nd}ztj | jj¡}| d¡}W n ty:   d}Y nw |  	t
||||j¡I dH }| |¡}|dusSJ ‚|S )zBGet rabit context arguments from data distribution in DaskDMatrix.Nr   z/:)Ú
CollConfigrF   r{   ÚcommÚget_address_hostrß   r~   ÚstripÚ	ExceptionÚrun_on_schedulerrv   Útracker_timeoutÚget_comm_config)	r‰   rX   r  r  r]   r^   Ú	user_addrÚ
sched_addrrm   rr   rr   rs   Ú_get_rabit_argsƒ  s(   €
ÿ
ÿ
r  c                   C   s   t jjdd d�S )NÚxgboost)Údefault)rØ   r   ró   rr   rr   rr   rs   Ú_get_dask_config¬  ó   r   ÚdtrainÚevalsc                 C   s~   t | j ¡ ƒ}|r;|D ]/}t|ƒdksJ ‚t|d tƒr#t|d tƒs%J ‚|d | u r,qt |d j ¡ ƒ}| |¡}qt|ƒS )Nr   r   rA   )	Úsetr¬   rÓ   ri   ra   rO   rc   Úunionr«   )r"  r#  ÚX_worker_maprq   r¬   rr   rr   rs   Ú_get_workers_from_data·  s    r'  c                 Ã   s@   �|j  ¡ I d H }|d  ¡ }t| ƒ| }|rtd|› �ƒ‚d S )Nrñ   zMissing required workers: )rß   ÚidentityrÓ   r$  ÚRuntimeError)rñ   r‰   ÚinfoÚcurrent_workersÚmissing_workersrr   rr   rs   Ú_check_workers_are_aliveÆ  s   €ÿr-  Úglobal_configÚparamsÚnum_boost_roundÚobjÚearly_stopping_roundsÚverbose_evalÚ	xgb_modelÚ	callbacksÚcustom_metricc                 ƒ   sP  �t ||ƒ}t|| ƒI d H  t| t|ƒ||d�I d H }t|ƒ dtdtttttf f dtdt	t dt	t dt
dt
d	tt f‡ ‡‡‡‡‡‡‡fd
d„}t || ¡4 I d H šF |d urpdd„ |D ƒ}dd„ |D ƒ}dd„ |D ƒ}ng }g }g }t| |||t|ƒ||g|g| ¢R d|iŽI d H }|W  d   ƒI d H  S 1 I d H s¡w   Y  d S )N)r  r  Ú
parametersÚ	coll_argsÚtrain_idÚ
evals_nameÚevals_idÚ	train_refrû   r[   c                    s  t  ¡ }|  ¡ }t||ƒ}	| |	|	dœ¡ i }
ˆ d|	i¡ tdi |¤Ž�H tjdi ˆ¤Ž�0 t||g|¢R |||	ˆdœŽ\}}t	||ˆ|
t
|ƒdkrM|nd ˆˆˆˆˆˆ d�}W d   ƒ n1 saw   Y  W d   ƒ n1 spw   Y  | ¡ dkr‚||
dœ}|S d }|S )N)ÚnthreadÚn_jobsr=  )r;  r:  Ú	n_threadsÚmodelr   )r/  r"  r0  Úevals_resultr#  r1  r6  r2  r3  r4  r5  rK   rr   )r{   r|   ÚcopyrG   rg   rN   r   Úconfig_contextrB   Úworker_trainri   Únum_row)r7  r8  r9  r:  r;  r<  rû   r€   Úlocal_paramr?  Úlocal_historyÚXyr#  rL   rŽ   ©r5  r6  r2  r.  r0  r1  r3  r4  rr   rs   Údo_trainê  sR   	
"þýøõô€ þÿz_train_async.<locals>.do_trainc                 S   s   g | ]\}}|‘qS rr   rr   ©rÎ   r»   Únrr   rr   rs   rÒ   "  rÑ   z _train_async.<locals>.<listcomp>c                 S   s   g | ]\}}|‘qS rr   rr   rK  rr   rr   rs   rÒ   #  rÑ   c                 S   s   g | ]}t |ƒ‘qS rr   )r  )rÎ   r»   rr   rr   rs   rÒ   $  rÑ   rñ   )r'  r-  r  ri   r0   r
   rc   r   rª   r   Údictr   rJ   r{   Ú	MultiLockr  r  )r‰   r.  r  r/  r"  r0  r#  r1  r2  r3  r4  r5  r6  r  rñ   r8  rJ  Ú
evals_datar:  r;  rã   rr   rI  rs   Ú_train_asyncÐ  s^   €
ÿÿþýüûúùø6ø	÷õ0érP  é
   T)r#  r1  r2  r4  r3  r5  r6  r  c                C   s(   t | ƒ} | jtft ¡ tƒ dœtƒ ¤ŽS )a½  Train XGBoost model.

    .. versionadded:: 1.0.0

    .. note::

        Other parameters are the same as :py:func:`xgboost.train` except for
        `evals_result`, which is returned as part of function return value instead of
        argument.

    Parameters
    ----------
    client :
        Specify the dask client used for training.  Use default client returned from
        dask if it's set to None.

    coll_cfg :
        Configuration for the communicator used during training. See
        :py:class:`~xgboost.collective.Config`.

    Returns
    -------
    results: dict
        A dictionary containing trained booster and evaluation history.  `history` field
        is the same as `eval_result` from `xgboost.train`.

        .. code-block:: python

            {'booster': xgboost.Booster,
             'history': {'train': {'logloss': ['0.48253', '0.35953']},
                         'eval': {'logloss': ['0.480385', '0.357756']}}}

    )r.  r  )r�   r®   rP  r   Ú
get_configr   Úlocals)r‰   r/  r"  r0  r#  r1  r2  r4  r3  r5  r6  r  rr   rr   rs   r@   :  s   1ÿýür@   Úis_dfÚoutput_shapec                 C   s   | ot |ƒdkS ©Nr   r¸   )rT  rU  rr   rr   rs   Ú_can_output_dft  r!  rW  rœ   Ú
predictionÚcolumnsc                 C   sš   t ||jƒrKt| ddƒ}t| ƒr.ddl}|jdkr"|ji |tjd�S |j||tj|d�}|S ddl	}|jdkrA|ji |tj|d�S |j||tj|d�}|S )z0Return dataframe for prediction when applicable.ÚindexNr   )rY  Údtype)rY  r[  rZ  )
rW  r¢   Úgetattrr(   ÚcudfÚsizer¥   rŸ   Úfloat32Úpandas)rœ   rX  rY  rT  rZ  r]  Úpdrr   rr   rs   Ú_maybe_dataframex  s&   

ÿö

ÿ
ÿrb  Úmapped_predictrL   zdistributed.Futurer‘   .rÂ   c                 Ã   s€  �t | ¡ ƒ}t|ƒdkrt|tjƒrtdƒ‚tt|tjƒ|ƒrR|d ur/t|tj	ƒr/| 
¡ }n|}tj| ||d||tj |¡d�}t|ƒdkrP|jd d …df }|S |d urdt|tjtjfƒrd| ¡ }	n|}	t|ƒdkrrdg}
g }n g }
t|tjƒr…ttt|ƒd ƒƒ}ndd	„ tt|ƒd ƒD ƒ}t|ƒdkr¬t|jƒ}t|tƒs¤J ‚|d f|d< nd }tj| ||d
||	||
|tjd�
}|S )Né   zGUse `da.Array` or `DaskDMatrix` when output has more than 2 dimensions.T)rÂ   rA   r   r   c                 S   s   g | ]}|d  ‘qS )r   rr   )rÎ   rë   rr   rr   rs   rÒ   Ò  rÑ   z(_direct_predict_impl.<locals>.<listcomp>F)ÚchunksÚ	drop_axisÚnew_axisr[  )rb   rÓ   ri   ra   r¤   r¥   r£   rW  r¦   r§   Úto_dask_dataframeÚmap_partitionsÚutilsÚ	make_metaÚilocr¨   Úto_dask_arrayr«   rÔ   re  Ú
map_blocksrŸ   r_  )rc  rL   rœ   r‘   rU  rÂ   rY  Úbase_margin_dfÚpredictionsÚbase_margin_arrayrf  rg  re  rr   rr   rs   Ú_direct_predict_impl™  sj   €	ÿÿ
ù
,
Öÿ

örr  ÚfeaturesÚinplacerý   c                 K   s¶   t |tƒsJ ‚tj d¡}| d|¡}|r$| ¡ }| d¡dkr$d|d< t|dd�}| j	|fdd	i|¤Ž}t
|jƒdkrA|jd nd}	i }
t||jƒrVt|	ƒD ]}d
|
|< qO|j|
fS )z@Create a dummy test sample to infer output shape for prediction.iÊ  rA   Úpredict_typeÚmarginTÚoutput_margin)r›   Úvalidate_featuresFÚf4)ra   rª   rŸ   ÚrandomÚRandomStateÚrandnrB  Úpopr,   rV   ri   r¢   rW  rÔ   )rL   rs  rT  rt  rý   ÚrngÚtest_sampleÚmÚ
test_predtÚ	n_columnsrÂ   rë   rr   rr   rs   Ú_infer_predict_outputî  s   

rƒ  r@  c                 Ã   s”   �t |tƒr| j|dd�I d H }|S t |tƒr%| j|d dd�I d H }|S t |tjƒr=|}|j}|tur;td|› �ƒ‚|S tttttjgt|ƒƒƒ‚)NF)ÚhashrL   z9Underlying type of model future should be `Booster`, got )	ra   r+   ÚscatterrM  r{   r   r‹   r�   r2   )r‰   r@  rL   Útrr   rr   rs   Ú_get_model_future  s    €

õ
÷ÿÿr‡  rw  r’   Ú	pred_leafÚpred_contribsÚapprox_contribsÚpred_interactionsrx  Úiteration_rangeÚstrict_shapec       	   %      ƒ   sò  �t | |ƒI d H }t|ttjtjfƒs!ttttjtjgt	|ƒƒƒ‚dt
dtdtdtt dtdtf‡ ‡‡‡‡‡‡	‡
‡‡f
dd„}t|tjtjfƒrt|  | jt||jd	 t|tjƒd
ˆˆ
ˆˆ ˆ	ˆd�¡I d H \}}t|||d ||d�I d H S |  | jt|| ¡ d
d
ˆˆ
ˆˆ ˆ	ˆd�¡I d H \}}|j‰|j‰|j‰|j‰dt
dtttf dtjf‡ ‡‡‡‡‡‡‡‡	‡
‡‡fdd„}g }g }g }g }t|j ¡ ƒ}|D ]"}|j| }|  |¡ |  t!|ƒ|g ¡ |  ‡fdd„|D ƒ¡ qÈt"||ƒD ]\}}| jdd„ ||gd�}| #|¡ qðtt"||||ƒƒ}t$|dd„ d�}dd„ |D ƒ}dd„ |D ƒ}dd„ |D ƒ}g }t"||ƒD ]\}}| j||||gd�} | #| ¡ �q2g }!|  %|¡I d H }t&|ƒD ]\}"}#|! #tj'||" |#f|d	d …  tj(d�¡ �qUtj)|!dd�}$|$S )NrL   Ú	partitionrT  rY  Ú_r[   c                    sp   t jdi ˆ¤Ž�& t|ˆdd�}| j|ˆˆˆˆ ˆˆ	ˆˆd�	}t||||ƒ}|W  d   ƒ S 1 s1w   Y  d S )NT)rœ   r’   r›   )	rœ   rw  rˆ  r‰  rŠ  r‹  rx  rŒ  r�  rr   )r   rC  r,   rV   rb  )rL   rŽ  rT  rY  r�  r€  Úpredt)
rŠ  r.  rŒ  r’   rw  r‰  r‹  rˆ  r�  rx  rr   rs   rc  0  s(   ý÷$îz&_predict_async.<locals>.mapped_predictrA   F)	rs  rT  rt  rw  rˆ  r‰  rŠ  r‹  r�  ©rc  rL   rœ   r‘   rU  rÂ   )
rL   rs  rT  rt  rw  rˆ  r‰  rŠ  r‹  r�  rÏ   c                    s|   |d }|  dd ¡}tjdi ˆ¤Ž�" t|ˆ|ˆˆdd�}| j|ˆˆ	ˆˆ ˆˆˆˆ
d�	}|W  d   ƒ S 1 s7w   Y  d S )Nrœ   r‘   T)r’   r‘   r”   r•   r›   )rw  rˆ  r‰  rŠ  r‹  rx  rŒ  r�  rr   )ró   r   rC  r,   rV   )rL   rÏ   rœ   r‘   r€  r�  )rŠ  r”   r•   r.  rŒ  r’   rw  r‰  r‹  rˆ  r�  rx  rr   rs   Údispatched_predictv  s0   ú÷$ìz*_predict_async.<locals>.dispatched_predictc                    s   g | ]}ˆ |j  ‘qS rr   rË   rÍ   )rÝ   rr   rs   rÒ   ˜  s    z"_predict_async.<locals>.<listcomp>c                 S   s   | d j d S )Nrœ   r   )r¢   )rÏ   rr   rr   rs   Ú<lambda>š  s    z _predict_async.<locals>.<lambda>)rñ   c                 S   s   | d S rV  rr   )Úprr   rr   rs   r“  ž  s    rË   c                 S   s   g | ]\}}}}|‘qS rr   rr   ©rÎ   rÏ   r¢   ÚorderÚwrr   rr   rs   rÒ   Ÿ  ó    c                 S   s   g | ]\}}}}|‘qS rr   rr   r•  rr   rr   rs   rÒ      r˜  c                 S   s   g | ]\}}}}|‘qS rr   rr   r•  rr   rr   rs   rÒ   ¡  r˜  )r¢   r[  r   ©Úaxis)*r‡  ra   rO   r¦   r§   r¤   r¥   r�   r2   r‹   r+   r   r÷   r   rª   rÚ   r  rƒ  r¢   rr  rõ   rÝ   r”   r•   r’   r
   rc   rŸ   Úndarrayr«   r¬   rÓ   Úextendri   ÚziprÖ   ÚsortedÚgatherrÞ   r  r_  Úconcatenate)%r‰   r.  r@  rœ   rw  r’   rˆ  r‰  rŠ  r‹  rx  rŒ  r�  Ú_boosterrc  Ú_output_shaperÂ   rU  r�  r’  Ú	all_partsÚ
all_ordersÚ
all_shapesÚall_workersÚworkers_addressrò   Úlist_of_partsr—  rÏ   ÚsÚparts_with_orderr	  ÚfÚarraysrë   Úrowsrp  rr   )rŠ  r”   r•   r.  rŒ  r’   rw  rÝ   r‰  r‹  rˆ  r�  rx  rs   Ú_predict_async  sº   €ÿÿÿÿÿ þ
õÿú	õÿ:

ÿÿr®  F)r   r   )	rw  r’   rˆ  r‰  rŠ  r‹  rx  rŒ  r�  c       	         C   ó$   t | ƒ} | jtfdt ¡ itƒ ¤ŽS )aÿ  Run prediction with a trained booster.

    .. note::

        Using ``inplace_predict`` might be faster when some features are not needed.
        See :py:meth:`xgboost.Booster.predict` for details on various parameters.  When
        output has more than 2 dimensions (shap value, leaf with strict_shape), input
        should be ``da.Array`` or ``DaskDMatrix``.

    .. versionadded:: 1.0.0

    Parameters
    ----------
    client:
        Specify the dask client used for training.  Use default client
        returned from dask if it's set to None.
    model:
        The trained model.  It can be a distributed.Future so user can
        pre-scatter it onto all workers.
    data:
        Input data used for prediction.  When input is a dataframe object,
        prediction output is a series.
    missing:
        Used when input data is not DaskDMatrix.  Specify the value
        considered as missing.

    Returns
    -------
    prediction: dask.array.Array/dask.dataframe.Series
        When input data is ``dask.array.Array`` or ``DaskDMatrix``, the return value is
        an array, when input data is ``dask.dataframe.DataFrame``, return value can be
        ``dask.dataframe.Series``, ``dask.dataframe.DataFrame``, depending on the output
        shape.

    r.  )r�   r®   r®  r   rR  rS  )r‰   r@  rœ   rw  r’   rˆ  r‰  rŠ  r‹  rx  rŒ  r�  rr   rr   rs   rV   ¶  s   3rV   ru  c        
         ƒ   s  �t | ƒ} t| |ƒI d H }
t|tjtjfƒs#tttjtjgt	|ƒƒƒ‚|d urAt|tjtjtj
fƒsAtttjtjtj
gt	|ƒƒƒ‚dtdtdtdtt dtdtf‡ ‡‡‡‡‡fdd„}|  | jt|
|jd	 t|tjƒd
ˆˆˆd�¡I d H \}}t||
||||d�I d H S )NrL   rŽ  rT  rY  r‘   r[   c              
      sZ   t jdi ˆ ¤Ž� | j|ˆˆˆ|ˆˆd�}W d   ƒ n1 sw   Y  t||||ƒ}|S )N)rŒ  ru  r’   r‘   rx  r�  rr   )r   rC  rW   rb  )rL   rŽ  rT  rY  r‘   rX  ©r.  rŒ  r’   ru  r�  rx  rr   rs   rc    s   ùÿ
z._inplace_predict_async.<locals>.mapped_predictrA   T)rs  rT  rt  ru  rŒ  r�  r‘  )r�   r‡  ra   r¦   r§   r¤   r¥   r�   r2   r‹   r¨   r+   r   r÷   r   rª   rÚ   r  rƒ  r¢   rr  )r‰   r.  r@  rœ   rŒ  ru  r’   rx  r‘   r�  rL   rc  r¢   rÂ   rr   r°  rs   Ú_inplace_predict_asyncí  sT   €
ÿÿþýüûú
øÿúr±  rí   )rŒ  ru  r’   rx  r‘   r�  c          	      C   r¯  )a±  Inplace prediction. See doc in :py:meth:`xgboost.Booster.inplace_predict` for
    details.

    .. versionadded:: 1.1.0

    Parameters
    ----------
    client:
        Specify the dask client used for training.  Use default client
        returned from dask if it's set to None.
    model:
        See :py:func:`xgboost.dask.predict` for details.
    data :
        dask collection.
    iteration_range:
        See :py:meth:`xgboost.Booster.predict` for details.
    predict_type:
        See :py:meth:`xgboost.Booster.inplace_predict` for details.
    missing:
        Value in the input data which needs to be present as a missing
        value. If None, defaults to np.nan.
    base_margin:
        See :py:obj:`xgboost.DMatrix` for details.

        .. versionadded:: 1.4.0

    strict_shape:
        See :py:meth:`xgboost.Booster.predict` for details.

        .. versionadded:: 1.4.0

    Returns
    -------
    prediction :
        When input data is ``dask.array.Array``, the return value is an array, when
        input data is ``dask.dataframe.DataFrame``, return value can be
        ``dask.dataframe.Series``, ``dask.dataframe.DataFrame``, depending on the output
        shape.

    r.  )r�   r®   r±  r   rR  rS  )	r‰   r@  rœ   rŒ  ru  r’   rx  r‘   r�  rr   rr   rs   rW   .  s   5ÿÿÿrW   ÚdeviceÚtree_methodr  c           
      ‹   s    �dt t dtdtf‡ ‡‡‡fdd„}td
d|i|¤Ž\}}|I dH }|du r+||fS g }|D ]}	|	d |u r=| |	¡ q/| |	d I dH |	d	 f¡ q/||fS )z(A switch function for async environment.r
  rý   r[   c                    s2   t ˆˆƒrtdˆ | ˆdœ|¤ŽS tddˆ i|¤ŽS )N)r‰   r
  r  r‰   rr   )r9   rP   rO   )r
  rý   ©r‰   r²  r  r³  rr   rs   Ú	_dispatchu  s   
ÿÿz2_async_wrap_evaluation_matrices.<locals>._dispatchÚcreate_dmatrixNr   rA   rr   )r   rO   r   r=   rÖ   )
r‰   r²  r³  r  rý   rµ  Útrain_dmatrixr#  Úawaitedrq   rr   r´  rs   Ú_async_wrap_evaluation_matricesl  s   €$	

r¹  ÚDaskScikitLearnBasec                 c   s$   � z|| _ | V  W d| _ dS d| _ w )z-Temporarily set the client for sklearn model.NrÀ   )r@  r‰   rr   rr   rs   Ú_set_worker_client‰  s
   €r»  c                       s0  e Zd ZdZdZddœdee deddf‡ fdd„Zd	e	d
e
de
dee dee defdd„Zedddddœde	d
e
de
dee dee defdd„ƒZ	d&de	dee defdd„Z	d&de	dee defdd„Zdee fdd„Zdefdd„Zed'dd „ƒZejd(d"d „ƒZd#ededefd$d%„Z‡  ZS ))rº  z<Base class for implementing scikit-learn interface with DaskN)r  r  rý   r[   c                   s   t ƒ jdi |¤Ž || _d S )Nrr   )ry   rz   r  )r   r  rý   r�   rr   rs   rz   š  s   
zDaskScikitLearnBase.__init__rœ   rw  rx  r‘   rŒ  c             
   Ã   sr   �|   |¡}|  ¡ sJ ‚t| j|  ¡ |||rdnd| j||d�I d H }t|tjƒr7| 	¡ }t
ƒ r7tƒ s7| ¡ }|S )Nrv  rí   )r‰   r@  rœ   rŒ  ru  r’   r‘   rx  )Ú_get_iteration_rangeÚ_can_use_inplace_predictrW   r‰   Úget_boosterr’   ra   r¤   r¥   rm  rD   rE   r¼   )r   rœ   rw  rx  r‘   rŒ  Úpredtsrr   rr   rs   r®  Ÿ  s$   €
	
ø
	z"DaskScikitLearnBase._predict_asyncFT©rw  rx  r‘   rŒ  rÆ   c                C   s   | j j| j|||||d�S )NrÀ  )r‰   r®   r®  )r   rÆ   rw  rx  r‘   rŒ  rr   rr   rs   rV   Æ  s   
úzDaskScikitLearnBase.predictc                 Ã   sJ   �|   |¡}t| j|| j| jd�I d H }t| j|  ¡ |d|d�I d H }|S )N)rœ   r’   r•   T)r@  rœ   rˆ  rŒ  )r¼  rO   r‰   r’   r•   rV   r¾  )r   rÆ   rŒ  Útest_dmatrixr¿  rr   rr   rs   Ú_apply_asyncÙ  s    €
üûz DaskScikitLearnBase._apply_asyncc                 C   s   | j j| j||d�S )N)rŒ  )r‰   r®   rÂ  )r   rÆ   rŒ  rr   rr   rs   Úapplyî  s   zDaskScikitLearnBase.applyc                    s$   dt t f‡ fdd„}ˆ  |¡ ¡ S )Nr[   c                   “   s   �ˆ S r±   rr   rr   r³   rr   rs   r�  ÷  s   €z(DaskScikitLearnBase.__await__.<locals>._)r   r   Ú_client_syncr²   )r   r�  rr   r³   rs   r²   õ  s   zDaskScikitLearnBase.__await__c                 C   s   | j  ¡ }d|v r|d= |S )NÚ_client)Ú__dict__rB  )r   Úthisrr   rr   rs   Ú__getstate__ü  s   
z DaskScikitLearnBase.__getstate__rŠ   c                 C   s   t | jƒ}|S )zûThe dask client used in this model.  The `Client` object can not be
        serialized for transmission, so if task is launched from a worker instead of
        directly from the client process, this attribute needs to be set at that worker.

        )r�   rÅ  )r   r‰   rr   rr   rs   r‰     s   
zDaskScikitLearnBase.clientÚcltc                 C   s   |d ur|j nd| _|| _d S )NF)ÚasynchronousÚ_asynchronousrÅ  )r   rÉ  rr   rr   rs   r‰     s   
rú   c              	   K   sæ   | j du rct| ddƒ}zt ¡  d}W n ty   d}Y nw |rct ¡ �6}t| |ƒ�}|jj|fi |¤d|i¤Ž}|W  d  ƒ W  d  ƒ S 1 sMw   Y  |W  d  ƒ S 1 s^w   Y  | jj|fi |¤d| jj	i¤ŽS )z‰Get the correct client, when method is invoked inside a worker we
        should use `worker_client' instead of default client.

        NrË  FTrÊ  )
rÅ  r\  r{   r|   r£   Úworker_clientr»  r‰   r®   rÊ  )r   rú   rý   rÊ  Ú	in_workerr‰   rÇ  rŽ   rr   rr   rs   rÄ    s2   
ÿ
ÿÿÿüÿ ú z DaskScikitLearnBase._client_syncr±   )r[   rŠ   )rÉ  rŠ   r[   N)rƒ   r„   r…   r†   rÅ  r   r  r   rz   rI   r÷   rH   r#   r®  r1   rV   rÂ  rÃ  r   r²   r
   rÈ  Úpropertyr‰   Úsetterr	   rÄ  rˆ   rr   rr   r�   rs   rº  •  st    $þüûúù
ø'ùþüûúùøýþý
üýþý
ü
z3Implementation of the Scikit-Learn API for XGBoost.Ú
estimatorsc                   @   s  e Zd ZdZdededee dee deeeeef   deee  deee  d	e	e
ef d
ee	eef  dee defdd„Zedddddddddœdededee dee deeeeef   d	ee	e
ef  d
ee	eeef  deee  deee  dee dd fdd„ƒZdS )rQ   zAdummy doc string to workaround pylint, replaced by the decorator.rÆ   ÚyÚsample_weightr‘   Úeval_setÚsample_weight_eval_setÚbase_margin_eval_setÚverboser4  rš   r[   c                Ã   s2  �|   ¡ }|  |	||
¡\}}}}
tdi d| j“d| j“d| j“d| j“d|“d|“dd “dd “d	|“d
|“d|
“d|“d|“d|“dd “dd “d| j“d| j“d| j	“ŽI d H \}}t
| jƒret| jƒ}nd }| jjtfd| jt ¡ tƒ |||  ¡ ||||| j| j| j|dœŽI d H }|d | _|  |d ¡ | S )Nr‰   r²  r³  r  rÆ   rÑ  r–   r—   rÒ  r‘   rš   rÓ  rÔ  rÕ  Ú
eval_groupÚeval_qidr’   r›   r•   T©rÊ  r‰   r.  r  r/  r"  r0  r#  r1  r6  r3  r2  r5  r  r4  rL   rM   rr   )Úget_xgb_paramsÚ_configure_fitr¹  r‰   r²  r³  r  r’   r›   r•   ÚcallableÚ	objectiver<   r®   rP  r   rR  r   Úget_num_boosting_roundsr2  r5  r  Ú_BoosterÚ_set_evaluation_result©r   rÆ   rÑ  rÒ  r‘   rÓ  rÔ  rÕ  rÖ  r4  rš   r/  r@  Úmetricr"  r#  r1  r  rr   rr   rs   Ú
_fit_async3  sŠ   €ÿÿþýüûúùø	÷
öõôóòñðïîí
ÿð
zDaskXGBRegressor._fit_asyncNT©rÒ  r‘   rÓ  rÖ  r4  rÔ  rÕ  rš   c                C   ó(   dd„ t ƒ  ¡ D ƒ}| j| jfi |¤ŽS )Nc                 S   ó   i | ]\}}|d vr||“qS ©)r   r‚   rr   ©rÎ   Úkr  rr   rr   rs   rÐ   †  ó    z(DaskXGBRegressor.fit.<locals>.<dictcomp>©rS  rÕ   rÄ  rã  ©r   rÆ   rÑ  rÒ  r‘   rÓ  rÖ  r4  rÔ  rÕ  rš   rw   rr   rr   rs   Úfitw  s   zDaskXGBRegressor.fit)rƒ   r„   r…   r†   rI   rH   r   r   r   r   rª   r÷   r+   r5   rã  r1   rc   rí  rr   rr   rr   rs   rQ   -  sr    þýûúù
ø
	÷

öõô
óDôþýûúùø	÷

ö
õôórQ   zBImplementation of the scikit-learn API for XGBoost classification.c                       s¦  e Zd Zdededee dee deeeeef   deee  deee  dee	e
f d	eeeef  d
ee dd fdd„Zdddddddddœdededee dee deeeeef   deee	e
f  d	eeeeef  deee  deee  d
ee dd fdd„Zdede
dee dee def
‡ fdd„Z			ddede
dee dee def
dd„Zejje_dede
de
dee dee def‡ fdd„Z‡  ZS )rR   rÆ   rÑ  rÒ  r‘   rÓ  rÔ  rÕ  rÖ  r4  rš   r[   c                Ã   sî  �|   ¡ }|  |	||
¡\}}}}
t| jfi d| j“d| j“d| j“d|“d|“dd “dd “d|“d	|“d
|
“d|“d|“d|“dd “dd “d| j“d| j“d| j	“ŽI d H \}}t
|tjƒrl| j t |¡¡I d H | _n| j | ¡ ¡I d H | _t| jƒrƒ| j ¡ | _t| jƒrŽ| j ¡ | _t | j¡| _t| jƒ| _| jdkrªd|d< | j|d< nd|d< t| jƒr¹t| jƒ}nd }| jjtfd| jt ¡ t ƒ |||  !¡ ||||| j"| j#| j$|dœŽI d H }|d | _%t| jƒsî|d | _|  &|d ¡ | S )Nr²  r³  r  rÆ   rÑ  r–   r—   rÒ  r‘   rš   rÓ  rÔ  rÕ  r×  rØ  r’   r›   r•   r   zmulti:softprobrÝ  Ú	num_classzbinary:logisticTrÙ  rL   rM   )'rÚ  rÛ  r¹  r‰   r²  r³  r  r’   r›   r•   ra   r¦   r§   rÚ   ÚuniqueÚclasses_Údrop_duplicatesr)   Úto_cupyr*   ró   rŸ   r   ri   Ú
n_classes_rÜ  rÝ  r<   r®   rP  r   rR  r   rÞ  r2  r5  r  rß  rà  rá  rr   rr   rs   rã  �  s¨   €ÿÿþýüûúùø	÷
öõôóòñðïîí



ÿð


zDaskXGBClassifier._fit_asyncNTrä  c                C   rå  )Nc                 S   ræ  rç  rr   rè  rr   rr   rs   rÐ   ö  rê  z)DaskXGBClassifier.fit.<locals>.<dictcomp>rë  rì  rr   rr   rs   rí  è  s   zDaskXGBClassifier.fitrx  rŒ  c                 ƒ   sZ   �| j dkr
tdƒ‚tƒ j|d|||d�I d H }tttjdd�tjƒ}tt	| ddƒ||ƒS )	Nzmulti:softmaxzSmulti:softmax doesn't support `predict_proba`.  Switch to `multi:softproba` insteadF)rœ   rw  rx  r‘   rŒ  T)Úallow_unknown_chunksizesró  r   )
rÝ  r£   ry   r®  r   r   r¦   Úvstackr;   r\  )r   rÆ   rx  r‘   rŒ  r¿  rõ  r�   rr   rs   Ú_predict_proba_asyncù  s    €
ÿûÿz&DaskXGBClassifier._predict_proba_asyncc                 C   s   | j | j||||d�S )N)rÆ   rx  r‘   rŒ  )rÄ  rö  )r   rÆ   rx  r‘   rŒ  rr   rr   rs   Úpredict_proba  s   ûzDaskXGBClassifier.predict_probarœ   rw  c          	      ƒ   sŽ   �t ƒ j|||||d�I d H }|r|S t|jƒdkr#|dk t¡}|S t|jƒdks,J ‚t|tjƒs4J ‚dt	dt	fdd„}tj
||dd	�}|S )
NrÀ  rA   g      à?r   Úxr[   c                 S   s   | j dd�S )NrA   r™  )Úargmax)rø  rr   rr   rs   Ú_argmax>  s   z1DaskXGBClassifier._predict_async.<locals>._argmax)rf  )ry   r®  ri   r¢   Úastyperª   ra   r¦   r§   r   rn  )	r   rœ   rw  rx  r‘   rŒ  Ú
pred_probsÚpredsrú  r�   rr   rs   r®  #  s$   €	û÷z DaskXGBClassifier._predict_async)TNN)rƒ   r„   r…   rI   rH   r   r   r   r   rª   r÷   r+   r5   rã  rc   rí  r#   rö  r   r÷  r3   r†   r®  rˆ   rr   rr   r�   rs   rR   Š  s¼    þýûúù
ø
	÷

öõô
ó]ôþýûúùø	÷

ö
õô
óþýüûúûþýüû
ú
þüûúùørR   zZImplementation of the Scikit-Learn API for XGBoost Ranking.

    .. versionadded:: 1.4.0

aq  
    allow_group_split :

        .. versionadded:: 3.0.0

        Whether a query group can be split among multiple workers. When set to `False`,
        inputs must be Dask dataframes or series. If you have many small query groups,
        this can significantly increase the fragmentation of the data, and the internal
        DMatrix construction can take longer.

zf
        .. note::

            For the dask implementation, group is not supported, use qid instead.
)Úextra_parametersÚend_notec                        s²  e Zd Zeddddœdededee ded	df
‡ fd
d„ƒZ	d	e
e f‡ fdd„Zdededee dee dee deeeeef   deee  deee  deee  deeef deeeef  dee d	d fdd„Zedddddddddddddœdededee dee dee dee deeeeef   deee  deee  deeeef  deeeeef  deee  deee  dee d	d fdd „ƒZejje_‡  ZS )!rS   z	rank:ndcgFN)rÝ  Úallow_group_splitr  rÝ  r   r  rý   r[   c                   s2   t |ƒrtdƒ‚|| _tƒ jd||dœ|¤Ž d S )Nz5Custom objective function not supported by XGBRanker.)rÝ  r  rr   )rÜ  r£   r   ry   rz   )r   rÝ  r   r  rý   r�   rr   rs   rz   ^  s   	zDaskXGBRanker.__init__c                    s   t ƒ  ¡ }| d¡ |S )Nr   )ry   Ú_wrapper_paramsÚadd©r   r/  r�   rr   rs   r  l  ó   

zDaskXGBRanker._wrapper_paramsrÆ   rÑ  r—   rÒ  r‘   rÓ  rÔ  rÕ  rØ  rÖ  r4  rš   c       
         Ã   s  �|   ¡ }|  |||¡\}}}}t| jfi d| j“d| j“d| j“d|“d|“dd “d|“d|“d	|“d
|“d|“d|“d|“dd “d|	“d| j“d| j“d| j	“ŽI d H \}}| jj
tfd| jt ¡ tƒ |||  ¡ |d ||
| j| j|| jdœŽI d H }|d | _|d | _| S )Nr²  r³  r  rÆ   rÑ  r–   r—   rÒ  r‘   rš   rÓ  rÔ  rÕ  r×  rØ  r’   r›   r•   T)rÊ  r‰   r.  r  r/  r"  r0  r#  r1  r6  r3  r2  r5  r4  r  rL   rM   )rÚ  rÛ  r¹  r‰   r²  r³  r  r’   r›   r•   r®   rP  r   rR  r   rÞ  r2  r5  r  rß  Úevals_result_)r   rÆ   rÑ  r—   rÒ  r‘   rÓ  rÔ  rÕ  rØ  rÖ  r4  rš   r/  r@  râ  r"  r#  r  rr   rr   rs   rã  q  s„   €ÿÿþýüûúùø	÷
öõôóòñðïîíÿð

zDaskXGBRanker._fit_async)r–   r—   rÒ  r‘   rÓ  r×  rØ  rÖ  r4  rÔ  rÕ  rš   r–   r×  c                C   sf  d}|d u r
|d u st |ƒ‚|d u rt dƒ‚dtdttj fdd„}dtt dtdtttj  fd	d
„}| j�s ||ƒrP||dƒrP||dƒrP||dƒrP||dƒsRJ ‚|d urZ|d us\J ‚t	|ƒ}t
| j|||||d�\}}}}}|d u�r g }g }g }g }|	s�J ‚t|ƒD ]ˆ\}\}}|r‘|| nd }|r™|| nd }||ƒs¡J ‚|	s¥J ‚|	| }|	r¿||dƒr¿||dƒr¿||dƒr¿||dƒsÁJ ‚|d urÉ|d usËJ ‚t	|ƒ|krát
| j|||||ƒ\}}}}}n|||||f\}}}}}| ||f¡ | |¡ |d u�r| |¡ |d u�r| |¡ q…|}|}	|�r|nd }|�r|nd }| j| j|||||||	|
||||d�S )Nz=Use the `qid` instead of the `group` with the dask interface.z`qid` is required for ranking.rÆ   r[   c                 S   s   t | tjƒs
tdƒ‚dS )NzJWhen `allow_group_split` is set to False, X is required to be a dataframe.T)ra   r¤   r¥   r�   )rÆ   rr   rr   rs   Úcheck_dfÊ  s
   ÿz#DaskXGBRanker.fit.<locals>.check_dfr—   r}   c                 S   s(   t | tjƒs| d urtd|› d�ƒ‚dS )Nz*When `allow_group_split` is set to False, z is required to be a series.T)ra   r¤   r¨   r�   )r—   r}   rr   rr   rs   Ú	check_serÒ  s
   
ÿz$DaskXGBRanker.fit.<locals>.check_serrÑ  rÒ  r‘   )rÑ  rÒ  r‘   )rÆ   rÑ  r—   rÒ  r‘   rÓ  rØ  rÖ  r4  rÔ  rÕ  rš   )r£   rH   r   r¤   r¥   r   rc   r¨   r   r  rC   r²  rÞ   rÖ   rÄ  rã  )r   rÆ   rÑ  r–   r—   rÒ  r‘   rÓ  r×  rØ  rÖ  r4  rÔ  rÕ  rš   r¹   r  r  ÚX_idÚnew_eval_setÚnew_eval_qidÚnew_sample_weight_eval_setÚnew_base_margin_eval_setrë   ÚXeÚyeÚweÚbeÚqerr   rr   rs   rí  ±  s´   ÿÿ
þ
ÿþýüûú
	ÿþýüûÿ




€ÿÿózDaskXGBRanker.fit)rƒ   r„   r…   r1   rc   r÷   r   r  r   rz   r   r  rI   rH   r   r   r   rª   r5   r+   rã  rí  r6   r†   rˆ   rr   rr   r�   rs   rS   E  s²    ûýüûúùþýûúùø
	÷

ö
õ
ôóò
ñ@ðþýûúùø	÷

ö
õôó
ò
ñðï{rS   zjImplementation of the Scikit-Learn API for XGBoost Random Forest Regressor.

    .. versionadded:: 1.4.0

rÝ  zI
    n_estimators : int
        Number of trees in random forest to fit.
)rþ  c                       ó  e Zd Zeddddddœdee dee dee d	ee d
ee deddf‡ fdd„ƒZde	e
ef f‡ fdd„Zdefdd„Zdddddddddœdededee dee deeeeef   deeeef  deeee
ef  deee  deee  dee dd f‡ fdd „Z‡  ZS )!rT   rA   çš™™™™™é?çñhãˆµøä>N©Úlearning_rateÚ	subsampleÚcolsample_bynodeÚ
reg_lambdar  r  r  r  r  r  rý   r[   c                   ó"   t ƒ jd|||||dœ|¤Ž d S ©Nr  rr   ©ry   rz   ©r   r  r  r  r  r  rý   r�   rr   rs   rz   =  ó   û
úzDaskXGBRFRegressor.__init__c                    ó   t ƒ  ¡ }| j|d< |S ©NÚnum_parallel_tree©ry   rÚ  Ún_estimatorsr  r�   rr   rs   rÚ  Q  r  z!DaskXGBRFRegressor.get_xgb_paramsc                 C   ó   dS ©NrA   rr   r³   rr   rr   rs   rÞ  V  ó   z*DaskXGBRFRegressor.get_num_boosting_roundsTrä  rÆ   rÑ  rÒ  r‘   rÓ  rÖ  r4  rÔ  rÕ  rš   c                   ó8   dd„ t ƒ  ¡ D ƒ}t| j| jƒ tƒ jdi |¤Ž | S )Nc                 S   ræ  rç  rr   rè  rr   rr   rs   rÐ   h  rê  z*DaskXGBRFRegressor.fit.<locals>.<dictcomp>rr   ©rS  rÕ   r:   r2  r5  ry   rí  rì  r�   rr   rs   rí  Z  ó   zDaskXGBRFRegressor.fit©rƒ   r„   r…   r1   r   rö   r  r   rz   r
   rc   rÚ  rª   rÞ  rI   rH   r   r   r   r÷   r+   r5   rí  rˆ   rr   rr   r�   rs   rT   0  ón    ùýüûúùø	÷	ôþýûúùø	÷

ö
õôórT   zkImplementation of the Scikit-Learn API for XGBoost Random Forest Classifier.

    .. versionadded:: 1.4.0

c                       r  )!rU   rA   r  r  Nr  r  r  r  r  r  rý   r[   c                   r  r  r  r  r�   rr   rs   rz   {  r  zDaskXGBRFClassifier.__init__c                    r  r   r"  r  r�   rr   rs   rÚ  �  r  z"DaskXGBRFClassifier.get_xgb_paramsc                 C   r$  r%  rr   r³   rr   rr   rs   rÞ  ”  r&  z+DaskXGBRFClassifier.get_num_boosting_roundsTrä  rÆ   rÑ  rÒ  r‘   rÓ  rÖ  r4  rÔ  rÕ  rš   c                   r'  )Nc                 S   ræ  rç  rr   rè  rr   rr   rs   rÐ   ¦  rê  z+DaskXGBRFClassifier.fit.<locals>.<dictcomp>rr   r(  rì  r�   rr   rs   rí  ˜  r)  zDaskXGBRFClassifier.fitr*  rr   rr   r�   rs   rU   n  r+  rU   )NN)rQ  )�r†   ÚloggingÚcollectionsr   Ú
contextlibr   Ú	functoolsr   r   Ú	threadingr   Útypingr   r   r	   r
   r   r   r   r   r   r   r   r   r   r   r   r   r   rØ   r{   rŸ   r   r¦   r   r  r   r¤   Údask.delayedr   r   Ú r   r   Ú_data_utilsr    Ú_typingr!   r"   r#   Úcallbackr$   r%   r  r&   ÚCollArgsr'   r‡   Úcompatr(   r)   r*   Úcorer+   r,   r-   r.   r/   r0   r1   r2   Úsklearnr3   r4   r5   r6   r7   r8   r9   r:   r;   r<   r=   r>   Útrackerr?   Útrainingr@   rD  rœ   rB   rC   rj  rD   rE   rF   rG   r§   r¥   r¨   rH   Ú__annotations__rI   rJ   Ú__all__Ú	getLoggerrj   rª   rc   rl   rv   rN   r�   rO   rø   rù   r  rP   r  r   r'  r-  r÷   rP  rW  rb  rr  rƒ  r‡  rö   r®  r    rV   r±  rW   r¹  r»  rº  rQ   rR   rS   rT   rU   rr   rr   rr   rs   Ú<module>   s  4L(8
þþ
#ÿþý
ü0ÿþýü
û tÿ
þýü
ûG?üÿþýü
û)ÿÿ
þÿÿ
þ
þ
ýü
ûúùø	÷
ö
õô
óòñ
ðjüóÿ
þýüúùø	÷

ö
õôóò9ÿÿÿÿ
þ!þýüû
ú
ù
øUÿÿÿÿÿ
þÿÿ
þÿ
þýüúùø	÷
öõôóò
ñ óÿþ
ýûúùø	÷
öõôóò6þ
ýüûúùø	÷
öõ
ôAöÿþýûúùø	÷
öõ=ÿþýüû
úÿÿþ ÿZþ 8î Tù2ù