
    ^j              	       4   d 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 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Z"dedee!   de	e!ef   fdZ# G d de      Zd1de!de$de$defdZ%dede$defdZ&dede!fdZ'd2de!d e
ee!e$f      dejP                  fd!Z)de!de
e$   fd"Z*d#ede$fd$Z+d#ede,fd%Z-d#ede,fd&Z.d'ede$fd(Z/d)e!d*eg ef   defd+Z0d,ede!fd-Z1d)e!defd.Z2d/e
e!   de,fd0Z3y)3z;Xgboost pyspark integration submodule for helper functions.    N)Thread)AnyCallableDictOptionalSetTypeUnion)BarrierTaskContextTaskContext)SparkSession   )CommunicatorContext)Config)_Args)_ArgVals)Booster)XGBModel)RabitTrackerclsreturnc                 8    | j                    d| j                   S )zReturn the class name..)
__module____name__)r   s    ^/var/www/ramen.bs-engineer-server.com/venv/lib/python3.12/site-packages/xgboost/spark/utils.pyget_class_namer      s    nnQs||n--    funcunsupported_setc                     t        j                  |       }i }|j                  j                         D ]C  }|j                  |j
                  us|j                  |vs+|j                  ||j                  <   E |S )zReturns 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	parameters        r   _get_default_params_from_funcr,      su     

D
!C^^**, E	 Y__4o53<3D3D 0E  r   c                   0     e Zd ZdZdededdf fdZ xZS )r   z&Context with PySpark specific task ID.contextargsr   Nc                 \    t        |j                               |d<   t        |   di | y )Ndmlc_task_id )strpartitionIdsuper__init__)selfr.   r/   	__class__s      r   r6   zCommunicatorContext.__init__3   s+    "7#6#6#89^ 4 r   )r   r   __qualname____doc__r   CollArgsValsr6   __classcell__)r8   s   @r   r   r   0   s&    0! 2 !L !T ! !r   r   host	n_workersportc                     d|i}t        || d|      }|j                          t        |j                        }d|_        |j                          |j                  |j                                |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threads         r   _start_trackerrL   8   s`    !9-DYVRVWGMMO7++,FFM
LLNKK##%&Kr   confc                     | j                   J | j                  dn| j                  }t        | j                   ||      }|S )z3Get rabit context arguments to send to each worker.r   )tracker_host_iptracker_portrL   )rM   r>   r?   envs       r   _get_rabit_argsrR   D   sE    +++!!)1t/@/@D
--y$
?CJr   r.   c                     | j                         D cg c]   }|j                  j                  d      d   " }}|d   S c c}w )zLGets the hostIP for Spark. This essentially gets the IP of the first worker.:r   )getTaskInfosaddresssplit)r.   infotask_ip_lists      r   _get_host_iprZ   L   sB    ;B;O;O;QR4DLL&&s+A.RLR? Ss   %?r(   levelc                    t        j                  |       }||j                  |       n<|j                  t         j                  k(  r|j                  t         j
                         |j                  sxt        j                         j                  sZt        j                  t        j                        }t        j                  d      }|j                  |       |j                  |       |S )zGGets a logger by name, or creates and configures it for the first time.z<%(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	formatters        r   
get_loggerrl   R   s    t$F <<7>>)OOGLL)??7#4#4#6#?#?''

3%%J
	 	Y''"Mr   c                     t        j                  |       }|j                  t         j                  k(  rdS |j                  S )z+Get the logger level for the given log nameN)r]   r^   r[   r`   )r(   ri   s     r   get_logger_levelrn   f   s0    t$F<<7>>14Cv||Cr   spark_sessionc                    t        |       rt        j                  S | j                  j                  j                         j                  | j                  j                  j                         j                         j                  d            S )z0Gets the current max number of concurrent tasks.r   )	_is_connectrd   maxsizesparkContext_jscscmaxNumConcurrentTasksresourceProfileManagerresourceProfileFromIdro   s    r   _get_max_num_concurrent_tasksrz   l   sl     =!{{ %%**--/EE""''**,			!		q	! r   c                     	 t        | t        j                  j                  j                  j
                        S # t        $ r Y yw xY w)NF)
isinstancepysparksqlconnectsessionr   AttributeErrorry   s    r   rq   rq   }   s<    -)<)<)D)D)Q)QRR s   7: 	AAc                 v    | j                   j                  dd      }|duxr |dk(  xs |j                  d      S )zWhether it is Spark local modezspark.masterNlocalzlocal[)rM   get
startswith)ro   masters     r   	_is_localr      sA     ##ND9FT6W#4#S8I8I(8STr   task_contextc                     | t        d      | j                         }d|vrt        d      t        |d   j                  d   j	                               S )z&Get the gpu id from the task resourcesz3_get_gpu_id should not be invoked from driver side.gpuzDCouldn't get the gpu id, Please check the GPU resource configurationr   )RuntimeError	resourcesint	addressesstrip)r   r   s     r   _get_gpu_idr      s`    PQQ&&(IIR
 	
 y))!,22455r   modelxgb_model_creatorc                 f     |       }|j                  t        | j                  d                   |S )zH
    Deserialize an xgboost.XGBModel instance from the input model.
    utf-8)
load_model	bytearrayencode)r   r   	xgb_models      r   deserialize_xgb_modelr      s.     "#I5<<#89:r   boosterc                 B    | j                  d      j                  d      S )z
    Serialize the input booster to a string.

    Parameters
    ----------
    booster:
        an xgboost.core.Booster instance
    jsonr   )save_rawdecode)r   s    r   serialize_boosterr      s      F#**733r   c                 l    t               }|j                  t        | j                  d                   |S )zN
    Deserialize an xgboost.core.Booster from the input ser_model_string.
    r   )r   r   r   r   )r   r   s     r   deserialize_boosterr      s,     iGyg!678Nr   devicec                 
    | dv S )z&Whether xgboost is using CUDA workers.)cudar   r2   )r   s    r   use_cudar      s    _$$r   )r   )N)4r:   r"   r]   rd   	threadingr   typingr   r   r   r   r   r	   r
   r}   r   r   pyspark.sqlr   
collectiver   CCtxr   r   CollArgsr   r;   corer   sklearnr   rJ   r   r3   r   r,   r   rL   rR   rZ   Loggerrl   rn   rz   boolrq   r   r   r   r   r   r   r2   r   r   <module>r      s   A   
  B B B  3 $ 4  * 1   ". . .
 
 %(X 	#s(^ &!$ !	 	 	C 	 	& S X ,  S %S/!: gnn (D3 D8C= D # "|  U\ Ud U6k 6c 6#+BL#9	4w 	43 	4s w %Xc] %t %r   