U
    Åmœd¸  ã                   @   sh   d dl Zd dlmZ d dlZd dœdd„Zejedœdd„Zej	d	d
�dd„ ƒZ
ej	d	d
�dd„ ƒZdS )é    N)Úsparse©Úaxisc                C   sv   t  | ¡rt| |d�\}}n6tj| |tjd�}t | | ¡j|tjd�}||d  }|| j| | j| d  9 }||fS )Nr   )r   Údtypeé   é   )r   ÚissparseÚsparse_mean_variance_axisÚnpÚmeanÚfloat64ÚmultiplyÚshape)ÚXr   r   ÚvarZmean_sq© r   úT/home/sam/Atlas/atlas_env/lib/python3.8/site-packages/scanpy/preprocessing/_utils.pyÚ_get_mean_var   s    
r   )Úmtxr   c                 C   s’   |dkst ‚t| tjƒr$d}| j}n*t| tjƒrFd}| jddd… }ntdƒ‚||krtt| j| j	| j
f|tjfžŽ S t| j| j	f|tjfžŽ S dS )a^  
    This code and internal functions are based on sklearns
    `sparsefuncs.mean_variance_axis`.

    Modifications:
    * allow deciding on the output type, which can increase accuracy when calculating the mean and variance of 32bit floats.
    * This doesn't currently implement support for null values, but could.
    * Uses numba not cython
    )r   r   r   r   Néÿÿÿÿz7This function only works on sparse csr and csc matrices)ÚAssertionErrorÚ
isinstancer   Z
csr_matrixr   Z
csc_matrixÚ
ValueErrorÚsparse_mean_var_major_axisÚdataÚindicesÚindptrr
   r   Úsparse_mean_var_minor_axis)r   r   Zax_minorr   r   r   r   r	      s$    
  ÿ ÿr	   T)Úcachec                 C   s  |j d }tj||d�}tj||d�}tj|tjd�}t|ƒD ] }	||	 }
||
  | |	 7  < q>t|ƒD ]}	||	  |  < qht|ƒD ]@}	||	 }
| |	 ||
  }||
  || 7  < ||
  d7  < q†t|ƒD ]8}	||	  |||	  ||	 d  7  < ||	  |  < qÐ||fS )zª
    Computes mean and variance for a sparse matrix for the minor axis.

    Given arrays for a csr matrix, returns the means and variances for each
    column back.
    r   ©r   r   r   )r   r
   ÚzerosÚ
zeros_likeZint64Úrange)r   r   Ú	major_lenÚ	minor_lenr   Znon_zeroÚmeansÚ	variancesÚcountsÚiZcol_indÚdiffr   r   r   r   -   s$    
$r   c                 C   sæ   t j||d�}t j||d�}t|ƒD ]¸}|| }	||d  }
|
|	 }t|	|
ƒD ]}||  | | 7  < qN||  |  < t|	|
ƒD ](}| | ||  }||  || 7  < q‚||  || || d  7  < ||  |  < q$||fS )z¦
    Computes mean and variance for a sparse array for the major axis.

    Given arrays for a csr matrix, returns the means and variances for each
    row back.
    r   r   r   )r
   r    r!   r"   )r   r   r   r#   r$   r   r%   r&   r(   ZstartptrZendptrr'   Újr)   r   r   r   r   P   s     r   )Únumpyr
   Zscipyr   Znumbar   ZspmatrixÚintr	   Znjitr   r   r   r   r   r   Ú<module>   s   

"
