o
    Ú­jÈ  ã                	   @   sd  d Z ddlZddlZddlZddlZddlZddlmZ ddlm	Z	m
Z
mZmZmZmZmZ ddlZddlmZmZmZ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# ddl$m%Z% dede&fdd„Z'de
dee& dee&e	f fdd„Z(G dd„ deƒZdDde&de)de)defdd„Z*dede)defdd „Z+d!ede&fd"d#„Z,defd$d%„Z-dEd&e&d'eee&e)f  dej.fd(d)„Z/d&e&dee) fd*d+„Z0d,ede)fd-d.„Z1d,ede2fd/d0„Z3dede2fd1d2„Z4d3ede)fd4d5„Z5de&fd6d7„Z6d8e&d9e
g e#f de#fd:d;„Z7d<e!de&fd=d>„Z8d8e&de!fd?d@„Z9dAee& de2fdBdC„Z:dS )Fz;Xgboost pyspark integration submodule for helper functions.é    N)ÚThread)ÚAnyÚCallableÚDictÚOptionalÚSetÚTypeÚUnion)ÚBarrierTaskContextÚ	SparkConfÚSparkContextÚ
SparkFilesÚTaskContext)ÚSparkSessioné   )ÚCommunicatorContext)ÚConfig)Ú_Args)Ú_ArgVals)ÚBooster)ÚXGBModel)ÚRabitTrackerÚclsÚreturnc                 C   s   | j › d| j› �S )zReturn the class name.Ú.)Ú
__module__Ú__name__)r   © r   úP/var/www/html/CropPilot/venv/lib/python3.10/site-packages/xgboost/spark/utils.pyÚget_class_name   s   r   ÚfuncÚunsupported_setc                 C   sD   t  | ¡}i }|j ¡ D ]}|j|jur|j|vr|j||j< q|S )z�Returns a dictionary of parameters and their default value of function fn.  Only
    the parameters with a default value will be included.

    )ÚinspectÚ	signatureÚ
parametersÚvaluesÚdefaultÚemptyÚname)r    r!   ÚsigÚfiltered_params_dictÚ	parameterr   r   r   Ú_get_default_params_from_func   s   

€r,   c                       s.   e Zd ZdZdededdf‡ fdd„Z‡  ZS )r   z&Context with PySpark specific task ID.ÚcontextÚargsr   Nc                    s&   t | ¡ ƒ|d< tƒ jdi |¤Ž d S )NÚdmlc_task_idr   )ÚstrÚpartitionIdÚsuperÚ__init__)Úselfr-   r.   ©Ú	__class__r   r   r3   5   s   zCommunicatorContext.__init__)r   r   Ú__qualname__Ú__doc__r
   ÚCollArgsValsr3   Ú__classcell__r   r   r5   r   r   2   s    "r   ÚhostÚ	n_workersÚportc                 C   sL   d|i}t || d|d�}| ¡  t|jd�}d|_| ¡  | | ¡ ¡ |S )z"Start Rabit tracker with n_workersr<   Útask)r<   Úhost_ipÚsortbyr=   )ÚtargetT)r   Ústartr   Úwait_forÚdaemonÚupdateÚworker_args)r;   r<   r=   r.   ÚtrackerÚthreadr   r   r   Ú_start_tracker:   s   rI   Úconfc                 C   s4   | j dusJ ‚| jdu rdn| j}t| j ||ƒ}|S )z3Get rabit context arguments to send to each worker.Nr   )Útracker_host_ipÚtracker_portrI   )rJ   r<   r=   Úenvr   r   r   Ú_get_rabit_argsF   s   rN   r-   c                 C   s   dd„ |   ¡ D ƒ}|d S )zLGets the hostIP for Spark. This essentially gets the IP of the first worker.c                 S   s   g | ]
}|j  d ¡d ‘qS )ú:r   )ÚaddressÚsplit)Ú.0Úinfor   r   r   Ú
<listcomp>P   s    z _get_host_ip.<locals>.<listcomp>r   )ÚgetTaskInfos)r-   Útask_ip_listr   r   r   Ú_get_host_ipN   s   rW   c                   C   s    t j ¡ durtdƒ‚tj ¡ S )z`Get or create spark session. Note: This function can only be invoked from driver
    side.

    Nz<_get_spark_session should not be invoked from executor side.)Úpysparkr   ÚgetÚRuntimeErrorr   ÚbuilderÚgetOrCreater   r   r   r   Ú_get_spark_sessionT   s
   ÿ
r]   r(   Úlevelc                 C   st   t  | ¡}|dur| |¡ n|jt jkr| t j¡ |js8t  ¡ js8t  tj	¡}t  
d¡}| |¡ | |¡ |S )zGGets a logger by name, or creates and configures it for the first time.Nz<%(asctime)s %(levelname)s %(name)s: %(funcName)s %(message)s)ÚloggingÚ	getLoggerÚsetLevelr^   ÚNOTSETÚINFOÚhandlersÚStreamHandlerÚsysÚstderrÚ	FormatterÚsetFormatterÚ
addHandler)r(   r^   ÚloggerÚhandlerÚ	formatterr   r   r   Ú
get_loggera   s   
ÿ

rn   c                 C   s    t  | ¡}|jt jkrdS |jS )z+Get the logger level for the given log nameN)r_   r`   r^   rb   )r(   rk   r   r   r   Úget_logger_levelu   s   
ro   Úspark_contextc                 C   s@   | j  ¡  ¡ dkr| j  ¡  | j  ¡  ¡  d¡¡S | j  ¡  ¡ S )z0Gets the current max number of concurrent tasks.z3.1r   )Ú_jscÚscÚversionÚmaxNumConcurrentTasksÚresourceProfileManagerÚresourceProfileFromId©rp   r   r   r   Ú_get_max_num_concurrent_tasks{   s
   
ÿrx   c                 C   s   | j  ¡  ¡ S )zWhether it is Spark local mode)rq   rr   ÚisLocalrw   r   r   r   Ú	_is_local†   s   rz   c                 C   s&   |   d¡}|d uo| d¡p| d¡S )Nzspark.masterzspark://zlocal-cluster)rY   Ú
startswith)rJ   Úmasterr   r   r   Ú_is_standalone_or_localclusterŒ   s   
ÿr}   Útask_contextc                 C   s>   | du rt dƒ‚|  ¡ }d|vrt dƒ‚t|d jd  ¡ ƒS )z&Get the gpu id from the task resourcesNz3_get_gpu_id should not be invoked from driver side.ÚgpuzDCouldn't get the gpu id, Please check the GPU resource configurationr   )rZ   Ú	resourcesÚintÚ	addressesÚstrip)r~   r€   r   r   r   Ú_get_gpu_id“   s   ÿr„   c                  C   s0   t  ¡ } tj | d¡}tj |¡st |¡ |S )Nzxgboost-tmp)r   ÚgetRootDirectoryÚosÚpathÚjoinÚexistsÚmakedirs)Úroot_dirÚxgb_tmp_dirr   r   r   Ú_get_or_create_tmp_dir¡   s
   
r�   ÚmodelÚxgb_model_creatorc                 C   s   |ƒ }|  t|  d¡ƒ¡ |S )zH
    Deserialize an xgboost.XGBModel instance from the input model.
    úutf-8)Ú
load_modelÚ	bytearrayÚencode)rŽ   r�   Ú	xgb_modelr   r   r   Údeserialize_xgb_model©   s   r•   Úboosterc                 C   s^   t j tƒ t ¡ › d�¡}|  |¡ t|dd��}| ¡ }W d  ƒ |S 1 s(w   Y  |S )z‡
    Serialize the input booster to a string.

    Parameters
    ----------
    booster:
        an xgboost.core.Booster instance
    ú.jsonr�   ©ÚencodingN)	r†   r‡   rˆ   r�   ÚuuidÚuuid4Ú
save_modelÚopenÚread)r–   Útmp_file_nameÚfÚser_model_stringr   r   r   Úserialize_booster´   s   



ÿþr¢   c                 C   sf   t ƒ }tj tƒ t ¡ › d�¡}t|ddd��}| | ¡ W d  ƒ n1 s'w   Y  | 	|¡ |S )zN
    Deserialize an xgboost.core.Booster from the input ser_model_string.
    r—   Úwr�   r˜   N)
r   r†   r‡   rˆ   r�   rš   r›   r�   Úwriter‘   )rŽ   r–   rŸ   r    r   r   r   Údeserialize_boosterÅ   s   ÿ
r¥   Údevicec                 C   s   | dv S )z&Whether xgboost is using CUDA workers.)Úcudar   r   )r¦   r   r   r   Úuse_cudaÒ   s   r¨   )r   )N);r8   r"   r_   r†   rf   rš   Ú	threadingr   Útypingr   r   r   r   r   r   r	   rX   r
   r   r   r   r   Úpyspark.sql.sessionr   Ú
collectiver   ÚCCtxr   r   ÚCollArgsr   r9   Úcorer   Úsklearnr   rG   r   r0   r   r,   r�   rI   rN   rW   r]   ÚLoggerrn   ro   rx   Úboolrz   r}   r„   r�   r•   r¢   r¥   r¨   r   r   r   r   Ú<module>   s`    $ÿÿ

þ&ÿ
ÿ
þ