U
    ½mœda(  ã                   @   sL  d dl 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
 d dlmZmZ d dlmZ d d	lmZ d d
lmZ e ¡ Zdd„ Zdd„ Zdd„ Zdd„ Zdd„ Zdd„ Zdd„ Zdd„ Zdd„ Zdd„ Zej  dd ¡d!d"„ ƒZ!d#d$„ Z"d%d&„ Z#ej  d'ej$ej%g¡ej  d(eeg¡d)d*„ ƒƒZ&ej  d(eeg¡d+d,„ ƒZ'dS )-é    )ÚlogN)Úassert_array_almost_equal)Úassert_almost_equal)Úassert_array_less)Úcheck_random_state)ÚBayesianRidgeÚARDRegression)ÚRidge)Údatasets)Úfast_logdetc                  C   s@   t jt j } }tdd�}| | |¡ |jj|jd fks<t‚dS )zCheck scores attribute shapeT©Úcompute_scoreé   N)	ÚdiabetesÚdataÚtargetr   ÚfitÚscores_ÚshapeZn_iter_ÚAssertionError©ÚXÚyÚclf© r   ú^/home/sam/Atlas/atlas_env/lib/python3.8/site-packages/sklearn/linear_model/tests/test_bayes.pyÚtest_bayesian_ridge_scores   s    
r   c               	   C   s  t jt j } }| jd }t tj¡j}dt |¡|  }d}d}d}d}d}	|t	|ƒ |	|  }
|
|t	|ƒ ||  7 }
d| t 
|¡ d| t | | j¡  }tj ||¡}|
dt|ƒt |j|¡ |t	dtj ƒ   7 }
t||||	dddd	�}| | |¡ t|jd |
d
d� dS )aÀ  Check value of score on toy example.

    Compute log marginal likelihood with equation (36) in Sparse Bayesian
    Learning and the Relevance Vector Machine (Tipping, 2001):

    - 0.5 * (log |Id/alpha + X.X^T/lambda| +
             y^T.(Id/alpha + X.X^T/lambda).y + n * log(2 * pi))
    + lambda_1 * log(lambda) - lambda_2 * lambda
    + alpha_1 * log(alpha) - alpha_2 * alpha

    and check equality with the score computed during training.
    r   ç      ð?çš™™™™™¹?g      à¿é   r   FT)Úalpha_1Úalpha_2Úlambda_1Úlambda_2Ún_iterZfit_interceptr   é	   ©ÚdecimalN)r   r   r   r   ÚnpZfinfoÚfloat64ÚepsÚvarr   ÚeyeÚdotÚTZlinalgZsolver   Úpir   r   r   r   )r   r   Ú	n_samplesr*   Úalpha_Úlambda_r    r!   r"   r#   ÚscoreÚMZM_inv_dot_yr   r   r   r   Ú test_bayesian_ridge_score_values"   s6    
(&ÿù	r5   c               
   C   sš   t  ddgddgddgddgddgddgddgg¡} t  ddddd	ddg¡j}td
d� | |¡}t|j|j d� | |¡}t|j	|j	ƒ t
|j|jƒ d S )Nr   é   é   é   é   r   é   é
   r   Tr   ©Úalpha©r(   Úarrayr.   r   r   r	   r2   r1   r   Úcoef_r   Z
intercept_)r   r   Úbr_modelÚrr_modelr   r   r   Útest_bayesian_ridge_parameterU   s    4rC   c               
   C   s¼   t  ddgddgddgddgddgddgddgg¡} t  ddddd	ddg¡j}t  dddddddg¡j}td
d�j| ||d�}t|j|j d�j| ||d�}t|j	|j	ƒ t
|j|jƒ d S )Nr   r6   r7   r8   r9   r   r:   r;   r   Tr   )Zsample_weightr<   r>   )r   r   ÚwrA   rB   r   r   r   Útest_bayesian_sample_weightsb   s    4  ÿrE   c                  C   st   t  dgdgdgdgdgg¡} t  dddddg¡}tdd�}| | |¡ dgdgd	gg}t| |¡ddd	gdƒ d S )
Nr   r   r:   é   r;   Tr   r6   r7   )r(   r?   r   r   r   Úpredict©r   ÚYr   Útestr   r   r   Útest_toy_bayesian_ridge_objectr   s    
rK   c                  C   sX   t  t  ddd¡d¡} t  dddddg¡}tddd�}| | |¡ | |¡}t|dƒ d S )	Nr   r7   r8   ç        r   ç      ð¿gü©ñÒMbP?)Z
alpha_initZlambda_init)r(   ZvanderZlinspacer?   r   r   r3   r   )r   r   ÚregÚr2r   r   r   Útest_bayesian_initial_params~   s
    rP   c            	      C   sˆ   d} d}t dƒ}| ¡ }| | |f¡}tj| |t |¡jd�}tj| |t |¡jd�}tƒ tƒ fD ] }| 	||¡ 
|¡}t||ƒ qbd S )Nr7   r8   é*   ©Údtype)r   ÚrandÚrandom_sampler(   Úfullr?   rS   r   r   r   rG   r   )	r0   Ú
n_featuresÚrandom_stateÚconstant_valuer   r   Úexpectedr   Zy_predr   r   r   Ú6test_prediction_bayesian_ridge_ard_with_constant_input‹   s    r[   c            
      C   s|   d} d}t dƒ}| ¡ }| | |f¡}tj| |t |¡jd�}d}tƒ tƒ fD ](}| 	||¡j
|dd�\}}	t|	|ƒ qNd S )Nr;   r8   rQ   rR   ç{®Gáz„?T©Z
return_std)r   rT   rU   r(   rV   r?   rS   r   r   r   rG   r   )
r0   rW   rX   rY   r   r   Zexpected_upper_boundaryr   Ú_Úy_stdr   r   r   Ú/test_std_bayesian_ridge_ard_with_constant_input›   s    r`   c                  C   s\   t  ddgddgg¡} t  ddg¡}tdd�}| | |¡ |jjdksJt‚|j| dd� d S )Nr   r   )r$   )r   r   Tr]   )r(   r?   r   r   Úsigma_r   r   rG   r   r   r   r   Útest_update_of_sigma_in_ard¬   s    
rb   c                  C   sh   t  dgdgdgg¡} t  dddg¡}tdd�}| | |¡ dgdgdgg}t| |¡dddgdƒ d S )Nr   r   r6   Tr   r7   )r(   r?   r   r   r   rG   rH   r   r   r   Útest_toy_ard_objectº   s    
rc   zn_samples, n_features))r;   éd   )rd   r;   c                 C   sZ   t j | ¡jdd�}|d d …df }tƒ }| ||¡ t  d|jd  ¡}|dk sVt‚d S )N)éú   r6   )Úsizer   g»½×Ùß|Û=)	r(   ÚrandomÚRandomStateÚnormalr   r   Úabsr@   r   )Úglobal_random_seedr0   rW   r   r   Z	regressorZabs_coef_errorr   r   r   Ú!test_ard_accuracy_on_easy_problemÆ   s    rl   c                     sè   ‡ ‡fdd„‰‡fdd„} d}d}d}t  dd	dd
d	g¡‰d‰ t j ||f¡}t j ||f¡}tdddgƒD ]v\}}| ||ƒ}tƒ }	|	 ||¡ |	j|dd�\}
}t|||d� tƒ }| ||¡ |j|dd�\}}t|||d� qld S )Nc                    s   t  | ˆ¡ˆ  S )N)r(   r-   )r   )ÚbrD   r   r   ÚfÖ   s    ztest_return_std.<locals>.fc                    s   ˆ | ƒt j | jd ¡|  S )Nr   )r(   rg   Úrandnr   )r   Ú
noise_mult)rn   r   r   Úf_noiseÙ   s    z test_return_std.<locals>.f_noiser8   é2   r;   r   rL   rM   r   r   r\   Tr]   r&   )	r(   r?   rg   Ú	enumerater   r   rG   r   r   )rq   ÚdZn_trainZn_testr   ZX_testr'   rp   r   Úm1Zy_mean1Zy_std1Úm2Zy_mean2Zy_std2r   )rm   rn   rD   r   Útest_return_stdÔ   s&    
rw   c                 C   s|   t j | ¡}d }}| ||¡}d}t  d|d ¡}t  dg| ¡}tƒ }| ||||¡}	| ||||¡}
t j	 
|	|
¡ d S )Nr;   r   T)r(   rg   rh   ro   Zaranger?   r   Z_update_sigmaZ_update_sigma_woodburyÚtestingÚassert_allclose)rk   Úrngr0   rW   r   r=   ZlmbdaZkeep_lambdarN   ÚsigmaZsigma_woodburyr   r   r   Útest_update_sigmaô   s    r|   rS   Ú	Estimatorc           	   	   C   sÂ   t jddgddgddgddgddgddgddgg| d	�}t  ddddd
ddg¡j}|ƒ }| ||¡ ddg}|D ]}t||ƒj|jkspt‚qp|j|dd�\}}|j|jks®t‚|j|jks¾t‚d S )Nr   r6   r7   r8   r9   r   r:   r;   rR   r   r@   ra   Tr]   )r(   r?   r.   r   ÚgetattrrS   r   rG   )	rS   r}   r   r   ÚmodelÚ
attributesÚ	attributeZy_meanr_   r   r   r   Útest_dtype_match  s    8r‚   c              
   C   s–   t  ddgddgddgddgddgddgddgg¡}t  ddddd	ddg¡j}| ƒ }| | t j¡|¡j}| | t j¡|¡j}t jj	||d
d� d S )Nr   r6   r7   r8   r9   r   r:   r;   r   g-Cëâ6?)Zrtol)
r(   r?   r.   r   ZastypeÚfloat32r@   r)   rx   ry   )r}   r   r   r   Zcoef_32Zcoef_64r   r   r   Útest_dtype_correctness  s    4r„   )(Úmathr   Únumpyr(   ZpytestZsklearn.utils._testingr   r   r   Zsklearn.utilsr   Zsklearn.linear_modelr   r   r	   Zsklearnr
   Zsklearn.utils.extmathr   Zload_diabetesr   r   r5   rC   rE   rK   rP   r[   r`   rb   rc   ÚmarkZparametrizerl   rw   r|   rƒ   r)   r‚   r„   r   r   r   r   Ú<module>   s<   
3
 