o
    Ú­j¬0  ã                   @   s
  d Z ddlmZmZ ddl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 ddlmZ ddlmZ dd	lmZmZmZ dd
lmZ ddlmZ dej dej!fdd„Z"eddƒZ#e#ddddddƒZ$de
eej!  de
ej! fdd„Z%deej& deej&e'e(gdf ddfdd„Z)G d d!„ d!eƒZ*d"ej&defd#d$„Z+d%ee'e	ej! f d&e
e, d'ee'ef d(e
e d)ee'ef defd*d+„Z-deej& d,e
ee'  d&e
e, d-e(d.ee'ef d/e(d0e(deee
e f fd1d2„Z.		3d9d4ed%ed5e
e d6e(dej!f
d7d8„Z/dS ):z*Utilities for processing spark partitions.é    )ÚdefaultdictÚ
namedtuple)	ÚAnyÚCallableÚDictÚIteratorÚListÚOptionalÚSequenceÚTupleÚUnionN)Ú
csr_matrixé   )Ú	ArrayLike©Úconcat)ÚDataIterÚDMatrixÚQuantileDMatrix)ÚXGBModelé   )Ú
get_loggerÚseriesÚreturnc                 C   s   | j dd�}t |¡}|S )zStack a series of arrays.F)Úcopy)Úto_numpyÚnpÚstack)r   Úarray© r   úO/var/www/html/CropPilot/venv/lib/python3.10/site-packages/xgboost/spark/data.pyÚstack_series   s   
r!   ÚAlias)ÚdataÚlabelÚweightÚmarginÚvalidÚqidÚvaluesr$   r%   Ú
baseMarginÚvalidationIndicatorr(   Úseqc                 C   s   | rt | ƒS dS )z&Concatenate the data if it's not None.Nr   )r,   r   r   r    Úconcat_or_none   s   r-   ÚiteratorÚappendc                    s¸   dt jdtddf‡ fdd„}d}| D ]G}|du rtj|jv }|du r*tj|jv s*J ‚|rF|j|tj  dd…f }|j|tj dd…f }n|d}}||dƒ |durY||dƒ qdS )	znExtract partitions from pyspark iterator. `append` is a user defined function for
    accepting new partition.ÚpartÚis_validr   Nc                    sJ   ˆ | t j|ƒ ˆ | t j|ƒ ˆ | t j|ƒ ˆ | t j|ƒ ˆ | t j|ƒ d S )N)Úaliasr#   r$   r%   r&   r(   )r0   r1   ©r/   r   r    Ú	make_blob,   s
   z#cache_partitions.<locals>.make_blobTF)ÚpdÚ	DataFrameÚboolr2   r'   ÚcolumnsÚloc)r.   r/   r4   Úhas_validationr0   Útrainr'   r   r3   r    Úcache_partitions&   s    


€òr<   c                       s|   e Zd ZdZdeeef dee de	ddf‡ fdd„Z
deeej  deej fd	d
„Zdedefdd„Zddd„Z‡  ZS )ÚPartIterz7Iterator for creating Quantile DMatrix from partitions.r#   Ú	device_idÚkwargsr   Nc                    s*   d| _ || _|| _|| _tƒ jdd� d S )Nr   T)Úrelease_data)Ú_iterÚ
_device_idÚ_dataÚ_kwargsÚsuperÚ__init__)Úselfr#   r>   r?   ©Ú	__class__r   r    rF   I   s
   zPartIter.__init__c                 C   sL   |sd S | j d ur!dd l}dd l}|jj | j ¡ | || j ¡S || j S ©Nr   )rB   ÚcudfÚcupyÚcudaÚruntimeÚ	setDevicer6   rA   )rG   r#   rK   Úcpr   r   r    Ú_fetchS   s   

zPartIter._fetchÚ
input_datac                 C   sž   | j t| jtj ƒkrdS |d|  | jtj ¡|  | j tjd ¡¡|  | j tjd ¡¡|  | j tj	d ¡¡|  | j tj
d ¡¡dœ| j¤Ž |  j d7  _ dS )NF©r#   r$   r%   Úbase_marginr(   r   Tr   )rA   ÚlenrC   r2   r#   rQ   Úgetr$   r%   r&   r(   rD   )rG   rR   r   r   r    Únextb   s   ûúzPartIter.nextc                 C   s
   d| _ d S rJ   )rA   )rG   r   r   r    Úresetp   s   
zPartIter.reset)r   N)Ú__name__Ú
__module__Ú__qualname__Ú__doc__r   Ústrr   r	   Úintr   rF   r
   r5   r6   rQ   r   r7   rW   rX   Ú__classcell__r   r   rH   r    r=   F   s    
ÿÿÿþ"
r=   r0   c                 C   sê   g dgg }}}d}t | j| j| j| jƒD ]B\}}}}|dkr)t|ƒ}	|}
|}nt|ƒ}	tj|	tj	d�}
|}|dkr=|	}||	ksCJ ‚| 
|
¡ | 
|d t|
ƒ ¡ | 
|¡ qt |¡}t |¡}t |¡}t|||ft| ƒ|fd�S )Nr   )Údtypeéÿÿÿÿ)Úshape)ÚzipÚfeatureVectorTypeÚfeatureVectorSizeÚfeatureVectorIndicesÚfeatureVectorValuesr^   rU   r   ÚarangeÚint32r/   r   Úconcatenater   )r0   Úcsr_indices_listÚcsr_indptr_listÚcsr_values_listÚ
n_featuresÚvec_typeÚ	vec_size_Úvec_indicesÚ
vec_valuesÚvec_sizeÚcsr_indicesÚ
csr_valuesÚcsr_indptr_arrÚcsr_indices_arrÚcsr_values_arrr   r   r    Ú)_read_csr_matrix_from_unwrapped_spark_vect   s6   ü



ÿry   r#   Údev_ordinalÚmetaÚrefÚparamsc                 C   sD   | st t d¡|d�S t| |fi |¤Ž}t |fi |¤d|i¤Ž}|S )z+Handle empty partition for QuantileDMatrix.©r   r   )r|   r|   )r   r   Úemptyr=   )r#   rz   r{   r|   r}   ÚitÚmr   r   r    Úmake_qdmŸ   s
   r‚   Úfeature_colsÚuse_qdmr?   Úenable_sparse_data_optimÚhas_validation_colc              	      sÊ  t tƒ‰t tƒ‰d‰dtjdtdtddf‡ ‡‡‡fdd„}dtjdtdtddf‡‡‡fd	d
„}dttttj	 f dttt
f dtfdd„}	|rV|}
dˆv rSˆd dksUJ ‚n|}
dtttt
f ttttttf f f f‡fdd„}|ƒ \}}ˆ dur‹|r‹t| |
ƒ tˆ||d|ƒ}n/ˆ durœ|sœt| |
ƒ |	ˆˆƒ}nˆ du r°|r°t| |
ƒ tˆ||d|ƒ}n
t| |
ƒ |	ˆˆƒ}|rÑ|rÇtˆ||||ƒ}n|rÎ|	ˆˆƒnd}nd}|durá| ¡ | ¡ ksáJ ‚||fS )a~  Create DMatrix from spark data partitions.

    Parameters
    ----------
    iterator :
        Pyspark partition iterator.
    feature_cols:
        A sequence of feature names, used only when rapids plugin is enabled.
    dev_ordinal:
        Device ordinal, used when GPU is enabled.
    use_qdm :
        Whether QuantileDMatrix should be used instead of DMatrix.
    kwargs :
        Metainfo for DMatrix.
    enable_sparse_data_optim :
        Whether sparse data should be unwrapped
    has_validation:
        Whether there's validation data.

    Returns
    -------
    Training DMatrix and an optional validation DMatrix.
    r   r0   Únamer1   r   Nc                    sâ   |t jks
|| jv ro|t jkr!ˆ d ur!| ˆ  jd dkr!| ˆ  }n| | jd dkr8| | }|t jkr7t|ƒ}nd }|t jkrU|d urUˆdkrL|jd ‰ˆ|jd ksUJ ‚|d u r[d S |rfˆ|  |¡ d S ˆ|  |¡ d S d S ©Nr   r   )r2   r#   r8   rb   r!   r/   ©r0   r‡   r1   r   )rƒ   rn   Ú
train_dataÚ
valid_datar   r    Úappend_mÕ   s*   


€
æz0create_dmatrix_from_partitions.<locals>.append_mc                    s€   |t jks
|| jv r>|t jkr&t| ƒ}ˆ dkr|jd ‰ ˆ |jd ks%J ‚n| | }|r5ˆ|  |¡ d S ˆ|  |¡ d S d S rˆ   )r2   r#   r8   ry   rb   r/   r‰   )rn   rŠ   r‹   r   r    Úappend_m_sparseó   s   

ôz7create_dmatrix_from_partitions.<locals>.append_m_sparser)   r?   c                 S   s¢   t | ƒdkrtdƒ d¡ tddt d¡i|¤ŽS t| tj ƒ}t|  	tj
d ¡ƒ}t|  	tjd ¡ƒ}t|  	tjd ¡ƒ}t|  	tjd ¡ƒ}td|||||dœ|¤ŽS )Nr   ÚXGBoostPySparkz_Detected an empty partition in the training data. Consider to enable repartition_random_shuffler#   r~   rS   r   )rU   r   Úwarningr   r   r   r-   r2   r#   rV   r$   r%   r&   r(   )r)   r?   r#   r$   r%   r&   r(   r   r   r    Úmake  s   ÿ
ÿÿz,create_dmatrix_from_partitions.<locals>.makeÚmissingg        c                     s@   d} i }i }ˆ   ¡ D ]\}}|| v r|||< q
|||< q
||fS )N)Úmax_binr‘   ÚsilentÚnthreadÚenable_categorical)Úitems)Únon_data_keysÚnon_data_paramsr{   ÚkÚv)r?   r   r    Úsplit_params  s   

z4create_dmatrix_from_partitions.<locals>.split_params)r   Úlistr5   r6   r]   r7   r   r   r   Úndarrayr   r   r   r   r^   Úfloatr<   r‚   Únum_col)r.   rƒ   rz   r„   r?   r…   r†   rŒ   r�   r�   Ú	append_fnr›   r{   r}   ÚdtrainÚdvalidr   )rƒ   r?   rn   rŠ   r‹   r    Úcreate_dmatrix_from_partitions®   sB   "&$,4






ÿr£   FÚmodelrT   Ústrict_shapec              	   C   sB   |   d¡}t||| j| j| j| j| jd�}|  ¡ j|dd||d�S )z4Predict contributions with data with the full model.N)rT   r‘   r”   Úfeature_typesÚfeature_weightsr•   TF)Úpred_contribsÚvalidate_featuresÚiteration_ranger¥   )	Ú_get_iteration_ranger   r‘   Ún_jobsr¦   r§   r•   Úget_boosterÚpredict)r¤   r#   rT   r¥   rª   Údata_dmatrixr   r   r    r¨   V  s"   
ù	ûr¨   )NF)0r\   Úcollectionsr   r   Útypingr   r   r   r   r   r	   r
   r   r   Únumpyr   Úpandasr5   Úscipy.sparser   Ú_typingr   Úcompatr   Úcorer   r   r   Úsklearnr   Úutilsr   ÚSeriesr�   r!   r"   r2   r-   r6   r]   r7   r<   r=   ry   r^   r‚   r£   r¨   r   r   r   r    Ú<module>   sˆ   ,
	"ÿÿ
þ .+ÿþ
ýü
û
úþ
ýüû
úùø	
÷ ,üÿþýüû