a
    d-                     @   s  U d dl mZmZmZmZ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 e
dZe
dZe
dZe
dZeZeZeegeeef f Zeeegef ZG d	d
 d
eZi Zeee ef ed< eeeddddZeeef eee ef dddZee eeeef dddZ ee eee ef dddZ!ee eee dddZ"eedf eee ef dddZ#ee eeedf dddZ$eeee ef dddZ%ee eedd d!Z&d"eee ef dd#d$Z'ee ed"dd%d&Z(ee)ee  ee*e!e" ee+e#e$ eee%e& eee'e( ee,d'd(d)Z-eed'd*d+Z.ee,d'd,d-Z/eG d.d/ d/Z0G d0d1 d1e0Z1eeee e0f d'd2d3Z2ee e0ed4d5d6Z3eeed7d8d9Z4eee ee f Z5eee ee ee f Z6eee eee df f Z7eeeeef gef Z8eeeef gef Z9eegef Z:eegef Z;eegeegef f Z<ee5eef e<e9eeef  d:d;d<Z=eee e<e:eef  d:d=d<Z=ee7e<e;e  d:d>d<Z=e7e<e;e  d:d?d<Z=eee e:eef eed@dAdBZ>ee5eef e9eeef eed@dCdBZ>ee6eeef e8eeeef eed@dDdBZ>e7e;e eed@dEdBZ>eege,f ee,dFdGdHZ?eege,f ee,dFdIdJZ@eee e:ee,f ee,dKdLdMZAee5eef e9eee,f ee,dKdNdMZAee6eeef e8eeee,f ee,dKdOdMZAe7e;e, ee,dKdPdMZAeee e:ee,f ee,dKdQdRZBee5eef e9eee,f ee,dKdSdRZBe7e;e, ee,dKdTdRZBee0e	ee  dUdVdWZCdS )X    )
NamedTupleCallableAnyTupleListDictTypecastOptionalTypeVaroverloadUnionN)
namedtupleOrderedDict)	dataclassTSURc                   @   s   e Zd ZU eed< eed< dS )NodeDef
flatten_fnunflatten_fnN)__name__
__module____qualname__FlattenFunc__annotations__UnflattenFunc r   r   \/var/www/html/stable-diffusion-webui/venv/lib/python3.9/site-packages/torch/utils/_pytree.pyr   )   s   
r   SUPPORTED_NODES)typr   r   returnc                 C   s   t ||t| < d S N)r   r    )r!   r   r   r   r   r   _register_pytree_node/   s    r$   )dr"   c                 C   s   t |  t |  fS r#   listvalueskeysr%   r   r   r   _dict_flatten2   s    r+   )r(   contextr"   c                 C   s   dd t || D S )Nc                 S   s   i | ]\}}||qS r   r   .0keyvaluer   r   r   
<dictcomp>6       z#_dict_unflatten.<locals>.<dictcomp>)zipr(   r,   r   r   r   _dict_unflatten5   s    r5   c                 C   s   | d fS r#   r   r*   r   r   r   _list_flatten8   s    r6   c                 C   s   t | S r#   r'   r4   r   r   r   _list_unflatten;   s    r8   .c                 C   s   t | d fS r#   r7   r*   r   r   r   _tuple_flatten>   s    r9   c                 C   s   t | S r#   )tupler4   r   r   r   _tuple_unflattenA   s    r;   c                 C   s   t | t| fS r#   )r'   typer*   r   r   r   _namedtuple_flattenD   s    r=   c                 C   s   t t||  S r#   )r	   r   r4   r   r   r   _namedtuple_unflattenG   s    r>   zOrderedDict[Any, Any]c                 C   s   t |  t |  fS r#   r&   r*   r   r   r   _odict_flattenJ   s    r?   c                 C   s   t dd t|| D S )Nc                 s   s   | ]\}}||fV  qd S r#   r   r-   r   r   r   	<genexpr>N   r2   z#_odict_unflatten.<locals>.<genexpr>)r   r3   r4   r   r   r   _odict_unflattenM   s    rA   )pytreer"   c                 C   sV   t | }|j}t|dks&|d tkr*dS t|dd }t|tsDdS tdd |D S )N   r   F_fieldsc                 s   s   | ]}t |tkV  qd S r#   )r<   str)r.   entryr   r   r   r@   a   r2   z*_is_namedtuple_instance.<locals>.<genexpr>)r<   	__bases__lenr:   getattr
isinstanceall)rB   r!   basesfieldsr   r   r   _is_namedtuple_instanceY   s    
rN   c                 C   s   t | rtS t| S r#   )rN   r   r<   rB   r   r   r   _get_node_typec   s    rP   c                 C   s   t | t vS r#   )rP   r    r)   rO   r   r   r   _is_leafi   s    rQ   c                   @   sJ   e Zd ZU eed< eed< ed  ed< ddddZdee	d	d
dZ
dS )TreeSpecr<   r,   children_specsNr"   c                 C   s   t dd | jD | _d S )Nc                 S   s   g | ]
}|j qS r   )
num_leaves)r.   specr   r   r   
<listcomp>y   r2   z*TreeSpec.__post_init__.<locals>.<listcomp>)sumrS   rU   selfr   r   r   __post_init__x   s    zTreeSpec.__post_init__r   indentr"   c                    s   d| j j d| j d}d}t| jr t|7  || jd  7 }|t| jdkrZdnd7 }|d fdd	| jdd  D 7 }| d
}|| S )Nz	TreeSpec(z, z, [ r   rC   ,c                    s"   g | ]}d d   |   qS )
 )__repr__)r.   childr]   r   r   rW      r2   z%TreeSpec.__repr__.<locals>.<listcomp>z]))r<   r   r,   rH   rS   rb   join)rZ   r]   Zrepr_prefixZchildren_specs_strZrepr_suffixr   rd   r   rb   {   s    
&
zTreeSpec.__repr__)r   )r   r   r   r   r   Contextr   r[   intrE   rb   r   r   r   r   rR   r   s
   
rR   c                       s4   e Zd Zdd fddZd	eedddZ  ZS )
LeafSpecNrT   c                    s   t  d d g  d| _d S )NrC   )super__init__rU   rY   	__class__r   r   rj      s    zLeafSpec.__init__r   r\   c                 C   s   dS )N*r   )rZ   r]   r   r   r   rb      s    zLeafSpec.__repr__)r   )r   r   r   rj   rg   rE   rb   __classcell__r   r   rk   r   rh      s   rh   c           
      C   sr   t | r| gt fS t| }t| j}|| \}}g }g }|D ]"}t|\}}	||7 }||	 q>|t|||fS )zkFlattens a pytree into a list of values and a TreeSpec that can be used
    to reconstruct the pytree.
    )rQ   rh   rP   r    r   tree_flattenappendrR   )
rB   	node_typer   child_pytreesr,   resultrS   rc   flat
child_specr   r   r   ro      s    
ro   )r(   rV   r"   c                 C   s   t |tstdt| dt| |jkrNtdt|  d|j d| dt |tr`| d S t|j j}d}d}g }|j	D ]*}||j7 }|
t| || | |}q~|||jS )zqGiven a list of values and a TreeSpec, builds a pytree.
    This is the inverse operation of `tree_flatten`.
    z^tree_unflatten(values, spec): Expected `spec` to be instance of TreeSpec but got item of type .z2tree_unflatten(values, spec): `values` has length z, but the spec refers to a pytree that holds z items (z).r   )rJ   rR   
ValueErrorr<   rH   rU   rh   r    r   rS   rp   tree_unflattenr,   )r(   rV   r   startendrr   ru   r   r   r   rx      s2    



rx   )fnrB   r"   c                    s$   t |\}}t fdd|D |S )Nc                    s   g | ]} |qS r   r   )r.   ir{   r   r   rW      r2   ztree_map.<locals>.<listcomp>)ro   rx   )r{   rB   	flat_argsrV   r   r}   r   tree_map   s    r   )tyr"   c                 C   s   d S r#   r   r   r   r   r   map_only   s    r   c                 C   s   d S r#   r   r   r   r   r   r      s    c                 C   s   d S r#   r   r   r   r   r   r      s    c                    s,   t tgtf t tgtf d fdd}|S )a  
    Suppose you are writing a tree_map over tensors, leaving everything
    else unchanged.  Ordinarily you would have to write:

        def go(t):
            if isinstance(t, Tensor):
                return ...
            else:
                return t

    With this function, you only need to write:

        @map_only(Tensor)
        def go(t):
            return ...

    You can also directly use 'tree_map_only'
    )fr"   c                    s$   t  ttd fdd}|S )N)xr"   c                    s   t | r | S | S d S r#   rJ   )r   )r   r   r   r   inner   s    
z%map_only.<locals>.deco.<locals>.inner)	functoolswrapsr   r   )r   r   r   )r   r   deco   s    zmap_only.<locals>.deco)r   r   r   )r   r   r   r   r   r      s    ()r   r{   rB   r"   c                 C   s   d S r#   r   r   r{   rB   r   r   r   tree_map_only   s    r   c                 C   s   d S r#   r   r   r   r   r   r     s    c                 C   s   d S r#   r   r   r   r   r   r     s    c                 C   s   t t| ||S r#   )r   r   r   r   r   r   r   	  s    )predrB   r"   c                 C   s   t |\}}tt| |S r#   )ro   rK   mapr   rB   r~   _r   r   r   tree_all  s    r   c                 C   s   t |\}}tt| |S r#   )ro   anyr   r   r   r   r   tree_any  s    r   )r   r   rB   r"   c                 C   s   d S r#   r   r   r   rB   r   r   r   tree_all_only  s    r   c                 C   s   d S r#   r   r   r   r   r   r     s    c                 C   s   d S r#   r   r   r   r   r   r     s    c                    s$   t |\}}t fdd|D S )Nc                 3   s    | ]}t |r |V  qd S r#   r   r.   r   r   r   r   r   r@   "  r2   z tree_all_only.<locals>.<genexpr>)ro   rK   r   r   rB   r~   r   r   r   r   r      s    c                 C   s   d S r#   r   r   r   r   r   tree_any_only$  s    r   c                 C   s   d S r#   r   r   r   r   r   r   (  s    c                    s$   t |\}}t fdd|D S )Nc                 3   s    | ]}t |r |V  qd S r#   r   r   r   r   r   r@   .  r2   z tree_any_only.<locals>.<genexpr>)ro   r   r   r   r   r   r   ,  s    )rB   rV   r"   c           
      C   s   t |tsJ t| r"| g|j S t |tr0d S t| }||jkrFd S t| j}|| \}}t	|t	|j
ksx||jkr|d S g }t||j
D ]*\}}t||}	|	d ur||	7 }q d S q|S r#   )rJ   rR   rQ   rU   rh   rP   r<   r    r   rH   rS   r,   r3   _broadcast_to_and_flatten)
rB   rV   rq   r   rr   ctxrs   rc   ru   rt   r   r   r   r   8  s&    




r   )Dtypingr   r   r   r   r   r   r   r	   r
   r   r   r   r   collectionsr   r   Zdataclassesr   r   r   r   r   rf   ZPyTreer   r   r   r    r   r$   r+   r5   r6   r8   r9   r;   r=   r>   r?   rA   dictr'   r:   boolrN   rP   rQ   rR   rh   ro   rx   r   ZType2ZType3ZTypeAnyZFn3ZFn2ZFnZFnAnyZ	MapOnlyFnr   r   r   r   r   r   r   r   r   r   r   <module>   s   :$ $
	(""(,"(,"(