U
    vIÀd¹  ã                   @   sò   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m	Z	m
Z
 d dlZd dlZddlmZ d	d
lmZmZ z$d dlZd dlmZmZmZmZ W n( ek
r¼   eeeef\ZZZZY nX G dd„ deƒZdd„ Zdd„ ZG dd„ deƒZdS )é    )Úissparse)Úceil)Úcopy)Úpartial)ÚDictÚUnionÚSequenceNé   )ÚAnnDataé   )ÚAnnCollectionÚ_ConcatViewMixin)ÚSamplerÚBatchSamplerÚDatasetÚ
DataLoaderc                   @   s&   e Zd Zd	dd„Zdd„ Zdd„ ZdS )
ÚBatchIndexSamplerFc                 C   s(   || _ ||k r|n|| _|| _|| _d S ©N)Ún_obsÚ
batch_sizeÚshuffleÚ	drop_last)Úselfr   r   r   r   © r   ú`/home/sam/Atlas/atlas_env/lib/python3.8/site-packages/anndata/experimental/pytorch/_annloader.pyÚ__init__   s    zBatchIndexSampler.__init__c                 c   sx   | j rtj | j¡ ¡ }ntt| jƒƒ}td| j| jƒD ]:}||t	|| j | jƒ… }t
|ƒ| jk rl| jrlq8|V  q8d S )Nr   )r   ÚnpÚrandomZpermutationr   ÚtolistÚlistÚranger   ÚminÚlenr   )r   ÚindicesÚiÚbatchr   r   r   Ú__iter__   s    zBatchIndexSampler.__iter__c                 C   s(   | j r| j| j }nt| j| j ƒ}|S r   )r   r   r   r   )r   Úlengthr   r   r   Ú__len__-   s    zBatchIndexSampler.__len__N)FF)Ú__name__Ú
__module__Ú__qualname__r   r&   r(   r   r   r   r   r      s   
r   c                 C   s†   t | tjƒr(|r|  ¡ } q‚|r‚|  ¡ } nZ| jjdkr‚t | jtj	¡r‚t
| ƒrT|  ¡ } |rhtj| dd�} nt | ¡} |r~|  ¡ n| } | S )NÚcategoryÚcuda)Zdevice)Ú
isinstanceÚtorchZTensorr-   Ú
pin_memoryZdtypeÚnamer   Z
issubdtypeÚnumberr   ZtoarrayZtensor)ÚarrÚuse_cudar0   r   r   r   Údefault_converter7   s    


r5   c                    sz   ˆ d krˆ}nht ˆ ƒr*‡ ‡fdd„}|}nLi }|D ]B}|ˆ krHˆ||< q2t|tƒrXd }n|| }tˆ | ˆ|ƒ||< q2|S )Nc                    s   ˆˆ | ƒƒS r   r   )r3   ©ÚconvertÚtop_convertr   r   Úcompose_convertM   s    z(_convert_on_top.<locals>.compose_convert)Úcallabler.   r   Ú_convert_on_top)r7   r8   Ú
attrs_keysZnew_convertr9   ÚattrZas_ksr   r6   r   r;   H   s    

r;   c                       sD   e Zd ZdZdeee eeef f e	e
e
e
dœ‡ fdd„Z‡  ZS )	Ú	AnnLoadera³      PyTorch DataLoader for AnnData objects.

    Builds DataLoader from a sequence of AnnData objects, from an
    :class:`~anndata.experimental.AnnCollection` object or from an `AnnCollectionView` object.
    Takes care of the required conversions.

    Parameters
    ----------
    adatas
        `AnnData` objects or an `AnnCollection` object from which to load the data.
    batch_size
        How many samples per batch to load.
    shuffle
        Set to `True` to have the data reshuffled at every epoch.
    use_default_converter
        Use the default converter to convert arrays to pytorch tensors, transfer to
        the default cuda device (if `use_cuda=True`), do memory pinning (if `pin_memory=True`).
        If you pass an AnnCollection object with prespecified converters, the default converter
        won't overwrite these converters but will be applied on top of them.
    use_cuda
        Transfer pytorch tensors to the default cuda device after conversion.
        Only works if `use_default_converter=True`
    **kwargs
        Arguments for PyTorch DataLoader. If `adatas` is not an `AnnCollection` object, then also
        arguments for `AnnCollection` initialization.
    é   FT)Úadatasr   r   Úuse_default_converterr4   c                    sÜ  t |tƒr|g}t |tƒs.t |tƒs.t |tƒrª| dd¡}| dd ¡}| dd ¡}	| dd ¡}
| dd ¡}| dd ¡}| dd	¡}| d
d	¡}t||||	|
||||d�	}nt |tƒr¾t|ƒ}nt	dƒ‚|rþ| dd¡}t
t||d�}t|j|t|jg d�ƒ|_d|k}d|k}d|k�o"|d d k	}d|k�o8|d dk}|�pB|}|d k	�r¾|dk�r¾|�s¾|�s¾| dd¡}|�r�| d¡}t|||d�}ntt|ƒ|||ƒ}tƒ j|fd |dœ|—Ž ntƒ j|f||dœ|—Ž d S )NÚjoin_obsÚinnerÚ	join_obsmÚlabelÚkeysÚindex_uniquer7   Úharmonize_dtypesTÚindices_strict)rB   rD   rE   rF   rG   r7   rH   rI   z1adata should be of type AnnData or AnnCollection.r0   F)r4   r0   )ÚXÚsamplerZbatch_samplerZworker_init_fnZnum_workersr   r?   r   )r   r   )r   rK   )r   r   )r.   r
   r   ÚtupleÚdictÚpopr   r   r   Ú
ValueErrorr   r5   r;   r7   r<   r   r   r"   Úsuperr   )r   r@   r   r   rA   r4   ÚkwargsrB   rD   rE   rF   rG   r7   rH   rI   Zdatasetr0   Z
_converterZhas_samplerZhas_batch_samplerZhas_worker_init_fnZhas_workersZuse_parallelr   rK   ©Ú	__class__r   r   r   }   s�    	
ÿþý÷

  ÿ  ÿÿ
ÿþýü
  ÿ   ÿzAnnLoader.__init__)r?   FTF)r)   r*   r+   Ú__doc__r   r   r
   r   ÚstrÚintÚboolr   Ú__classcell__r   r   rR   r   r>   `   s       úúr>   )Zscipy.sparser   Úmathr   r   Ú	functoolsr   Útypingr   r   r   Únumpyr   ÚwarningsZ_core.anndatar
   Zmulti_files._anncollectionr   r   r/   Ztorch.utils.datar   r   r   r   ÚImportErrorÚobjectr   r5   r;   r>   r   r   r   r   Ú<module>   s"    