a
    dJ                     @   s@  d dl 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	m
Z
mZmZmZmZmZmZmZ d dlmZ d dlmZmZmZmZmZ d dlmZ d d	lmZ d d
lmZ d dl m!Z! G dd deZ"ee#dddZ$eeeef e#dddZ%eeeef edddZ&d1eeeef e#eee	 ee'e	f f dddZ(d2eeeef e#ee'e	f dddZ)ee*ddddZ+dde	e	ee e	e	d d!d"Z,d3e
ee' e
d#d$d%Z-e
e"e
d&d'd(Z.ed4eee' ed) d*d+d,Z/e'e	ee	d-f ee'e	f ee'e	f ee'd-f ee#ee	d-f ee'e	f f d.d/d0Z0dS )5    N)OrderedDict)contextmanager)partial)	AnyCallableDict	GeneratorIterableOptionalTupleTypeUnion)get_all_subclasses)BatchSampler
DataLoaderDatasetIterableDatasetSampler)LightningEnum)MisconfigurationException)rank_zero_warn)pl_worker_init_functionc                   @   s$   e Zd ZdZdZeddddZdS )_WrapAttrTagsetdelN)argsreturnc                 G   s   | | j krt}nt}|| S N)SETsetattrdelattr)selfr   fn r#   h/var/www/html/stable-diffusion-webui/venv/lib/python3.9/site-packages/lightning_fabric/utilities/data.py__call__$   s    
z_WrapAttrTag.__call__)__name__
__module____qualname__r   DELr   r%   r#   r#   r#   r$   r       s   r   )
dataloaderr   c                 C   s   t | dot| jtS )Ndataset)hasattr
isinstancer+   r   )r*   r#   r#   r$   has_iterable_dataset-   s    r.   c              	   C   sf   z(t | dkr"td| jj d d}W n ttfyB   d}Y n0 |rbt| trbt| rbtd |S )z}Checks if a given Dataloader has ``__len__`` method implemented i.e. if it is a finite dataloader or
    infinite dataloader.r   `z>` returned 0 length. Please make sure this was your intention.TFzYour `IterableDataset` has `__len__` defined. In combination with multi-process data loading (when num_workers > 1), `__len__` could be inaccurate if each worker is not configured independently to avoid having duplicate data.)	lenr   	__class__r&   	TypeErrorNotImplementedErrorr-   r   r.   )r*   has_lenr#   r#   r$   r4   1   s    
r4   )r*   samplerr   c                 C   s(   t | |\}}t| g|R i |} | S r   )$_get_dataloader_init_args_and_kwargs_reinstantiate_wrapped_cls)r*   r5   dl_args	dl_kwargsr#   r#   r$   _update_dataloaderH   s    r:   F)r*   r5   disallow_batch_samplerr   c              
      s  t | tstd|  dt| d}|rB| j}| j| j | j}n(dd t| 	 D d }| j
d< d tt| jj}tdd	 | D }|r|r|d
d ttjj	 D  n |ttjj |dd  |sfdd|	 D d fdd	 D d}d|}t |trHd d< d d< nt| ||  fdd| D }	|	rt|	}
| jj}ddd	 |
D }td| d|
 d| d| d	|stt B |  }|rt|}| jj}td| d| d| d|fS )NzThe dataloader z0 needs to subclass `torch.utils.data.DataLoader`__pl_saved_argsc                 S   s    i | ]\}}| d s||qS )_)
startswith.0kvr#   r#   r$   
<dictcomp>^       z8_get_dataloader_init_args_and_kwargs.<locals>.<dictcomp>multiprocessing_contextr#   c                 s   s   | ]}|j |ju V  qd S r   )kindVAR_KEYWORDr@   pr#   r#   r$   	<genexpr>h   rD   z7_get_dataloader_init_args_and_kwargs.<locals>.<genexpr>c                 S   s"   i | ]\}}|j |jur||qS r#   )defaultemptyr?   r#   r#   r$   rC   p   rD   r!   c                    s*   h | ]"\}}| v r|j  | ur|qS r#   )rK   )r@   namerI   )attrsr#   r$   	<setcomp>x   rD   z7_get_dataloader_init_args_and_kwargs.<locals>.<setcomp>r+   c                    s   i | ]\}}| v r||qS r#   r#   r?   )non_defaultsr#   r$   rC   }   rD   batch_samplerr5   c                    sD   h | ]<}|j |j|jfv r|j|ju r|jvr|j vr|jqS r#   )rF   POSITIONAL_ONLYPOSITIONAL_OR_KEYWORDrK   rL   rM   rH   )	arg_namesr9   r#   r$   rO      s   

z, c                 s   s   | ]}d | dV  qdS )z`self.r/   Nr#   )r@   Zarg_namer#   r#   r$   rJ      rD   z,Trying to inject custom `Sampler` into the `z` instance. This would fail as some of the `__init__` arguments are not available as instance attributes. The missing attributes are z. If you instantiate your `zZ` inside a `*_dataloader` hook of your module, we will do this for you. Otherwise, define z inside your `__init__`.z&Trying to inject parameters into the `z{` instance. This would fail as it doesn't expose all its attributes in the `__init__` signature. The missing arguments are z. HINT: If you wrote the `zA` class, add the `__init__` arguments or allow passing `**kwargs`) r-   r   
ValueErrorr,   r<   __pl_saved_kwargs__pl_saved_arg_namesZ	__datasetvarsitemsrE   dictinspect	signature__init__
parametersanyvaluesupdatepopaddgetr   '_dataloader_init_kwargs_resolve_samplersortedr1   r&   joinr   r   keysr2   )r*   r5   r;   Zwas_wrappedr8   Zoriginal_datasetparamsZhas_variadic_kwargsr+   Zrequired_argsZsorted_required_argsZdataloader_cls_nameZmissing_args_messageZmissing_kwargsZsorted_missing_kwargsr#   )rT   rN   r9   rP   r$   r6   N   sx    




	r6   c              
   C   sF  t | d}|dur:|rDt|tu r:|j|kr:| j|jksBtdnt|tur:t|}t|dr|j}|j}|j	}|j
}td|||||\}	}}|	std|j dt|g|R i |}nhz|||j|jd}W nP ty( }
 z6d	dl}|d
t|
}|s
 td|
W Y d}
~
n
d}
~
0 0 dd|dddS |dddS )zThis function is used to handle the sampler, batch_sampler arguments associated within a DataLoader for its
    re-instantiation.rQ   NzdIt is not possible to have a batch sampler in your dataloader, when running on multiple IPU devices.r<   r5   zYTrying to inject a modified sampler into the batch sampler; however, it seems the class `z` does not have an argument called `sampler.` To mitigate this, expose an argument `sampler` in the `__init__` method of your custom class.)
batch_size	drop_lastr   z:.*__init__\(\) (got multiple values)|(missing \d required)zWe tried to re-instantiate your custom batch sampler and failed. To mitigate this, either follow the API of `BatchSampler` or instantiate your custom batch sampler inside `*_dataloader` hooks of your module.F   )r5   shufflerQ   rj   rk   )r5   rm   rQ   )getattrtyper   r5   rj   r   r,   r<   rV   __pl_saved_default_kwargsrW   _replace_value_in_saved_argsr2   r(   r7   rk   rematchstr)r*   r5   r;   rQ   Zbatch_sampler_clsr   kwargsdefault_kwargsrT   successerr   rs   r#   r#   r$   re      sh    






re   )r*   rankr   c                 C   s.   t tjddr*| jd u r*tt|d| _d S )NZPL_SEED_WORKERSr   )ry   )intosenvironrd   Zworker_init_fnr   r   )r*   ry   r#   r#   r$   _auto_add_worker_init_fn   s    r}   )explicit_cls)orig_objectr   r~   ru   r   c                O   s   |d u rt | n|}z||i |}W nx ty } z`dd l}|dt|}|sT | d }	d|j d|	 d|	 d|	 d	}
t|
|W Y d }~n
d }~0 0 t| dt	 }|D ]\}}||g|R   q|S )	Nr   z-.*__init__\(\) got multiple values .* '(\w+)'zThe zd implementation has an error where more than one `__init__` argument can be passed to its parent's `zr=...` `__init__` argument. This is likely caused by allowing passing both a custom argument that will map to the `zc` argument as well as `**kwargs`. `kwargs` should be filtered to make sure they don't contain the `zR` key. This argument was automatically passed to your object by PyTorch Lightning.__pl_attrs_record)
ro   r2   rr   rs   rt   groupsr&   r   rn   list)r   r~   r   ru   constructorresultrx   rr   rs   argumentmessageattrs_recordr"   r#   r#   r$   r7      s,    
 r7   )initstore_explicit_argr   c                    s(   t  tttdd fdd}|S )zWraps the ``__init__`` method of classes (currently :class:`~torch.utils.data.DataLoader` and
    :class:`~torch.utils.data.BatchSampler`) in order to enable re-instantiation of custom subclasses.N)objr   ru   r   c                    s"  t | dd}t| dd tj}tdd | D }t|d t	|  fdd|
 D }t| dst| d| t| d	  t| d
 t| d| d urv rt| d |  n  v rt| d    | g|R i   t| d| d S )N__pl_inside_initFTc                 s   s6   | ].}|j d kr|j|j|jfvr|j |jfV  qdS )r!   N)rM   rF   VAR_POSITIONALrG   rK   )r@   paramr#   r#   r$   rJ   &  s   z5_wrap_init_method.<locals>.wrapper.<locals>.<genexpr>c                    s2   i | ]*\}}| vr|vr|t jjkr||qS r#   )r[   	ParameterrL   )r@   rM   valueru   Zparam_namesr#   r$   rC   .  s   z6_wrap_init_method.<locals>.wrapper.<locals>.<dictcomp>r<   rV   rW   rp   __)rn   object__setattr__r[   r\   r^   r   r`   tupler0   rY   r,   index)r   r   ru   Zold_inside_initri   Zparameters_defaultsrv   r   r   r   r$   wrapper  s,    
 z"_wrap_init_method.<locals>.wrapper	functoolswrapsr   )r   r   r   r#   r   r$   _wrap_init_method  s    'r   )methodtagr   c                    s&   t  ttdd fdd}|S )zWraps the ``__setattr__`` or ``__delattr__`` method of classes (currently
    :class:`~torch.utils.data.DataLoader` and :class:`~torch.utils.data.BatchSampler`) in order to enable re-
    instantiation of custom subclasses.N)r   r   r   c                    s   |^}}t | dd\}}||ko&|k }t| d|f  | g|R   |rt | ddst | dt }||f t| d| t| d||f d S )NZ__pl_current_call)Nr   r   Tr   )rn   r   r   r   append)r   r   rM   r=   Zprev_call_nameZprev_call_methodZ
first_callr   r   r   r#   r$   r   N  s    z"_wrap_attr_method.<locals>.wrapperr   )r   r   r   r#   r   r$   _wrap_attr_methodI  s    r   )NNN)base_clsr   r   c              	   c   s   t | | hB }|D ]}d|jv r6|j|_t|j||_dtjfdtjffD ]N\}}||jv sd|| u rJd| }t||t	|| t||t
t	||| qJqdV  |D ]F}dD ]<}d| |jv rt||t	|d|  t|d|  qqdS )zThis context manager is used to add support for re-instantiation of custom (subclasses) of `base_cls`.

    It patches the ``__init__``, ``__setattr__`` and ``__delattr__`` methods.
    r]   r   __delattr__Z__oldN)r   r   r]   )r   __dict__r]   Z__old__init__r   r   r   r)   r   rn   r   r    )r   r   classesclsZpatch_fn_namer   Z
saved_nameZpatched_namer#   r#   r$   _replace_dunder_methodse  s     

r   .)replace_keyreplace_valuer   ru   rv   rT   r   c                 C   sj   | |v r>| | }|d| |f ||d d  }d||fS | |v sN| |v r`||| < d||fS d||fS )zTries to replace an argument value in a saved list of args and kwargs.

    Returns a tuple indicating success of the operation and modified saved args and kwargs
    Nrl   TF)r   )r   r   r   ru   rv   rT   Zreplace_indexr#   r#   r$   rq     s    
"

rq   )F)F)N)N)1r   r[   r{   collectionsr   
contextlibr   r   typingr   r   r   r   r	   r
   r   r   r   Z$lightning_utilities.core.inheritancer   Ztorch.utils.datar   r   r   r   r   Z lightning_fabric.utilities.enumsr   Z%lightning_fabric.utilities.exceptionsr   Z$lightning_fabric.utilities.rank_zeror   Zlightning_fabric.utilities.seedr   r   boolr.   r4   r:   rt   r6   re   rz   r}   r7   r   r   r   rq   r#   r#   r#   r$   <module>   sX   ,	 
a 

J /



