o
    Ú­j‘ ã                   @   s>  d 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 ddl	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mZmZmZ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&m'Z'm(Z( ddl)m*Z*m+Z+m,Z,m-Z-m.Z.m/Z/m0Z0 ddl1m2Z2m3Z3m4Z4m5Z5m6Z6m7Z7 ddl8m9Z9m:Z: ddl;m<Z<m=Z= ddl>m?Z?m@Z@mAZAmBZBmCZC ddlDmEZEmFZFmGZGmHZHmIZImJZJmKZKmLZL ddlMmNZNmOZO ddlPmQZQ ddlRmSZS ddlTmUZUmVZVmWZW ddlXmYZYmZZZ ddl[m\Z\m]Z]m^Z^ ddl_m`Z`maZambZbmcZc ddldmeZf ddlgmhZhmiZimjZjmkZkmlZl ddlmmnZnmoZompZpmqZqmrZrmsZs ddltmuZu ddlvmwZwmxZxmyZymzZzm{Z{m|Z|m}Z}m~Z~mZm€Z€m�Z�m‚Z‚mƒZƒm„Z„m…Z…m†Z† g d¢Z‡g d ¢Zˆd!d"d#d$d%d&d'd(œZ‰d)d*„ e‰ Š¡ D ƒZ‹g d+¢ZŒh d,£Z�d-d.hZŽh d/£Z�d0d1iZ�ed2d3ƒZ‘e‘d4d5d6d7ƒZ’d8Z“d9Z”G d:d;„ d;e*e+e0e,e/eneoereqesepƒZ•d<e=d=ee– d>ee< fd?d@„Z—d<e=dAe–d>e<fdBdC„Z˜d>eee<e–f ge<f fdDdE„Z™dFe<d>ee< fdGdH„ZšedIdJƒZ›dKZœG dLdM„ dMee•e4e6ƒZ�G dNdO„ dOee•e4e6ƒZžG dPdQ„ dQeže-e.epƒZŸG dRdS„ dSƒZ G dTdU„ dUe7ƒZ¡G dVdW„ dWe5ƒZ¢G dXdY„ dYe7ƒZ£G dZd[„ d[e5ƒZ¤dS )\z4XGBoost pyspark integration submodule for core code.é    N)Ú
namedtuple)Úasdict)
ÚAnyÚCallableÚDictÚIteratorÚListÚOptionalÚTupleÚTypeÚUnionÚcast)ÚRDDÚ	SparkConfÚSparkContextÚcloudpickle)Ú	EstimatorÚModel)Úarray_to_vectorÚvector_to_array)Ú	VectorUDT)ÚParamÚParamsÚTypeConverters)ÚHasFeaturesColÚHasLabelColÚHasPredictionColÚHasProbabilityColÚHasRawPredictionColÚHasValidationIndicatorColÚHasWeightCol)ÚDefaultParamsReaderÚDefaultParamsWriterÚ
MLReadableÚMLReaderÚ
MLWritableÚMLWriter)ÚResourceProfileBuilderÚTaskResourceRequests)ÚColumnÚ	DataFrame)ÚcolÚcountDistinctÚ
pandas_udfÚrandÚstruct)Ú	ArrayTypeÚBooleanTypeÚ
DoubleTypeÚ	FloatTypeÚIntegerTypeÚIntegralTypeÚLongTypeÚ	ShortType)ÚexpitÚsoftmaxé   )Ú	ArrayLike)ÚConfig)Úimport_cupyÚis_cudf_availableÚis_cupy_available)Úconfig_contextÚ
get_config)ÚBoosterÚ_check_distributed_paramsÚ_py_version)ÚDEFAULT_N_ESTIMATORSÚXGBClassifierÚXGBModelÚ_can_use_qdm)Útrainé   )Ú)_read_csr_matrix_from_unwrapped_spark_vecÚaliasÚcreate_dmatrix_from_partitionsÚpred_contribsÚstack_series)ÚHasArbitraryParamsDictÚHasBaseMarginColÚHasContribPredictionColÚHasEnableSparseDataOptimÚHasFeaturesColsÚHasQueryIdCol)ÚXGBoostTrainingSummary)ÚCommunicatorContextÚ_get_default_params_from_funcÚ_get_gpu_idÚ_get_host_ipÚ_get_max_num_concurrent_tasksÚ_get_rabit_argsÚ_get_spark_sessionÚ	_is_localÚ_is_standalone_or_localclusterÚdeserialize_boosterÚdeserialize_xgb_modelÚget_class_nameÚ
get_loggerÚget_logger_levelÚserialize_boosterÚuse_cuda)ÚfeaturesColÚlabelColÚ	weightColÚrawPredictionColÚpredictionColÚprobabilityColÚvalidationIndicatorColÚbase_margin_colÚarbitrary_params_dictÚforce_repartitionÚnum_workersÚfeature_namesÚfeatures_colsÚenable_sparse_data_optimÚqid_colÚrepartition_random_shuffleÚpred_contrib_colÚlaunch_tracker_on_driverÚcoll_cfg)ÚmissingÚn_estimatorsÚfeature_typesÚfeature_weightsrg   rh   ri   rj   rk   rl   rm   )Úfeatures_colÚ	label_colÚ
weight_colÚraw_prediction_colÚprediction_colÚprobability_colÚvalidation_indicator_colc                 C   s   i | ]\}}||“qS © r…   ©Ú.0ÚkÚvr…   r…   úO/var/www/html/CropPilot/venv/lib/python3.10/site-packages/xgboost/spark/core.pyÚ
<dictcomp>Ž   s    r‹   )Úenable_categoricalÚn_jobsÚnthread>	   ÚqidÚgroupÚeval_qidÚeval_setÚ
eval_groupÚbase_marginÚsample_weightÚbase_margin_eval_setÚsample_weight_eval_setÚevalsÚevals_result>   r”   Úoutput_marginÚvalidate_featuresrŒ   z—`xgboost.spark` estimators do not have 'enable_categorical' param, but you can set `feature_types` param and mark categorical features with 'c' string.ÚPred)Ú
predictionÚraw_predictionÚprobabilityÚpred_contribr�   ÚrawPredictionrŸ   ÚpredContribzinit_booster.jsonzXGBoost-PySparkc                	   @   sÐ  e Zd Zee ¡ ddejƒZee ¡ ddej	ƒZ
ee ¡ ddejƒZee ¡ ddejƒZee ¡ d	d
ejƒZee ¡ ddejƒZee ¡ ddejƒZdedd fdd„Zdedd fdd„Zedee fdd„ƒZedeeef fdd„ƒZd9dd„Z	d:dedeeef fdd„Z edeeef fd d!„ƒZ!d9d"d#„Z"deeef fd$d%„Z#edeeef fd&d'„ƒZ$d9d(d)„Z%deeef fd*d+„Z&	d:d,ed-e'd.eddfd/d0„Z(d9d1d2„Z)defd3d4„Z*d5d6defd7d8„Z+dS );Ú_SparkXGBParamsrq   zQThe number of XGBoost workers. Each XGBoost worker corresponds to one spark task.ÚdevicezÒThe device type for XGBoost executors. Available options are `cpu`,`cuda` and `gpu`. Set `device` to `cuda` or `gpu` if the executors are running on GPU instances. Currently, only one GPU per task is supported.rp   z÷A boolean variable. Set force_repartition=true if you want to force the input dataset to be repartitioned before XGBoost training.Note: The auto repartitioning judgement is not fully accurate, so it is recommendedto have force_repartition be True.rv   z’A boolean variable. Set repartition_random_shuffle=true if you want to random shuffle dataset when repartitioning is required. By default is True.rr   z'A list of str to specify feature names.rx   z¨A boolean variable. Set launch_tracker_on_driver to true if you want the tracker to be launched on the driver side; otherwise, it will be launched on the executor side.ry   z8xgboost.collective.Config. The collective configuration.ÚvalueÚreturnc                 C   s    t |tƒsJ ‚|  | j|¡ | S )zSet collective configuration)Ú
isinstancer<   Úsetry   ©Úselfr¥   r…   r…   rŠ   Úset_coll_cfg   s   z_SparkXGBParams.set_coll_cfgc                 C   s*   t d|iƒ |dv sJ ‚|  | j|¡ | S )z*Set device, optional value: cpu, cuda, gpur¤   )ÚcpuÚcudaÚgpu)rC   r¨   r¤   r©   r…   r…   rŠ   Ú
set_device  s   z_SparkXGBParams.set_devicec                 C   ó   t ƒ ‚)zi
        Subclasses should override this method and
        returns an xgboost.XGBModel subclass
        ©ÚNotImplementedError©Úclsr…   r…   rŠ   Ú_xgb_cls  ó   z_SparkXGBParams._xgb_clsc                    s0   |   ¡ ƒ }| ¡ ‰ ‡ fdd„ˆ D ƒ}t|d< |S )zGGet the xgboost.sklearn.XGBModel default parameters and filter out somec                    s   i | ]}|t vr|ˆ | “qS r…   )Ú_unsupported_xgb_params)r‡   rˆ   ©Úparams_dictr…   rŠ   r‹     ó    z;_SparkXGBParams._get_xgb_params_default.<locals>.<dictcomp>r{   )rµ   Ú
get_paramsrE   )r´   Úxgb_model_defaultÚfiltered_params_dictr…   r¸   rŠ   Ú_get_xgb_params_default  s   

ÿz'_SparkXGBParams._get_xgb_params_defaultNc                 C   ó   |   ¡ }| jdi |¤Ž dS )z,Set xgboost parameters into spark parametersNr…   )r¾   Ú_setDefault©rª   r½   r…   r…   rŠ   Ú_set_xgb_params_default!  ó   z'_SparkXGBParams._set_xgb_params_defaultFÚgen_xgb_sklearn_estimator_paramc                 C   sz   i }t tƒ|  ¡  ¡ B |  ¡  ¡ B }|s|t tƒO }|  ¡ D ]}|j|vr-|  |¡||j< q|  |  	d¡¡}| 
|¡ |S )zIGenerate the xgboost parameters which will be passed into xgboost libraryro   )r¨   Ú_pyspark_specific_paramsÚ_get_fit_params_defaultÚkeysÚ_get_predict_params_defaultÚ_non_booster_paramsÚextractParamMapÚnameÚgetOrDefaultÚgetParamÚupdate)rª   rÄ   Ú
xgb_paramsÚnon_xgb_paramsÚparamro   r…   r…   rŠ   Ú_gen_xgb_params_dict&  s$   
ÿ
þÿ
€ÿ
z$_SparkXGBParams._gen_xgb_params_dictc                 C   ó   t |  ¡ jtƒ}|S )z+Get the xgboost.XGBModel().fit() parameters)rX   rµ   ÚfitÚ_unsupported_fit_params)r´   Ú
fit_paramsr…   r…   rŠ   rÆ   =  ó   
ÿz'_SparkXGBParams._get_fit_params_defaultc                 C   r¿   )zLGet the xgboost.XGBModel().fit() parameters and set them to spark parametersNr…   )rÆ   rÀ   rÁ   r…   r…   rŠ   Ú_set_fit_params_defaultE  rÃ   z'_SparkXGBParams._set_fit_params_defaultc                 C   ó<   |   ¡  ¡ }i }|  ¡ D ]}|j|v r|  |¡||j< q|S )zBGenerate the fit parameters which will be passed into fit function)rÆ   rÇ   rÊ   rË   rÌ   )rª   Úfit_params_keysrÖ   rÑ   r…   r…   rŠ   Ú_gen_fit_params_dictJ  ó   
€z$_SparkXGBParams._gen_fit_params_dictc                 C   rÓ   )z4Get the parameters from xgboost.XGBModel().predict())rX   rµ   ÚpredictÚ_unsupported_predict_params)r´   Úpredict_paramsr…   r…   rŠ   rÈ   S  r×   z+_SparkXGBParams._get_predict_params_defaultc                 C   r¿   )z_Get the parameters from xgboost.XGBModel().predict() and
        set them into spark parametersNr…   )rÈ   rÀ   rÁ   r…   r…   rŠ   Ú_set_predict_params_default[  s   z+_SparkXGBParams._set_predict_params_defaultc                 C   rÙ   )zRGenerate predict parameters which will be passed into xgboost.XGBModel().predict())rÈ   rÇ   rÊ   rË   rÌ   )rª   Úpredict_params_keysrß   rÑ   r…   r…   rŠ   Ú_gen_predict_params_dicta  rÜ   z(_SparkXGBParams._gen_predict_params_dictÚspark_versionÚconfÚis_localc                 C   sÖ   |   ¡ re|rt| jjƒ d|  | j¡¡ dS | d¡}|du r#tdƒ‚| d¡}|dur<t	|ƒdkr<t| jjƒ d|¡ |dk sQd|  krJd	k rgn dS t
|ƒsi|durat	|ƒdk r_td
ƒ‚dS tdƒ‚dS dS dS )z2Validate the gpu parameters and gpu configurationsz_You have enabled GPU in spark local mode. Please make sure your local node has at least %d GPUsú"spark.executor.resource.gpu.amountNzIThe `spark.executor.resource.gpu.amount` is required for training on GPU.úspark.task.resource.gpu.amountç      ð?z’The configuration assigns %s GPUs to each Spark task, but each XGBoost training task only utilizes 1 GPU, which will lead to unnecessary GPU wasteú3.4.0ú3.5.1a  XGBoost doesn't support GPU fractional configurations. Please set `spark.task.resource.gpu.amount=spark.executor.resource.gpu.amount`. To enable GPU fractional configurations, you can try standalone/localcluster with spark 3.4.0+ andYARN/K8S with spark 3.5.1+zEThe `spark.task.resource.gpu.amount` is required for training on GPU.)Ú_run_on_gpurc   Ú	__class__Ú__name__ÚwarningrÌ   rq   ÚgetÚ
ValueErrorÚfloatr_   )rª   rã   rä   rå   Úexecutor_gpusÚgpu_per_taskr…   r…   rŠ   Ú_validate_gpu_paramsj  s@   
ý
ÿ
üÿÿÿ	ÿÐ#z$_SparkXGBParams._validate_gpu_paramsc                 C   sd  |   d¡}|d urt|tƒstdƒ‚|   | j¡dk r&td|   | j¡› d�ƒ‚|   |  d¡¡}|dkr6tdƒ‚|   d	¡d urIt|   d	¡tƒsItd
ƒ‚d}|   |¡d urrt|   |¡tƒsrt|   |¡tƒrntdd„ |   |¡D ƒƒsrtdƒ‚|   d¡d urƒ|  	| j
¡sƒtdƒ‚|   | j¡rž|   d¡dkr”tdƒ‚|   | j¡ržtdƒ‚tƒ }|j}|  |j| ¡ t|ƒ¡ d S )NÚ	xgb_modelzGThe xgb_model param must be set with a `xgboost.core.Booster` instance.rJ   zNumber of workers was z(.It cannot be less than 1 [Default is 1]Útree_methodÚexactzAThe `exact` tree method is not supported for distributed systems.Ú	objectivez.Only string type 'objective' param is allowed.Úeval_metricc                 s   s   � | ]}t |tƒV  qd S ©N)r§   Ústr)r‡   Úmetricr…   r…   rŠ   Ú	<genexpr>Ã  s
   € ÿ
ÿz3_SparkXGBParams._validate_params.<locals>.<genexpr>zGOnly string type or list of string type 'eval_metric' param is allowed.Úearly_stopping_roundszbIf 'early_stopping_rounds' param is set, you need to set 'validation_indicator_col' param as well.rz   g        zIIf enable_sparse_data_optim is True, missing param != 0 is not supported.z¢If enable_sparse_data_optim is True, you cannot set multiple feature columns but you should set one feature column with values of `pyspark.ml.linalg.Vector` type.)rÌ   r§   rB   rð   rq   rÍ   rû   r   ÚallÚ_col_is_defined_not_emptyrm   rt   rs   r]   ÚsparkContextrô   ÚversionÚgetConfr^   )rª   Ú
init_modelrö   rù   ÚssÚscr…   r…   rŠ   Ú_validate_params¤  s^   
ÿÿÿÿýþü
ÿÿ	ÿÿz _SparkXGBParams._validate_paramsc                 C   s   t |  | j¡ƒS )z<If train or transform on the gpu according to the parameters)rf   rÌ   r¤   ©rª   r…   r…   rŠ   rë   ì  s   z_SparkXGBParams._run_on_gpurÑ   z
Param[str]c                 C   s   |   |¡o|  |¡dvS )N)NÚ )Ú	isDefinedrÌ   )rª   rÑ   r…   r…   rŠ   r   ñ  s   z)_SparkXGBParams._col_is_defined_not_empty©r¦   N)F),rí   Ú
__module__Ú__qualname__r   r   Ú_dummyr   ÚtoIntrq   ÚtoStringr¤   Ú	toBooleanrp   rv   ÚtoListrr   rx   Úidentityry   r<   r«   rû   r¯   Úclassmethodr   rG   rµ   r   r   r¾   rÂ   ÚboolrÒ   rÆ   rØ   rÛ   rÈ   rà   râ   r   rô   r  rë   r   r…   r…   r…   rŠ   r£   ¿   sš    üø
ù	ûüûü

ÿÿ

þ
	

ÿÿÿÿ
þ
:Hr£   ÚdatasetÚfeatures_col_namesr¦   c                 C   sn   g }|D ]0}t | j| jtƒr| t|ƒ tƒ ¡ |¡¡ qt | j| jtt	fƒr1| t|ƒ¡ qt
dƒ‚|S )zFValues in feature columns must be integral types or float/double typeszGValues in feature columns must be integral types or float/double types.)r§   ÚschemaÚdataTyper2   Úappendr+   r   r3   rL   r5   rð   )r  r  Úfeature_colsÚcr…   r…   rŠ   Ú3_validate_and_convert_feature_col_as_float_col_listõ  s   ÿr  Úfeatures_col_namec                 C   sŒ   | j | j}t|ƒ}t|tƒr1t|jtttt	t
fƒs#td|j› d�ƒ‚| ttƒ ƒ¡ tj¡}|S t|tƒrBt|dd� tj¡}|S tdƒ‚)zQIt handles
    1. Convert vector type to array type
    2. Cast to Array(Float32)zGIf feature column is array type, its elements must be number type, got Ú.Úfloat32)Údtypezãfeature column must be array type or `pyspark.ml.linalg.Vector` type, if you want to use multiple numetric columns as features, please use `pyspark.ml.transform.VectorAssembler` to assemble them into a vector type column first.)r  r  r+   r§   r0   ÚelementTyper2   r3   r6   r4   r7   rð   r   rL   Údatar   r   )r  r  Úfeatures_col_datatyper~   Úfeatures_array_colr…   r…   rŠ   Ú._validate_and_convert_feature_col_as_array_col  s,   
þÿÿ
õÿ
úÿr&  c               
   C   s\   z	ddl m}  | W S  ty   Y nw z	ddlm} |W S  ty- } ztdƒ|‚d }~ww )Nr   )Ú
unwrap_udtzfCannot import pyspark `unwrap_udt` function. Please install pyspark>=3.4 or run on Databricks Runtime.)Úpyspark.sql.functionsr'  ÚImportErrorÚ pyspark.databricks.sql.functionsÚRuntimeError)r'  Údatabricks_unwrap_udtÚexcr…   r…   rŠ   Ú_get_unwrap_udt_fn&  s"   ÿÿý€ÿr.  Úfeature_colc                 C   sF   t ƒ }|| ƒ}|j d¡|j d¡|j d¡|j ttƒ ƒ¡ d¡gS )NÚfeatureVectorTypeÚfeatureVectorSizeÚfeatureVectorIndicesÚfeatureVectorValues)	r.  ÚtyperL   ÚsizeÚindicesÚvaluesr   r0   r3   )r/  r'  Úfeatures_unwrapped_vec_colr…   r…   rŠ   Ú_get_unwrapped_vec_cols9  s   


ÿúr9  ÚFeatureProp)rt   Úhas_validation_colÚfeatures_cols_namesi  @ c                	       sÂ  e Zd ZU eeef ed< d3‡ fdd„Zdeddfdd„Ze	de
d	 fd
d„ƒZdededd	fdd„Zdededefdd„Zdedefdd„Zdedeeef fdd„Ze	deeef deeeef eeef f fdd„ƒZdedeee ef fdd„Zdedeeef fdd„Zdedeeeef eeef eeef f fd d!„Zd"ed#edefd$d%„Zd&edefd'd(„Z deeeeef f fd)d*„Z!dedd	fd+d,„Z"d4d.d/„Z#e	d5d1d2„ƒZ$‡  Z%S )6Ú_SparkXGBEstimatorÚ_input_kwargsr¦   Nc                    sP   t ƒ  ¡  |  ¡  |  ¡  |  ¡  | jddddd d d i dd�	 t| jjƒ| _	d S )NrJ   r¬   FT)	rq   r¤   rp   rv   rr   r|   r}   ro   rx   )
ÚsuperÚ__init__rÂ   rØ   rà   rÀ   rc   rì   rí   Úloggerr  ©rì   r…   rŠ   r@  _  s    
÷z_SparkXGBEstimator.__init__Úkwargsc                 K   sV  i }d|v r
t dƒ‚| ¡ D ]†\}}|| jjkr t d|› d�ƒ‚|tv r.t dt| › d�ƒ‚|tv rL|t| jj krFt|tƒrF| jj}|}nt| }|}|  	|¡rr|dkret|tƒre| j
di d|i¤Ž q| j
di t|ƒ|i¤Ž q|tv s‚|tv s‚|tv s‚|tv r�t |d|› d	�¡}t |ƒ‚|||< qt|ƒ |  | j¡}| j
i |¥|¥d
� dS )z/
        Set params for the estimator.
        ro   z,Invalid param name: 'arbitrary_params_dict'.zUnsupported param 'z"' please use features_col instead.zPlease use param name z	 instead.r~   rs   z'.)ro   Nr…   )rð   Úitemsrs   rË   Ú _inverse_pyspark_param_alias_mapÚ_pyspark_param_alias_maprg   r§   ÚlistÚhasParamÚ_setrû   r·   rÕ   rÞ   Ú_unsupported_train_paramsÚ _unsupported_params_hint_messagerï   rC   rÌ   ro   )rª   rC  Ú_extra_paramsrˆ   r‰   Úreal_kÚerr_msgÚ_existing_extra_paramsr…   r…   rŠ   Ú	setParamsu  sL   
ÿÿÿþ
ÿ
z_SparkXGBEstimator.setParamsÚ_SparkXGBModelc                 C   r°   )zf
        Subclasses should override this method and
        returns a _SparkXGBModel subclass
        r±   r³   r…   r…   rŠ   Ú_pyspark_model_cls§  r¶   z%_SparkXGBEstimator._pyspark_model_clsrõ   Útraining_summaryc                 C   s   |   ¡ ||ƒS rú   )rR  )rª   rõ   rS  r…   r…   rŠ   Ú_create_pyspark_model¯  s   z(_SparkXGBEstimator._create_pyspark_modelÚboosterÚconfigc                 C   s8   | j dd�}|  ¡ di |¤Ž}| |¡ |j |¡ |S )NT©rÄ   r…   )rÒ   rµ   Ú
load_modelÚ_BoosterÚload_config)rª   rU  rV  Úxgb_sklearn_paramsÚsklearn_modelr…   r…   rŠ   Ú_convert_to_sklearn_model´  s   ÿ
z,_SparkXGBEstimator._convert_to_sklearn_modelr  c                 C   s0   |   | j¡rdS |   | j¡}|j ¡ }||k S )zn
        We repartition the dataset if the number of workers is not equal to the number of
        partitions.T)rÌ   rp   rq   ÚrddÚgetNumPartitions)rª   r  rq   Únum_partitionsr…   r…   rŠ   Ú_repartition_needed½  s
   

z&_SparkXGBEstimator._repartition_neededc                 C   s¢   |   ¡ }|  ¡ }| dd¡}| |¡ ||d< |  ¡ tk}|rAt| tt	j
ƒ¡ ¡ d d ƒ}|dkr8d|d< nd|d< ||d	< n|  d¡|d< |  d
¡|d< |S )zQ
        This just gets the configuration params for distributed xgboost
        ÚverboseNÚverbose_evalr   r:   zbinary:logisticrø   zmulti:softprobÚ	num_classr{   Únum_boost_round)rÒ   rÛ   ÚpoprÎ   rµ   rF   ÚintÚselectr,   rL   ÚlabelÚcollectrÌ   )rª   r  ÚparamsrÖ   rc  ÚclassificationÚnum_classesr…   r…   rŠ   Ú_get_distributed_train_paramsÇ  s"   
ÿ

z0_SparkXGBEstimator._get_distributed_train_paramsÚtrain_paramsc                 C   sZ   t ttƒ}i i }}| ¡ D ]\}}||v r|||< q|||< qdd„ | ¡ D ƒ}||fS )Nc                 S   s   i | ]\}}|t vr||“qS r…   )rÉ   r†   r…   r…   rŠ   r‹   ó  rº   z?_SparkXGBEstimator._get_xgb_train_call_args.<locals>.<dictcomp>)rX   Úworker_trainrJ  rD  )r´   ro  Úxgb_train_default_argsÚbooster_paramsÚkwargs_paramsÚkeyr¥   r…   r…   rŠ   Ú_get_xgb_train_call_argså  s   ÿ


ÿz+_SparkXGBEstimator._get_xgb_train_call_argsc                 C   s~  t |  | j¡ƒ tj¡}|g}d }|  | j¡}|r8|  | j¡}|j| j}t	|t
ƒs.tdƒ‚| tt |ƒƒ¡ n%|  | j¡rO|  | j¡}t||ƒ}| |¡ nt||  | j¡ƒ}	| |	¡ |  | j¡rr| t |  | j¡ƒ tj¡¡ d}
|  | j¡r‹| t |  | j¡ƒ tj¡¡ d}
|  | j¡r | t |  | j¡ƒ tj¡¡ |  | j¡rµ| t |  | j¡ƒ tj¡¡ t||
|ƒ}||fS )NzgIf enable_sparse_data_optim is True, the feature column values must be `pyspark.ml.linalg.Vector` type.FT)r+   rÌ   rh   rL   ri  rt   rg   r  r  r§   r   rð   Úextendr9  rs   r  r&  r  r   ri   Úweightrm   Úvalidrn   Úmarginru   r�   r:  )rª   r  r   Úselect_colsr<  rt   r  r$  rs   r%  r;  Úfeature_propr…   r…   rŠ   Ú'_prepare_input_columns_and_feature_propø  sT   
ÿÿÿ
ÿÿÿÿz:_SparkXGBEstimator._prepare_input_columns_and_feature_propc                 C   sè   |   |¡\}}|j|Ž }|  | j¡}tƒ j}t|ƒ}|jr-|jt	j
 j}t|tƒs-tdƒ‚||kr;t| jjƒ d|¡ |  |¡rb|  | j¡rN| |t	j¡}n|  | j¡r]| |tdƒ¡}n| |¡}|  | j¡rp|jt	jdd�}||fS )zAPrepare the input including column pruning, repartition and so onz.The validation indicator must be boolean type.zÔThe num_workers %s set for xgboost distributed training is greater than current max number of concurrent spark task slots, you need wait until more task slots available or you need increase spark cluster workers.rJ   T)Ú	ascending)r|  rh  rÌ   rq   r]   r  r[   r;  r  rL   rx  r  r§   r1   Ú	TypeErrorrc   rì   rí   rî   ra  r   ru   ÚrepartitionByRanger�   rv   Úrepartitionr.   ÚsortWithinPartitions)rª   r  rz  r{  rq   r  Úmax_concurrent_tasksr!  r…   r…   rŠ   Ú_prepare_input2  s2   ÿ

û

z!_SparkXGBEstimator._prepare_inputc                 C   s¸   |   |¡}|  |¡\}}ttƒ j ¡  dd¡ƒ}||  d¡|  d¡|  d¡t|  d¡ƒdœ}|d d ur8d|d	< ||d
< dd„ | 	¡ D ƒ}dd„ | 	¡ D ƒ}dd„ | 	¡ D ƒ}|||fS )Nzspark.task.cpusÚ1r|   rr   r}   rz   )rŽ   r|   rr   r}   rz   TrŒ   rŽ   c                 S   ó   i | ]\}}|d ur||“qS rú   r…   r†   r…   r…   rŠ   r‹   z  ó    z:_SparkXGBEstimator._get_xgb_parameters.<locals>.<dictcomp>c                 S   r…  rú   r…   r†   r…   r…   rŠ   r‹   {  rº   c                 S   r…  rú   r…   r†   r…   r…   rŠ   r‹   ~  r†  )
rn  ru  rg  r]   r  r  rï   rÌ   rñ   rD  )rª   r  ro  rr  Útrain_call_kwargs_paramsÚcpu_per_taskÚdmatrix_kwargsr…   r…   rŠ   Ú_get_xgb_parametersc  s,   
ÿÿûÿ
z&_SparkXGBEstimator._get_xgb_parametersrã   rä   c                 C   sð   |   ¡ rv|dk r| j d¡ dS d|  krdk r)n nt|ƒs)| j d|¡ dS | d¡}| d¡}|du s;|du rC| j d	¡ dS t|ƒd
krQ| j d¡ dS t|ƒd
kr_| j d¡ dS | d¡}|du rjdS t|ƒt|ƒkrtdS dS dS )zaCheck if stage-level scheduling is not needed,
        return true to skip stage-level schedulingré   z?Stage-level scheduling in xgboost requires spark version 3.4.0+Trê   zYFor %s, Stage-level scheduling in xgboost requires spark standalone or local-cluster modeúspark.executor.coresræ   NznStage-level scheduling in xgboost requires spark.executor.cores, spark.executor.resource.gpu.amount to be set.rJ   zDStage-level scheduling in xgboost requires spark.executor.cores > 1 zYStage-level scheduling in xgboost will not work when spark.executor.resource.gpu.amount>1rç   F)rë   rA  Úinfor_   rï   rg  rñ   )rª   rã   rä   Úexecutor_coresrò   Útask_gpu_amountr…   r…   rŠ   Ú_skip_stage_level_scheduling‚  sL   ÿÿý

ÿÿÿ
z/_SparkXGBEstimator._skip_stage_level_schedulingr^  c                 C   sâ   t ƒ }|j ¡ }t|jƒs|  |j|¡r|S | d¡}|dus!J ‚|j dd¡}|dus.J ‚|j dd¡}|dus;J ‚d|v rId| ¡ krIt	|ƒnt	|ƒd d	 }d
}t
ƒ  |¡ d|¡}	tƒ  |	¡j}
| j d||¡ | |
¡S )z$Try to enable stage-level schedulingr‹  Nzspark.pluginsú zspark.rapids.sql.enabledÚtruezcom.nvidia.spark.SQLPluginr:   rJ   rè   r®   z>XGBoost training tasks require the resource(cores=%s, gpu=%s).)r]   r  r  r^   r�  r  rï   rä   Úlowerrg  r(   ÚcpusÚresourcer'   ÚrequireÚbuildrA  rŒ  ÚwithResources)rª   r^  r  rä   r�  Úspark_pluginsÚspark_rapids_sql_enabledÚ
task_coresÚ	task_gpusÚtreqsÚrpr…   r…   rŠ   Ú_try_stage_level_schedulingÅ  s4   
ÿ
þüý
z._SparkXGBEstimator._try_stage_level_schedulingc                 C   sÊ   |   | j¡}i }|rAtƒ }|  | j¡r |   | j¡}t|tƒs J ‚|jdu r/tƒ j 	¡  
d¡|_|   | j¡}| t||ƒ¡ ||fS |  | j¡ra|   | j¡}t|tƒsTJ ‚|jduratd|j› �ƒ‚||fS )z@Start the tracker and return the tracker envs on the driver sideNzspark.driver.hostz>You must enable launch_tracker_on_driver to use tracker host: )rÌ   rx   r<   r
  ry   r§   Útracker_host_ipr]   r  r  rï   rq   rÎ   r\   rð   )rª   rx   Ú
rabit_argsrä   rq   r…   r…   rŠ   Ú_get_tracker_argsò  s.   
ÿ
ø
ÿÿz$_SparkXGBEstimator._get_tracker_argsc           	         sP  ˆ  ¡  ˆ ˆ¡\‰‰ˆ ˆ¡\‰‰‰ˆ ¡ ‰ttƒ jƒ‰ˆ ˆj¡‰	ˆ 	¡ \‰‰
ˆ 
ˆj¡r5ˆ ˆj¡nd ‰ttƒ‰tƒ d ‰dttj dttj f‡‡‡‡‡‡‡‡	‡
‡‡‡fdd„‰ dttttf f‡ ‡‡fdd„}ttƒ dtƒ ˆ	ˆˆˆ¡ |ƒ \}}}ttƒ d	¡ ˆ t|d
ƒ|¡}t t |¡¡}ˆ ||¡}| ˆj¡ ˆ  |¡S )NÚuse_rmmÚpandas_df_iterr¦   c                 3   s¾  � ddl m} | ¡ }| ¡  d}tˆ  dd¡ˆ  dd¡ƒ}ˆ  dd¡}d}ˆ	rMˆr.| ¡ nt|ƒ}d	t|ƒ ˆ d< |o>tƒ }d
ˆ d › d|rIdnd› �}|r]ˆ  dd¡dur]ˆ d ˆd< ˆ}| ¡ dkr‚ˆszˆdurmˆnt	ƒ }t
|ƒ|_t|ˆƒ}ttˆƒ |¡ d|i}	ˆsŒ||	d< |jt |	¡d�}
ttdd„ |
D ƒƒƒdkr¦tdƒ‚ˆs±t |
d ¡d }i }t|ˆd��N t|fi |¤Ž�6 t| ˆj||ˆˆjˆjd�\}}|durà|df|dfg}n|dfg}tdˆ |||dœˆ
¤Ž}W d  ƒ n1 sûw   Y  W d  ƒ n	1 �sw   Y  | ¡  | ¡ dk�r[t dt t |ƒ¡gi¡V  | !¡ }t d|gi¡V  | "d¡ #d¡}t$dt|ƒt%ƒD ]}|||t% … }t d|gi¡V  �qFdS dS )zˆTakes in an RDD partition and outputs a booster for that partition after
            going through the Rabit Ring protocol

            r   )ÚBarrierTaskContextNrö   r¤   Ú	verbosityrJ   zTraining on CPUsúcuda:zLeveraging z to train with QDM: ÚonÚoffÚmax_binÚuse_qdmÚ	rabit_msg)Úmessagec                 s   s   � | ]
}t  |¡d  V  qdS )rª  N)ÚjsonÚloads)r‡   Úxr…   r…   rŠ   rý   \  s   € zB_SparkXGBEstimator._fit.<locals>._train_booster.<locals>.<genexpr>z1The workers' cudf environments are in-consistent )r¥  r¢  )Úiteratorr  Údev_ordinalrª  rC  rt   r;  ÚtrainingÚ
validation)rk  Údtrainr˜   r™   r#  r­  úutf-8r…   )&Úpysparkr¤  rï   ÚbarrierrH   ÚpartitionIdrY   rû   r>   r<   rZ   rŸ  r\   rc   Ú_LOG_TAGrŒ  Ú	allGatherr­  ÚdumpsÚlenr¨   r+  r®  r@   rW   rM   r<  rt   r;  rp  Úpdr*   ÚdictÚsave_configÚsave_rawÚdecodeÚrangeÚ_MODEL_CHUNK_SIZE)r£  r¤  Úcontextr±  rª  r¥  ÚmsgÚ_rabit_argsÚ_confÚworker_messageÚmessagesr™   r´  ÚdvalidÚdvalrU  rV  Úbooster_jsonÚoffsetÚbooster_chunk)rr  rä   r‰  r{  rå   rx   Ú	log_levelrq   r   Ú
run_on_gpur‡  r¢  r…   rŠ   Ú_train_booster(  sœ   €

þÿ

ÿÿ

ÿÿþ
ù	
üûð€ øz/_SparkXGBEstimator._fit.<locals>._train_boosterc                     s^   ˆj ˆ dd�j ¡  dd„ ¡} ˆ | ¡}| ¡ }dd„ |D ƒ}|d |d d	 |d
d … ¡fS )Nzdata string)r  c                 S   s   | S rú   r…   )r¯  r…   r…   rŠ   Ú<lambda>�  s    z;_SparkXGBEstimator._fit.<locals>._run_job.<locals>.<lambda>c                 S   s   g | ]}|d  ‘qS )r   r…   )r‡   r‰   r…   r…   rŠ   Ú
<listcomp>‘  s    z=_SparkXGBEstimator._fit.<locals>._run_job.<locals>.<listcomp>r   rJ   r	  r:   )ÚmapInPandasr^  r·  ÚmapPartitionsrž  rj  Újoin)r^  Úrdd_with_resourceÚretr#  )rÑ  r  rª   r…   rŠ   Ú_run_job†  s   þ
ú
 z)_SparkXGBEstimator._fit.<locals>._run_jobzkRunning xgboost-%s on %s workers with
	booster params: %s
	train_call_kwargs_params: %s
	dmatrix_kwargs: %szFinished xgboost training!rµ  )!r  rƒ  rŠ  rë   r^   r]   r  rÌ   rq   r¡  ÚisSetry   rd   r¹  rA   r   r½  r*   r
   rû   rc   rŒ  rD   r]  Ú	bytearrayrV   Úfrom_metricsr­  r®  rT  Ú	_resetUidÚuidÚ_copyValues)	rª   r  rÙ  r™   rV  rU  Úresult_xgb_modelrS  Úspark_modelr…   )rÑ  rr  rä   r  r‰  r{  rå   rx   rÏ  rq   r   rÐ  rª   r‡  r¢  rŠ   Ú_fit  sL   üÿ
ÿ$þ ^÷
ÿ
z_SparkXGBEstimator._fitÚSparkXGBWriterc                 C   ó   t | ƒS )z=
        Return the writer for saving the estimator.
        )rã  r  r…   r…   rŠ   Úwrite¬  ó   z_SparkXGBEstimator.writeÚSparkXGBReaderc                 C   rä  )z>
        Return the reader for loading the estimator.
        )rç  r³   r…   r…   rŠ   Úread²  ó   z_SparkXGBEstimator.readr  )r¦   rã  )r¦   rç  )&rí   r  r  r   rû   r   Ú__annotations__r@  rP  r  r   rR  rG   rV   rT  rÛ  r]  r*   r  ra  rn  r
   ru  r   r)   r:  r|  rƒ  rŠ  r   r�  r   rž  r¡  râ  rå  rè  Ú__classcell__r…   r…   rB  rŠ   r=  \  sR   
 2ÿÿ
þ	

ÿþÿ
þ:1ÿ$
þC- 
 r=  c                
       s4  e Zd Z		d%dee dee ddf‡ fdd„Zedee fdd„ƒZ	de
fd	d
„Z	d&dedeeeeee f f fdd„Zd'dd„Zed(dd„ƒZdedeee eee  f fdd„Zdee fdd„Zdeeef fdd„Zdefdd„Zdededefdd „Zdef‡ fd!d"„Zdedefd#d$„Z‡  Z S ))rQ  NÚxgb_sklearn_modelrS  r¦   c                    s   t ƒ  ¡  || _|| _d S rú   )r?  r@  Ú_xgb_sklearn_modelrS  )rª   rì  rS  rB  r…   rŠ   r@  »  s   

z_SparkXGBModel.__init__c                 C   r°   rú   r±   r³   r…   r…   rŠ   rµ   Ä  s   z_SparkXGBModel._xgb_clsc                 C   s   | j dusJ ‚| j  ¡ S )z=
        Return the `xgboost.core.Booster` instance.
        N)rí  Úget_boosterr  r…   r…   rŠ   rî  È  s   
z_SparkXGBModel.get_boosterrw  Úimportance_typec                 C   s   |   ¡ j|d�S )a�  Get feature importance of each feature.
        Importance type can be defined as:

        * 'weight': the number of times a feature is used to split the data across all trees.
        * 'gain': the average gain across all splits the feature is used in.
        * 'cover': the average coverage across all splits the feature is used in.
        * 'total_gain': the total gain across all splits the feature is used in.
        * 'total_cover': the total coverage across all splits the feature is used in.

        Parameters
        ----------
        importance_type: str, default 'weight'
            One of the importance types defined above.
        )rï  )rî  Ú	get_score)rª   rï  r…   r…   rŠ   Úget_feature_importancesÏ  s   z&_SparkXGBModel.get_feature_importancesÚSparkXGBModelWriterc                 C   rä  )z9
        Return the writer for saving the model.
        )rò  r  r…   r…   rŠ   rå  â  ræ  z_SparkXGBModel.writeÚSparkXGBModelReaderc                 C   rä  )z:
        Return the reader for loading the model.
        )ró  r³   r…   r…   rŠ   rè  è  ré  z_SparkXGBModel.readr  c                 C   sŠ   |   | j¡rd}tt|   | j¡ƒƒ}||fS |   | j¡}g }|r3t|ƒ t|jƒ¡r3t	||ƒ}||fS d}| 
t||   | j¡ƒ¡ ||fS )z¸XGBoost model trained with features_cols parameter can also predict
        vector or array feature type. But first we need to check features_cols
        and then featuresCol
        N)rÌ   rt   r9  r+   rg   rs   r¨   ÚissubsetÚcolumnsr  r  r&  )rª   r  Úfeature_col_namesr~   r…   r…   rŠ   Ú_get_feature_colï  s(   ÿÿúÿÿz_SparkXGBModel._get_feature_colc                 C   s    d}|   | j¡r|  | j¡}|S )z$Return the pred_contrib_col col nameN)r   rw   rÌ   )rª   Úpred_contrib_col_namer…   r…   rŠ   Ú_get_pred_contrib_col_name  s   z)_SparkXGBModel._get_pred_contrib_col_namec                 C   s(   |   ¡ durdtj› dtj› d�fS dS )zÆReturn the bool to indicate if it's a single prediction, true is single prediction,
        and the returned type of the user-defined function. The value must
        be a DDL-formatted type string.NFú	 double, ú array<double>)TÚdouble)rù  Úpredr�   r    r  r…   r…   rŠ   Ú_out_schema  s   z_SparkXGBModel._out_schemac              
      sD   |   ¡ ‰|  ¡ ‰ dtdtdtt dttjtjf f‡ ‡fdd„}|S )zNReturn the true prediction function which will be running on the executor sideÚmodelÚXr”   r¦   c                    sj   i }| j |f|ddœˆ¤Ž}t |¡|tj< ˆ d ur0t| ||ƒ}t t|ƒ¡|tj< tj|d�S |tj S )NF)r”   r›   ©r#  )	rÝ   r½  ÚSeriesrý  r�   rN   rG  r    r*   )rÿ  r   r”   r#  ÚpredsÚcontribs©rø  rß   r…   rŠ   Ú_predict+  s   ÿýü
z2_SparkXGBModel._get_predict_func.<locals>._predict)	râ   rù  rG   r;   r	   r   r½  r*   r  ©rª   r  r…   r  rŠ   Ú_get_predict_func%  s   ÿÿÿþz _SparkXGBModel._get_predict_funcÚpred_colc                 C   s–   |   | j¡}|  ¡ \}}|r|r| ||¡}|S d}| ||¡}|r.| |tt|ƒtjƒ¡}|  ¡ }|durD| |t	tt|ƒtj
ƒƒ¡}| |¡}|S )zPost process of transformÚ_prediction_structN)rÌ   rk   rþ  Ú
withColumnÚgetattrr+   rý  r�   rù  r   r    Údrop)rª   r  r	  Úprediction_col_nameÚsingle_predÚ_Úpred_struct_colrø  r…   r…   rŠ   Ú_post_transform@  s(   ðÿþ
z_SparkXGBModel._post_transformc                    sN   t ƒ  ¡ }ttƒ jƒr|S tƒ j ¡  d¡}|du r%|r#ttƒ 	d¡ dS |S )z`If gpu is used to do the prediction according to the parameters
        and spark configurationsrç   NzADo the prediction on the CPUs since no gpu configurations are setF)
r?  rë   r^   r]   r  r  rï   rc   r¹  rî   )rª   Úuse_gpu_by_paramsró   rB  r…   rŠ   rë   [  s   
ýÿz_SparkXGBModel._run_on_gpuc              
      sø   | j ‰d }|  | j¡rt|  | j¡ƒ tj¡}|d u‰|  |¡\}‰|  | j¡‰ |  	¡ ‰|  
¡ \}}ttƒ jƒ‰|  ¡ ‰ttƒ‰t|ƒdttj dttj f‡ ‡‡‡‡‡‡‡fdd„ƒ}ˆrp|d usdJ ‚|tg |¢|‘R Ž ƒ}n|t|Ž ƒ}|  ||¡S )Nr°  r¦   c                 3   sR  � ˆd usJ ‚ˆ}ddl m} | ¡ }|d usJ ‚d‰ d}ˆr[tƒ rYtƒ rYˆr=tƒ }|jj ¡ }|dkr<| 	¡ }|| ‰ nt
|ƒ‰ ˆ dkrVdtˆ ƒ }d| }|j|d� nd}nd	}| 	¡ dkrittˆƒ |¡ d
tdtf‡ fdd„}	| D ]0}
ˆrt|
ƒ}nˆd urˆ|
ˆ }nt|
tj ƒ}|	|ƒ}ˆr�|	|
tj ƒ}nd }ˆ|||ƒV  qvd S )Nr   )ÚTaskContextéÿÿÿÿzDo the inference on the CPUsr¦  zDo the inference with device: )r¤   zCCouldn't get the correct gpu id, fallback the inference on the CPUsz?CUDF or Cupy is unavailable, fallback the inference on the CPUsr#  r¦   c                    s:   ˆ dkrddl }ddl}|jj ˆ ¡ | | ¡}~ |S | S )z Move the data to gpu if possibler   N)ÚcudfÚcupyr­   ÚruntimeÚ	setDevicer*   )r#  r  ÚcpÚdf©r±  r…   rŠ   Úto_gpu_if_possible·  s   
zJ_SparkXGBModel._transform.<locals>.predict_udf.<locals>.to_gpu_if_possible)r¶  r  rï   r>   r?   r=   r­   r  ÚgetDeviceCountr¸  rY   rû   Ú
set_paramsrc   r¹  rŒ  r;   rK   rO   rL   r#  ry  )r°  rÿ  r  rÄ  rÅ  r  Ú
total_gpusÚpartition_idr¤   r  r#  r   Útmpr”   ©rt   rö  Úhas_base_marginrå   rÏ  Úpredict_funcrÐ  rì  r  rŠ   Úpredict_udf�  sN   €€

ñz._SparkXGBModel._transform.<locals>.predict_udf)rí  r   rn   r+   rÌ   rL   ry  r÷  rt   r  rþ  r^   r]   r  rë   rd   r¹  r-   r   r½  r*   r  r/   r  )rª   r  rn   r~   r  r  r&  r	  r…   r#  rŠ   Ú
_transformw  s*   ÿ2Ez_SparkXGBModel._transform)NN)rw  )r¦   rò  )r¦   ró  )!rí   r  r  r	   rG   rV   r@  r  r   rµ   rB   rî  rû   r   r   rñ   r   rñ  rå  rè  r*   r
   r)   r÷  rù  r  rþ  r   r  r  rë   r'  rë  r…   r…   rB  rŠ   rQ  º  sB    ýþýü	ÿÿ
þ
ÿ
þ$
rQ  c                   @   sJ   e Zd ZdZdeeef fdd„Zdefdd„Z	de
dede
fd	d
„ZdS )Ú_ClassificationModelzu
    The model returned by :func:`xgboost.spark.SparkXGBClassifier.fit`

    .. Note:: This API is experimental.
    r¦   c                 C   sB   t j› dt j› dt j› d�}|  ¡ d ur|› dt j› d�}d|fS )Nz array<double>, rú  rû  z, z array<array<double>>F)rý  rž   r�   rŸ   rù  r    )rª   r  r…   r…   rŠ   rþ  è  s   ÿÿz _ClassificationModel._out_schemac              
      sh   |   ¡ ‰|  ¡ ‰ dtjdttjtjf fdd„‰dtdtdttj dtt	j
t	jf f‡ ‡‡fdd	„}|S )
NÚmarginsr¦   c                 S   s`   | j dkr$t| ƒ}d| }t |  | f¡ ¡ }t ||f¡ ¡ }||fS | }t|dd�}||fS )NrJ   rè   ©Úaxis)Úndimr8   ÚnpÚvstackÚ	transposer9   )r)  Úclassone_probsÚclasszero_probsÚ	raw_predsÚclass_probsr…   r…   rŠ   Útransform_marginø  s   
þz@_ClassificationModel._get_predict_func.<locals>.transform_marginrÿ  r   r”   c           	   	      s    | j |f|dddœˆ¤Ž}ˆ|ƒ\}}tj|dd�}tjt t|ƒ¡tjt |¡tj	t t|ƒ¡i}ˆ d urJt
| ||dd�}t t| ¡ ƒ¡|tj< tj|d�S )NTF)r”   rš   r›   rJ   r*  )Ústrict_shaper  )rÝ   r-  Úargmaxrý  rž   r½  r  rG  r�   rŸ   rN   Útolistr    r*   )	rÿ  r   r”   r)  r2  r3  r  Úresultr  ©rø  rß   r4  r…   rŠ   r    s&   ÿüûýz8_ClassificationModel._get_predict_func.<locals>._predict)râ   rù  r-  Úndarrayr
   rG   r;   r	   r   r½  r*   r  r  r…   r9  rŠ   r  ô  s    ÿÿÿþz&_ClassificationModel._get_predict_funcr  r	  c                 C   sÂ   d}|  ||¡}|  | j¡}|r|  |ttt|ƒtjƒƒ¡}|  | j¡}|r2|  |tt|ƒtj	ƒ¡}|  | j
¡}|rH|  |ttt|ƒtjƒƒ¡}|  ¡ }|d ur\|  |tt|ƒtjƒ¡}| |¡S )Nr
  )r  rÌ   rj   r   r  r+   rý  rž   rk   r�   rl   rŸ   rù  r    r  )rª   r  r	  r  Úraw_prediction_col_namer  Úprobability_col_namerø  r…   r…   rŠ   r  "  s4   þÿþþ
z$_ClassificationModel._post_transformN)rí   r  r  Ú__doc__r
   r  rû   rþ  r   r  r*   r)   r  r…   r…   r…   rŠ   r(  ß  s
    .r(  c                   @   s˜   e Zd Ze	ddeeef dedede	j
deeeef  ddfdd	„ƒZed
eee ee f dedede	j
deeeef eeef f f
dd„ƒZdS )Ú_SparkXGBSharedReadWriteNÚinstanceÚpathr  rA  ÚextraMetadatar¦   c                 C   s  |   ¡  g d¢}i }| j ¡ D ]\}}|j|vr|||j< q|p!i }|  d¡}	|	dur?| d¡ t t 	|	¡¡ 
d¡}
|
|d< |  d¡}|durLt|d< |  d	¡r`|  d	¡}|dur`t|ƒ|d	< tj| ||||d
� |dur‰t|ƒ}tj |t¡}tƒ  |fgdg¡j |¡ dS dS )zs
        Save the metadata of an xgboost.spark._SparkXGBEstimator or
        xgboost.spark._SparkXGBModel.
        )Ú	callbacksrõ   ry   rB  NzœThe callbacks parameter is saved using cloudpickle and it is not a fully self-contained format. It may fail to load with different versions of dependencies.ÚasciiÚserialized_callbacksrõ   Úinit_boosterry   )rA  ÚparamMap)r  Ú	_paramMaprD  rË   rÌ   rî   Úbase64Úencodebytesr   r»  rÁ  Ú_INIT_BOOSTER_SAVE_PATHr
  r   r"   ÚsaveMetadatare   Úosr@  rÖ  r]   ÚcreateDataFramerå  Úparquet)r?  r@  r  rA  rA  Ú
skipParamsÚ
jsonParamsÚpr‰   rB  rD  rE  rä   Úser_init_boosterÚ	save_pathr…   r…   rŠ   rK  E  sJ   

€
ÿÿþ



ÿ
ÿûz%_SparkXGBSharedReadWrite.saveMetadataÚpyspark_xgb_clsc              
   C   s  t j||t| ƒd�}| ƒ }t  ||¡ d|v rK|d }zt t | d¡¡¡}| 	|j
|¡ W n tyJ } z| d|› d�¡ W Y d}~nd}~ww d|v r[| tdi |d ¤Ž¡ d|v rtj ||d ¡}	tƒ j |	¡ ¡ d	 j}
t|
ƒ}| 	|j|¡ | |d
 ¡ ||fS )z¶
        Load the metadata and the instance of an xgboost.spark._SparkXGBEstimator or
        xgboost.spark._SparkXGBModel.

        :return: a tuple of (metadata, instance)
        )ÚexpectedClassNamerD  rC  z)Fails to load the callbacks param due to zC. Please set the callbacks param manually for the loaded estimator.Nry   rE  r   rÞ  r…   )r!   ÚloadMetadatarb   ÚgetAndSetParamsr   r®  rH  ÚdecodebytesÚencoder¨   rB  Ú	Exceptionrî   r«   r<   rL  r@  rÖ  r]   rè  rN  rj  rE  r`   rõ   rÝ  )rT  r@  r  rA  ÚmetadataÚpyspark_xgbrD  rB  ÚeÚ	load_pathrR  rE  r…   r…   rŠ   ÚloadMetadataAndInstancew  s8   
ÿÿ
ÿ€ÿÿz0_SparkXGBSharedReadWrite.loadMetadataAndInstancerú   )rí   r  r  Ústaticmethodr   r=  rQ  rû   r   ÚloggingÚLoggerr	   r   r   rK  r   r
   r_  r…   r…   r…   rŠ   r>  D  s8    û
ÿþýüûú1ÿþýüûr>  c                       s4   e Zd ZdZd‡ fdd„Zdeddfd	d
„Z‡  ZS )rã  z)
    Spark Xgboost estimator writer.
    r?  r=  r¦   Nc                    ó&   t ƒ  ¡  || _t| jjdd�| _d S ©NÚWARN)Úlevel©r?  r@  r?  rc   rì   rí   rA  ©rª   r?  rB  r…   rŠ   r@  ª  ó   
zSparkXGBWriter.__init__r@  c                 C   s   t  | j|| j| j¡ dS )z
        save model.
        N)r>  rK  r?  r  rA  )rª   r@  r…   r…   rŠ   ÚsaveImpl¯  s   zSparkXGBWriter.saveImpl)r?  r=  r¦   N)rí   r  r  r=  r@  rû   rj  rë  r…   r…   rB  rŠ   rã  ¥  s    rã  c                       ó@   e Zd ZdZded ddf‡ fdd„Zdeddfd	d
„Z‡  ZS )rç  z)
    Spark Xgboost estimator reader.
    r´   r=  r¦   Nc                    rc  rd  ©r?  r@  r´   rc   rì   rí   rA  ©rª   r´   rB  r…   rŠ   r@  »  ri  zSparkXGBReader.__init__r@  c                 C   s$   t  | j|| j| j¡\}}td|ƒS )z
        load model.
        r=  )r>  r_  r´   r  rA  r   )rª   r@  r  r\  r…   r…   rŠ   ÚloadÀ  s   ÿ
zSparkXGBReader.load©	rí   r  r  r=  r   r@  rû   rn  rë  r…   r…   rB  rŠ   rç  ¶  ó    rç  c                       s<   e Zd ZdZdeddf‡ fdd„Zdeddfdd	„Z‡  ZS )
rò  z%
    Spark Xgboost model writer.
    r?  r¦   Nc                    rc  rd  rg  rh  rB  r…   rŠ   r@  Ï  ri  zSparkXGBModelWriter.__init__r@  c                 C   s–   | j j}|dus
J ‚t | j || j| j¡ tj |d¡}| 	¡  
d¡ d¡}g }tdt|ƒtƒD ]}| |||t … ¡ q0tƒ j |d¡ |¡ dS )z›
        Save metadata and model for a :py:class:`_SparkXGBModel`
        - save metadata to path/metadata
        - save model to path/model.json
        Nrÿ  r­  rµ  r   rJ   )r?  rí  r>  rK  r  rA  rL  r@  rÖ  rî  rÀ  rÁ  rÂ  r¼  rÃ  r  r]   r  ÚparallelizeÚsaveAsTextFile)rª   r@  rõ   Úmodel_save_pathrU  Úbooster_chunksrÍ  r…   r…   rŠ   rj  Ô  s   ÿzSparkXGBModelWriter.saveImpl)	rí   r  r  r=  rQ  r@  rû   rj  rë  r…   r…   rB  rŠ   rò  Ê  s    rò  c                       rk  )ró  z%
    Spark Xgboost model reader.
    r´   rQ  r¦   Nc                    rc  rd  rl  rm  rB  r…   rŠ   r@  î  ri  zSparkXGBModelReader.__init__r@  c                    sz   t  ˆ j|ˆ jˆ j¡\}}td|ƒ}|jdd�‰tj 	|d¡}d 	t
ƒ j |¡ ¡ ¡}d‡ ‡fdd	„}t||ƒ}||_|S )z—
        Load metadata and model for a :py:class:`_SparkXGBModel`

        :return: SparkXGBRegressorModel or SparkXGBClassifierModel instance
        rQ  TrW  rÿ  r	  r¦   rG   c                      s   ˆ j  ¡ di ˆ¤ŽS )Nr…   )r´   rµ   r…   ©rª   r[  r…   rŠ   Úcreate_xgb_model  s   z2SparkXGBModelReader.load.<locals>.create_xgb_modelN)r¦   rG   )r>  r_  r´   r  rA  r   rÒ   rL  r@  rÖ  r]   r  ÚtextFilerj  ra   rí  )rª   r@  r  Úpy_modelÚmodel_load_pathÚser_xgb_modelrv  rõ   r…   ru  rŠ   rn  ó  s   ÿ
ÿÿ
zSparkXGBModelReader.loadro  r…   r…   rB  rŠ   ró  é  rp  ró  )¥r=  rH  r­  ra  rL  Úcollectionsr   Údataclassesr   Útypingr   r   r   r   r   r	   r
   r   r   r   Únumpyr-  Úpandasr½  r¶  r   r   r   r   Ú
pyspark.mlr   r   Úpyspark.ml.functionsr   r   Úpyspark.ml.linalgr   Úpyspark.ml.paramr   r   r   Úpyspark.ml.param.sharedr   r   r   r   r   r   r    Úpyspark.ml.utilr!   r"   r#   r$   r%   r&   Úpyspark.resourcer'   r(   Úpyspark.sqlr)   r*   r(  r+   r,   r-   r.   r/   Úpyspark.sql.typesr0   r1   r2   r3   r4   r5   r6   r7   Úscipy.specialr8   r9   Ú_typingr;   Ú
collectiver<   Úcompatr=   r>   r?   rV  r@   rA   ÚcorerB   rC   rD   ÚsklearnrE   rF   rG   rH   r²  rI   rp  r#  rK   rL   rM   rN   rO   rk  rP   rQ   rR   rS   rT   rU   ÚsummaryrV   ÚutilsrW   rX   rY   rZ   r[   r\   r]   r^   r_   r`   ra   rb   rc   rd   re   rf   rÅ   rÉ   rF  rD  rE  r·   rÕ   rJ  rÞ   rK  rœ   rý  rJ  r¹  r£   rû   r  r&  r.  r9  r:  rÃ  r=  rQ  r(  r>  rã  rç  rò  ró  r…   r…   r…   rŠ   Ú<module>   sÜ    0$ 	(
 Hù
þ	ÿÿ
õ  8ÿÿ
þÿÿ
þ  þ    b  
'ÿea