a
    d:                     @   s.  d dl mZmZ d dlmZ d dlmZmZmZm	Z	m
Z
mZmZmZ d dlZd dlmZmZ d dlmZ d dlmZ d dlmZm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ddZ$G dd deZ%eedddZ&G dd de%Z'G dd de'Z(G dd de	Z)G dd de%Z*dS )    )ABCabstractmethod)deepcopy)AnyCallableIterableIteratorListOptionalSizedTupleN)apply_to_collectionapply_to_collections)
DataLoader)has_len)CombinedLoaderCycleIterator)_add_capture_metadata_collate_patch_dataloader_get_iterators"_teardown_dataloader_get_iteratorsIteratorStateMergedIteratorStatepatch_dataloader_iterator)MisconfigurationException)_fault_tolerant_trainingreturnc                   C   s   d S N r   r   r   m/var/www/html/stable-diffusion-webui/venv/lib/python3.9/site-packages/pytorch_lightning/utilities/fetching.py_profile_nothing%   s    r    c                   @   sP  e Zd ZdZeedddZeddddZeddd	Zeedd
ddZ	ddddZ
d0eddddZeeddddZeedddZeeddddZddddZeee ddddZeedd d!Zeedd"d#Zeee dd$d%Zddd&d'Zd dd(d)Zedd*d+Zddd,d-Zddd.d/Z dS )1AbstractDataFetchera  This base class should be used to implement a fault tolerant ``DataFetcher``. It is required to override the
    ``fetching_function`` with fetching logic.

    Example::

        class SimpleDataFetcher(AbstractDataFetcher):
            def fetching_function(self):
                while True:
                    try:
                        return next(self.dataloader_iter), False
                    except StopIteration:
                        return None, True
    r   c                 C   s   dS )z&Override with your own fetching logic.Nr   selfr   r   r   fetching_function9   s    z%AbstractDataFetcher.fetching_functionNc                 C   s   dS )z*Override with your own pre-fetching logic.Nr   r"   r   r   r   prefetching=   s    zAbstractDataFetcher.prefetchingc                 C   s   dS )z=Hook to override to handle the logic before fetching a batch.Nr   r"   r   r   r   on_fetch_startA   s    z"AbstractDataFetcher.on_fetch_startbatchstart_outputr   c                 C   s   dS z>Hook to extend which handles the logic after fetching a batch.Nr   r#   r(   r)   r   r   r   on_fetch_endD   s    z AbstractDataFetcher.on_fetch_endc                 C   s   dS )zDHook to override to indicate the `DataFetcher` to wait for an event.Nr   r"   r   r   r   waitG   s    zAbstractDataFetcher.waitr   prefetch_batchesr   c                 C   s>   |dk rt d|| _d | _d | _d| _d| _t| _t| _d S )Nr   z(`prefetch_batches` should at least be 0.F)	r   r/   _dataloaderdataloader_iterfetcheddoner    _start_profiler_stop_profilerr#   r/   r   r   r   __init__J   s    zAbstractDataFetcher.__init__)
dataloaderkwargsr   c                 K   s"   |  | || _t  |   d S r   )r   r0   r   _attach_data_fetcher)r#   r8   r9   r   r   r   setupU   s    
zAbstractDataFetcher.setupc                 C   s$   | j d u rtd| jj d| j S )N`z6` should have been `setup` with a dataloader iterable.)r0   r   	__class____name__r"   r   r   r   r8   [   s
    
zAbstractDataFetcher.dataloader)r8   r   c                 C   s2   t | ttfsd S t | tr"| j} t| tt d S r   )
isinstancer   r   loadersr   r   )r8   r   r   r   r   c   s
    
z1AbstractDataFetcher._add_capture_metadata_collatec                    s0   t td d fdd}t j jtt f| d S )N)loaderiteratorr   c                    s<   t | tr| j} |j}t | tr8t r8 | _t| |  d S r   )r?   r   rA   Z_loader_iterr   r   _lightning_fetcherr   )rA   rB   r"   r   r   _apply_patch_fnn   s    
z9AbstractDataFetcher._apply_patch.<locals>._apply_patch_fn)r   r   r   r@   loader_iters)r#   rD   r   r"   r   _apply_patchm   s    
z AbstractDataFetcher._apply_patch)r1   dataloader_iter_statesr   c                 C   s   t |dd d u ri |_t |dd d u r.t |_|D ].}|j}||jvrPg |j|< |j| | q2| j| jkr|D ]>}t|jrt	|j|_
|j}|j| d}|j|| qrd S )Ncache_statesstater   )getattrrH   r   rI   nameappendr2   r/   lenr   previous_statepopupdate)r#   r1   rG   Z
iter_stateZ	iter_namerI   r   r   r   _store_dataloader_iter_statez   s     


z0AbstractDataFetcher._store_dataloader_iter_statec                 C   s   t | jtr| jjS | jS r   )r?   r8   r   r@   r"   r   r   r   r@      s    zAbstractDataFetcher.loadersc                 C   s,   | j d u rtdt| jtr&| j jS | j S )NzCThe `dataloader_iter` isn't available outside the __iter__ context.)r1   r   r?   r8   r   rE   r"   r   r   r   rE      s
    
z AbstractDataFetcher.loader_itersc                 C   s   t tddd}t| jt |S )NrB   r   c                 S   s   | j S r   )rI   )rB   r   r   r   collect_state   s    z0AbstractDataFetcher.state.<locals>.collect_state)r   r   r   rE   )r#   rS   r   r   r   rI      s    zAbstractDataFetcher.statec                    s*   t d d fdd}t jt tf| d S )N)rA   r   c                    s*   t | tr| j} t | tr&t r& | _d S r   )r?   r   rA   r   r   rC   )rA   r"   r   r   _attach_data_fetcher_fn   s    
zIAbstractDataFetcher._attach_data_fetcher.<locals>._attach_data_fetcher_fn)r   r   r@   r   )r#   rT   r   r"   r   r:      s    z(AbstractDataFetcher._attach_data_fetcherc                 C   s(   |    t| j| _|   |   | S r   )resetiterr8   r1   rF   r%   r"   r   r   r   __iter__   s
    zAbstractDataFetcher.__iter__c                 C   s   |   S r   )r$   r"   r   r   r   __next__   s    zAbstractDataFetcher.__next__c                 C   s   d| _ d| _d S )Nr   F)r2   r3   r"   r   r   r   rU      s    zAbstractDataFetcher.resetc                 C   sF   |    t| jtr| j   t| jtr6t| j d | _t  d S r   )rU   r?   r0   r   r   Z$_shutdown_workers_and_reset_iteratorr1   r   r"   r   r   r   teardown   s    
zAbstractDataFetcher.teardown)r   )!r>   
__module____qualname____doc__r   r   r$   r%   r&   r,   r-   intr7   r   r;   propertyr8   staticmethodr   rF   r   r	   r   rQ   r@   rE   r   rI   r:   rW   rX   rU   rY   r   r   r   r   r!   )   s:   	
r!   r(   r   c                 C   s   | S r   r   )r(   r   r   r   _no_op_batch_to_device   s    ra   c                       s   e Zd ZdZdeedd fddZdeee	e
ge
f  dd fd	d
Ze
dddZe
e
ddddZddddZe
dddZeddddZe
e
dddZdd fddZ  ZS )DataFetcherao  This class is used to control batch fetching flow.

    Args:
        prefetch_batches: Number of batches to pre-fetch. Pre-fetching at least 1 batch is necessary to properly track
            whether a batch is the last one (available with :attr:`self.done`) under any training setup.
        store_on_device: Whether to store the pre-fetched batches on device.
       TN)r/   store_on_devicer   c                    s*   t  j|d || _t| _g | _d| _d S )N)r/   F)superr7   rd   ra   batch_to_devicebatches_has_len)r#   r/   rd   r=   r   r   r7      s
    zDataFetcher.__init__)r8   rf   r   c                    s(   t  | t|| _|d ur$|| _d S r   )re   r;   r   rh   rf   )r#   r8   rf   ri   r   r   r;      s    
zDataFetcher.setupr   c                 C   s   |    d S r   )r4   r"   r   r   r   r&      s    zDataFetcher.on_fetch_startr'   c                 C   s   |    | j| dS r*   )r5   rg   rL   r+   r   r   r   r,      s    zDataFetcher.on_fetch_endc              	   C   sT   | j }|d usJ t| jD ]2}z| | W q tyL   d| _Y  qPY q0 qd S )NT)r1   ranger/   _fetch_next_batchStopIterationr3   )r#   rB   _r   r   r   r%      s    zDataFetcher.prefetchingc              
   C   s   | j d usJ | jrP| jd}z| | j  W q tyL   | j | _Y q0 nX| jsz| | j  | jd}W q ty } zd| _|W Y d }~qd }~0 0 nt|   | |S )Nr   T)r1   rg   rO   rk   rl   r3   r-   move_to_device)r#   r(   er   r   r   r$      s"    zDataFetcher.fetching_functionrR   c              
   C   s   |   }zt|}W n0 tyD } z|   |W Y d }~n
d }~0 0 |  jd7  _| js| jr| j}t|t	stJ | jt
|k| _| || d S )Nrc   )r&   nextrl   r5   r2   r/   rh   r8   r?   r   rM   r3   r,   )r#   rB   r)   r(   ro   r8   r   r   r   rk     s    zDataFetcher._fetch_next_batchr`   c                 C   s   | j r| |}|S r   )rd   rf   r#   r(   r   r   r   rn   $  s    
zDataFetcher.move_to_devicec                    s   t    g | _d S r   )re   rU   rg   r"   ri   r   r   rU   )  s    
zDataFetcher.reset)rc   T)N)r>   rZ   r[   r\   r]   boolr7   r   r
   r   r   r;   r&   r,   r%   r$   r   rk   rn   rU   __classcell__r   r   ri   r   rb      s   
 
rb   c                       sp   e Zd ZdZeedd fddZeed fddZd	d
ddZeej	j
ddddZdd
ddZ  ZS )InterBatchParallelDataFetchera  This class implements inter-batch parallelism, which aims at hiding the latency of host-to-device copy of
    input batches behind computationally intensive operations.

    code-block::

        Without parallelization:

        batch 0: [HtoD][forward][backward]
        batch 1:                          [HtoD][forward][backward]
        batch 2:                                                   [HtoD][forward][backward]

        With parallelization, the latency of HtoD copy can be hidden:

        batch 0: [HtoD][forward][backward]
        batch 1:       [HtoD]             [forward][backward]
        batch 2:             [HtoD]                          [forward][backward]
    N)argsr9   r   c                    s(   t  j|i | tj | _g | _d S r   )re   r7   torchcudaZStreamcuda_streamevents)r#   ru   r9   ri   r   r   r7   B  s    z&InterBatchParallelDataFetcher.__init__r`   c                    s>   t j| j t |W  d    S 1 s00    Y  d S r   )rv   rw   streamrx   re   rn   rq   ri   r   r   rn   G  s    z,InterBatchParallelDataFetcher.move_to_deviceztorch.cuda.Eventr   c                 C   s   t j }|   |S r   )rv   rw   Eventr4   r#   eventr   r   r   r&   K  s    
z,InterBatchParallelDataFetcher.on_fetch_start)r(   r}   r   c                 C   s,   |    | j| |  | j| d S r   )r5   rg   rL   recordry   )r#   r(   r}   r   r   r   r,   Q  s    z*InterBatchParallelDataFetcher.on_fetch_endc                 C   s   | j d}|  d S )Nr   )ry   rO   r-   r|   r   r   r   r-   W  s    z"InterBatchParallelDataFetcher.wait)r>   rZ   r[   r\   r   r7   rn   r&   rv   rw   r{   r,   r-   rs   r   r   ri   r   rt   .  s   rt   c                   @   s0   e Zd ZdZeeddddZedddZdS )	StepFuncDataLoaderIterztThis class is a wrapper to keep track of dataloader iterator fetching event while left entirely to user
    control.N)rB   data_fetcherr   c                 C   s   || _ || _d S r   )rB   r   )r#   rB   r   r   r   r   r7   b  s    zStepFuncDataLoaderIter.__init__r   c              
   C   sj   z4| j   t| j}| j   | j  jd7  _|W S  tyd } zd| j _|W Y d }~n
d }~0 0 d S )Nrc   T)r   r4   rp   rB   r5   r2   rl   r3   )r#   dataro   r   r   r   rX   f  s    


zStepFuncDataLoaderIter.__next__)	r>   rZ   r[   r\   r   r!   r7   r   rX   r   r   r   r   r   ]  s   r   c                       sN   e Zd ZdZdedd fddZdddd	Zeeef dd
dZ	  Z
S )DataLoaderIterDataFetchera  This class is used to return directly the `dataloader_iter` to the ``LightningModule`` training_step for
    users to implement their own pre-fetching logic. This feature can be activated as follows:

    Example::

        Class MyModel(LightningModule):

            def __init__(self):
                self.automatic_optimization = False

            def training_step(self, dataloader_iter: Iterator, batch_idx: int) -> None:
                # it is the user responsibility to fetch and move the batch to the right device.
                batch = next(dataloader_iter)
                batch = batch.to(self.device)
                ...
    r   Nr.   c                    s   t    d| _d S )NF)re   r7   rd   r6   ri   r   r   r7     s    
z"DataLoaderIterDataFetcher.__init__r   c                 C   s&   | j }|d usJ tt|| | _d S r   )r1   rV   r   rB   )r#   rB   r   r   r   r%     s    z%DataLoaderIterDataFetcher.prefetchingc                 C   s   | j s| j| jfS td S r   )r3   r2   rB   rl   r"   r   r   r   r$     s    z+DataLoaderIterDataFetcher.fetching_function)r   )r>   rZ   r[   r\   r]   r7   r%   r   r   r$   rs   r   r   ri   r   r   r  s   r   )+abcr   r   copyr   typingr   r   r   r   r	   r
   r   r   rv   Z#lightning_utilities.core.apply_funcr   r   Ztorch.utils.data.dataloaderr   Zlightning_fabric.utilities.datar   Z$pytorch_lightning.trainer.supportersr   r   Z(pytorch_lightning.utilities.auto_restartr   r   r   r   r   r   Z&pytorch_lightning.utilities.exceptionsr   Z#pytorch_lightning.utilities.importsr   r    r!   ra   rb   rt   r   r   r   r   r   r   <module>   s$   (   b/