a
    d$                     @   s   d dl Z d dlZd dlZd dlZd dlZd dlmZmZmZm	Z	m
Z
mZ d dlmZ d dlmZ 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mZ d d	lmZ erd d
lmZ dZG dd deZ eeedddZ!G dd deZ"dS )    N)AnyCallableOptionalSizedTupleUnion)	HTTPError)warn)Tensor)
DataLoaderDatasetrandom_split)_IS_WINDOWS)LightningDataModule)_TORCHVISION_AVAILABLE)
transformsz./datac                       s   e Zd ZdZdZdZdZdZd%ee	e
e	edd	 fd
dZeeeef dddZedddZeedddZee	dddZd&e	ddddZeddddZed'eeeeeef ddd Zed(eeeef eeef ed"d#d$Z  ZS ))_MNISTzCarbon copy of ``tests_pytorch.helpers.datasets.MNIST``.

    We cannot import the tests as they are not distributed with the package.
    See https://github.com/Lightning-AI/lightning/pull/7614#discussion_r671183652 for more context.
    )zChttps://pl-public-data.s3.amazonaws.com/MNIST/processed/training.ptz?https://pl-public-data.s3.amazonaws.com/MNIST/processed/test.ptztraining.ptztest.ptZcompleteTg_)Ǻ?gGr?N)roottrain	normalizedownloadkwargsreturnc                    sZ   t    || _|| _|| _| | | jr2| jn| j}| t	j
| j|\| _| _d S N)super__init__r   r   r   prepare_dataTRAIN_FILE_NAMETEST_FILE_NAME	_try_loadospathjoincached_folder_pathdatatargets)selfr   r   r   r   r   Z	data_file	__class__ q/var/www/html/stable-diffusion-webui/venv/lib/python3.9/site-packages/pytorch_lightning/demos/mnist_datamodule.pyr   5   s    

z_MNIST.__init__)idxr   c                 C   sV   | j |  d}t| j| }| jd urNt| jdkrN| j|g| jR  }||fS )Nr      )r%   floatZ	unsqueezeintr&   r   lennormalize_tensor)r'   r,   imgtargetr*   r*   r+   __getitem__B   s
    z_MNIST.__getitem__r   c                 C   s
   t | jS r   )r0   r%   r'   r*   r*   r+   __len__K   s    z_MNIST.__len__c                 C   s   t j| jd| jS )NMNIST)r!   r"   r#   r   cache_folder_namer6   r*   r*   r+   r$   N   s    z_MNIST.cached_folder_path)data_folderr   c                 C   s4   d}| j | jfD ]}|o,tjtj||}q|S )NT)r   r   r!   r"   isfiler#   )r'   r:   existingfnamer*   r*   r+   _check_existsR   s    z_MNIST._check_exists)r   r   c                 C   s4   |r|  | js| | j |  | js0tdd S )NzDataset not found.)r>   r$   	_downloadRuntimeError)r'   r   r*   r*   r+   r   X   s    z_MNIST.prepare_datac                 C   sR   t j|dd | jD ]8}td|  t j|t j|}tj	
|| qd S )NT)exist_okzDownloading )r!   makedirs	RESOURCESlogginginfor"   r#   basenameurllibrequesturlretrieve)r'   r:   urlZfpathr*   r*   r+   r?   ^   s
    
z_MNIST._download         ?)	path_datatrialsdeltar   c                 C   s   d\}}|sJ dt j| s.J d|  t|D ]R}zt| }W n: ty } z"|}t|t		   W Y d}~q6d}~0 0  qq6|dusJ |dur||S )zHResolving loading from the same time from multiple concurrent processes.)NNz!at least some trial has to be setzmissing file: N)
r!   r"   r;   rangetorchload	Exceptiontimesleeprandom)rM   rN   rO   res	exception_exr*   r*   r+   r    e   s    (z_MNIST._try_load        )tensormeanstdr   c                 C   s8   t j|| j| jd}t j|| j| jd}| ||S )N)dtypedevice)rQ   Z	as_tensorr_   r`   subdiv)r\   r]   r^   r*   r*   r+   r1   z   s    z_MNIST.normalize_tensor)Tr   T)T)rK   rL   )r[   rL   )__name__
__module____qualname____doc__rC   r   r   r9   strbooltupler   r   r/   r   r
   r4   r7   propertyr$   r>   r   r?   staticmethodr.   r    r   r1   __classcell__r*   r*   r(   r+   r   %   s*    
	 r   )argsr   r   c               
   O   s   t tdd }|rlzddlm} |tdd W n8 tyj } z td| d d}W Y d }~n
d }~0 0 |s|td	 t}|| i |S )
NZPL_USE_MOCKED_MNISTFr   )r8   T)r   zError z) downloading `torchvision.datasets.MNIST`zD`torchvision.datasets.MNIST` not available. Using our hosted version)	rh   r!   getenvZtorchvision.datasetsr8   _DATASETS_PATHr   printr   )rm   r   Ztorchvision_mnist_availabler8   er*   r*   r+   r8      s    r8   c                       s   e Zd ZdZdZedddddfeeeeeee	e	dd		 fd
dZ
eedddZddddZeddddZedddZedddZedddZeee dddZ  ZS )MNISTDataModulezStandard MNIST, train, val, test splits and transforms.

    >>> MNISTDataModule()  # doctest: +ELLIPSIS
    <...mnist_datamodule.MNISTDataModule object at ...>
    Zmnisti     F*       N)	data_dir	val_splitnum_workersr   seed
batch_sizerm   r   r   c           	         sV   t  j|i | |r.tr.td| d d}|| _|| _|| _|| _|| _|| _	dS )an  
        Args:
            data_dir: where to save/load the data
            val_split: how many of the training images to use for the validation split
            num_workers: how many workers to use for loading data
            normalize: If true applies image normalize
            seed: starting seed for RNG.
            batch_size: desired batch size.
        zYou have requested num_workers=zA on Windows, but currently recommended is 0, so we set it for your   N)
r   r   r   r	   rv   rw   rx   r   ry   rz   )	r'   rv   rw   rx   r   ry   rz   rm   r   r(   r*   r+   r      s    
zMNISTDataModule.__init__r5   c                 C   s   dS )N
   r*   r6   r*   r*   r+   num_classes   s    zMNISTDataModule.num_classesc                 C   s$   t | jddd t | jddd dS )zSaves MNIST files to `data_dir`Tr   r   FN)r8   rv   r6   r*   r*   r+   r      s    zMNISTDataModule.prepare_data)stager   c                 C   sf   | j rt| j dni }t| jfddd|}t|ts<J t|}t||| j | jg\| _	| _
dS )z"Split the train and valid dataset.Z	transformTFr}   N)default_transformsdictr8   rv   
isinstancer   r0   r   rw   dataset_traindataset_val)r'   r~   extradatasetZtrain_lengthr*   r*   r+   setup   s
    zMNISTDataModule.setupc                 C   s   t | j| jd| jddd}|S )z7MNIST train set removes a subset to use for validation.Trz   shufflerx   Z	drop_lastZ
pin_memory)r   r   rz   rx   r'   loaderr*   r*   r+   train_dataloader   s    z MNISTDataModule.train_dataloaderc                 C   s   t | j| jd| jddd}|S )z?MNIST val set uses a subset of the training set for validation.FTr   )r   r   rz   rx   r   r*   r*   r+   val_dataloader   s    zMNISTDataModule.val_dataloaderc                 C   sJ   | j rt| j dni }t| jfddd|}t|| jd| jddd}|S )z#MNIST test set uses the test split.r   Fr}   Tr   )r   r   r8   rv   r   rz   rx   )r'   r   r   r   r*   r*   r+   test_dataloader   s    zMNISTDataModule.test_dataloaderc                 C   s8   t sd S | jr,tt tjdddg}nt }|S )N)g      ?)r]   r^   )r   r   transform_libZComposeZToTensorZ	Normalize)r'   Zmnist_transformsr*   r*   r+   r      s    z"MNISTDataModule.default_transforms)rc   rd   re   rf   namero   rg   r/   rh   r   r   rj   r|   r   r   r   r   r   r   r   r   r   rl   r*   r*   r(   r+   rr      s8   $rr   )#rD   r!   rV   rT   rG   typingr   r   r   r   r   r   urllib.errorr   warningsr	   rQ   r
   Ztorch.utils.datar   r   r   Z"lightning_fabric.utilities.importsr   Zpytorch_lightningr   Z#pytorch_lightning.utilities.importsr   Ztorchvisionr   r   ro   r   r8   rr   r*   r*   r*   r+   <module>   s&    \