U
    ½mœdÿ	 ã                   @   sò
  d dl Zd dlmZ d dlmZ 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 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" 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. d dl,m/Z/ d d l,m0Z0 d d!l1m2Z2 d d"l1m3Z3 d d#l1m4Z4 d d$l1m5Z5 d d%l1m6Z6 d d&l7m8Z8 d d'lm9Z9 d(Z:d)Z;d*Z<e =¡ Z>e>j?e>j@ ZAZBe CeAjDd  ¡ZEejF Gd ¡ZHeH IeE¡ eEdd+… ZEeAeE eBeE  ZAZBe J¡ ZKe LeKj?¡ZMeKj@ZNd,d-„ ZOd.d/„ ZPd0d1„ ZQd2d3„ ZRe	jSd4d5gd6�d7d8„ ƒZTe	jU Vd9e:¡e	jU Vd:d;d<g¡d=d>„ ƒƒZWe	jU Vd9e:¡e	jU Vd:d;d<g¡d?d@„ ƒƒZXe	jU Vd9e:¡e	jU Vd:d;d<g¡dAdB„ ƒƒZYe	jU Vd9e:¡e	jU Vd:d;d<g¡dCdD„ ƒƒZZe	jU Vd9e:¡e	jU Vd:d;d<g¡dEdF„ ƒƒZ[e	jU Vd9e:¡e	jU Vd:d;d<g¡dGdH„ ƒƒZ\e	jU Vd9e:¡e	jU Vd:d;d<g¡e	jU VdId;d<g¡e	jU VdJdKdLg¡dMdN„ ƒƒƒƒZ]dOdP„ Z^dQdR„ Z_dSdT„ Z`dUdV„ ZadWdX„ ZbdYdZ„ Zce	jU Vd[d\d]d^g¡d_d`„ ƒZde	jU Vdadbdcdddedfg¡e	jU Vdgd;d<g¡dhdi„ ƒƒZee	jU Vdadbdcdddedfg¡e	jU Vdgd;d<g¡djdk„ ƒƒZf�ddrds„Zge	jU Vdtdudv„ edwdxdydzd{d|gd<d;gƒD ƒ¡e	jU Vd}d~dd€g¡e	jU Vd�e Cd‚¡¡dƒd„„ ƒƒƒZhe	jU Vd…d†d‡g¡e	jU VdˆejiejLg¡e	jU Vd‰dŠd‹g¡e	jU Vd:d;d<g¡e	jU VdŒd�dŽd�g¡d�d‘„ ƒƒƒƒƒZjd’d“„ Zke	jU Vd…d†d‡g¡e	jU VdˆejiejLg¡e	jU Vd”d•d–g¡e	jU Vd—d˜d™dšd›g¡dœd�„ ƒƒƒƒZle	jU Vdžd;d<g¡e	jU VdŸd d¡d¢d£g¡d¤d¥„ ƒƒZmd¦d§„ Znd¨d©„ Zoe	jU Vdªe#d<d«�e.fe%d<d«�e/fg¡d¬d­„ ƒZpe	jU Vdªe#ƒ e.fe%ƒ e/fg¡e	jU Vd®dd‚g¡d¯d°„ ƒƒZqd±d²„ Zrd³d´„ Zsdµd¶„ Ztd·d¸„ Zue	jU Vd¹ddºeQg¡e	jU Vd®de3d»ƒg¡e	jU Vd¼eOePg¡d½d¾„ ƒƒƒZve	jU Vd®de3d»ƒg¡e	jU Vd¼eOePg¡d¿dÀ„ ƒƒZwdÁdÂ„ ZxdÃdÄ„ Zye	jU VdÅeneoeseteuexf¡dÆdÇ„ ƒZzdÈdÉ„ Z{e	jU VdÊe$e%f¡dËdÌ„ ƒZ|dÍdÎ„ Z}e	jU Vd¹ddÏeRg¡dÐdÑ„ ƒZ~e	jU Vd¹ddºeQg¡dÒdÓ„ ƒZe	jU VdÔe#e%g¡dÕdÖ„ ƒZ€d×dØ„ Z�dÙdÚ„ Z‚dÛdÜ„ ZƒdÝdÞ„ Z„e	jU VdÔe#e%g¡e	jU Vdßdàdáie…dâfdàdãie…däfdàdåie†dæfg¡dçdè„ ƒƒZ‡e	jU VdÔe#e%g¡dédê„ ƒZˆdëdì„ Z‰dídî„ ZŠedïdð„ ƒZ‹e	jU Vd9dzdydñdòg¡e	jU Vdód;d<g¡dôdõ„ ƒƒZŒe	jU Vd9d{d†dwg¡död÷„ ƒZ�e	jU Vdód;d<g¡dødù„ ƒZŽe	jU Vdúd<d;g¡e	jU Vdûde �dü¡g¡e	jU Vdýej�ejLg¡e	jU Vd9dòdydwdzdxd{dñg¡dþdÿ„ ƒƒƒƒZ‘e	jU Vd9d†dydwdzdxd{dñg¡�d �d„ ƒZ’�d�d„ Z“e	jU Vd9d†dwdzdydxd{dñg¡e	jU Vd�e”doƒ¡�d�d„ ƒƒZ•�d�d„ Z–e	jU V�de$i fe%d®dife%d®d‚ifg¡�d	�d
„ ƒZ—e	jU Vd9dòdñg¡e	jU Vd:d;d<g¡e	jU VdJ�ddL�ddKg¡�d�d„ ƒƒƒZ˜e	jU Vd:d;d<g¡e	jU VdJ�ddL�ddKg¡�d�d„ ƒƒZ™e	jU Vd9d†dwdzdydxd{g¡�d�d„ ƒZše	jU VdJ�ddL�ddKg¡�d�d„ ƒZ›e	jU VdJ�ddL�ddKg¡�d�d„ ƒZœ�d�d„ Z�e	jU Vd9dwdzdyd†dxd{dñg¡�d�d„ ƒZždS (  é    N)Úlinalg)Úproduct)Ú	_IS_32BIT)Úassert_almost_equal)Úassert_allclose)Úassert_array_almost_equal)Úassert_array_equal)Úignore_warnings)Úcheck_sample_weights_invariance)ÚConvergenceWarning)Údatasets©Úmean_squared_error)Úmake_scorer)Ú
get_scorer)ÚLinearRegression)Úridge_regression)ÚRidge)Ú	_RidgeGCV)ÚRidgeCV)ÚRidgeClassifier)ÚRidgeClassifierCV)Ú_solve_cholesky)Ú_solve_cholesky_kernel)Ú
_solve_svd)Ú_solve_lbfgs)Ú_check_gcv_mode)Ú_X_CenterStackOp)Úmake_low_rank_matrix)Úmake_regression)Úmake_classification)Úmake_multilabel_classification)ÚGridSearchCV)ÚKFold)Ú
GroupKFold)Úcross_val_predict)ÚLeaveOneOut)Úminmax_scale)Úcheck_random_state)ÚsvdÚ	sparse_cgÚcholeskyÚlsqrÚsagÚsaga)r*   r-   )r*   r+   r,   r-   r.   éÈ   c                 C   s   | S ©N© ©ÚXr1   r1   ú^/home/sam/Atlas/atlas_env/lib/python3.8/site-packages/sklearn/linear_model/tests/test_ridge.pyÚDENSE_FILTERF   s    r5   c                 C   s
   t  | ¡S r0   )ÚspÚ
csr_matrixr2   r1   r1   r4   ÚSPARSE_FILTERJ   s    r8   c                 C   s   t  | |k¡S r0   )ÚnpÚmean©Zy_testÚy_predr1   r1   r4   Ú_accuracy_callableN   s    r=   c                 C   s   | | d   ¡ S )Né   )r:   r;   r1   r1   r4   Ú_mean_squared_error_callableR   s    r?   ÚlongZwide)Úparamsc                 C   s°  |j dkrd\}}nd\}}t||ƒ}tj | ¡}t||||d�}d|dd…df< t |¡\}}}	t |dk¡stt	‚|dd…d|…f |dd…|d…f  }
}|	d|…dd…f |	|d…dd…f  }}|j dk�r
|j
d	d
|d�}|| }|||j|| d�d  7 }n.|j
d	d
|d�}|jt d| ¡ |
j | }d}|t |¡ }d|d< t |j| | |j| ¡}|||  }|||  }tj |¡tj |¡k �s¤t	‚||||fS )aD  Dataset with OLS and Ridge solutions, well conditioned X.

    The construction is based on the SVD decomposition of X = U S V'.

    Parameters
    ----------
    type : {"long", "wide"}
        If "long", then n_samples > n_features.
        If "wide", then n_features > n_samples.

    For "wide", we return the minimum norm solution w = X' (XX')^-1 y:

        min ||w||_2 subject to X w = y

    Returns
    -------
    X : ndarray
        Last column of 1, i.e. intercept.
    y : ndarray
    coef_ols : ndarray of shape
        Minimum norm OLS solutions, i.e. min ||X w - y||_2_2 (with mininum ||w||_2 in
        case of ambiguity)
        Last coefficient is intercept.
    coef_ridge : ndarray of shape (5,)
        Ridge solution with alpha=1, i.e. min ||X w - y||_2_2 + ||w||_2^2.
        Last coefficient is intercept.
    r@   )é   é   )rC   rB   )Ú	n_samplesÚ
n_featuresÚeffective_rankÚrandom_stateé   Néÿÿÿÿçü©ñÒMbP?éöÿÿÿé
   ©ÚlowÚhighÚsize©rP   r>   r   )rI   rI   )ÚparamÚminr9   ÚrandomÚRandomStater   r   r)   ÚallÚAssertionErrorÚuniformÚnormalÚTZdiagÚidentityZsolveÚnorm)Úglobal_random_seedÚrequestrD   rE   ÚkÚrngr3   ÚUÚsZVtZU1ZU2ZVt1Ú_Zcoef_olsÚyÚalphaÚdZ
coef_ridgeZR_OLSZR_Ridger1   r1   r4   Úols_ridge_datasetV   s<    


   ÿ**rg   ÚsolverÚfit_interceptTFc                 C   sl  |\}}}}d}t |d| | dkr$dnd|d�}	|t |¡ }
|||  }dt |d ¡t |
d ¡  }tf |	Ž}|d	d	…d	d
…f }|r”|d
 }n ||jdd� }|| ¡  }d}| ||¡ |d	d
… }|jt |¡ksàt	‚t
|j|ƒ | ||¡t |¡k�st	‚tf |	Žj||t |jd ¡d�}|jt |¡k�s@t	‚t
|j|ƒ | ||¡t |¡k�sht	‚d	S )zˆTest that Ridge converges for all solvers to correct solution.

    We work with a simple constructed data set with known solution.
    ç      ð?T©r-   r.   çVçž¯Ò<ç»½×Ùß|Û=©re   ri   rh   ÚtolrG   rH   r>   NrI   r   ©Úaxis©Úsample_weight)Údictr9   r:   Úsumr   ÚfitÚ
intercept_ÚpytestÚapproxrW   r   Úcoef_ÚscoreÚonesÚshape)rh   ri   rg   r]   r3   rd   rc   Úcoefre   rA   Zres_nullZ	res_RidgeZR2_RidgeÚmodelÚ	interceptr1   r1   r4   Útest_ridge_regression�   s8    û	 

"r�   c                 C   sü   |\}}}}|j \}}	d}
t|
d || | dkr2dnd|d�}|dd…dd…f }d	tj||fd
d� }tj |¡t||	d
 ƒks„t‚|r’|d }n ||jdd� }|| ¡  }d}| 	||¡ |dd… }|j
t |¡ksÞt‚t|jtj||f dd� dS )a  Test that Ridge converges for all solvers to correct solution on hstacked data.

    We work with a simple constructed data set with known solution.
    Fit on [X] with alpha is the same as fit on [X, X]/2 with alpha/2.
    For long X, [X, X] is a singular matrix.
    rj   r>   rk   rl   rm   rn   NrI   ç      à?rH   rp   r   ç:Œ0âŽyE>©Úatol)r}   r   r9   Úconcatenater   Úmatrix_rankrS   rW   r:   rv   rw   rx   ry   r   rz   Úr_©rh   ri   rg   r]   r3   rd   rc   r~   rD   rE   re   r   r€   r1   r1   r4   Ú test_ridge_regression_hstacked_XÉ   s,    
û
rŠ   c                 C   sø   |\}}}}|j \}}	d}
td|
 || | dkr2dnd|d�}|dd…dd…f }tj||fd	d
�}tj |¡t||	ƒks|t‚tj||f }|r˜|d }n ||j	d	d
� }|| 	¡  }d	}| 
||¡ |dd… }|jt |¡ksät‚t|j|dd� dS )aJ  Test that Ridge converges for all solvers to correct solution on vstacked data.

    We work with a simple constructed data set with known solution.
    Fit on [X] with alpha is the same as fit on [X], [y]
                                                [X], [y] with 2 * alpha.
    For wide X, [X', X'] is a singular matrix.
    rj   r>   rk   rl   rm   rn   NrI   r   rp   rƒ   r„   )r}   r   r9   r†   r   r‡   rS   rW   rˆ   r:   rv   rw   rx   ry   r   rz   r‰   r1   r1   r4   Ú test_ridge_regression_vstacked_Xñ   s.    
û
r‹   c                 C   s8  |\}}}}|j \}}	d}
t|
|| | dkr.dnd|d�}tf |Ž}|rp|dd…dd…f }|d }|dd… }nd}| ||¡ ||	ksŒ|s®|jt |¡ks t‚t|j	|ƒ n†t| 
|¡|ƒ t|| | |ƒ tj tj|j|j	f ¡tj tj||f ¡k�st‚tjdd	� |jt |¡k�s(t‚t|j	|ƒ dS )
a  Test that unpenalized Ridge = OLS converges for all solvers to correct solution.

    We work with a simple constructed data set with known solution.
    Note: This checks the minimum norm solution for wide X, i.e.
    n_samples < n_features:
        min ||w||_2 subject to X w = y
    r   rk   rl   rm   rn   NrI   ú1Ridge does not provide the minimum norm solution.©Úreason)r}   rt   r   rv   rw   rx   ry   rW   r   rz   Úpredictr9   r   r\   rˆ   Úxfail)rh   ri   rg   r]   r3   rd   r~   rc   rD   rE   re   rA   r   r€   r1   r1   r4   Ú!test_ridge_regression_unpenalized  s8    
û
ÿr‘   c                 C   sr  |\}}}}|j \}}	d}
t|
|| | dkr.dnd|d�}|rf|dd…dd…f }|d }|dd… }nd}dtj||fd	d
� }tj |¡t||	ƒksšt‚| ||¡ ||	ks²|sî|j	t
 |¡ksÆt‚| dkrÖt
 ¡  t|jtj||f ƒ n€t| |¡|ƒ tj tj|j	|jf ¡tj tj|||f ¡k�s6t‚t
jdd� |j	t
 |¡k�sXt‚t|jtj||f ƒ dS )a^  Test that unpenalized Ridge = OLS converges for all solvers to correct solution.

    We work with a simple constructed data set with known solution.
    OLS fit on [X] is the same as fit on [X, X]/2.
    For long X, [X, X] is a singular matrix and we check against the minimum norm
    solution:
        min ||w||_2 subject to min ||X w - y||_2
    r   rk   rl   rm   rn   NrI   r‚   rH   rp   r+   rŒ   r�   )r}   r   r9   r†   r   r‡   rS   rW   rv   rw   rx   ry   Úskipr   rz   rˆ   r�   r\   r�   ©rh   ri   rg   r]   r3   rd   r~   rc   rD   rE   re   r   r€   r1   r1   r4   Ú,test_ridge_regression_unpenalized_hstacked_XR  s<    
ûÿr”   c                 C   sV  |\}}}}|j \}}	d}
t|
|| | dkr.dnd|d�}|rf|dd…dd…f }|d }|dd… }nd}tj||fdd�}tj |¡t||	ƒks–t‚tj||f }| 	||¡ ||	ks¼|sÞ|j
t |¡ksÐt‚t|j|ƒ ntt| |¡|ƒ tj tj|j
|jf ¡tj tj||f ¡k�s$t‚tjd	d
� |j
t |¡k�sFt‚t|j|ƒ dS )aˆ  Test that unpenalized Ridge = OLS converges for all solvers to correct solution.

    We work with a simple constructed data set with known solution.
    OLS fit on [X] is the same as fit on [X], [y]
                                         [X], [y].
    For wide X, [X', X'] is a singular matrix and we check against the minimum norm
    solution:
        min ||w||_2 subject to X w = y
    r   rk   rl   rm   rn   NrI   rp   rŒ   r�   )r}   r   r9   r†   r   r‡   rS   rW   rˆ   rv   rw   rx   ry   r   rz   r�   r\   r�   r“   r1   r1   r4   Ú,test_ridge_regression_unpenalized_vstacked_X‰  s:    
ûÿr•   ÚsparseXre   rj   ç{®Gáz„?c                 C   s<  |r.|r| t krt ¡  n|s.| tkr.t ¡  |\}}}}	|j\}
}tjdd|
d�}t||| | dkrhdndd|d�}|d	d	…d	d
…f }tj	||fdd�}tj
||f }tj
|d| f | }|rÌ|	d
 }n ||jdd� }|| ¡  }d}|rút |¡}|j|||d� |	d	d
… }	|jt |¡k�s,t‚t|j|	ƒ d	S )zÝTest that Ridge with sample weights gives correct results.

    We use the following trick:
        ||y - Xw||_2 = (z - Aw)' W (z - Aw)
    for z=[y, y], A' = [X', X'] (vstacked), and W[:n/2] + W[n/2:] = 1, W=diag(W)
    r   rH   rM   rk   rl   rm   é † )re   ri   rh   ro   Úmax_iterrG   NrI   rp   rr   )ÚSPARSE_SOLVERS_WITH_INTERCEPTrx   r’   Ú SPARSE_SOLVERS_WITHOUT_INTERCEPTr}   r`   rX   r   r9   r†   rˆ   r:   r6   r7   rv   rw   ry   rW   r   rz   )rh   ri   r–   re   rg   r]   r3   rd   rc   r~   rD   rE   Úswr   r€   r1   r1   r4   Ú$test_ridge_regression_sample_weightsÀ  s>    

ú

r�   c                  C   sX   t  dd¡} tt| dgd�}t ttj¡}t|| dgd�}t tj|¡j}t||ƒ d S )NrI   rH   r—   ©re   )	Ú
y_diabetesÚreshaper   Ú
X_diabetesr9   ÚdotrZ   r   r   )rd   r~   ÚKZ	dual_coefZcoef2r1   r1   r4   Útest_primal_dual_relationshipñ  s    r¤   c               
   C   sZ   t j d¡} |  d¡}|  dd¡}d}tjt|d�� t||dddd d	d
� W 5 Q R X d S )Nr   é   rL   z3sparse_cg did not converge after [0-9]+ iterations.©Úmatchrj   r*   ç        rH   )re   rh   ro   r™   Úverbose)r9   rT   rU   Úrandnrx   Úwarnsr   r   )r`   rd   r3   Zwarning_messager1   r1   r4   Ú&test_ridge_regression_convergence_failú  s    
      ÿr¬   c                  C   sX  t j d¡} d\}}|  ||¡}|  |¡}|d d …t jf }t j|d| f }tƒ }| ||¡ |jj	|fksrt
‚|jj	dks‚t
‚t|jt jƒs”t
‚t|jtƒs¤t
‚| ||¡ |jj	d|fksÄt
‚|jj	dksÔt
‚t|jt jƒsæt
‚t|jt jƒsøt
‚| ||¡ |jj	d|fk�st
‚|jj	dk�s,t
‚t|jt jƒ�s@t
‚t|jt jƒ�sTt
‚d S )Nr   ©r¥   rL   rH   r1   ©rH   r>   )r>   )r9   rT   rU   rª   ÚnewaxisÚc_r   rv   rz   r}   rW   rw   Ú
isinstanceZndarrayÚfloat)r`   rD   rE   r3   rd   ZY1ÚYÚridger1   r1   r4   Útest_ridge_shapes_type  s,    
rµ   c                  C   sˆ   t j d¡} d\}}|  ||¡}|  |¡}t j|d| f }tƒ }| ||¡ |j}| ||¡ t|jd |ƒ t|jd |d ƒ d S )Nr   r­   rj   rH   )	r9   rT   rU   rª   r°   r   rv   rw   r   )r`   rD   rE   r3   rd   r³   r´   r€   r1   r1   r4   Útest_ridge_intercept#  s    
r¶   c                  C   s�   t j d¡} d\}}|  |¡}|  ||¡}tddd�}tdd�}| ||¡ | ||¡ t|j|jƒ | ||¡ | ||¡ t|j|jƒ d S )Nr   )r¥   rC   r¨   F©re   ri   ©ri   )	r9   rT   rU   rª   r   r   rv   r   rz   )r`   rD   rE   rd   r3   r´   Zolsr1   r1   r4   Útest_ridge_vs_lstsq5  s    

r¹   c            	   	      sÂ   t j d¡} d\}}}|  ||¡‰ |  ||¡‰t  |¡‰t  ‡ fdd„tˆˆjƒD ƒ¡}‡ ‡‡fdd„dD ƒ}|D ]}t||ƒ qrt	ˆd d… d�}d	}t
jt|d
�� | ˆ ˆ¡ W 5 Q R X d S )Né*   )é   rL   r¥   c                    s&   g | ]\}}t |d d� ˆ |¡j‘qS )r+   ©re   rh   ©r   rv   rz   )Ú.0re   Útargetr2   r1   r4   Ú
<listcomp>V  s   ÿz3test_ridge_individual_penalties.<locals>.<listcomp>c                    s$   g | ]}t ˆ|d d� ˆ ˆ¡j‘qS )çê-�™—q=)re   rh   ro   r½   )r¾   rh   ©r3   Z	penaltiesrd   r1   r4   rÀ   \  s   ÿ)r)   r*   r,   r+   r-   r.   rI   rž   zCNumber of targets and number of penalties do not correspond: 4 != 5r¦   )r9   rT   rU   rª   ÚarangeÚarrayÚziprZ   r   r   rx   ÚraisesÚ
ValueErrorrv   )	r`   rD   rE   Ú	n_targetsÚcoef_choleskyZcoefs_indiv_penZcoef_indiv_penr´   Úerr_msgr1   rÂ   r4   Útest_ridge_individual_penaltiesJ  s&    



þÿþrË   Ún_colr1   r®   )é   c           	      C   sÀ   t j d¡}| dd¡}| d¡}| t|ƒ¡}|jd| žŽ }|jd| žŽ }tt |¡||ƒ}t  ||d d …d f |  |d d …d f g¡}t	| 
|¡| 
|¡ƒ t	|j 
|¡|j 
|¡ƒ d S )Nr   é   é   é	   )rÎ   )rÐ   )r9   rT   rU   rª   Úlenr   r6   r7   Zhstackr   r¢   rZ   )	rÌ   r`   r3   ZX_mÚsqrt_swr³   ÚAÚoperatorZreference_operatorr1   r1   r4   Útest_X_CenterStackOpj  s    
.rÕ   r}   )rL   rH   )é   rÐ   )rÍ   é   )r>   r>   )r»   r»   Úuniform_weightsc                 C   sÆ   t j d¡}|j| Ž }|r,t  |jd ¡}n| d| d ¡}t  |¡}t j|d|d�}|| |d d …d f  }| 	|j
¡}t ||d d …d f  ¡}	tdd�}
|
 |	|¡\}}t||ƒ t||ƒ d S ©Nr   rH   )rq   ÚweightsTr¸   )r9   rT   rU   rª   r|   r}   Ú	chisquareÚsqrtÚaverager¢   rZ   r6   r7   r   Z_compute_gramr   )r}   rØ   r`   r3   rœ   rÒ   ÚX_meanÚ
X_centeredZ	true_gramÚX_sparseÚgcvZcomputed_gramÚcomputed_meanr1   r1   r4   Útest_compute_gramx  s    



rã   c                 C   sÆ   t j d¡}|j| Ž }|r,t  |jd ¡}n| d| d ¡}t  |¡}t j|d|d�}|| |d d …d f  }|j	 
|¡}t ||d d …d f  ¡}	tdd�}
|
 |	|¡\}}t||ƒ t||ƒ d S rÙ   )r9   rT   rU   rª   r|   r}   rÛ   rÜ   rÝ   rZ   r¢   r6   r7   r   Z_compute_covariancer   )r}   rØ   r`   r3   rœ   rÒ   rÞ   rß   Ztrue_covariancerà   rá   Zcomputed_covrâ   r1   r1   r4   Útest_compute_covarianceŒ  s    



rä   éd   r‚   rL   rH   ç      *@ç      >@c                 C   sÔ   t | ||||||d|d�	\}}}|dkr4t |g¡}||7 }tj |¡ d||j¡dk}| ¡ }d|| < d||< || |¡8 }|
r®|| t 	|¡d | ¡7 }t 	|¡d }|dkr¾|d }|	rÌ|||fS ||fS )NT)	rD   rE   Ún_informativerÈ   ÚbiasÚnoiseÚshuffler~   rG   rH   r   r¨   )
r   r9   ÚasarrayrT   rU   Zbinomialr}   Úcopyr¢   Úabs)rD   rE   Úproportion_nonzerorè   rÈ   ré   ÚX_offsetrê   rë   r~   ÚpositiverG   r3   rd   ÚcÚmaskZ	removed_Xr1   r1   r4   Ú_make_sparse_offset_regression   s8    ÷ÿ

rô   zsolver, sparse_Xc                 c   s&   | ]\}}|r|d ks||fV  qdS ))r*   ÚridgecvNr1   )r¾   rh   Úsparse_Xr1   r1   r4   Ú	<genexpr>Ï  s    ûr÷   r+   r-   r*   r,   r.   rõ   z"n_samples,dtype,proportion_nonzero)r»   Úfloat32çš™™™™™¹?)é(   rø   rj   )r»   Úfloat64çš™™™™™É?ÚseedrÍ   c                 C   sÎ   d}|dkrdnd}t dd||||d�\}}	t|ƒ}td|d	� ||	¡}
|j|d
d�}|	j|d
d�}	|rrt |¡}| dkrˆt|gd�}nt| d|d�}| ||	¡ t|j	|
j	ddd� t|j
|
j
ddd� d S )Nrj   gÍÌÌÌÌÌì?g      I@g     @@rL   é   )ré   rE   rï   rê   rG   rD   r)   )rh   re   F)rí   rõ   ©Úalphasrm   )rh   ro   re   rJ   ©r…   Úrtol)rô   r'   r   rv   Úastyper6   r7   r   r   rz   rw   )rh   rï   rD   Údtyperö   rý   re   rê   r3   rd   Z	svd_ridger´   r1   r1   r4   Útest_solver_consistencyÍ  s,    ú

r  Úgcv_moder)   ÚeigenÚX_constructorÚX_shape)rÎ   rÏ   )rÎ   r»   zy_shape, noise)©rÎ   rj   )©rÎ   rH   rç   )©rÎ   rÍ   ç     Àb@c              	   C   sÎ   |\}}t |ƒdkr|d nd}t|||dd|dd�\}	}
|
 |¡}
dd	d
ddg}t|||dd�}t| ||d�}| |	|
¡ ||	ƒ}| ||
¡ |jt |j¡ks¦t‚t	|j
|j
dd� t	|j|jdd� d S )Nr>   rI   rH   r   Fr¥   ©rD   rE   rÈ   rG   rë   rê   rè   rJ   rù   rj   ç      $@ç     @�@Úneg_mean_squared_error©Úcvri   r   Úscoring)r  ri   r   ©r  )rÑ   rô   r    r   rv   Úalpha_rx   ry   rW   r   rz   rw   )r  r  r	  Úy_shaperi   rê   rD   rE   rÈ   r3   rd   r   Ú	loo_ridgeÚ	gcv_ridgeÚX_gcvr1   r1   r4   Útest_ridge_gcv_vs_ridge_loo_cvþ  s<    ù
	
üýr  c            	   	   C   s¬   d} d\}}d}t |||ddddd�\}}dd	d
ddg}t|d|| d�}td|| d�}| ||¡ | ||¡ |jt |j¡ks„t‚t|j|jdd� t|j	|j	dd� d S )NZexplained_variance)rL   r¥   rH   r   Fr¥   r  rJ   rù   rj   r  r  Tr  )ri   r   r  r  )
rô   r   rv   r  rx   ry   rW   r   rz   rw   )	r  rD   rE   rÈ   r3   rd   r   r  r  r1   r1   r4   Útest_ridge_loo_cv_asym_scoring1  s2    ù

   ÿr  rE   rÏ   r»   zy_shape, fit_intercept, noise)r
  Trj   )r  Tg      4@)r  Tr  )r  Frç   c                    s  dddddg}t j d¡}t|ƒdkr.|d nd	}td
||dd|d�\}	}
|
 |¡}
d| t|	ƒ¡ }|| ¡  d	  t	¡}t  
t  |	jd ¡|¡‰ | t¡}|	ˆ  |
ˆ   }}t|	jd d�}|j||ˆ d�}t||d|d�}| ||¡ t|j|d�}|j||ˆ d�}t||||d�}|| d ‰‡ ‡fdd„t  |	jd ¡D ƒ‰t  ˆ¡‰||	ƒ}t|d| |d�}|j||
|d� t|ƒdk�r¨|jd d …d d …| |j¡f }n|jd d …| |j¡f }|jt |j¡k�sÚt‚t|ˆdd� t|j|jdd� t|j|jdd� d S )NrJ   rù   rj   r  r  r   r>   rI   rH   rÎ   F)rD   rE   rÈ   rG   rë   rê   rÍ   )Zn_splits)Úgroupsr  )r   r  r  ri   r·   ©r  c                    s"   g | ]}t jˆˆ |k d d�‘qS )r   rp   )r9   ru   )r¾   Úi©ÚindicesZkfold_errorsr1   r4   rÀ     s    z1test_ridge_gcv_sample_weights.<locals>.<listcomp>T)r   Ústore_cv_valuesr  ri   rr   r  )r9   rT   rU   rÑ   rô   r    rª   rS   r  ÚintÚrepeatrÃ   r}   r²   r$   Úsplitr   rv   r   r  r%   rì   Ú
cv_values_Úindexrx   ry   rW   r   rz   rw   )r  r  ri   rE   r  rê   r   r`   rÈ   r3   rd   rs   ZX_tiledZy_tiledr  ZsplitsZkfoldZ	ridge_regZpredictionsr  r  Z
gcv_errorsr1   r   r4   Útest_ridge_gcv_sample_weightsO  sb    ú


üÿ
ü"r(  Úsparsez2mode, mode_n_greater_than_p, mode_p_greater_than_n)Nr)   r  )Úautor)   r  )r  r  r  )r)   r)   r)   c                 C   sH   t ddd�\}}| rt |¡}t||ƒ|ks0t‚t|j|ƒ|ksDt‚d S )Nr¥   r>   )rD   rE   )r   r6   r7   r   rW   rZ   )r)  ÚmodeZmode_n_greater_than_pZmode_p_greater_than_nr3   rc   r1   r1   r4   Útest_check_gcv_mode_choice—  s
    
r,  c                 C   s¦  t jd }g }| tk}t|d�}| | t ƒt¡ |j}| |¡ t}t	t
dd�}td|d�}||jƒ| t ƒtƒ |jt |¡ks„t‚dd„ }	t	|	ƒ}td|d�}
||
jƒ| t ƒtƒ |
jt |¡ksÈt‚tdƒ}td|d�}| | t ƒt¡ |jt |¡k�st‚| tk�r<|j| t ƒtt |¡d	� |jt |¡k�s<t‚t ttf¡j}| | t ƒ|¡ | | t ƒ¡}| | t ƒt¡ | | t ƒ¡}tt ||f¡j|d
d� |S )Nr   r¸   F)Zgreater_is_better)ri   r  c                 S   s   t | |ƒ S r0   r   )Úxrd   r1   r1   r4   ÚfuncÁ  s    z_test_ridge_loo.<locals>.funcr  rr   çñhãˆµøä>r  )r¡   r}   r5   r   rv   rŸ   r  Úappendr	   r   r   r   rx   ry   rW   r   r9   r|   ÚvstackrZ   r�   r   )Úfilter_rD   Úretri   Z	ridge_gcvr  Úfr  Z
ridge_gcv2r.  Z
ridge_gcv3ZscorerZ
ridge_gcv4r³   ÚY_predr<   r1   r1   r4   Ú_test_ridge_loo«  s>    



r6  c                 C   sª   t ƒ }| | tƒt¡ | | tƒ¡ t|jjƒdks8t‚t	|j
ƒtjksLt‚tdƒ}|j|d� | | tƒt¡ | | tƒ¡ t|jjƒdks’t‚t	|j
ƒtjks¦t‚d S )NrH   r¥   r  )r   rv   r¡   rŸ   r�   rÑ   rz   r}   rW   Útyperw   r9   rû   r#   Ú
set_params)r2  Úridge_cvr  r1   r1   r4   Ú_test_ridge_cvá  s    r:  zridge, make_dataset)r"  c                 C   s.   |ddd�\}}|   ||¡ t| dƒr*t‚d S )Né   rº   ©rD   rG   r&  )rv   ÚhasattrrW   )r´   Úmake_datasetr3   rd   r1   r1   r4   Ú#test_ridge_gcv_cv_values_not_storedò  s    	r?  r  c                 C   sL   |ddd�\}}| j d|d� |  ||¡ t| dƒs8t‚t| jtƒsHt‚d S )Nr;  rº   r<  F)r"  r  Úbest_score_)r8  rv   r=  rW   r±   r@  r²   )r´   r>  r  r3   rd   r1   r1   r4   Útest_ridge_best_score   s
    rA  c               	      s¼  t j d¡} d\}}}|  ||¡}t  |d d …dgf t  d|f¡¡t  |d d …dgf dt  d|f¡ ¡ t  |d d …dgf dt  d|f¡ ¡ |  ||¡ ‰ d‰‡ ‡fd	d
„|jD ƒ}tˆdd� ˆ |¡}t	||j
ƒ tt|j
d� ˆ |¡j|jƒ tˆddd� ˆ |¡}|j
j|fk�s$t‚|jj|fk�s8t‚|jj|tˆƒ|fk�sTt‚tdddd� ˆ |¡}|j
j|fk�s~t‚|jj|fk�s’t‚|jj||dfk�sªt‚tˆddd� ˆ |d d …df ¡}t  |j
¡�sÞt‚t  |j¡�sðt‚|jj|tˆƒfk�s
t‚tˆddd� ˆ |¡}t	||j
ƒ tt|j
d� ˆ |¡j|jƒ tˆtƒ dd�}d}tjt|d�� | ˆ |¡ W 5 Q R X tˆddd�}tjt|d�� | ˆ |¡ W 5 Q R X d S )Nrº   )r»   r¥   rÍ   r   rH   gš™™™™™©?r>   rJ   )rH   rå   éè  c                    s    g | ]}t ˆd � ˆ |¡j‘qS )rÿ   )r   rv   r  )r¾   r¿   ©r3   r   r1   r4   rÀ   !  s     z6test_ridge_cv_individual_penalties.<locals>.<listcomp>T)r   Úalpha_per_targetrž   )r   rD  r"  Úr2)r   rD  r  )r   r  rD  z3cv!=None and alpha_per_target=True are incompatibler¦   r;  )r9   rT   rU   rª   r¢   r|   rZ   r   rv   r   r  r   r   rz   r}   rW   r@  r&  rÑ   Zisscalarr&   rx   rÆ   rÇ   )r`   rD   rE   rÈ   rd   Zoptimal_alphasr9  Úmsgr1   rC  r4   Ú"test_ridge_cv_individual_penalties  sd    
"&ÿ&þ
ýÿ ÿ ÿ ÿ ÿrG  c                 C   s2   t dd�}| | tƒt¡ t | | tƒt¡d¡S )NFr¸   r¥   )r   rv   r¡   rŸ   r9   Úroundr{   )r2  r´   r1   r1   r4   Ú_test_ridge_diabetesU  s    
rI  c                 C   s’   t  ttf¡j}tjd }tdd�}| | tƒ|¡ |jjd|fksHt	‚| 
| tƒ¡}| | tƒt¡ | 
| tƒ¡}tt  ||f¡j|dd� d S )NrH   Fr¸   r>   rÍ   ©Údecimal)r9   r1  rŸ   rZ   r¡   r}   r   rv   rz   rW   r�   r   )r2  r³   rE   r´   r5  r<   r1   r1   r4   Ú_test_multi_ridge_diabetes[  s    

rL  c                 C   s¾   t  t¡jd }tjd }tƒ tƒ fD ]L}| | tƒt¡ |jj||fksNt	‚| 
| tƒ¡}t  t|k¡dks&t	‚q&tdƒ}t|d�}| | tƒt¡ | 
| tƒ¡}t  t|k¡dksºt	‚d S )Nr   rH   gHáz®Gé?r¥   r  gš™™™™™é?)r9   ÚuniqueÚy_irisr}   ÚX_irisr   r   rv   rz   rW   r�   r:   r#   )r2  Ú	n_classesrE   Úregr<   r  r1   r1   r4   Ú_test_ridge_classifiersi  s    

rR  r  Zaccuracyr¥   r2  c                 C   s>   t |ƒrt|ƒn|}t||d�}| | tƒt¡ | tƒ¡ d S )N)r  r  )Úcallabler   r   rv   rO  rN  r�   )r2  r  r  Úscoring_Úclfr1   r1   r4   Ú"test_ridge_classifier_with_scoringy  s    rV  c                 C   sj   dd„ }t jdddd�}t|t|ƒ|d�}| | tƒt¡ |jt 	d¡ksNt
‚|jt 	|d	 ¡ksft
‚d S )
Nc                 S   s   dS )Nçáz®GáÚ?r1   r;   r1   r1   r4   Ú_dummy_scoreŒ  s    z:test_ridge_regression_custom_scoring.<locals>._dummy_scoreéþÿÿÿr>   r¥   )Únum)r   r  r  rW  r   )r9   Zlogspacer   r   rv   rO  rN  r@  rx   ry   rW   r  )r2  r  rX  r   rU  r1   r1   r4   Ú$test_ridge_regression_custom_scoring†  s    r[  c                 C   sh   t ddd�}| | tƒt¡ | | tƒt¡}t ddd�}| | tƒt¡ | | tƒt¡}||ksdt‚d S )Nr/  F)ro   ri   rJ   )r   rv   r¡   rŸ   r{   rW   )r2  r´   r{   Zridge2Zscore2r1   r1   r4   Ú_test_tolerance—  s    r\  c                 C   s2   | t ƒ}| tƒ}|d k	r.|d k	r.t||dd� d S )NrÍ   rJ  )r5   r8   r   )Ú	test_funcZ	ret_denseZ
ret_sparser1   r1   r4   Úcheck_dense_sparse£  s    r^  r]  c                 C   s   t | ƒ d S r0   )r^  )r]  r1   r1   r4   Útest_dense_sparse­  s    r_  c                  C   sd  t  ddgddgddgddgddgg¡} dddddg}td d�}| | |¡ t| d	dgg¡t  dg¡ƒ tdd
id�}| | |¡ t| d	dgg¡t  dg¡ƒ tdd�}| | |¡ t| d	dgg¡t  dg¡ƒ t  ddgddgddgddgg¡} ddddg}td d�}| | |¡ tdd�}| | |¡ t|jƒdk�sDt‚t	|j
|j
ƒ t	|j|jƒ d S )Nç      ð¿r   çš™™™™™é¿rj   r¨   rH   rI   ©Úclass_weightrü   rJ   Úbalancedr>   )r9   rÄ   r   rv   r   r�   rÑ   Zclasses_rW   r   rz   rw   )r3   rd   rQ  Zregar1   r1   r4   Útest_class_weights¼  s(    (

"

re  rQ  c                 C   sø   | ƒ }|  tjtj¡ | dd�}|  tjtj¡ t|j|jƒ t tjj¡}|tjdk  d9  < ddddœ}| ƒ }|  tjtj|¡ | |d�}|  tjtj¡ t|j|jƒ | ƒ }|  tjtj|d ¡ | |d�}|  tjtj|¡ t|j|jƒ d	S )
z5Check class_weights resemble sample_weights behavior.rd  rb  rH   rå   rj   g      Y@)r   rH   r>   r>   N)	rv   ÚirisÚdatar¿   r   rz   r9   r|   r}   )rQ  Zreg1Zreg2rs   rc  r1   r1   r4   Ú"test_class_weight_vs_sample_weightß  s$    


rh  c                  C   sš   t  ddgddgddgddgddgg¡} dddddg}td dd	dgd
�}| | |¡ tddidd	ddgd
�}| | |¡ t| ddgg¡t  dg¡ƒ d S )Nr`  r   ra  rj   r¨   rH   rI   r—   rù   )rc  r   rJ   rL   gš™™™™™É¿r>   )r9   rÄ   r   rv   r   r�   )r3   rd   rQ  r1   r1   r4   Útest_class_weights_cvü  s    (ri  r  c              	   C   sê   t j d¡}d}d}| ||¡}dddg}t|ƒ}t| ƒrBt| ƒn| }t|d d|d�}| |¡}	| ||	¡ |j	j
||fks€t‚d	}
| ||
¡}	| ||	¡ |j	j
||
|fks²t‚td	d| d
�}tjtdd�� | ||	¡ W 5 Q R X d S )Nrº   rÏ   r¥   rù   rj   r  T©r   r  r"  r  rÍ   )r  r"  r  zcv!=None and store_cv_valuesr¦   )r9   rT   rU   rª   rÑ   rS  r   r   rv   r&  r}   rW   rx   rÆ   rÇ   )r  r`   rD   rE   r-  r   Ún_alphasrT  Úrrd   rÈ   r1   r1   r4   Útest_ridgecv_store_cv_values  s$    

rm  c           	   	   C   s  t  ddgddgddgddgddgg¡}t  dddddg¡}|jd }ddd	g}t|ƒ}t| ƒrht| ƒn| }t|d d
|d�}d}| ||¡ |jj|||fks¢t	‚t  dddddgdddddgdddddgg¡ 
¡ }|jd }| ||¡ |jj|||fk�st	‚d S )Nr`  r   ra  rj   r¨   rH   rI   rù   r  Trj  )r9   rÄ   r}   rÑ   rS  r   r   rv   r&  rW   Z	transpose)	r  r-  rd   rD   r   rk  rT  rl  rÈ   r1   r1   r4   Ú(test_ridge_classifier_cv_store_cv_values+  s*    (

   ÿ&ÿ
rn  Ú	Estimatorc                 C   sŽ   t j d¡}d}d\}}| tkr,| |¡}n| dd|¡}| ||¡}| |d�}|j|ksltd| j› d�ƒ‚| 	||¡ t
|jt  |¡ƒ d S )Nr   ©rù   rj   r  ©r¥   r¥   r>   rÿ   z`alphas` was mutated in `z
.__init__`)r9   rT   rU   r   rª   Úrandintr   rW   Ú__name__rv   r   rì   )ro  r`   r   rD   rE   rd   r3   Z	ridge_estr1   r1   r4   Útest_ridgecv_alphas_conversionH  s    
ÿþrt  c                  C   s´   t j d¡} d}dD ]š\}}|  |¡}|  ||¡}d|  |¡ }tdƒ}t||d�}|j|||d� d|i}	tt	ƒ |	|d	�}
|
j|||d� |j
|
jjksžt‚t|j|
jjƒ qd S )
Nr   rp  )©r;  r¥   r­   rj   r¥   )r   r  rr   re   r  )r9   rT   rU   rª   Úrandr#   r   rv   r"   r   r  Zbest_estimator_re   rW   r   rz   )r`   r   rD   rE   rd   r3   rs   r  rõ   Ú
parametersÚgsr1   r1   r4   Útest_ridgecv_sample_weight]  s    
ry  c               
      s(  ddg} ddg}t j d¡}t| |ƒD ]ü\}}| ||¡‰ | |¡‰| |¡d d }d}d}|d d …t jf ‰|t jd d …f ‰tdd�‰ˆ ˆ ˆ|¡ ˆ ˆ ˆ|¡ ˆ ˆ ˆ|¡ ‡ ‡‡‡fdd	„}‡ ‡‡‡fd
d„}	d}
tj	t
|
d�� |ƒ  W 5 Q R X d}
tj	t
|
d�� |	ƒ  W 5 Q R X q&d S )Nr>   rÍ   rº   rH   rj   g       @rž   c                      s   ˆ  ˆ ˆˆ¡ d S r0   ©rv   r1   )r3   r´   Úsample_weights_not_OKrd   r1   r4   Úfit_ridge_not_ok�  s    zStest_raises_value_error_if_sample_weights_greater_than_1d.<locals>.fit_ridge_not_okc                      s   ˆ  ˆ ˆˆ¡ d S r0   rz  r1   )r3   r´   Úsample_weights_not_OK_2rd   r1   r4   Úfit_ridge_not_ok_2�  s    zUtest_raises_value_error_if_sample_weights_greater_than_1d.<locals>.fit_ridge_not_ok_2z)Sample weights must be 1D array or scalarr¦   )r9   rT   rU   rÅ   rª   r¯   r   rv   rx   rÆ   rÇ   )Ú
n_samplessÚn_featuressr`   rD   rE   Zsample_weights_OKZsample_weights_OK_1Zsample_weights_OK_2r|  r~  rÊ   r1   )r3   r´   r{  r}  rd   r4   Ú9test_raises_value_error_if_sample_weights_greater_than_1du  s.    

r�  c                  C   sÐ   ddg} ddg}t j d¡}tjtjtjtjtjg}t	ddd�}t	ddd�}t
| |ƒD ]t\}}| ||¡}| |¡}	| |¡d d }
|D ]>}||ƒ}|j||	|
d� |j||	|
d� t|j|jd	d
� qŠqVd S )Nr>   rÍ   rº   rj   Fr·   rH   rr   r;  rJ  )r9   rT   rU   r6   Z
coo_matrixr7   Z
csc_matrixZ
lil_matrixZ
dok_matrixr   rÅ   rª   rv   r   rz   )r  r€  r`   Zsparse_matrix_convertersÚsparse_ridgeÚdense_ridgerD   rE   r3   rd   Zsample_weightsZsparse_converterrà   r1   r1   r4   Ú&test_sparse_design_with_sample_weightsœ  s(    û
r„  c                  C   sP   t  ddgddgddgddgddgg¡} dddddg}tdd	�}| | |¡ d S )
Nr`  r   ra  rj   r¨   rH   rI   )rH   rL   rå   rÿ   )r9   rÄ   r   rv   )r3   rd   r´   r1   r1   r4   Útest_ridgecv_int_alphas»  s    (
r…  zparams, err_type, err_msgr   )rH   rI   iœÿÿÿz alphas\[1\] == -1, must be > 0.0)gš™™™™™¹¿r`  g      $Àz"alphas\[0\] == -0.1, must be > 0.0)rH   rj   Ú1z1alphas\[2\] must be an instance of float, not strc              	   C   sR   d\}}t  ||¡}t  dd|¡}tj||d�� | f |Ž ||¡ W 5 Q R X dS )z?Check the `alphas` validation in RidgeCV and RidgeClassifierCV.rq  r   r>   r¦   N)r`   rª   rr  rx   rÆ   rv   )ro  rA   Zerr_typerÊ   rD   rE   r3   rd   r1   r1   r4   Útest_ridgecv_alphas_validationÄ  s
    r‡  c                 C   sL   d\}}t  ||¡}| tkr(t  |¡}nt  dd|¡}| dd� ||¡ dS )zÉCheck the case when `alphas` is a scalar.
    This case was supported in the past when `alphas` where converted
    into array in `__init__`.
    We add this test to ensure backward compatibility.
    rq  r   r>   rH   rÿ   N)r`   rª   r   rr  rv   )ro  rD   rE   r3   rd   r1   r1   r4   Útest_ridgecv_alphas_scalarà  s    rˆ  c                      s&   d‰t ‰ dˆ ‰‡ ‡‡‡fdd„‰d S )Nz5This is not a solver (MagritteSolveCV QuantumBitcoin)zQKnown solvers are 'sparse_cg', 'cholesky', 'svd' 'lsqr', 'sag' or 'saga'. Got %s.c               	      sH   t  d¡} t  d¡}t| |dˆd� tjˆ ˆd�� ˆƒ  W 5 Q R X d S )NrÍ   rj   r¼   r¦   )r9   Úeyer|   r   rx   rÆ   ©r3   rd   ©Ú	exceptionr.  ÚmessageZwrong_solverr1   r4   r.  þ  s
    

z=test_raises_value_error_if_solver_not_supported.<locals>.func)rÇ   r1   r1   r‹  r4   Ú/test_raises_value_error_if_solver_not_supportedò  s    ÿÿrŽ  c                  C   s6   t ddd�} |  tt¡ | jjd tjd ks2t‚d S )Nr*   rH   )rh   r™   r   )r   rv   r¡   rŸ   rz   r}   rW   )rQ  r1   r1   r4   Útest_sparse_cg_max_iter  s    r�  c                  C   sž   d} t t }}t || df¡j}tddƒD ]<}dD ]2}t||dd�}| ||¡ t|j	t || ¡ƒ q2q*dD ],}t|ddd�}| ||¡ |j	d kslt
‚qld S )	Nr>   rH   rC   )r-   r.   r,   rÁ   )rh   r™   ro   )r*   r)   r+   rù   )r¡   rŸ   r9   ZtilerZ   Úranger   rv   r   Zn_iter_rW   )rÈ   r3   rd   Zy_nr™   rh   rQ  r1   r1   r4   Útest_n_iter  s    
r‘  Úlbfgsr*  Úwith_sample_weightc                 C   sº   | dk}t d||d�\}}d}|rDtj |¡}d|j|jd d� }| dkrPd	n| }t|d
|d�}	t| d
|d�}
|	j|||d� |
jt 	|¡||d� t
|	j|
jƒ t
|	j|
jdd� dS )aó  Check that ridge finds the same coefs and intercept on dense and sparse input
    in the presence of sample weights.

    For now only sparse_cg and lbfgs can correctly fit an intercept
    with sparse X with default tol and max_iter.
    'sag' is tested separately in test_ridge_fit_intercept_sparse_sag because it
    requires more iterations and should raise a warning if default max_iter is used.
    Other solvers raise an exception, as checked in
    test_ridge_fit_intercept_sparse_error
    r’  r»   )rE   rG   rñ   Nrj   r   rQ   r*  r*   rÁ   )rh   ro   rñ   rr   g�íµ ÷Æ >r  )rô   r9   rT   rU   rX   r}   r   rv   r6   r7   r   rw   rz   )rh   r“  r]   rñ   r3   rd   rs   r`   Zdense_solverrƒ  r‚  r1   r1   r4   Útest_ridge_fit_intercept_sparse   s"      ÿ
r”  c              	   C   sX   t ddd�\}}t |¡}t| d�}d | ¡}tjt|d�� | ||¡ W 5 Q R X d S )Nr»   r   )rE   rG   ©rh   zsolver='{}' does not supportr¦   )	rô   r6   r7   r   Úformatrx   rÆ   rÇ   rv   )rh   r3   rd   ÚX_csrr‚  rÊ   r1   r1   r4   Ú%test_ridge_fit_intercept_sparse_errorF  s    


r˜  c           
   	   C   s
  t dd|dd�\}}| r<tj |¡}d|j|jd d� }nd }t |¡}tddd	d
dd�}t	f |Ž}t	f |Ž}	|j
|||d� t ¡ �" t dt¡ |	j
|||d� W 5 Q R X t|j|	jdd� t|j|	jdd� tjtdd�� t	dd	dd d� 
||¡ W 5 Q R X d S )Nr¥   r»   g      @)rE   rD   rG   rð   rj   r   rQ   r-   Trm   r˜   )re   rh   ri   ro   r™   rr   Úerrorç-Cëâ6?r  z"sag" solver requires.*r¦   rJ   )rh   ri   ro   r™   )rô   r9   rT   rU   rX   r}   r6   r7   rt   r   rv   ÚwarningsÚcatch_warningsÚsimplefilterÚUserWarningr   rw   rz   rx   r«   )
r“  r]   r3   rd   r`   rs   r—  rA   rƒ  r‚  r1   r1   r4   Ú#test_ridge_fit_intercept_sparse_sagP  s8       ÿ

    ÿ


rŸ  Úreturn_interceptrs   rB  Úarr_typec                 C   sþ   t dƒ}| dd¡}dddg}t ||¡}d}| r6d}||7 }||ƒ}	d	\}
}trVd
nd}|dk}|dkr¤| r¤tjtdd�� t|	||
||| ||d� W 5 Q R X dS t|	||
|||| |d�}| rê|\}}t	||d|d� t	||d|d� nt	||d|d� dS )z=check if all combinations of arguments give valid estimationsrº   rB  rÍ   rH   r>   rù   r¨   g     ˆÃ@)rJ   ç�íµ ÷Æ°>rJ   rš  r’  )r-   r*  zIn Ridge, only 'sag' solverr¦   )re   rh   rs   r   rñ   ro   N)re   rh   rs   rñ   r   ro   r   ©r  r…   )
r(   rv  r9   r¢   r   rx   rÆ   rÇ   r   r   )r   rs   r¡  rh   r`   r3   Z
true_coefsrd   Ztrue_interceptZ	X_testingre   ro   r…   rñ   Úoutr~   r€   r1   r1   r4   Ú.test_ridge_regression_check_arguments_validityk  sP    
ø
ør¥  c                 C   s  t j d¡}d}| dk}d\}}| ||¡}| |¡}| t j¡}| t j¡}	dt  t j¡j }
t|| d|
|d�}| 	||	¡ |j
}t|| d|
|d�}| 	||¡ |j
}|j|jks¸t‚|j|jksÈt‚| |¡j|jksÞt‚| |¡j|jksôt‚t|j
|j
dd	d
� d S )Nr   rj   r’  ru  r>   éô  )re   rh   r™   ro   rñ   rš  gü©ñÒMb@?r£  )r9   rT   rU   rª   r  rø   ZfinfoÚ
resolutionr   rv   rz   r  rW   r�   r   )rh   r`   re   rñ   rD   rE   ÚX_64Úy_64ÚX_32Úy_32ro   Úridge_32Úcoef_32Úridge_64Úcoef_64r1   r1   r4   Útest_dtype_match¨  s@    
    ÿ    ÿr°  c                  C   sò   t j d¡} t  ddg¡}d\}}}|  ||¡}|  ||¡}| t j¡}| t j¡}t|dd�}	|	 ||¡ |	j	}
t|dd�}| ||¡ |j	}|
j
|j
ks t‚|j
|j
ks°t‚|	 |¡j
|j
ksÆt‚| |¡j
|j
ksÜt‚t|	j	|j	dd� d S )	Nr   rj   r‚   )r;  r×   r>   r+   r¼   r¥   rJ  )r9   rT   rU   rÄ   rª   r  rø   r   rv   rz   r  rW   r�   r   )r`   re   rD   rE   Zn_targetr¨  r©  rª  r«  r¬  r­  r®  r¯  r1   r1   r4   Útest_dtype_match_choleskyÍ  s$    
r±  c                 C   sð   t j |¡}d\}}| ||¡}| |¡}t  ||¡d| |¡  }d}| dk}	tƒ }
| dkrbdnd}t jt jfD ]2}t| 	|¡| 	|¡|| |d |	dd	d
d
d�|
|< qr|
t j j
t jks¼t‚|
t j j
t jksÒt‚t|
t j |
t j |d� d S )Nru  r—   rj   r’  r*   rJ   r/  r¦  rm   F)	re   rh   rG   rs   rñ   r™   ro   Zreturn_n_iterr   r„   )r9   rT   rU   rª   r¢   rt   rø   rû   r   r  r  rW   r   )rh   rý   rG   rD   rE   r3   r~   rd   re   rñ   Úresultsr…   Zcurrent_dtyper1   r1   r4   Ú%test_ridge_regression_dtype_stabilityë  s4    
õr³  c                  C   sR   t dd�\} }t | ¡} | d d d…d d …f } |d d d… }tdd� | |¡ d S )Nrº   ©rG   r>   r-   r•  )r   r9   Zasfortranarrayr   rv   rŠ  r1   r1   r4   Útest_ridge_sag_with_X_fortran  s
    
rµ  zClassifier, paramsc                 C   s’   t ddd�\}}| dd¡}tj||gdd�}| f |Ž ||¡}| |¡}|j|jksZt‚t|dd…df |dd…df ƒ t	dd� ||¡ dS )	zRCheck that multilabel classification is supported and give meaningful
    results.rH   r   )rP  rG   rI   rp   Nr-   r•  )
r!   r    r9   r†   rv   r�   r}   rW   r   r   )Ú
ClassifierrA   r3   rd   r³   rU  r5  r1   r1   r4   Útest_ridgeclassifier_multilabel  s    
"r·  rJ   rù   c                 C   s†   t  ddgddgddgddgg¡}t  dd	g¡}|rHd
}| |¡| }n
| |¡}t|d| |d�}| ||¡ t  |jdk¡s‚t‚dS )z:Test that positive Ridge finds true positive coefficients.rH   r>   rÍ   rC   r¥   r;  r×   rÏ   rK   r»   T©re   rñ   rh   ri   r   N)r9   rÄ   r¢   r   rv   rV   rz   rW   )rh   ri   re   r3   r~   r€   rd   r   r1   r1   r4   Ú#test_ridge_positive_regression_test/  s    "
   ÿr¹  c           
      C   s¬   t j d¡}| dd¡}|jdd|jd d�}| rDd}|| | }n|| }||j|jd d�d	 7 }g }d
D ](}t||| dd�}	| |	 	||¡j
¡ qnt|dddœŽ dS )z¸Test that Ridge w/wo positive converges to the same solution.

    Ridge with positive=True and positive=False must give the same
    when the ground truth coefs are all positive.
    rº   é,  rå   rù   rj   rH   rQ   r   r—   )TFrm   )re   rñ   ri   ro   r¢  r  N)r9   rT   rU   rª   rX   r}   rY   r   r0  rv   rz   r   )
ri   re   r`   r3   r~   r€   rd   r²  rñ   r   r1   r1   r4   Ú%test_ridge_ground_truth_positive_testC  s$       ÿr»  c              	   C   sœ   d}t  ddgddgg¡}t  ddg¡}|| }t|d| dd	�}tjtd
d�� | ||¡ W 5 Q R X tjtdd�� t|||d| dd�\}}W 5 Q R X dS )z5Test input validation for positive argument in Ridge.rù   rH   r>   rÍ   rC   rI   TFr¸  zdoes not support positiver¦   zonly 'lbfgs' solver can be used)rñ   rh   r   N)r9   rÄ   r   rx   rÆ   rÇ   rv   r   )rh   re   r3   r~   rd   r   rc   r1   r1   r4   Útest_ridge_positive_error_test^  s          ÿr¼  c           	         s˜   t dddd�\‰ ‰d‰d}d‡ ‡‡fdd	„	}tˆd
� ˆ ˆ¡}tˆdd� ˆ ˆ¡}||ƒ}||ƒ}||ksnt‚t|ƒD ]}|||d�}||ksvt‚qvdS )z?Check ridge loss consistency when positive argument is enabled.rº  rº   ©rD   rE   rG   rù   rå   Nrƒ   c                    sp   | j }|d k	r6tj |¡}| j|jd|| jjd� }n| j}dt ˆˆ |  | d ¡ dˆ t |d ¡  S )Nr   rQ   r‚   r>   )rw   r9   rT   rU   rz   rX   r}   ru   )r   rG   Znoise_scaler€   r`   r~   ©r3   re   rd   r1   r4   Ú
ridge_lossy  s    &ÿz,test_positive_ridge_loss.<locals>.ridge_lossrž   T)re   rñ   r´  )Nrƒ   )r   r   rv   rW   r�  )	re   Zn_checksr¿  r   Zmodel_positiveZlossZloss_positiverG   Zloss_perturbedr1   r¾  r4   Útest_positive_ridge_lossr  s    rÀ  c                 C   sf   t dddd�\}}t |d¡}t | g¡} ddddœ}t||| f|Ž}t||| ƒ}t||d	d
d� dS )zETest that LBGFS gets almost the same coef of svd when positive=False.rº  rº   r½  rH   Fg¼‰Ø—²Òœ<i ¡ )rñ   ro   r™   rš  r   r  N)r   r9   Zexpand_dimsrì   r   r   r   )re   r3   rd   ÚconfigZ
coef_lbfgsrÉ   r1   r1   r4   Útest_lbfgs_solver_consistency˜  s    ýrÂ  c               	   C   sb   t  ddgddgg¡} t  ddg¡}tddddd	dd
�}tjtdd�� | | |¡ W 5 Q R X dS )z1Test that LBFGS solver raises ConvergenceWarning.rH   rI   g    _ Âg    _ Br—   r’  FrÁ   T)re   rh   ri   ro   rñ   r™   zlbfgs solver did not converger¦   N)r9   rÄ   r   rx   r«   r   rv   )r3   rd   r   r1   r1   r4   Útest_lbfgs_solver_error©  s    úrÃ  c                 C   s   t d| d| dkd�}tf |Ž}|jj}t||dd� t||dd� tj d¡}td	d
dd|d�\}}|j	dd|j
d d�}tj||gdd�}tj||gdd�}	tj||gdd�}
tf |Žj||d| d�}tf |Žj||	|
d�}t|j|jƒ t|j|jƒ dS )z›Test that Ridge fulfils sample weight invariance.

    Note that this test is stricter than the common test
    check_sample_weights_invariance alone.
    rj   rÁ   r’  )re   rh   ro   rñ   r|   )ÚkindZzerosrº   rå   rº  rL   é2   )rD   rE   rF   rè   rG   r—   r>   r   rM   rp   rr   N)rt   r   Ú	__class__rs  r
   r9   rT   rU   r   rX   r}   r†   rv   r   rz   rw   )rh   rA   rQ  Únamer`   r3   rd   rœ   ZX_dupZy_dupZsw_dupZ	ridge_2swZ	ridge_dupr1   r1   r4   Ú#test_ridge_sample_weight_invarianceº  s4    	ü
û
rÈ  )rå   rå   r‚   rL   rH   ræ   rç   rç   TFFN)ŸÚnumpyr9   Zscipy.sparser)  r6   Zscipyr   Ú	itertoolsr   rx   r›  Zsklearn.utilsr   Zsklearn.utils._testingr   r   r   r   r	   Zsklearn.utils.estimator_checksr
   Zsklearn.exceptionsr   Zsklearnr   Zsklearn.metricsr   r   r   Zsklearn.linear_modelr   r   r   Zsklearn.linear_model._ridger   r   r   r   r   r   r   r   r   r   Zsklearn.datasetsr   r   r    r!   Zsklearn.model_selectionr"   r#   r$   r%   r&   Zsklearn.preprocessingr'   r(   ZSOLVERSrš   r›   Zload_diabetesZdiabetesrg  r¿   r¡   rŸ   rÃ   r}   ÚindrT   rU   r`   rë   Z	load_irisrf  r7   rO  rN  r5   r8   r=   r?   Zfixturerg   ÚmarkZparametrizer�   rŠ   r‹   r‘   r”   r•   r�   r¤   r¬   rµ   r¶   r¹   rË   rÕ   rã   rä   rô   r  rì   r  r  r(  r,  r6  r:  r?  rA  rG  rI  rL  rR  rV  r[  r\  r^  r_  re  rh  ri  rm  rn  rt  ry  r�  r„  r…  rÇ   Ú	TypeErrorr‡  rˆ  rŽ  r�  r‘  r”  r˜  rŸ  r|   rÄ   r¥  r°  r±  r�  r³  rµ  r·  r¹  r»  r¼  rÀ  rÂ  rÃ  rÈ  r1   r1   r1   r4   Ú<module>   sN  

F*&(555-	 
           ô
-þþþþ!ýþ'üþ	<üþ	
6þþ
þ	G

úþ
#
 ÿ


'	ýýùþ

$
	
 ÿ7 ÿ" ÿ 


ýþ ÿ% ÿ