a
    d                     @   s~   d dl mZmZ d dlZd dlmZ ddlmZ ddlm	Z	mZ e ej
ej
dddZej
edd	d
ZG dd dZdS )    )ListTupleN   )_ddp)_get_registrycontractmodulereturnc                 K   s   t  j| fi | | S )zReplicates a module

    Args:
        module (torch.nn.Module): module to replicate

    Example::
        >>> # xdoctest: +REQUIRES(module:torch._C._distributed_c10d)
        >>> module = nn.Linear(3, 3)
        >>> replicate(module)
    )_ReplicateStatemark_modules)r	   kwargs r   p/var/www/html/stable-diffusion-webui/venv/lib/python3.9/site-packages/torch/distributed/_composable/replicate.py	replicate
   s    r   c                 C   s   dt | vS )z2Check if module is composable for `replicate` API.Zfully_shard)r   )r	   r   r   r   _can_compose   s    r   c                   @   s   e Zd ZddddZejddddZejddd	d
ZddddZeje	e
jdf ddddZeje	e
j e
je
jdddZdS )r   N)r
   c                 C   s    g | _ d| _t | _i | _d S )NF)moduleshas_initializednnZParameterList_param_listr   )selfr   r   r   __init__#   s    
z_ReplicateState.__init__)r   r
   c                 O   s`   |D ]P}t |std| j| | t|_dt|_|| j	 |
| j q|| _d S )NzGCannot apply `replicate()` on a Module already managed by `fully_shard`F)r   AssertionErrorr   appendr   stateZ_distributed_state_params_collectedZregister_forward_pre_hookforward_pre_hookZregister_forward_hookforward_post_hookr   )r   r   r   r	   r   r   r   r   )   s    z_ReplicateState.mark_modulesr   c                 C   sr   t |sd S tt|dr8t|jr,d S dt|_| jdd |jddD  | D ]}| 	| q^d S )Nr   Tc                 s   s   | ]}|j r|V  qd S N)Zrequires_grad).0paramr   r   r   	<genexpr>B   s   z<_ReplicateState._recursive_collect_params.<locals>.<genexpr>F)Zrecurse)
r   hasattrr   r   r   r   extend
parameterschildren_recursive_collect_params)r   r	   childr   r   r   r&   7   s    
z)_ReplicateState._recursive_collect_paramsc                 C   sB   | j r
d S d| _ | jD ]}| | qtj| jfi | j| _d S )NT)r   r   r&   r   ZDistributedDataParallelr   r   )r   r	   r   r   r   init_helperH   s    
z_ReplicateState.init_helper.)r	   inputr
   c                 C   s   |    | j  d S r   )r(   r   Zpre_forward)r   r	   r)   r   r   r   r   R   s    z _ReplicateState.forward_pre_hook)r	   r)   outputr
   c                 C   s   | j |S r   )r   Zpost_forward)r   r	   r)   r*   r   r   r   r   X   s    z!_ReplicateState.forward_post_hook)__name__
__module____qualname__r   r   Moduler   r&   r(   r   torchZTensorr   r   r   r   r   r   r   "   s   r   )typingr   r   r/   Ztorch.nnr    r   r   r   r.   r   boolr   r   r   r   r   r   <module>   s   