a
    dG                     @   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
 d dlmZ eeZdd Zd	d
 Zdd ZejdddZejdddZdS )    N)
eval_frame)counters)aot_module_simplified)
FakeTensor)_disable_current_modesc                     s   t jjd fdd}|S )N)gmc                    s6  dd l }tdr$d  d< d|jj_d|jj_td d  d7  < d}|rvt	d td d	  d7  < | S  fd
d}dpd  |d< ddl
m} z^| B t| |fi }td d  d7  < t|W  d    W S 1  s0    Y  W n* ty0   td d	  d7  <  Y n0 d S )Nr   decompositionsTaot_autogradtotal   Fz5Unable to use AOT Autograd because graph has mutationZnot_okc                     s   t t  | i |S N)r   disable)argskwargsbw_compiler f/var/www/html/stable-diffusion-webui/venv/lib/python3.9/site-packages/torch/_dynamo/backends/common.py_wrapped_bw_compiler$   s    z?aot_autograd.<locals>.compiler_fn.<locals>._wrapped_bw_compilerr   fw_compiler)enable_aot_loggingok)functorch.compilecallablegetcompileconfigZuse_functionalizeZuse_fake_tensorr   logdebugZtorch._inductor.debugr   r   r   r   	Exception)r   example_inputsZ	functorchZuse_fallbackr   r   Zcgr   r   r   compiler_fn   s.    


0z!aot_autograd.<locals>.compiler_fn)torchZfxZGraphModule)r   r"   r   r!   r   r	      s    (r	   c                 C   s0   ddl m}m}m} |||d}| r,||d< |S )Nr   )default_decompositions#min_cut_rematerialization_partition
ts_compile)r   r   Zpartition_fnr   )r   r$   r%   r&   )Zuse_decompsr$   r%   r&   r   r   r   r   mem_efficient_fusion_kwargs:   s    r'   c                    s$   dd  t  fdd}|S )zg
    Decorator for backends that need real inputs.  We swap out fake
    tensors for zero tensors.
    c                 S   sp   t | ts| S | jr:dd |  D }dd |  D }n|  }|  }tj||| j| j| j	d}|
  |S )Nc                 S   s   g | ]}|j j|j jqS r   nodeZ	shape_envZ	size_hintexpr.0sr   r   r   
<listcomp>X       z;fake_tensor_unsupported.<locals>.defake.<locals>.<listcomp>c                 S   s   g | ]}|j j|j jqS r   r(   r+   r   r   r   r.   Y   r/   )dtypedevicerequires_grad)
isinstancer   Z_has_symbolic_sizes_stridessizestrider#   Zempty_stridedr0   r1   r2   Zzero_)xr4   r5   yr   r   r   defakeT   s     
z'fake_tensor_unsupported.<locals>.defakec                    sJ   t  0 tt |}| |fi |W  d    S 1 s<0    Y  d S r   )r   listmap)modelinputsr   r8   fnr   r   wrapperg   s    z(fake_tensor_unsupported.<locals>.wrapper)	functoolswraps)r>   r?   r   r=   r   fake_tensor_unsupportedN   s    rB   )returnc                 C   s"   | D ]}t |dr|j  S qd S )Nr1   )hasattrr1   r    r6   r   r   r   device_from_inputsp   s    
rF   c                 C   s"   | D ]}t |dr|j  S qd S )Nr0   )rD   r0   rE   r   r   r   dtype_from_inputsv   s    
rG   )r@   loggingr#   Ztorch._dynamor   Ztorch._dynamo.utilsr   Ztorch._functorch.aot_autogradr   Ztorch._subclassesr   Ztorch.utils._python_dispatchr   	getLogger__name__r   r	   r'   rB   r1   rF   r0   rG   r   r   r   r   <module>   s   
,"