U
    Åmœd34  ã                   @   s’  d Z ddl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dddd	œeeeee f eee ee ee ee e
jd
œdd„Zd$e
je
jed ee ee
j eeee ee ee f dœdd„Ze
jee ed edœdd„Zd%ddddœeee eeeef  eeee
jdœdd„Zd&ddœeee eeeef  ee
jdœdd„Zdddddœd d!„Zdddddœd"d#„ZdS )'z9This module contains helper functions for accessing data.é    )ÚOptionalÚIterableÚTupleÚUnionÚListN)Úspmatrix)ÚAnnDataé   )ÚLiteralZrank_genes_groups)ÚkeyÚpval_cutoffÚ
log2fc_minÚ
log2fc_maxÚgene_symbols)ÚadataÚgroupr   r   r   r   r   Úreturnc                   s�  t ˆtƒrˆg‰ˆdkr.tˆ jˆ d jjƒ‰dddddg}‡ ‡‡fdd„|D ƒ}tj|d	dd
g|d�}|jd	d� 	¡ }tj
|d
 ˆd�|d
< | d
dg¡jdd�}|dk	r¼||d |k  }|dk	rÔ||d |k }|dk	rì||d |k  }|dk	�r
|jˆ j| dd�}dddœ ¡ D ]N\}	}
|	ˆ jˆ k�rˆ jˆ |	 ˆ jdd� 	¡ jdd
|
d�}| |¡}�qtˆƒd	k�r„|jd
dd� |j	dd�S )aÜ      :func:`scanpy.tl.rank_genes_groups` results in the form of a
    :class:`~pandas.DataFrame`.

    Params
    ------
    adata
        Object to get results from.
    group
        Which group (as in :func:`scanpy.tl.rank_genes_groups`'s `groupby`
        argument) to return results from. Can be a list. All groups are
        returned if groups is `None`.
    key
        Key differential expression groups were stored under.
    pval_cutoff
        Return only adjusted p-values below the  cutoff.
    log2fc_min
        Minimum logfc to return.
    log2fc_max
        Maximum logfc to return.
    gene_symbols
        Column name in `.var` DataFrame that stores gene symbols. Specifying
        this will add that column to the returned dataframe.

    Example
    -------
    >>> import scanpy as sc
    >>> pbmc = sc.datasets.pbmc68k_reduced()
    >>> sc.tl.rank_genes_groups(pbmc, groupby="louvain", use_raw=True)
    >>> dedf = sc.get.rank_genes_groups_df(pbmc, group="0")
    NÚnamesZscoresZlogfoldchangesZpvalsZ	pvals_adjc                    s$   g | ]}t  ˆ jˆ | ¡ˆ ‘qS © )ÚpdÚ	DataFrameÚuns)Ú.0Úc©r   r   r   r   úG/home/sam/Atlas/atlas_env/lib/python3.8/site-packages/scanpy/get/get.pyÚ
<listcomp>@   s     z(rank_genes_groups_df.<locals>.<listcomp>é   r   )Úaxisr   Úkeys)Úlevel)Ú
categoriesZlevel_0)Úcolumns)ÚonZpct_nz_groupZpct_nz_reference)ÚptsZpts_rest©Úindex)Zid_varsÚvar_nameZ
value_nameT)r"   Zinplace)Údrop)Ú
isinstanceÚstrÚlistr   Zdtyper   r   ÚconcatÚstackZreset_indexZCategoricalZsort_valuesr(   ÚjoinÚvarÚitemsZrename_axisZmeltÚmergeÚlen)r   r   r   r   r   r   r   ZcolnamesÚdr$   ÚnameZpts_dfr   r   r   Úrank_genes_groups_df   s@    )

ÿ
  ýÿr5   F©Úobsr/   )Údim_dfÚ	alt_indexÚdimr   Úalias_indexÚuse_rawr   c                 C   sÆ  |r
d}nd}d|dk }d}|dk	rLt j||d�}	|j}|› d|› d�}
nt j||d�}	|› d	�}
g }g }g }g }| jjs¤| j| j ¡   ¡ }td
|› d|› �ƒ‚|jsÊt|› d|› d|› d|› d�ƒ‚t 	|¡D ]¶}|| jk�r| 
|¡ ||	jk�rŠtd|› d|› d|› d|
› d�	ƒ‚qÔ||	jk�r€|	| }t|t jƒ�rj|dk	�sNt‚td|› d|› d|
› d�ƒ‚| 
|¡ | 
|¡ qÔ| 
|¡ qÔt|ƒdk�r¼td|› d|› d|› d|
› d�	ƒ‚|||fS )z8Common logic for checking indices for obs_df and var_df.z	adata.rawr   r6   r7   Nr%   z['z']Z_nameszadata.z_ contains duplicated columns. Please rename or remove these columns first.
`Duplicated columns Ú.z5_names contains duplicated items
Please rename these z& names first for example using `adata.z_names_make_unique()`z	The key 'z' is found in both adata.z and zFound duplicate entries for 'z' in r   zCould not find keys 'z' in columns of `adata.z` or in )r   ZSeriesr4   r"   Z	is_uniqueZ
duplicatedÚtolistÚ
ValueErrorÚnpÚuniqueÚappendr&   ÚKeyErrorr)   ÚAssertionErrorr2   )r8   r9   r:   r   r;   r<   Zalt_reprZalt_dimZ
alias_nameZ	alt_namesZalt_search_reprZcol_keysZ
index_keysZindex_aliasesÚ	not_foundZdup_colsr   Úvalr   r   r   Ú_check_indices`   s\    	
ÿÿ
ÿÿ
ÿrG   )r   r   )Ú	dim_namesr   r   Úbackedc                 C   s”   t d ƒt d ƒg}| |¡}|r`t |¡}| ¡ }|| ||< t |¡||< | t|ƒ t|ƒ }	n|||< | t|ƒ }	ddlm}
 |
|	ƒr�|	 ¡ }	|	S )Nr   )Úissparse)	ÚsliceZget_indexerr@   ZargsortÚcopyÚtupleÚscipy.sparserJ   Útoarray)ÚXrH   r   r   rI   Zmutable_idxerÚidxZ	idx_orderZ	rev_idxerÚmatrixrJ   r   r   r   Ú_get_array_values¯   s    

rS   r   )Úlayerr   r<   )r   r   Ú	obsm_keysrT   r   r<   r   c                C   sŠ  |r|dkst dƒ‚| jj}n| j}|dk	r<t || ¡}nd}t| j|jd|||d�\}}	}
tj| j	d�}t
|	ƒdkr¸tt| ||d�|j|	d| jd	�}tj|tj||
| j	d
�gdd�}t
|ƒdkrÜtj|| j| gdd�}|rè|| }|D ]˜\}}|› d|› �}| j| }t|tjƒ�r6t |dd…|f ¡||< qìt|tƒ�rbt |dd…|f  ¡ ¡||< qìt|tjƒrì|jdd…|f ||< qì|S )a      Return values for observations in adata.

    Params
    ------
    adata
        AnnData object to get values from.
    keys
        Keys from either `.var_names`, `.var[gene_symbols]`, or `.obs.columns`.
    obsm_keys
        Tuple of `(key from obsm, column index of obsm[key])`.
    layer
        Layer of `adata` to use as expression values.
    gene_symbols
        Column of `adata.var` to search for `keys` in.
    use_raw
        Whether to get expression values from `adata.raw`.

    Returns
    -------
    A dataframe with `adata.obs_names` as index, and values specified by `keys`
    and `obsm_keys`.

    Examples
    --------
    Getting value for plotting:

    >>> pbmc = sc.datasets.pbmc68k_reduced()
    >>> plotdf = sc.get.obs_df(
            pbmc,
            keys=["CD8B", "n_genes"],
            obsm_keys=[("X_umap", 0), ("X_umap", 1)]
        )
    >>> plotdf.plot.scatter("X_umap0", "X_umap1", c="CD8B")

    Calculating mean expression for marker genes by cluster:

    >>> pbmc = sc.datasets.pbmc68k_reduced()
    >>> marker_genes = ['CD79A', 'MS4A1', 'CD8A', 'CD8B', 'LYZ']
    >>> genedf = sc.get.obs_df(
            pbmc,
            keys=["louvain", *marker_genes]
        )
    >>> grouped = genedf.groupby("louvain")
    >>> mean, var = grouped.mean(), grouped.var()
    Nz9Cannot specify use_raw=True and a layer at the same time.r7   )r;   r<   r%   r   )rT   r<   r   ©r   rI   ©r"   r&   ©r   ú-)rD   Úrawr/   r   ÚIndexrG   r7   r&   r   Ú	obs_namesr2   rS   Ú_get_obs_repÚisbackedr,   Úobsmr)   r@   ÚndarrayÚravelr   rO   Úloc)r   r   rU   rT   r   r<   r/   r;   Zobs_colsZvar_idx_keysZvar_symbolsÚdfrR   ÚkrQ   Úadded_krF   r   r   r   Úobs_dfÍ   sZ    7ÿþ
ú
ûþ
 rf   ©rT   )r   r   Ú	varm_keysrT   r   c                C   sD  t | j| jd|ƒ\}}}tj| jjd�}t|ƒdkrttt| |d�| j|d| j	d�j
}tj|tj||| jd�gdd�}t|ƒdkr˜tj|| j| gdd�}|r¤|| }|D ]–\}	}
|	› d	|
› �}| j|	 }t|tjƒrðt |d
d
…|
f ¡||< q¨t|tƒ�rt |d
d
…|
f  ¡ ¡||< q¨t|tjƒr¨|jd
d
…|
f ||< q¨|S )aË      Return values for observations in adata.

    Params
    ------
    adata
        AnnData object to get values from.
    keys
        Keys from either `.obs_names`, or `.var.columns`.
    varm_keys
        Tuple of `(key from varm, column index of varm[key])`.
    layer
        Layer of `adata` to use as expression values.

    Returns
    -------
    A dataframe with `adata.var_names` as index, and values specified by `keys`
    and `varm_keys`.
    r/   r%   r   rg   rV   rW   r   rX   rY   N)rG   r/   r\   r   r   r&   r2   rS   r]   r^   ÚTr,   Z	var_namesZvarmr)   r@   r`   ra   r   rO   rb   )r   r   rh   rT   Zvar_colsZobs_idx_keysÚ_rc   rR   rd   rQ   re   rF   r   r   r   Úvar_df?  s8    
ûþ
 rk   )r<   rT   r_   Úobspc          
      C   s®   t |tƒstdt|ƒ› d�ƒ‚|dk	}|dk	}|dk	}|dk	}t||||fƒ}	|	dksZt‚|	dkrh| jS |rv| j| S |r‚| jjS |r�| j	| S |rž| j
| S dsªtdƒ‚dS )z3
    Choose array aligned with obs annotation.
    z!use_raw expected to be bool, was r=   NFr   r   ú[That was unexpected. Please report this bug at:

	 https://github.com/scverse/scanpy/issues)r)   ÚboolÚ	TypeErrorÚtypeÚsumrD   rP   ÚlayersrZ   r_   rl   )
r   r<   rT   r_   rl   Úis_layerÚis_rawÚis_obsmÚis_obspÚchoices_mader   r   r   r]   €  s*    



ÿr]   c                C   sš   |dk	}|dk	}|dk	}|dk	}	t ||||	fƒ}
|
dks<t‚|
dkrL|| _nJ|r\|| j|< n:|rj|| j_n,|rz|| j|< n|	rŠ|| j|< nds–tdƒ‚dS )z(
    Set value for observation rep.
    NFr   r   rm   )rq   rD   rP   rr   rZ   r_   rl   )r   rF   r<   rT   r_   rl   rs   rt   ru   rv   rw   r   r   r   Ú_set_obs_repŸ  s&    
ÿrx   )NF)r   r   )r   r   )Ú__doc__Útypingr   r   r   r   r   Únumpyr@   Zpandasr   rN   r   Zanndatar   Z_compatr
   r*   Úfloatr   r5   r[   rn   rG   rS   Úintrf   rk   r]   rx   r   r   r   r   Ú<module>   s‚   ø÷T  úùQû   ýùøt  ýûúA