U
    Åmœd†N  ã                	   @   s  d dl mZ d dl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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 e d¡Zdeeeeeee f  eeee f eee edœdd„Z G dd„ deƒZ!G dd„ dƒZ"dS )é    )ÚMutableMapping)ÚIterableÚUnionÚOptionalN)Úversion)Úcheck_random_state)Úissparse)ÚAnnDataé   )Úsettings)Úlogging)Ú_rp_forest_generate)ÚNeighborsView)Úpkg_versionz0.7rc1©ÚumapÚpcaÚknnT)ÚadataÚ	adata_refÚobsÚembedding_methodÚlabeling_methodÚneighbors_keyÚinplacec                 K   s  t dƒ}|tk r,tdt› d|› dt› d�ƒ‚t d¡}	t|tƒrF|gn|}t|tƒrZ|gn|}t|tƒrn|gn|}t|ƒdkršt|p†g ƒdkrš|t|ƒ }t||ƒ}
|
 	| ¡ |D ]}|
 
|¡ q²|dk	rø|
jf |Ž t|ƒD ]\}}|
 ||| ¡ qÞtjd	|	d
� |
 |¡S )u4      Map labels and embeddings from reference data to new data.

    :tutorial:`integrating-data-using-ingest`

    Integrates embeddings and annotations of an `adata` with a reference dataset
    `adata_ref` through projecting on a PCA (or alternate
    model) that has been fitted on the reference data. The function uses a knn
    classifier for mapping labels and the UMAP package [McInnes18]_ for mapping
    the embeddings.

    .. note::

        We refer to this *asymmetric* dataset integration as *ingesting*
        annotations from reference data to new data. This is different from
        learning a joint representation that integrates both datasets in an
        unbiased way, as CCA (e.g. in Seurat) or a conditional VAE (e.g. in
        scVI) would do.

    You need to run :func:`~scanpy.pp.neighbors` on `adata_ref` before
    passing it.

    Parameters
    ----------
    adata
        The annotated data matrix of shape `n_obs` Ã— `n_vars`. Rows correspond
        to cells and columns to genes. This is the dataset without labels and
        embeddings.
    adata_ref
        The annotated data matrix of shape `n_obs` Ã— `n_vars`. Rows correspond
        to cells and columns to genes.
        Variables (`n_vars` and `var_names`) of `adata_ref` should be the same
        as in `adata`.
        This is the dataset with labels and embeddings
        which need to be mapped to `adata`.
    obs
        Labels' keys in `adata_ref.obs` which need to be mapped to `adata.obs`
        (inferred for observation of `adata`).
    embedding_method
        Embeddings in `adata_ref` which need to be mapped to `adata`.
        The only supported values are 'umap' and 'pca'.
    labeling_method
        The method to map labels in `adata_ref.obs` to `adata.obs`.
        The only supported value is 'knn'.
    neighbors_key
        If not specified, ingest looks adata_ref.uns['neighbors']
        for neighbors settings and adata_ref.obsp['distances'] for
        distances (default storage places for pp.neighbors).
        If specified, ingest looks adata_ref.uns[neighbors_key] for
        neighbors settings and
        adata_ref.obsp[adata_ref.uns[neighbors_key]['distances_key']] for distances.
    inplace
        Only works if `return_joint=False`.
        Add labels and embeddings to the passed `adata` (if `True`)
        or return a copy of `adata` with mapped embeddings and labels.

    Returns
    -------
    * if `inplace=False` returns a copy of `adata`
      with mapped embeddings and labels in `obsm` and `obs` correspondingly
    * if `inplace=True` returns `None` and updates `adata.obsm` and `adata.obs`
      with mapped embeddings and labels

    Example
    -------
    Call sequence:

    >>> import scanpy as sc
    >>> sc.pp.neighbors(adata_ref)
    >>> sc.tl.umap(adata_ref)
    >>> sc.tl.ingest(adata, adata_ref, obs='cell_type')

    .. _ingest PBMC tutorial: https://scanpy-tutorials.readthedocs.io/en/latest/integrating-pbmcs-using-ingest.html
    .. _ingest Pancreas tutorial: https://scanpy-tutorials.readthedocs.io/en/latest/integrating-pancreas-using-ingest.html
    Úanndataz*ingest only works correctly with anndata>=z (you have z) as prior to z4, `AnnData.concatenate` did not concatenate `.obsm`.zrunning ingesté   Nz    finished)Útime)r   ÚANNDATA_MIN_VERSIONÚ
ValueErrorÚloggÚinfoÚ
isinstanceÚstrÚlenÚIngestÚfitÚmap_embeddingÚ	neighborsÚ	enumerateÚ
map_labelsÚto_adata)r   r   r   r   r   r   r   ÚkwargsZanndata_versionÚstartZingÚmethodÚiÚcol© r1   úM/home/sam/Atlas/atlas_env/lib/python3.8/site-packages/scanpy/tools/_ingest.pyÚingest   s.    Vÿ
ÿÿ

r3   c                   @   sF   e Zd Zddd„Zdd„ Zdd„ Zd	d
„ Zdd„ Zdd„ Zdd„ Z	dS )Ú_DimDictr   Nc                 C   s(   i | _ || _|| _|d k	r$|  |¡ d S ©N)Ú_dataÚ_dimÚ_axisÚupdate)ÚselfÚdimÚaxisÚvalsr1   r1   r2   Ú__init__Ž   s
    z_DimDict.__init__c              
   C   sN   |j | j | jkr@td|› d|j | j › d| j› d| j› d�	ƒ‚|| j|< d S )NzValue passed for key 'z)' is of incorrect shape. Value has shape z for dimension z while it should have Ú.)Úshaper8   r7   r   r6   )r:   ÚkeyÚvaluer1   r1   r2   Ú__setitem__•   s
    (ÿz_DimDict.__setitem__c                 C   s
   | j | S r5   ©r6   ©r:   rA   r1   r1   r2   Ú__getitem__Ÿ   s    z_DimDict.__getitem__c                 C   s   | j |= d S r5   rD   rE   r1   r1   r2   Ú__delitem__¢   s    z_DimDict.__delitem__c                 C   s
   t | jƒS r5   )Úiterr6   ©r:   r1   r1   r2   Ú__iter__¥   s    z_DimDict.__iter__c                 C   s
   t | jƒS r5   )r$   r6   rI   r1   r1   r2   Ú__len__¨   s    z_DimDict.__len__c                 C   s   t | ƒj› d| j› d�S )Nú(ú))ÚtypeÚ__name__r6   rI   r1   r1   r2   Ú__repr__«   s    z_DimDict.__repr__)r   N)
rO   Ú
__module__Ú__qualname__r>   rC   rF   rG   rJ   rK   rP   r1   r1   r1   r2   r4   �   s   

r4   c                   @   sš   e Zd ZdZdd„ Zdd„ Zdd„ Zdd	„ Zd
d„ Zd)dd„Z	d*dd„Z
dd„ Zdd„ Zd+dd„Zdd„ Zdd„ Zdd„ Zd d!„ Zd,d#d$„Zd-d'd(„ZdS ).r%   uX      Class to map labels and embeddings from existing data to new data.

    You need to run :func:`~scanpy.pp.neighbors` on `adata` before
    initializing Ingest with it.

    Parameters
    ----------
    adata : :class:`~anndata.AnnData`
        The annotated data matrix of shape `n_obs` Ã— `n_vars`
        with embeddings and labels.
    c                 C   sV  dd l }| jsd|j_|j| j|jd d  dd¡d�| _| jj	| j_
| j| j_d | j_| j ¡  |jd | j_t| jƒ| j_| jjd dk | j_| j| j_| j| j_| j| j_| j�s| jd k	sÌ| jd k	rê| j| j_| j| j_| j| j_| jd k	rþ| j| j_| j| j_| j| j_n
| j| j_ |jd d d	 | j_!|jd d d
 | j_"d | j_#d S )Nr   Fr   ÚparamsÚrandom_state)ÚmetricrT   ÚX_umapi   ÚaÚb)$r   Ú_use_pynndescentZumap_Z_HAVE_PYNNDESCENTZUMAPÚ_metricÚunsÚgetÚ_umapZlearning_rateZ_initial_alphaÚ_repÚ	_raw_dataZ	knn_distsZ_validate_parametersÚobsmZ
embedding_r   Z_sparse_datar@   Z_small_dataÚ_metric_kwdsÚ_n_neighborsÚn_neighborsÚ_random_initÚ
_tree_initÚ_searchÚ
_dist_funcZ_input_distance_funcÚ
_rp_forestÚ_search_graphÚ_nnd_idxZ_knn_search_indexZ_aÚ_bZ_input_hash)r:   r   Úur1   r1   r2   Ú
_init_umap½   s<    þ











zIngest._init_umapc                    sö   ddl m} ddlm} ddlm} d | _d | _d | _d | _	d | _
|| j ‰tdƒt d¡k ržddlm}m} |ˆˆ ƒ\| _| _||| j| jd�}|ˆˆ ƒ}nHdd	lm}	 dd
lm}
 |	‡ ‡fdd„ƒ}|||d�}||
|d�}|| _
|| _|| _	d S )Nr   )Úpartial)Úinitialise_search)Únamed_distancesú
umap-learnz0.4.0)Úmake_initialisationsÚmake_initialized_nnd_search)Zinit_from_randomZinit_from_tree)Únjit)Úinitialized_nnd_searchc                    s   ˆ| |fˆ žŽ S r5   r1   )ÚxÚy©Ú	dist_argsZ	dist_funcr1   r2   Úpartial_dist_func  s    z3Ingest._init_dist_search.<locals>.partial_dist_func)Údist)Ú	functoolsrn   Zumap.nndescentro   Zumap.distancesrp   rd   re   Ú_initialise_searchrf   rg   rZ   r   r   Úparserr   rs   Znumbart   ru   )r:   ry   rn   ro   rp   rr   rs   r}   rf   rt   ru   rz   r1   rx   r2   Ú_init_dist_searchê   s<    
 ÿýzIngest._init_dist_searchc              	   C   sº   ddl m} d| _t |jd ¡d d …d f }t |t | ¡ j	¡f¡}|| j
| j| j| j|| jd�| _ddlm} t| jjƒ}|| jj| jj| jj| jj| jj|| jj| jjƒ| j_d S )Nr   )Ú	NNDescentT)ÚdatarU   Úmetric_kwdsrc   Z
init_graphrT   )Úmake_forest)Zpynndescentr€   rY   ÚnpZaranger@   ZhstackÚstackZtolilÚrowsr^   rZ   ra   rb   Ú_neigh_random_staterj   Zpynndescent.rp_treesrƒ   r   rT   r_   rc   Zn_search_treesZ	leaf_sizeÚ	rng_stateZn_jobsZ_angular_treesrh   )r:   Ú	distancesr€   Z	first_colZinit_indicesrƒ   Zcurrent_random_stater1   r1   r2   Ú_init_pynndescent  s0    ú
øzIngest._init_pynndescentc                 C   s¶  t ||ƒ}|d d | _d|d krR|d d | _| jdkrB|jn
|j| j | _nŒd|d kr’d| _|d d | _|jd d d …d | j…f | _nL|jtj	krÞd|j 
¡ krÞd| _|jd d d …d tj	…f | _| jjd | _d|d k�r
|d d | _t| j ¡ ƒ}n
i | _d	}|d d
 | _tdƒt d¡k �r’|  |¡ |d  ¡ }|jdk tj¡|_| | ¡ ¡| _d|k�rŠt|d ƒ| _nd | _n |d  dd¡| _|   |d ¡ d S )NrS   rc   Zuse_repÚXÚn_pcsÚX_pcar   r‚   r1   rU   rq   z0.5.0r‰   r   Z	rp_forestrT   )!r   rb   Ú_use_repr‹   r`   r^   Ú_n_pcsZn_varsr   ZN_PCSÚkeysr@   ra   ÚtupleÚvaluesrZ   r   r   r~   r   Úcopyr�   Úastyper„   Zint8ÚmaximumZ	transposeri   r   rh   r\   r‡   rŠ   )r:   r   r   r(   ry   Zsearch_graphr1   r1   r2   Ú_init_neighbors9  s:    
  

zIngest._init_neighborsc                 C   sr   |j d d d | _|j d d d | _| jrDd|j ¡ krDtdƒ‚| jrb|jd |jd  | _n|jd | _d S )Nr   rS   Zzero_centerZuse_highly_variableÚhighly_variablez*Did not find adata.var['highly_variable'].ÚPCs)r[   Ú_pca_centeredÚ_pca_use_hvgÚvarr�   r   ÚvarmÚ
_pca_basis©r:   r   r1   r1   r2   Ú	_init_pcab  s    zIngest._init_pcaNc                 C   s¤   |j | _d| _d | _|| _d | _d| _d|jkr:|  |¡ |d krFd}||jkr^|  	||¡ nt
d|› d�ƒ‚d|jkr‚|  |¡ d | _d | _d | _d | _d | _d S )Nr‹   Fr   r(   z*There is no neighbors data in `adata.uns["z"]`.
Please run pp.neighbors.rV   )r‹   r^   rŽ   r�   Ú
_adata_refÚ
_adata_newrY   r[   rŸ   r–   r   r`   rm   Ú_obsmÚ_obsZ_labelsÚ_indicesÚ
_distances)r:   r   r   r1   r1   r2   r>   n  s,    



ÿ

zIngest.__init__c                 C   sv   | j j}t|ƒr| ¡ n| ¡ }| jr>|d d …| jjd f }| jrT||j	dd�8 }t
 || jd d …d |…f ¡}|S )Nr—   r   ©r<   )r¡   r‹   r   Ztoarrayr“   rš   r    r›   r™   Zmeanr„   Údotr�   )r:   rŒ   r‹   r�   r1   r1   r2   Ú_pca“  s    zIngest._pcac                 C   sN   | j }| jd k	r|  | j¡S | jdkr,|jS | j|j ¡ krH|j| j S |jS )Nr‹   )r¡   r�   r¨   rŽ   r‹   r`   r�   rž   r1   r1   r2   Ú	_same_rep�  s    

zIngest._same_repc                 C   sf   | j jj ¡ }|jj ¡ }| |¡s,tdƒ‚tj|jj	d�| _
t|jdd�| _|| _|  ¡ | jd< dS )aØ          Map `adata_new` to the same representation as `adata`.

        This function identifies the representation which was used to
        calculate neighbors in 'adata' and maps `adata_new` to
        this representation.
        Variables (`n_vars` and `var_names`) of `adata_new` should be the same
        as in `adata`.

        `adata` refers to the :class:`~anndata.AnnData` object
        that is passed during the initialization of an Ingest instance.
        zNVariables in the new adata are different from variables in the reference adata)Úindexr   r¦   ÚrepN)r    Z	var_namesr#   ÚupperÚequalsr   ÚpdZ	DataFramer   rª   r£   r4   Zn_obsr¢   r¡   r©   )r:   Z	adata_newZref_var_namesZnew_var_namesr1   r1   r2   r&   §  s    
ÿz
Ingest.fité   çš™™™™™¹?r   c                 C   sö   ddl m}m} t|ƒ}| ||d¡ tj¡}| j}| j	d }	|dkrL| j
}| jrt|| j_| j |	||¡\| _| _n~ddlm}
 | j| j||	t|| ƒ|d�}|  || jj| jj||	¡}|
|ƒ\}}|dd…d|…f |dd…d|…f  | _| _dS )z´        Calculate neighbors of `adata_new` observations in `adata`.

        This function calculates `k` neighbors in `adata` for
        each observation of `adata_new`.
        r   )Ú	INT32_MAXÚ	INT32_MINé   r«   N)Údeheap_sort)rˆ   )Z
umap.umap_r±   r²   r   Úrandintr”   r„   Zint64r^   r¢   rb   rY   rj   Zsearch_rng_stateÚqueryr¤   r¥   Z
umap.utilsr´   r}   rh   Úintrf   ri   ZindptrÚindices)r:   ÚkZ
queue_sizeÚepsilonrT   r±   r²   rˆ   ÚtrainÚtestr´   ÚinitÚresultr¸   Údistsr1   r1   r2   r(   Ã  s6    
   
 ÿ    ÿzIngest.neighborsc                 C   s   | j  | jd ¡S )Nr«   )r]   Z	transformr¢   rI   r1   r1   r2   Ú_umap_transformç  s    zIngest._umap_transformc                 C   s<   |dkr|   ¡ | jd< n |dkr0|  ¡ | jd< ntdƒ‚dS )zá        Map embeddings of `adata` to `adata_new`.

        This function infers embeddings, specified by `method`,
        for `adata_new` from existing embeddings in `adata`.
        `method` can be 'umap' or 'pca'.
        r   rV   r   r�   z5Ingest supports only umap and pca embeddings for now.N)rÀ   r¢   r¨   ÚNotImplementedError)r:   r.   r1   r1   r2   r'   ê  s    ÿzIngest.map_embeddingc                    s8   | j j|  d¡‰ ‡ fdd„| jD ƒ}tj|ˆ jjd�S )NÚcategoryc                    s   g | ]}ˆ |   ¡ d  ‘qS )r   )Úmode)Ú.0Zinds©Z	cat_arrayr1   r2   Ú
<listcomp>ÿ  s     z(Ingest._knn_classify.<locals>.<listcomp>)r’   Ú
categories)r    r   r”   r¤   r®   ZCategoricalÚcatrÇ   )r:   Úlabelsr’   r1   rÅ   r2   Ú_knn_classifyû  s
    ÿzIngest._knn_classifyc                 C   s&   |dkr|   |¡| j|< ntdƒ‚dS )zÂ        Map labels of `adata` to `adata_new`.

        This function infers `labels` for `adata_new.obs`
        from existing labels in `adata.obs`.
        `method` can be only 'knn'.
        r   z%Ingest supports knn labeling for now.N)rÊ   r£   rÁ   )r:   rÉ   r.   r1   r1   r2   r*     s    zIngest.map_labelsFc                 C   sJ   |r
| j n| j  ¡ }|j | j¡ | jD ]}| j| |j|< q(|sF|S dS )aV          Returns `adata_new` with mapped embeddings and labels.

        If `inplace=False` returns a copy of `adata_new`
        with mapped embeddings and labels in `obsm` and `obs` correspondingly.
        If `inplace=True` returns nothing and updates `adata_new.obsm`
        and `adata_new.obs` with mapped embeddings and labels.
        N)r¡   r“   r`   r9   r¢   r£   r   )r:   r   r   rA   r1   r1   r2   r+     s    	
zIngest.to_adataÚbatchú-c                 C   sú   | j j| j|||d�}| j ¡ }||j| dk j|_|j |¡ | j	D ]2}|| j j
krHt | j j
| | j	| f¡|j
|< qH| jdkr¬t | j j
| j | j	d f¡|j
| j< d| j	krÈ| j jd |jd< d| j	krö| j jd |jd< | j jd	 |jd	< |S )
zõ        Returns concatenated object.

        This function returns the new :class:`~anndata.AnnData` object
        with concatenated existing embeddings and labels of 'adata'
        and inferred embeddings and labels for `adata_new`.
        )Ú	batch_keyÚbatch_categoriesÚindex_uniqueÚ1)r�   r‹   r«   rV   r   r�   r   r˜   )r    Zconcatenater¡   r£   r“   r   Z	obs_namesrª   r9   r¢   r`   r„   ZvstackrŽ   r[   rœ   )r:   rÍ   rÎ   rÏ   r   Z
obs_updaterA   r1   r1   r2   Úto_adata_joint"  s0    
ü

ÿ
ÿ

zIngest.to_adata_joint)N)N)Nr¯   r°   r   )F)rË   NrÌ   )rO   rQ   rR   Ú__doc__rm   r   rŠ   r–   rŸ   r>   r¨   r©   r&   r(   rÀ   r'   rÊ   r*   r+   rÑ   r1   r1   r1   r2   r%   ¯   s(   -/ )
%



$
     ÿr%   )Nr   r   NT)#Úcollections.abcr   Útypingr   r   r   Zpandasr®   Únumpyr„   Ú	packagingr   Zsklearn.utilsr   Zscipy.sparser   r   r	   Ú r   r   r    r(   r   Ú_utilsr   Z_compatr   r~   r   r#   Úboolr3   r4   r%   r1   r1   r1   r2   Ú<module>   s:   
     ùùy"