U
    Åmœds›  ã                   @   sŠ  d dl Zd dlmZ d dlmZmZmZ d dlm	Z	 d dl
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Zzd dlZW n& ek
rª   edƒ ed	ƒd‚Y nX eej d
¡d  ƒZedk rÚedƒ ed	ƒd‚zd dlZW n ek
�r   edƒ Y nX G dd„ deƒZdd„ Zd)dd„Zd*dd„Zd+dd„Z d,dd„Z!dd„ Z"d-dd„Z#d d!„ Z$d"d#„ Z%d.d%d&„Z&G d'd(„ d(ej'j(ƒZ)dS )/é    N)ÚUMAP)ÚwarnÚcatch_warningsÚfilterwarnings)ÚTypingError)Úspectral_layout)Úcheck_random_state)ÚKDTreea  The umap.parametric_umap package requires Tensorflow > 2.0 to be installed.
    You can install Tensorflow at https://www.tensorflow.org/install
    
    or you can install the CPU version of Tensorflow using 

    pip install umap-learn[parametric_umap]

    z/umap.parametric_umap requires Tensorflow >= 2.0Ú.é   aü   Global structure preservation in the umap.parametric_umap package requires 
        tensorflow_probability to be installed. You can install tensorflow_probability at
        https://www.tensorflow.org/probability, 
        
        or via

        pip install --upgrade tensorflow-probability

        Please ensure to install a version which is compatible to your tensorflow 
        installation. You can verify the correct release at 
        https://github.com/tensorflow/probability/releases.

        c                       s¨   e Zd Zdddddddejjjdd�dddddddi f‡ fd	d
„	Zd‡ fdd„	Zd‡ fdd„	Z	‡ fdd„Z
‡ fdd„Zdd„ Zdd„ Zdd„ Zdd„ Zddd„Z‡  ZS ) ÚParametricUMAPNTF)Zfrom_logitsç      ð?é
   é   r   c                    s  t ƒ jf |Ž || _|| _|| _|| _|| _|| _|	| _|| _	|
| _
|| _|| _dtjkrb|| _ntdƒ d| _|| _|| _d| _|| _|dkr¸|r¦tjj d¡| _q¾tjj d¡| _n|| _|rÔ|sÔtdƒ d| _| jdk	�r|jd jd	 | jk�rtd
 |jd jd	 | j¡ƒ‚dS )a  
        Parametric UMAP subclassing UMAP-learn, based on keras/tensorflow.
        There is also a non-parametric implementation contained within to compare
        with the base non-parametric implementation.

        Parameters
        ----------
        optimizer : tf.keras.optimizers, optional
            The tensorflow optimizer used for embedding, by default None
        batch_size : int, optional
            size of batch used for batch training, by default None
        dims :  tuple, optional
            dimensionality of data, if not flat (e.g. (32x32x3 images for ConvNet), by default None
        encoder : tf.keras.Sequential, optional
            The encoder Keras network
        decoder : tf.keras.Sequential, optional
            the decoder Keras network
        parametric_embedding : bool, optional
            Whether the embedder is parametric or non-parametric, by default True
        parametric_reconstruction : bool, optional
            Whether the decoder is parametric or non-parametric, by default False
        parametric_reconstruction_loss_fcn : bool, optional
            What loss function to use for parametric reconstruction, by default tf.keras.losses.BinaryCrossentropy
        parametric_reconstruction_loss_weight : float, optional
            How to weight the parametric reconstruction loss relative to umap loss, by default 1.0
        autoencoder_loss : bool, optional
            [description], by default False
        reconstruction_validation : array, optional
            validation X data for reconstruction loss, by default None
        loss_report_frequency : int, optional
            how many times per epoch to report loss, by default 1
        n_training_epochs : int, optional
            number of epochs to train for, by default 1
        global_correlation_loss_weight : float, optional
            Whether to additionally train on correlation of global pairwise relationships (>0), by default 0
        run_eagerly : bool, optional
            Whether to run tensorflow eagerly
        keras_fit_kwargs : dict, optional
            additional arguments for model.fit (like callbacks), by default {}
        Útensorflow_probabilityz˜tensorflow_probability not installed or incompatible to current                 tensorflow installation. Setting global_correlation_loss_weight to zero.r   Nçü©ñÒMbP?gš™™™™™¹?zpParametric decoding is not implemented with nonparametric             embedding. Turning off parametric decodingFéÿÿÿÿzNDimensionality of embedder network output ({}) doesnot match n_components ({}))ÚsuperÚ__init__ÚdimsÚencoderÚdecoderÚparametric_embeddingÚparametric_reconstructionÚ"parametric_reconstruction_loss_fcnÚ%parametric_reconstruction_loss_weightÚrun_eagerlyÚautoencoder_lossÚ
batch_sizeÚloss_report_frequencyÚsysÚmodulesÚglobal_correlation_loss_weightr   Úreconstruction_validationÚkeras_fit_kwargsÚparametric_modelÚn_training_epochsÚtfÚkerasÚ
optimizersZAdamÚ	optimizerÚoutputsÚshapeÚn_componentsÚ
ValueErrorÚformat)Úselfr*   r   r   r   r   r   r   r   r   r   r#   r   r&   r"   r   r$   Úkwargs©Ú	__class__© úM/home/sam/Atlas/atlas_env/lib/python3.8/site-packages/umap/parametric_umap.pyr   ?   sX    >ÿÿÿ
ÿÿÿ þþzParametricUMAP.__init__c                    s@   | j dkr.|d krtdƒ‚|| _tƒ  ||¡S tƒ  ||¡S d S ©NÚprecomputedzTPrecomputed distances must be supplied if metric                     is precomputed.)Úmetricr.   Ú_Xr   Úfit©r0   ÚXÚyZprecomputed_distancesr2   r4   r5   r:   ¾   s    
ÿzParametricUMAP.fitc                    s@   | j dkr.|d krtdƒ‚|| _tƒ  ||¡S tƒ  ||¡S d S r6   )r8   r.   r9   r   Úfit_transformr;   r2   r4   r5   r>   Ì   s    
ÿzParametricUMAP.fit_transformc                    s:   | j r"| jjt |¡| j| jd�S tdƒ tƒ  	|¡S dS )aw  Transform X into the existing embedded space and return that
        transformed output.
        Parameters
        ----------
        X : array, shape (n_samples, n_features)
            New data to be transformed.
        Returns
        -------
        X_new : array, shape (n_samples, n_components)
            Embedding of the new data in low-dimensional space.
        ©r   Úverbosez_Embedding new data is not supported by ParametricUMAP.                 Using original embedder.N)
r   r   ÚpredictÚnpÚ
asanyarrayr   r@   r   r   Ú	transform©r0   r<   r2   r4   r5   rD   Û   s      ÿÿzParametricUMAP.transformc                    s2   | j r"| jjt |¡| j| jd�S tƒ  |¡S dS )aš   "Transform X in the existing embedded space back into the input
        data space and return that transformed output.
        Parameters
        ----------
        X : array, shape (n_samples, n_components)
            New points to be inverse transformed.
        Returns
        -------
        X_new : array, shape (n_samples, n_features)
            Generated data points new data in data space.
        r?   N)	r   r   rA   rB   rC   r   r@   r   Úinverse_transformrE   r2   r4   r5   rF   ò   s      ÿz ParametricUMAP.inverse_transformc           
      C   sŽ  i }| j rštjjj| jdd�}tjjj| jdd�}||g}|  |¡}|  |¡}| jr˜| jrf|  	|¡}n|  	t 
|¡¡}tjjjdd„ dd�|ƒ}||d< n„tjjjdtjd	d
�}t t | j|d ¡¡}t t | j|d ¡¡}|  |¡dd…ddd…f }|  |¡dd…ddd…f }|g}tj||gdd�}	tjjjdd„ dd�|	ƒ}	|	|d< | jdk�r|tjjjdd„ dd�|ƒ|d< t||d�| _dS )zDefine the model in kerasÚto_x)r,   ÚnameÚfrom_xc                 S   s   | S ©Nr4   ©Úxr4   r4   r5   Ú<lambda>  ó    z.ParametricUMAP._define_model.<locals>.<lambda>Úreconstruction)rH   r   Úbatch_sample)r,   ÚdtyperH   r   Nr   ©Úaxisc                 S   s   | S rJ   r4   rK   r4   r4   r5   rM   5  rN   Úumapc                 S   s   | S rJ   r4   rK   r4   r4   r5   rM   <  rN   Úglobal_correlation)Úinputsr+   )r   r'   r(   ÚlayersZInputr   r   r   r   r   Zstop_gradientÚLambdaÚint32ÚsqueezeÚgatherÚheadÚtailÚconcatr"   ÚGradientClippedModelr%   )
r0   r+   rG   rI   rV   Úembedding_toÚembedding_fromZembedding_to_reconrP   Zembedding_to_fromr4   r4   r5   Ú_define_model  sR    

 ÿþ
  ÿÿ ÿþzParametricUMAP._define_modelc                 C   s    i }i }t | j| j| j| j| j| jƒ}||d< d|d< | jdkrjt|d< | j|d< | j	dkrjt
dƒ d| _	| jr„| j|d< | j|d< | jj| j||| j	d	� d
S )z2
        Compiles keras model with losses
        rT   r   r   rU   Fz>Setting tensorflow to run eagerly for global_correlation_loss.TrO   )r*   ÚlossÚloss_weightsr   N)Ú	umap_lossr   Únegative_sample_rateÚ_aÚ_bÚedge_weightr   r"   Údistance_loss_corrr   r   r   r   r   r%   Úcompiler*   )r0   Úlossesrd   Úumap_loss_fnr4   r4   r5   Ú_compile_modelD  s6    ú




üzParametricUMAP._compile_modelc              	   C   sF  | j dkr| j}| jd kr.t |¡d g| _n*t| jƒdkrXt |t|ƒgt| jƒ ¡}| jr‚t 	|¡dkszt 
|¡dk r‚tdƒ t|| j| j| j| j| j| jƒ\}| _}}}| _t t | tj¡d¡¡| _t t | tj¡d¡¡| _| jröd }	n t|| j| j| j| j | jdd	�}	t|ƒ}
t| j| j | j| j|
| j| j|	ƒ\| _| _ |  !¡  |  "¡  | j�rvt#|| j | j$ ƒ}nd
}| j�rÞ| j%d k	�rÞt| jƒdk�rÀt | j%t| j%ƒgt| jƒ ¡| _%| j%t &| j%¡fd| j%if}nd }| j'j(|f| j$| j) |d
|dœ| j*—Ž}|j+| _,| j�r.| jj-|| j.d�}n| jj/d  0¡ }|i fS )Nr7   r   r   r   ç        zMData should be scaled to the range 0-1 for cross-entropy reconstruction loss.r   Úspectral)Úinitéd   rO   )ZepochsÚsteps_per_epochZmax_queue_sizeÚvalidation_data)r@   )1r8   r9   r   rB   r,   ÚlenZreshapeÚlistr   ÚmaxÚminr   Úconstruct_edge_datasetÚgraph_Ún_epochsr   r   r"   ri   r'   ZconstantÚexpand_dimsÚastypeÚint64r\   r]   Úinit_embedding_from_graphr-   Úrandom_stateÚ_metric_kwdsÚprepare_networksr   r   rb   rn   Úintr   r#   Ú
zeros_liker%   r:   r&   r$   ÚhistoryZ_historyrA   r@   Útrainable_variablesÚnumpy)r0   r<   r{   rq   r€   Úedge_datasetZn_edgesr\   r]   Úinit_embeddingÚn_datars   rt   r…   Ú	embeddingr4   r4   r5   Ú_fit_embed_datai  s®    

"ÿùù
ùøÿÿþþ
þûÿ
ûú	zParametricUMAP._fit_embed_datac                 C   s   t dd„ | j ¡ D ƒƒS )Nc                 s   s,   | ]$\}}t ||ƒr|d kr||fV  qdS ))r*   r   r   r%   N)Úshould_pickle)Ú.0ÚkÚvr4   r4   r5   Ú	<genexpr>ã  s   
 þz.ParametricUMAP.__getstate__.<locals>.<genexpr>)ÚdictÚ__dict__Úitems)r0   r4   r4   r5   Ú__getstate__á  s    þzParametricUMAP.__getstate__c              
   C   s  | j d k	r6tj |d¡}| j  |¡ |r6td |¡ƒ | jd k	rltj |d¡}| j |¡ |rltd |¡ƒ | jd k	r¢tj |d¡}| j |¡ |r¢td |¡ƒ t	ƒ �b t
dƒ | j ¡ | _tj |d¡}t|d	ƒ�}t | |tj¡ W 5 Q R X |�rtd
 |¡ƒ W 5 Q R X d S )Nr   zKeras encoder model saved to {}r   zKeras decoder model saved to {}r%   zKeras full model saved to {}Úignoreú	model.pklÚwbz*Pickle of ParametricUMAP model saved to {})r   ÚosÚpathÚjoinÚsaveÚprintr/   r   r%   r   r   r*   Z
get_configÚ_optimizer_dictÚopenÚpickleÚdumpÚHIGHEST_PROTOCOL)r0   Úsave_locationr@   Úencoder_outputÚdecoder_outputÚparametric_model_outputÚmodel_outputÚoutputr4   r4   r5   rœ   é  s.    


zParametricUMAP.save)NN)NN)T)Ú__name__Ú
__module__Ú__qualname__r'   r(   rl   ZBinaryCrossentropyr   r:   r>   rD   rF   rb   rn   rŒ   r•   rœ   Ú__classcell__r4   r4   r2   r5   r   >   s8   ÿí?%xr   c                 C   sŒ   |   ¡ }| ¡  |jd }|dkr:|jd dkr6d}nd}d|j|j|j ¡ t|ƒ k < | ¡  ||j }|j}|j}|j}||||||fS )a=  
    gets elements of graphs, weights, and number of epochs per edge

    Parameters
    ----------
    graph_ : scipy.sparse.csr.csr_matrix
        umap graph of probabilities
    n_epochs : int
        maximum number of epochs per edge

    Returns
    -------
    graph scipy.sparse.csr.csr_matrix
        umap graph
    epochs_per_sample np.array
        number of epochs to train each sample for
    head np.array
        edge head
    tail np.array
        edge tail
    weight np.array
        edge weight
    n_vertices int
        number of verticies in graph
    r   Nr   é'  iô  éÈ   ro   )	ZtocooZsum_duplicatesr,   Údatarw   ÚfloatZeliminate_zerosÚrowÚcol)rz   r{   ÚgraphÚ
n_verticesÚepochs_per_sampler\   r]   Úweightr4   r4   r5   Úget_graph_elements  s    

r·   rp   c                 C   sD  |dkrt dƒ}t|tƒrF|dkrF|jdd|jd |fd� tj¡}nút|tƒr°|dkr°t| |||||d�}dt 	|¡ 
¡  }	||	  tj¡|jd	|jd |gd
� tj¡ }n�t |¡}
t|
jƒdk�r@tj|
dd�jd |
jd k �r<t|
ƒ}|j|
dd�\}}t |dd…df ¡}|
|jd| |
jd
� tj¡ }n|
}|S )a*  Initialize embedding using graph. This is for direct embeddings.

    Parameters
    ----------
    init : str, optional
        Type of initialization to use. Either random, or spectral, by default "spectral"

    Returns
    -------
    embedding : np.array
        the initialized embedding
    NÚrandomg      $Àg      $@r   )ÚlowÚhighÚsizerp   )r8   Zmetric_kwdsç-Cëâ6?)Úscaler»   r   rR   )r�   r   r   )r   Ú
isinstanceÚstrÚuniformr,   r}   rB   Úfloat32r   Úabsrw   ÚnormalÚarrayru   Úuniquer	   ÚqueryZmean)Z	_raw_datar³   r-   r€   r8   r�   rq   r‹   ZinitialisationZ	expansionZ	init_dataÚtreeÚdistÚindZnndistr4   r4   r5   r   B  sX      ÿþúÿ ÿýþ	
  ÿþr   r   c                 C   s   dd|| d|     S )a³  
     convert distance representation into probability,
        as a function of a, b params

    Parameters
    ----------
    distances : array
        euclidean distance between two points in embedding
    a : float, optional
        parameter based on min_dist, by default 1.0
    b : float, optional
        parameter based on min_dist, by default 1.0

    Returns
    -------
    float
        probability in embedding space
    r   r   r4   )Z	distancesÚaÚbr4   r4   r5   Úconvert_distance_to_probability|  s    rÌ   r¼   c                 C   sV   |  t j t  ||d¡¡ }d|   t j t  d| |d¡¡ | }|| }|||fS )a¼  
    Compute cross entropy between low and high probability

    Parameters
    ----------
    probabilities_graph : array
        high dimensional probabilities
    probabilities_distance : array
        low dimensional probabilities
    EPS : float, optional
        offset to to ensure log is taken of a positive number, by default 1e-4
    repulsion_strength : float, optional
        strength of repulsion between negative samples, by default 1.0

    Returns
    -------
    attraction_term: tf.float32
        attraction term for cross entropy loss
    repellant_term: tf.float32
        repellant term for cross entropy loss
    cross_entropy: tf.float32
        cross entropy umap loss

    r   )r'   ÚmathÚlogÚclip_by_value)Úprobabilities_graphÚprobabilities_distanceZEPSÚrepulsion_strengthZattraction_termZrepellant_termZCEr4   r4   r5   Úcompute_cross_entropy’  s    
ÿÿþÿrÓ   c                    s6   ˆst  |ˆd ¡‰tj‡ ‡‡‡‡‡‡fdd„ƒ}|S )a  
    Generate a keras-ccompatible loss function for UMAP loss

    Parameters
    ----------
    batch_size : int
        size of mini-batches
    negative_sample_rate : int
        number of negative samples per positive samples to train on
    _a : float
        distance parameter in embedding space
    _b : float float
        distance parameter in embedding space
    edge_weights : array
        weights of all edges from sparse UMAP graph
    parametric_embedding : bool
        whether the embeddding is parametric or nonparametric
    repulsion_strength : float, optional
        strength of repulsion vs attraction for cross-entropy, by default 1.0

    Returns
    -------
    loss : function
        loss function that takes in a placeholder (0) and the output of the keras network
    r   c              
      sÞ   t j|ddd�\}}t j|ˆdd�}t j|ˆdd�}t  |t j t  t  |¡d ¡¡¡}t jt j	|| dd�t j	|| dd�gdd�}t
|ˆ ˆƒ}t jt  ˆ¡t  ˆˆ ¡gdd�}	t|	|ˆd�\}
}}ˆsÔ|ˆ }t  |¡S )Nr   r   )Znum_or_size_splitsrS   r   rR   )rÒ   )r'   ÚsplitÚrepeatr[   r¸   ÚshuffleÚranger,   r^   ZnormrÌ   ZonesÚzerosrÓ   Úreduce_mean)Zplaceholder_yZembed_to_fromr`   ra   Zembedding_neg_toZ
repeat_negZembedding_neg_fromZdistance_embeddingrÑ   rÐ   Zattraction_lossZrepellant_lossZce_loss©rg   rh   r   rf   r   rÒ   Zweights_tiledr4   r5   rc   ã  sD      ÿ
 ÿþû	  ÿ ÿýzumap_loss.<locals>.loss)rB   Ztiler'   Úfunction)r   rf   rg   rh   Zedge_weightsr   rÒ   rc   r4   rÚ   r5   re   ¼  s
    #,re   c                 C   sò   t jj ¡ | ƒ} t jj ¡ |ƒ}dd„ }|| ƒ} ||ƒ}t  | dd¡} t  |dd¡}t jj| dd… | dd…  dd�}t jj|dd… |dd…  dd�}|t j |j	¡d	  }t  
tjjt  |d¡t  |d¡d
�¡}t j |¡rìtdƒ‚| S )z6Loss based on the distance between elements in a batchc                 S   s   | t  | ¡ t j | ¡ S rJ   )r'   rÙ   rÍ   Z
reduce_stdrK   r4   r4   r5   Úz_score  s    z#distance_loss_corr.<locals>.z_scoreiöÿÿÿr   r   Nr   rR   g»½×Ùß|Û=)rL   r=   z%NaN values found in correlation loss.)r'   r(   rW   ÚFlattenrÏ   rÍ   Zreduce_euclidean_normr¸   rÀ   r,   rZ   r   ÚstatsZcorrelationr|   Úis_nanr.   )rL   Zz_xrÜ   ZdxZdzZcorr_dr4   r4   r5   rj     s&    $$
 
ÿÿrj   c           	      C   s2  |rr| dkr¬t j t jjj|d�t jj ¡ t jjjddd�t jjjddd�t jjjddd�t jjj|dd�g¡} n:t jjj||dd	�}|jd
d� | 	|g¡ t j |g¡} |dk�r*|�r*t j t jjj|d�t jjjddd�t jjjddd�t jjjddd�t jjjt
 |¡ddd�t jj |¡g¡}| |fS )aÉ  
    Generates a set of keras networks for the encoder and decoder if one has not already
    been predefined.

    Parameters
    ----------
    encoder : tf.keras.Sequential
        The encoder Keras network
    decoder : tf.keras.Sequential
        the decoder Keras network
    n_components : int
        the dimensionality of the latent space
    dims : tuple of shape (dim1, dim2, dim3...)
        dimensionality of data
    n_data : number of elements in dataset
        # of elements in training dataset
    parametric_embedding : bool
        Whether the embedder is parametric or non-parametric
    parametric_reconstruction : bool
        Whether the decoder is parametric or non-parametric
    init_embedding : array (optional, default None)
        The initial embedding, for nonparametric embeddings

    Returns
    -------
    encoder: tf.keras.Sequential
        encoder keras network
    decoder: tf.keras.Sequential
        decoder keras network
    N)Zinput_shaperr   Zrelu)ÚunitsÚ
activationÚz)rà   rH   r   )Zinput_length©r   Zrecon)rà   rH   rá   )r'   r(   Z
SequentialrW   Z
InputLayerrÝ   ZDenseZ	EmbeddingÚbuildZset_weightsrB   ÚproductZReshape)	r   r   r-   r   rŠ   r   r   r‰   Zembedding_layerr4   r4   r5   r‚   7  sF    )
úÿ  ÿ
  ÿøÿr‚   c                    s�  ‡ fdd„‰ˆ j d dkrdnd‰‡ ‡‡fdd„}‡‡‡fd	d
„}dd„ }	t||ƒ\}
}}}}}ˆdkr„|r|t |dg¡‰nt|ƒ‰t || d¡¡t || d¡¡ }}tj t	t|ƒƒ¡}||  tj
¡}||  tj
¡}|�rJtjj ||f¡}| ¡ }| d¡}|jˆdd�}|j|tjjjd�}|j|tjjjd�}| d¡}n2|	ƒ }tjjj|tjtjft d¡t d¡fd�}|ˆt|ƒ|||fS )a  
    Construct a tf.data.Dataset of edges, sampled by edge weight.

    Parameters
    ----------
    X : array, shape (n_samples, n_features)
        New data to be transformed.
    graph_ : scipy.sparse.csr.csr_matrix
        Generated UMAP graph
    n_epochs : int
        # of epochs to train each edge
    batch_size : int
        batch size
    parametric_embedding : bool
        Whether the embedder is parametric or non-parametric
    parametric_reconstruction : bool
        Whether the decoder is parametric or non-parametric
    c                    s   ˆ |  S rJ   r4   )Úindex)r<   r4   r5   Úgather_index¢  s    z,construct_edge_dataset.<locals>.gather_indexg•Ö&è.>g      à?TFc                    sV   ˆr6t  ˆ| gt jg¡d }t  ˆ|gt jg¡d }nt  ˆ | ¡}t  ˆ |¡}||fS )Nr   )r'   Zpy_functionrÁ   r[   )Zedge_toZ	edge_fromÚedge_to_batchÚedge_from_batch)r<   rç   Úgather_indices_in_pythonr4   r5   Úgather_X©  s    z(construct_edge_dataset.<locals>.gather_Xc                    s8   dt  dˆ ¡i}ˆdkr | |d< ˆr,| |d< | |f|fS )NrT   r   rU   rO   )r'   rÕ   )rè   ré   r+   )r   r"   r   r4   r5   Úget_outputs³  s    z+construct_edge_dataset.<locals>.get_outputsc                  S   s   dd„ } | S )zº
        The sham generator is a placeholder when all data is already intrinsic to
        the model, but keras wants some input data. Used for non-parametric
        embedding.
        c                   s   s(   t jdt jd�t jdt jd�fV  q d S )Nr   )rQ   )r'   rØ   rY   r4   r4   r4   r5   Úsham_generatorÄ  s    zKconstruct_edge_dataset.<locals>.make_sham_generator.<locals>.sham_generatorr4   )rí   r4   r4   r5   Úmake_sham_generator½  s    z3construct_edge_dataset.<locals>.make_sham_generatorNiè  rƒ   r­   )Zdrop_remainder)Znum_parallel_callsr   r   rã   )Zoutput_shapes)Únbytesr·   rB   rx   ru   rÕ   r}   r¸   Zpermutationr×   r~   r'   r¯   ZDatasetZfrom_tensor_slicesrÖ   ÚbatchÚmapZexperimentalZAUTOTUNEZprefetchZfrom_generatorrY   ZTensorShape)r<   rz   r{   r   r   r   r"   rë   rì   rî   r³   rµ   r\   r]   r¶   r´   Zedges_to_expZedges_from_expZshuffle_maskrˆ   Úgenr4   )r<   r   rç   rê   r"   r   r5   ry   †  sT    

 ÿþÿ
 ÿ ÿ
ýry   c                 C   sÌ   z0t  t |¡d¡ ¡ }t t  | ¡ d¡¡}W n– tjtjj	t
tjjtjjtttfk
r† } ztd | |¡ƒ W Y ¢dS d}~X Y nB tk
rÆ } z$td| › d|› d|› �ƒ W Y ¢dS d}~X Y nX dS )	a  
    Checks if a dictionary item can be pickled

    Parameters
    ----------
    key : try
        key for dictionary element
    val : None
        element of dictionary

    Returns
    -------
    picklable: bool
        whether the dictionary item can be pickled
    Úbase64zDid not pickle {}: {}FNzFailed at pickling ú:z due to T)ÚcodecsÚencoder    ÚdumpsÚdecodeÚloadsÚPicklingErrorr'   ÚerrorsZInvalidArgumentErrorÚ	TypeErrorZInternalErrorZNotFoundErrorÚOverflowErrorr   ÚAttributeErrorr   r/   r.   )ÚkeyÚvalZpickledZ	unpickledÚer4   r4   r5   r�   û  s&    ø
r�   Tc           
      C   s.  t j | d¡}t t|dƒ¡}|r0td |¡ƒ |jd }t	t
jj|ƒ}| |j¡|_t j | d¡}t j |¡r’t
jj |¡|_|r’td |¡ƒ t j | d¡}t j |¡rÊt
jj |¡|_td |¡ƒ t|j|j|j|j|j|jƒ}t j | d	¡}	t j |	¡�r*t
jjj|	d
|id�|_td |	¡ƒ |S )aŒ  
    Load a parametric UMAP model consisting of a umap-learn UMAP object
    and corresponding keras models.

    Parameters
    ----------
    save_location : str
        the folder that the model was saved in
    verbose : bool, optional
        Whether to print the loading steps, by default True

    Returns
    -------
    parametric_umap.ParametricUMAP
        Parametric UMAP objects
    r—   Úrbz-Pickle of ParametricUMAP model loaded from {}rH   r   z"Keras encoder model loaded from {}r   z"Keras decoder model loaded from {}r%   rc   )Zcustom_objectszKeras full model loaded from {})r™   rš   r›   r    ÚloadrŸ   r�   r/   rž   Úgetattrr'   r(   r)   Úfrom_configr*   ÚexistsÚmodelsZ
load_modelr   r   re   r   rf   rg   rh   ri   r   r%   )
r£   r@   r§   ÚmodelÚ
class_nameZOptimizerClassr¤   r¥   rm   r¦   r4   r4   r5   Úload_ParametricUMAP#  s@    
ú
 ÿr
  c                   @   s   e Zd ZdZdd„ ZdS )r_   zg
    We need to define a custom keras model here for gradient clipping,
    to stabilize training.
    c           	   	   C   s˜   |\}}t  ¡ �$}| |dd�}| j||| jd�}W 5 Q R X | j}| ||¡}dd„ |D ƒ}dd„ |D ƒ}| j t||ƒ¡ | j	 
||¡ dd„ | jD ƒS )	NT)Ztraining)Zregularization_lossesc                 S   s   g | ]}t  |d d¡‘qS )g      Àg      @)r'   rÏ   ©rŽ   Zgradr4   r4   r5   Ú
<listcomp>w  s     z3GradientClippedModel.train_step.<locals>.<listcomp>c                 S   s(   g | ] }t  t j |¡t  |¡|¡‘qS r4   )r'   ÚwhererÍ   rß   r„   r  r4   r4   r5   r  x  s   ÿc                 S   s   i | ]}|j | ¡ “qS r4   )rH   Úresult)rŽ   Úmr4   r4   r5   Ú
<dictcomp>‚  s      z3GradientClippedModel.train_step.<locals>.<dictcomp>)r'   ZGradientTapeZcompiled_lossrl   r†   Zgradientr*   Zapply_gradientsÚzipZcompiled_metricsZupdate_stateZmetrics)	r0   r¯   rL   r=   ZtapeZy_predrc   Ztrainable_varsZ	gradientsr4   r4   r5   Ú
train_stepi  s    
þzGradientClippedModel.train_stepN)r©   rª   r«   Ú__doc__r  r4   r4   r4   r5   r_   c  s   r_   )rp   )r   r   )r¼   r   )r   )N)T)*r‡   rB   rT   r   Úwarningsr   r   r   Znumbar   r™   Zumap.spectralr   Zsklearn.utilsr   rõ   r    Zsklearn.neighborsr	   r    Z
tensorflowr'   ÚImportErrorrƒ   Ú__version__rÔ   ZTF_MAJOR_VERSIONr   r   r·   r   rÌ   rÓ   re   rj   r‚   ry   r�   r
  r(   ZModelr_   r4   r4   r4   r5   Ú<module>   s`   ÿ
ÿ

ÿ
   Q7 ÿ
:
   ÿ
1 ù
W, ø
Ou(
@