a
    d''                     @   s  d dl mZ d dlmZ d dlZd dlmZ d dlm  m	Z
 d dlmZmZ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mZmZmZmZ d d
lmZ d dlm Z m!Z! d dlm"Z" eG dd de#Z$edZ%dd Z&e%'ej(e%'ej)dd Z*e%'ej+e%'ej,dd Z-e%'edd Z.e%'edd Z/e%'ej0dd Z1dd Z2dd Z3e%'ej4j5j6j7d d! Z8e%9ej: e%9ej; e%9ej< dS )"    )	dataclass)partialN)StorageWeakRef)DispatchKeyDispatchKeySetExcludeDispatchKeyGuard)#_unwrap_all_tensors_from_functional_wrap_all_tensors_to_functionalfunctionalize)
PyOperator)FakeTensorMode)disable_proxy_modes_tracingProxyTorchDispatchModemake_fxtrack_tensor_treeunwrap_proxy)_extract_tensor_metadata)_get_current_dispatch_mode_pop_mode_temporarily)tree_flattenc                   @   s   e Zd ZU eed< dS )!UnsupportedAliasMutationExceptionreasonN)__name__
__module____qualname__str__annotations__ r   r   e/var/www/html/stable-diffusion-webui/venv/lib/python3.9/site-packages/functorch/experimental/_cond.pyr      s   
r   condc                 C   s  t |ttfsJ dtdd |D s0J dt ( t|| }t|| }W d    n1 sd0    Y  g }g }	|jjD ]}
|
jdkr~|	|
j
 q~|jjD ]}
|
jdkr|		|
j
 qt|\}}t|	\}}t|t|ksJ tdt|D ],}|| }|| }|jd |jd ksJ qd }d}|sdd| }t| jj|r\|d	7 }n|}q2|}d
| }t| jj|rJ | jj|| | jj|| ||||f}ttt| |}| jjd||i dd}|| }t||d | jdS )Nz0Cond operands must be a list or tuple of tensorsc                 s   s   | ]}t |tjV  qd S N)
isinstancetorchTensor).0or   r   r   	<genexpr>)       ztrace_cond.<locals>.<genexpr>z'Cond operands must be a list of tensorsoutputr   Ztensor_metaZtrue_graph_   Zfalse_graph_call_functionZconditional)name)Zconstanttracer)r!   listtupleallr   r   graphnodesopextendargspytreer   lenrangemetahasattrr,   rootZregister_moduleZtree_mapr   r   Zcreate_proxyr   )Z
proxy_modeZfunc_overloadpredtrue_fnfalse_fnoperandsZ
true_graphZfalse_graph	true_outsZ
false_outsnodeflat_true_outs_flat_false_outsitrue_out	false_outZ	next_name	candidateZ	true_nameZ
false_namer4   Z
proxy_argsZ	out_proxyoutr   r   r   
trace_cond'   sN    *





rI   c                 C   s.   t  }|d u sJ d| r"|| S || S d S )Nz-Mode should never be enabled for CPU/CUDA key)r   )r;   r<   r=   r>   moder   r   r   
cond_denseh   s
    rK   c                 G   sN   t ||g|g \}}tdd |D s,J tttj}t| ||g|R  S )Nc                 S   s    g | ]}t |tjr|j qS r   )r!   r"   r#   Zrequires_grad)r$   fr   r   r   
<listcomp>x   s   z!cond_autograd.<locals>.<listcomp>)r   r/   r   r   r   AutogradCPUr   )r;   r<   r=   r>   Zflat_operandsrB   Zguardr   r   r   cond_autograds   s    rO   c                 C   sR   t  }|d usJ dt "}t|t| |||}W d    n1 sD0    Y  |S )Nz5Mode should always be enabled for python fallback key)r   r   rI   r   )r;   r<   r=   r>   rJ   resr   r   r   inner   s
    0rQ   c                 C   s   || }t |\}}t || \}}t|t|kr@tdt||D ]4\}}	t|}
t|	}|
|krJtd|
 d| qJ|S )Nz1Unmatched number of outputs from cond() branches.z=Unmatched tensor metadata from cond() branches.
true branch: z, false branch: )r5   r   r6   RuntimeErrorzipr   )r;   r<   r=   r>   r?   rA   rB   rC   rE   rF   Z	true_metaZ
false_metar   r   r   cond_fake_tensor_mode   s    rT   c                  G   s   t ttj}t|  S r    )r   r   r   PythonDispatcherr   )r4   rB   r   r   r   cond_python_dispatcher   s    rV   c              
   C   s   zt | | }W n: ty$   Y dS  tyJ } z|W Y d}~n
d}~0 0 t }|jjD ]Z}|jdkrr|| |jdkrZ|j}t	|t
jjrZ|jjrZ|jD ]}||v r  dS qqZdS )z
    Dispatch-trace the branch with fake inputs and check if
    producing graph has mutable op on the input. This is
    bit restrictive as the branch must be traceable.
    TNplaceholderr*   F)r   r   	Exceptionsetr0   r1   r2   addtargetr!   r"   Z_opsZ
OpOverloadZ_schemaZ
is_mutabler4   )branchfake_inputsgmeZinput_nodesr@   r[   argr   r   r   $_has_potential_branch_input_mutation   s"    



ra   c           	   
   C   s   zt | | }W n: ty$   Y dS  tyJ } z|W Y d}~n
d}~0 0 t }|jjD ]&}|jdkrZ|t|j	d 
  qZt|| \}}|D ]&}t|tjrt|
 |v r dS qdS )z
    Dispatch-trace the branch with fake inputs and check if
    producing graph has output aliasing the branch input. This is
    bit restrictive as the branch must be traceable.
    TNrW   valF)r   r   rX   rY   r0   r1   r2   rZ   r   r8   Z_typed_storager5   r   r!   r"   r#   )	r\   r]   r^   r_   Zinput_storagesr@   ZoutsrB   rH   r   r   r   !_has_potential_branch_input_alias   s    
rc   c              	      s2  |   }|rdnd}t||d}t||d}t||d}	t||d}
|   t }| |	|
fD ]2} fdd}ttj||}t	||rbt
dqb||fD ]2} fdd}ttj||}t||rt
d	qW d
   n1 s0    Y  t||	|
|}t||  dW  d
   S 1 s$0    Y  d
S )z
    Functionalization implementation for torch.cond. Currently:
      1. We don't allow any input mutation inside the branches
      2. Our check for above condition is not exhaustive
    Zmutations_and_viewsZ	mutations)reapply_views)removec                    s      | S r    Zfake_tensor_converterxZft_moder   r   convert   s    z#cond_functionalize.<locals>.convertz6One of torch.cond branch might be modifying the input!c                    s      | S r    rf   rg   ri   r   r   rj      s    z5One of torch.cond branch might be aliasing the input!N)level)Zfunctionalize_add_back_viewsr   r
   lowerr   r5   Ztree_map_onlyr"   r#   ra   r   rc   r   r	   rk   )interpreterr;   r<   r=   inputsrd   rJ   Zunwrapped_inputsZunwrapped_predZfunctional_true_fnZfunctional_false_fnZfake_tensor_moder\   rj   r]   Zcond_returnr   ri   r   cond_functionalize   s*    



(ro   )=Zdataclassesr   	functoolsr   r"   Z torch.multiprocessing.reductionsr   Ztorch.utils._pytreeutilsZ_pytreer5   Ztorch._Cr   r   r   Z!torch._functorch.eager_transformsr   r	   r
   Z
torch._opsr   Ztorch._subclasses.fake_tensorr   Z"torch.fx.experimental.proxy_tensorr   r   r   r   r   Ztorch.fx.passes.shape_propr   Ztorch.utils._python_dispatchr   r   r   rR   r   r   rI   Zpy_implZCUDAZCPUrK   ZAutogradCUDArN   rO   rQ   rT   rU   rV   ra   rc   Z_CZ
_functorchZTransformTypeZFunctionalizero   ZfallthroughZPythonTLSSnapshotZADInplaceOrViewZBackendSelectr   r   r   r   <module>   sF   A

	







%