a
    d                     @   s|  d dl mZ d dl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mZ d d	lmZ d
dlmZmZm 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j,dd Z-e!#ej.j/j0j1dd Z2e!3ej4 e!3ej5 e!3ej6 dS )    )partialN)DispatchKeyDispatchKeySetExcludeDispatchKeyGuard)#_unwrap_all_tensors_from_functional_wrap_all_tensors_to_functionalfunctionalize)
PyOperator)FakeTensorMode)disable_proxy_modes_tracingmake_fxProxyTorchDispatchModetrack_tensor_treeunwrap_proxy)_get_current_dispatch_mode_pop_mode_temporarily)tree_flatten   )!_has_potential_branch_input_alias$_has_potential_branch_input_mutation!UnsupportedAliasMutationExceptionmapc                    s`  t |tjstdt|jdks0|jd dkr8tdtdd  D sRtdt ( t||d g R  W d    n1 s0    Y  d }d}|sd| }t	| j
j|r|d7 }q|}q| j
j| |g R }ttt| |}	| j
jd	||	i d
d}
 fdd|D }|d |jd g|d j}|t| t||
d | j
dS )Nzmap() must loop over a tensorr   zDmap() cannot be traced with scalar tensors or zero dimension tensorsc                 s   s   | ]}t |tjV  qd S N)
isinstancetorchTensor).0o r   d/var/www/html/stable-diffusion-webui/venv/lib/python3.9/site-packages/functorch/experimental/_map.py	<genexpr>        ztrace_map.<locals>.<genexpr>z3map() operands must be a list of tensors or modulesZbody_graph_r   call_functionr   )namec                    s   g | ]}|g R  qS r   r   r   xargsZ
body_graphr   r   
<listcomp>4   r!   ztrace_map.<locals>.<listcomp>)Zconstanttracer)r   r   r   
ValueErrorlenshapeallr   r   hasattrr)   rootZregister_modulepytreeZtree_mapr   r   Zcreate_proxy	new_emptyZcopy_stackr   )Z
proxy_modeZfunc_overloadfxsr'   Z	next_namei	candidateZ	node_argsZ
proxy_argsZ	out_proxyoutsoutr   r&   r   	trace_map   s2    6

 r9   c                    s0   t  }|d u sJ dt fdd|D S )Nz-Mode should never be enabled for CPU/CUDA keyc                    s   g | ]}|g R  qS r   r   r$   r'   r3   r   r   r(   C   r!   zmap_cpu.<locals>.<listcomp>)r   r   r2   )r3   r4   r'   moder   r:   r   map_cpu>   s    r<   c                 G   sH   t | ||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   r3   r   r   r   r(   K   s   z map_autograd.<locals>.<listcomp>)r   r-   r   r   r   AutogradCPUr   )r3   r4   r'   Zflat_operands_r   r   r   map_autogradF   s    r?   c                 G   sV   t  }|d usJ dt &}t|t| |g|R  }W d    n1 sH0    Y  |S )Nz5Mode should always be enabled for python fallback key)r   r   r9   r   )r3   r4   r'   r;   resr   r   r   map_proxy_torch_dispatch_modeR   s
    4rA   c                    s4    fdd|D }|d  |jd g|d jS )Nc                    s   g | ]}|g R  qS r   r   r$   r:   r   r   r(   ]   r!   z(map_fake_tensor_mode.<locals>.<listcomp>r   )r1   r,   )r3   r4   r'   r7   r   r:   r   map_fake_tensor_mode[   s    rB   c                  G   s   t ttj}t|  S r   )r   r   r   PythonDispatcherr   )r'   r>   r   r   r   map_python_dispatchera   s    rD   c              	      s   |   }|rdnd}t||d}t||d}t||d}|   t }	|	J  fdd}
|
||}t||rvtdt||rtdW d	   n1 s0    Y  t||g|R  }t	|| 
 d
W  d	   S 1 s0    Y  d	S )z
    Functionalization implementation for torch.map. Currently:
      1. We don't allow any input mutation inside the map function
      2. Our check for above condition is not exhaustive
    Zmutations_and_viewsZ	mutations)reapply_views)removec                    s2      | }ttj fdd|}|d f| S )Nc                    s      | S r   )fake_tensor_converter)r%   Zft_moder   r   <lambda>~   r!   z<map_functionalize.<locals>.get_fake_inputs.<locals>.<lambda>r   )rG   r0   Ztree_map_onlyr   r   )unwrapped_xsunwrapped_argsZfake_xsZ	fake_argsrH   r   r   get_fake_inputsz   s    
z*map_functionalize.<locals>.get_fake_inputsz torch.map is mutating the input!z torch.map is aliasing the input!N)level)Zfunctionalize_add_back_viewsr   r   lowerr
   r   r   r   r   r   rM   )interpreterr3   r4   r'   rE   r;   rJ   rK   Zfunctional_map_fnZfake_tensor_moderL   Zfake_inputsZ
map_returnr   rH   r   map_functionalizef   s(    
	


"rP   )7	functoolsr   r   Ztorch.utils._pytreeutilsZ_pytreer0   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.utils._python_dispatchr   r   r   _condr   r   r   r   r9   Zpy_implZCUDAZCPUr<   ZAutogradCUDAr=   r?   rA   rB   rC   rD   Z_CZ
_functorchZTransformTypeZFunctionalizerP   ZfallthroughZPythonTLSSnapshotZADInplaceOrViewZBackendSelectr   r   r   r   <module>   s:   #









+