U
    ½mœd§?  ã                   @   sL  d dl Z d dlZd dlm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 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mZ ej d ¡Zejdd�Zejdd�Zeejdd�dd…ejf  Zeejdd�dd…ejf  Zej  ddddg¡ej  ddddg¡ej  dd dg¡dd„ ƒƒƒZ!ej  dddg¡ej  ddddg¡ej  dd dg¡d d!„ ƒƒƒZ"d"d#„ Z#d$d%„ Z$d&d'„ Z%d(d)„ Z&d*d+„ Z'd,d-„ Z(d.d/„ Z)d0d1„ Z*d2d3„ Z+d4d5„ Z,d6d7„ Z-d8d9„ Z.d:d;„ Z/d<d=„ Z0d>d?„ Z1d@dA„ Z2dBdC„ Z3dDdE„ Z4ej  dFeeeeg¡dGdH„ ƒZ5dIdJ„ Z6dS )Ké    N)Ú
csr_matrix)Úassert_array_equal)Úassert_array_almost_equal)Úassert_allclose)Úkernel_metrics)Ú
RBFSampler)ÚAdditiveChi2Sampler)ÚSkewedChi2Sampler)ÚNystroem)ÚPolynomialCountSketch)Úmake_classification)Úpolynomial_kernelÚ
rbf_kernelÚchi2_kernel)é,  é2   ©Úsizeé   ©ZaxisÚgammaçš™™™™™¹?g      @zdegree, n_components)r   éô  )é   r   )é   iˆ  Úcoef0c           
      C   sœ   t tt| ||d�}t|| ||dd�}| t¡}| t¡}t ||j¡}|| }	t 	t 
|	¡¡dksft‚tj	|	|	d� t |	¡dks†t‚t 
|	¡dks˜t‚d S )N)r   Údegreer   é*   )Ún_componentsr   r   r   Úrandom_stateçš™™™™™©?©Úoutr   )r   ÚXÚYr   Úfit_transformÚ	transformÚnpÚdotÚTÚabsÚmeanÚAssertionErrorÚmax)
r   r   r   r   ÚkernelZps_transformÚX_transÚY_transÚkernel_approxÚerror© r3   ú`/home/sam/Atlas/atlas_env/lib/python3.8/site-packages/sklearn/tests/test_kernel_approximation.pyÚtest_polynomial_count_sketch   s     û

r5   ç      ð?r   r   r   c           	      C   sl   t d| ||dd�}| t¡}| t¡}t d| ||dd�}| ttƒ¡}| ttƒ¡}t||ƒ t||ƒ dS )zZCheck that PolynomialCountSketch results are the same for dense and sparse
    input.
    r   r   )r   r   r   r   r   N)r   r%   r#   r&   r$   r   r   )	r   r   r   Zps_denseZXt_denseZYt_denseZ	ps_sparseZ	Xt_sparseZ	Yt_sparser3   r3   r4   Ú)test_polynomial_count_sketch_dense_sparse9   s(        ÿ

    ÿ
r7   c                 C   s   t  | |j¡S )N)r'   r(   r)   )r#   r$   r3   r3   r4   Ú_linear_kernelP   s    r8   c               	   C   s´  t d d …tjd d …f } ttjd d …d d …f }d|  | | |  }|jdd�}tdd�}| t ¡}| t¡}t ||j	¡}t
||dƒ | tt ƒ¡}| ttƒ¡}	t||jƒ t||	jƒ t ¡ }
d|
d< d}tjt|d	�� | |
¡ W 5 Q R X td
d�}t d¡}tjt|d	�� | t ¡ W 5 Q R X dddg}|D ]:}t|d�}|jd k�sXt‚| t ¡ |jd k	�s:t‚�q:d}td
|d�}|j|k�s–t‚| t ¡ |j|k�s°t‚d S )Nr   r   r   ©Úsample_stepsr   éÿÿÿÿ©r   r   z!Negative values in data passed to©Úmatché   zHIf sample_steps is not in [1, 2, 3], you need to provide sample_intervalg333333Ó?)r:   Úsample_interval)r#   r'   Únewaxisr$   Úsumr   r%   r&   r(   r)   r   r   r   ÚAÚcopyÚpytestÚraisesÚ
ValueErrorÚreÚescapeÚfitr@   r,   Zsample_interval_)ZX_ZY_Zlarge_kernelr.   r&   r/   r0   r1   Z
X_sp_transZ
Y_sp_transÚY_negÚmsgZsample_steps_availabler:   r@   r3   r3   r4   Útest_additive_chi2_samplerT   sF    



ÿ



rM   c               	   C   s:  d} |  d t d< t|  d d …tjd d …f }t |  tjd d …d d …f }t |¡d t |¡d  t d¡ t || ¡ }t |jdd�¡}t| ddd�}| t¡}| 	t ¡}t 
||j¡}t||d	ƒ t |¡ ¡ sâtd
ƒ‚t |¡ ¡ søtdƒ‚t  ¡ }	|  d |	d< d}
tjt|
d�� | 	|	¡ W 5 Q R X d S )Ng¸…ëQ¸ž?g       @r<   r   r   éè  r   )Z
skewednessr   r   r   zNaNs found in the Gram matrixz)NaNs found in the approximate Gram matrixz2X may not contain entries smaller than -skewednessr=   )r$   r#   r'   rA   ÚlogÚexprB   r	   r%   r&   r(   r)   r   ÚisfiniteÚallr,   rD   rE   rF   rG   )ÚcZX_cZY_cZ
log_kernelr.   r&   r/   r0   r1   rK   rL   r3   r3   r4   Útest_skewed_chi2_sampler“   s&    2ÿ

rT   c               	   C   sl   t ƒ } t ¡ }d|d< tjtdd�� |  |¡ W 5 Q R X tjtdd�� |  t¡ |  |¡ W 5 Q R X dS )zEnsures correct error messager;   r<   zX in AdditiveChi2Sampler.fitr=   z"X in AdditiveChi2Sampler.transformN)r   r#   rD   rE   rF   rG   rJ   r&   )ZtransformerZX_negr3   r3   r4   Ú%test_additive_chi2_sampler_exceptions»   s    
rU   c                  C   s˜   d} t tt| d�}t| ddd�}| t¡}| t¡}t ||j¡}|| }t 	t 
|¡¡dksbt‚tj	||d� t |¡dks‚t‚t 
|¡d	ks”t‚d S )
Ng      $@©r   rN   r   )r   r   r   g{®Gáz„?r!   r   r    )r   r#   r$   r   r%   r&   r'   r(   r)   r*   r+   r,   r-   )r   r.   Zrbf_transformr/   r0   r1   r2   r3   r3   r4   Útest_rbf_samplerÇ   s    

rW   c                 C   sT   t ƒ }tjddgddgddgg| d�}| |¡ |jj| ks@t‚|jj| ksPt‚dS ©	zRCheck that the fitted attributes are stored accordingly to the
    data type of X.r   r   r   r?   é   é   ©ÚdtypeN)r   r'   ÚarrayrJ   Úrandom_offset_r\   r,   Úrandom_weights_)Úglobal_dtypeÚrbfr#   r3   r3   r4   Ú(test_rbf_sampler_fitted_attributes_dtypeÚ   s
     
rb   c                  C   sŒ   t dd�} tjddgddgddggtjd	�}|  |¡ t dd�}tjddgddgddggtjd	�}| |¡ t| j|jƒ t| j|jƒ d
S ©z?Check the equivalence of the results with 32 and 64 bits input.r   )r   r   r   r   r?   rY   rZ   r[   N)	r   r'   r]   Úfloat32rJ   Úfloat64r   r^   r_   )Zrbf32ZX32Zrbf64ZX64r3   r3   r4   Ú"test_rbf_sampler_dtype_equivalenceç   s    
"

"
rf   c                  C   sD   dgdggddg } }t dd�}| | |¡ |jt d¡ks@t‚dS )	z4Check the inner value computed when `gamma='scale'`.g        r6   r   r   ÚscalerV   r?   N)r   rJ   Z_gammarE   Zapproxr,   )r#   Úyra   r3   r3   r4   Útest_rbf_sampler_gamma_scaleõ   s    
ri   c                 C   sT   t ƒ }tjddgddgddgg| d�}| |¡ |jj| ks@t‚|jj| ksPt‚dS rX   )r	   r'   r]   rJ   r^   r\   r,   r_   )r`   Zskewed_chi2_samplerr#   r3   r3   r4   Ú0test_skewed_chi2_sampler_fitted_attributes_dtypeý   s
     
rj   c                  C   sŒ   t dd�} tjddgddgddggtjd	�}|  |¡ t dd�}tjddgddgddggtjd	�}| |¡ t| j|jƒ t| j|jƒ d
S rc   )	r	   r'   r]   rd   rJ   re   r   r^   r_   )Zskewed_chi2_sampler_32ZX_32Zskewed_chi2_sampler_64ZX_64r3   r3   r4   Ú*test_skewed_chi2_sampler_dtype_equivalence
  s    
"

"
 ÿ ÿrk   c                  C   sj   ddgddgddgg} t ƒ  | ¡ | ¡ tƒ  | ¡ | ¡ tƒ  | ¡ | ¡ t| ƒ} tƒ  | ¡ | ¡ d S )Nr   r   r   r?   rY   rZ   )r   rJ   r&   r	   r   r   )r#   r3   r3   r4   Útest_input_validation  s    rl   c                  C   sþ   t j d¡} | jdd�}t|jd d� |¡}t|ƒ}tt  	||j
¡|ƒ td| d�}| |¡ |¡}|j|jd dfks~t‚tdt| d�}| |¡ |¡}|j|jd dfks´t‚tƒ }|D ]:}td|| d�}| |¡ |¡}|j|jd dfks¾t‚q¾d S )Nr   ©é
   r?   r   ©r   r   ©r   r   )r   r.   r   )r'   ÚrandomÚRandomStateÚuniformr
   Úshaper%   r   r   r(   r)   rJ   r&   r,   r8   r   )Úrndr#   ÚX_transformedÚKZtransZkernels_availableÚkernr3   r3   r4   Útest_nystroem_approximation(  s     ry   c                  C   sŽ   t j d¡} | jdd�}tdd�}| |¡}t|d d�}t  ||j¡}t	||ƒ tddd�}| |¡}t
|d	d�}t  ||j¡}t	||ƒ d S )
Nr   rm   r   rn   ro   rV   Zchi2©r.   r   r   )r'   rq   rr   rs   r
   r%   r   r(   r)   r   r   )ru   r#   Únystroemrv   rw   ZK2r3   r3   r4   Ú test_nystroem_default_parametersC  s    



r|   c                  C   s†   t j d¡} |  dd¡}t  |gd ¡}d}t||jd d� |¡}| |¡}t	||d�}t
|t  ||j¡ƒ t  t  t¡¡s‚t‚d S )Nr   rn   é   r   éd   )r   r   rV   )r'   rq   rr   ZrandZvstackr
   rt   rJ   r&   r   r   r(   r)   rR   rQ   r$   r,   )Úrngr#   r   ÚNrv   rw   r3   r3   r4   Útest_nystroem_singular_kernelW  s    
r�   c                  C   s^   t j d¡} | jdd�}t|ddd�}td|jd ddd	�}| |¡}tt  	||j
¡|ƒ d S )
Né%   rm   r   gÍÌÌÌÌÌ@r   ©r   r   Z
polynomialr   )r.   r   r   r   )r'   rq   rr   rs   r   r
   rt   r%   r   r(   r)   )ru   r#   rw   r{   rv   r3   r3   r4   Ú test_nystroem_poly_kernel_paramsg  s       ÿ
r„   c            	   
   C   sÐ   t j d¡} d}| j|dfd�}dd„ }g }t|ƒ}t||d d|id	� |¡ t|ƒ||d  d
 kslt‚d}ddiddidd
if}|D ]@}tf t	|d dœ|—Ž}t
jt|d�� | |¡ W 5 Q R X qŠd S )Nr   rn   r?   r   c                 S   s   |  d¡ t | |¡ ¡ S )z&Histogram kernel that writes to a log.r   )Úappendr'   ÚminimumrB   )Úxrh   rO   r3   r3   r4   Úlogging_histogram_kernelz  s    
z8test_nystroem_callable.<locals>.logging_histogram_kernelr   rO   )r.   r   Zkernel_paramsr   ú-Don't pass gamma, coef0 or degree to Nystroemr   r   r   rz   r=   )r'   rq   rr   rs   Úlistr
   rJ   Úlenr,   r8   rE   rF   rG   )	ru   Ú	n_samplesr#   rˆ   Z
kernel_logrL   ÚparamsÚparamÚnyr3   r3   r4   Útest_nystroem_callablet  s(    ýür�   c            	   
   C   s¼   t j d¡} | jdd�}t|ddd�}td|jd d	�}| |¡}tt  	||j
¡|ƒ d
}ddiddiddif}|D ]B}tf d|jd d	œ|—Ž}tjt|d�� | |¡ W 5 Q R X qtd S )Né   rm   r   r   r   rƒ   Zprecomputedr   rz   r‰   r   r   r   r   r=   )r'   rq   rr   rs   r   r
   rt   r%   r   r(   r)   rE   rF   rG   rJ   )	ru   r#   rw   r{   rv   rL   r�   rŽ   r�   r3   r3   r4   Ú test_nystroem_precomputed_kernel‘  s    
r’   c                  C   s:   t ddd�\} }tddd�}| | ¡ |jjdks6t‚dS )	zÓCheck that `component_indices_` corresponds to the subset of
    training points used to construct the feature map.
    Non-regression test for:
    https://github.com/scikit-learn/scikit-learn/issues/20474
    r~   r}   )rŒ   Z
n_featuresrn   r   rp   )rn   N)r   r
   rJ   Zcomponent_indices_rt   r,   )r#   Ú_Zfeature_map_nystroemr3   r3   r4   Útest_nystroem_component_indices¥  s    þ
r”   Ú	Estimatorc                    sR   | ƒ   t¡}| t¡}| ¡ }| j ¡ ‰ ‡ fdd„t|jd ƒD ƒ}t||ƒ dS )zCheck get_feature_names_outc                    s   g | ]}ˆ › |› �‘qS r3   r3   )Ú.0Úi©Ú
class_namer3   r4   Ú
<listcomp>¾  s     z.test_get_feature_names_out.<locals>.<listcomp>r   N)	rJ   r#   r&   Úget_feature_names_outÚ__name__ÚlowerÚrangert   r   )r•   Zestr/   Ú	names_outÚexpected_namesr3   r˜   r4   Útest_get_feature_names_out´  s    

r¡   c                  C   s|   t j d¡} | jdd�}tdd� |¡}dddg}d	d
dddddddddddddg}|j|d�}dd„ |D ƒ}t||ƒ dS )z4Check get_feature_names_out for AdditiveChi2Sampler.r   )r   r   r   r   r9   Zf0Úf1Úf2Zf0_sqrtZf1_sqrtZf2_sqrtZf0_cos1Zf1_cos1Zf2_cos1Zf0_sin1Zf1_sin1Zf2_sin1Zf0_cos2Zf1_cos2Zf2_cos2Zf0_sin2Zf1_sin2Zf2_sin2)Zinput_featuresc                 S   s   g | ]}d |› �‘qS )Zadditivechi2sampler_r3   )r–   Úsuffixr3   r3   r4   rš   Ü  s     zBtest_additivechi2sampler_get_feature_names_out.<locals>.<listcomp>N)r'   rq   rr   Úrandom_sampler   rJ   r›   r   )r   r#   Zchi2_samplerZinput_namesÚsuffixesrŸ   r    r3   r3   r4   Ú.test_additivechi2sampler_get_feature_names_outÂ  s.    
ñr§   )7rH   Únumpyr'   Zscipy.sparser   rE   Zsklearn.utils._testingr   r   r   Zsklearn.metrics.pairwiser   Zsklearn.kernel_approximationr   r   r	   r
   r   Zsklearn.datasetsr   r   r   r   rq   rr   r   r¥   r#   r$   rB   rA   ÚmarkZparametrizer5   r7   r8   rM   rT   rU   rW   rb   rf   ri   rj   rk   rl   ry   r|   r�   r„   r�   r’   r”   r¡   r§   r3   r3   r3   r4   Ú<module>   sf   ?( 
ÿ
