a
    d8Z                     @   sj  d dl Z d dlZd dlZd dlZd dlZd dlZd dlmZmZm	Z	m
Z
 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 d	d
lmZmZ d	dlmZmZ d	dlmZm Z m!Z! d	dlm"Z"m#Z#m$Z$m%Z%m&Z&m'Z' d	dl(m)Z)m*Z*m+Z+m,Z,m-Z- d	dl.m/Z/m0Z0 d	dl1m2Z2m3Z3m4Z4m5Z5 d	dl6m7Z7 e 8e9Z:dd Z;G dd dej<j=Z>dS )    N)DictListOptionalSet)get_decompositions)dynamo_timed)ShapeEnv)no_dispatch   )config   )r   ir)CppWrapperCodeGenWrapperCodeGen)LoweringExceptionMissingOperatorWithDecompMissingOperatorWithoutDecomp)ConstantFixedLayoutInputBuffer	Pointwise	Reduction	TensorBox)FALLBACK_ALLOW_LISTlayout_constraints	loweringsmake_fallbackneeds_realized_inputs)CppSizeVarAllocatorSizeVarAllocator)convert_shape_to_inductorgather_originsget_dtype_sizesympy_product)Vc                 C   s,   t jt jt jt jt jt jt jt jh}| |v S N)	torchfloat32float64int64int32int16Zint8Zuint8bool)dtypeZsupported_dtype r.   ^/var/www/html/stable-diffusion-webui/venv/lib/python3.9/site-packages/torch/_inductor/graph.pysupported_dtype_of_cpp_wrapper/   s    r0   c                       s  e Zd ZejdddZejdddZdSejjd fdd	Z	d
d Z
edd ZedddZejdddZdd Ze fddZdd ZejdddZejdddZedd d!Zd"d# Zeejd$d%d&Zed' fd(d)Z fd*d+Zd,d- Zd.d/ Zd0d1 Z  fd2d3Z!d4d5 Z"ejj#d6 fd7d8Z$d9d: Z%d;d< Z&d=d> Z'd?d@ Z(dAdB Z)dCdD Z*dEdF Z+dGdH Z,dIdJ Z-edKdL Z.dMdN Z/dOdP Z0eddQdRZ1  Z2S )TGraphLowering)exc                 C   sx   | j rt| t| fS ddlm} |dt| jj }| j	||\}}}dd |D }dd |D }||fS )z
        Support dynamic shapes and dynamic strides by assigning variables
        to each dimension.  We duck-shape tensors, so if two tensors
        have the same size they get assigned the same symbolic variable.
        r   )ConstantSourceZ__unknown_tensor_c                 S   s$   g | ]}t |tjr|jjn|qS r.   
isinstancer&   ZSymIntnodeexpr.0ir.   r.   r/   
<listcomp>Y       z8GraphLowering.symbolic_sizes_strides.<locals>.<listcomp>c                 S   s$   g | ]}t |tjr|jjn|qS r.   r4   r8   r.   r.   r/   r;   Z   r<   )
reuse_shape_envr    sizestrideZtorch._dynamo.sourcer3   len
_shape_envZ
var_to_valZ,create_symbolic_sizes_strides_storage_offset)selfr2   r3   sourcer>   r?   _r.   r.   r/   symbolic_sizes_strides@   s     z$GraphLowering.symbolic_sizes_stridesc                 C   s,   dd |  D }dd | D }||fS )z+
        Primarily used to weights
        c                 S   s   g | ]}t |qS r.   sympyIntegerr8   r.   r.   r/   r;   a   r<   z6GraphLowering.static_sizes_strides.<locals>.<listcomp>c                 S   s   g | ]}t |qS r.   rF   r8   r.   r.   r/   r;   b   r<   )r>   r?   )rB   r2   r>   r?   r.   r.   r/   static_sizes_strides]   s    z"GraphLowering.static_sizes_stridesN)gmc                    s   t  | |d u r"t }d| _n|| _d| _|| _t|| _i | _i | _d | _	t
 | _g | _i | _t
 | _t
 | _d | _|| _t
 | _t
 | _td| _g | _i | _t | _d| _tj| _|| _d | _ dh| _!d S )NFTr   r1   zaten.convolution_backward)"super__init__r   r=   rA   r   sizevarsgraph_inputsgraph_inputs_originalgraph_outputssetdevice_typesbuffers	constantsremoved_buffersZinplaced_to_removewrapper_codenum_static_inputsZmutated_inputsZunaligned_buffersrG   rH   randomness_offsetrandomness_seedsname_to_buffertimeZcreation_timenamer   cpp_wrapper_can_use_cpp_wrappergraph_id	scheduler_warned_fallback)rB   rJ   Z	shape_envrW   r_   	__class__r.   r/   rL   e   s:    

zGraphLowering.__init__c                 C   s*   || j vr&| j | td|  d S )NzUsing FallbackKernel: )ra   addloginforB   r\   r.   r.   r/   warn_fallback   s    
zGraphLowering.warn_fallbackc                 C   s   t jS r%   )r$   	fake_moderB   r.   r.   r/   ri      s    zGraphLowering.fake_mode)buffer_namec                 C   sx   || j v r| j | jS || jv r.| j|  S || jv rF| j|  S td|}|rf| |dS td| d S )Nzas_strided\(([a-zA-Z0-9_]+),r   zcould not find )	rT   r-   rZ   	get_dtyperN   rematchgroupKeyError)rB   rk   mr.   r.   r/   rl      s    


zGraphLowering.get_dtype)devicec                 C   s`   d|j  d|j }|| jvrBtjd|tjd| j|< | j| tj	|tj
|tjg g ddS )a  
        Return a device-unique 1-element tensor storing our RNG seed.
        This will get initialized at the start of each graph in
        `wrapper.py`.

        Note this is only used by cuda backends.  The CPU backend handles
        RNG seeds as a sizevar.
        Zseed_rD   r.   )rr   r-   )rr   r-   r>   r?   )r\   Zlayout)typeindexrT   r&   zerosr)   rY   appendr   ZRandSeedBufferr   )rB   rr   r\   r.   r.   r/   random_seed_buffer   s    	
z GraphLowering.random_seed_bufferc                 C   s   | j }|| | _ |S )zX
        A global counter of how many random numbers we have handed out so far.
        )rX   )rB   Znumeloffsetr.   r.   r/   increment_randomness_offset   s    
z)GraphLowering.increment_randomness_offsetc                    s   t  j| S r%   )rK   run)rB   argsrb   r.   r/   rz      s    zGraphLowering.runc                 C   s   d| _ td| d S )NFz+Set _can_use_cpp_wrapper to False due to %s)r^   re   debug)rB   Zcondr.   r.   r/   disable_cpp_wrapper   s    z!GraphLowering.disable_cpp_wrapper)bufferc                 C   s&   t |tjr"t|dds"| d d S )NZ
cpp_kernelFExternKernel)r5   r   r   getattrr}   )rB   r~   r.   r.   r/   check_buffer_for_cpp_wrapper   s    z*GraphLowering.check_buffer_for_cpp_wrapperc                 C   s:   t jr| | dt| j }| j| || j|< |S )Nbuf)r   r]   r   r@   rS   rv   rZ   )rB   r~   r\   r.   r.   r/   register_buffer   s    

zGraphLowering.register_bufferr\   c              	      sb   t  tsJ  fdd| j D ]6\}}z| W q& tyZ   tjddd Y q&0 q&dS )z
        When a buffer is mutated we need to make sure all the reads to
        the old version are realized before the mutation happens.
        c                    sB   t | ttfr fdd| D S t | tjr>|  r>|   | S )Nc                    s   g | ]} |qS r.   r.   r9   x)visitr.   r/   r;      r<   zAGraphLowering.realize_users_of.<locals>.visit.<locals>.<listcomp>)r5   listtupler   IRNodeZ
is_user_ofrealize)valuer\   r   r.   r/   r      s    
z-GraphLowering.realize_users_of.<locals>.visitzerror in realize_users_ofT)exc_infoN)r5   strenvitems	Exceptionre   warning)rB   r\   keyr   r.   r   r/   realize_users_of   s    zGraphLowering.realize_users_ofc              
      s:    fdd}t t| t j jg R  S )Nc                     s   j  D ]X\} }  | kr
  | kr
 j|jkr
 j|jkr
t | r
|   S q
dt	j  }  j | < | S )NZconstant)
rT   r   r>   r?   r-   rr   r&   eqallr@   )r\   r   datarB   r.   r/   allocate   s    



z3GraphLowering.add_tensor_constant.<locals>.allocate)r   creater   ConstantBufferr   rr   r-   rI   )rB   r   r   r.   r   r/   add_tensor_constant   s    z!GraphLowering.add_tensor_constant)r\   device_overridec                 C   sZ   | j | j|ks|du r|S | d|j |jp0d }|| j vrV| j | || j |< |S )z
        We AOT copy constants to the devices they are needed on.
        If device_override doesn't match the constant's device, then
        copy it and return a different name.
        NrD   r   )rT   rr   rs   rt   to)rB   r\   r   Zalt_namer.   r.   r/   constant_name  s    
zGraphLowering.constant_name)targetc              	      s   t  |||}tjrBt| j| jk s,tjsB|j	sB| 
|\}}n| |\}}tt|t|j|j||}|| j|< |jj| j|< | j|jj |S r%   )rK   placeholderr   Zstatic_weight_shapesr@   rN   rW   dynamo_configZdynamic_shapesZ_has_symbolic_sizes_stridesrI   rE   r   r   r   r   rr   r-   r   rO   rR   rd   rs   )rB   r   r{   kwargsZexamplesizesstridestensorrb   r.   r/   r     s*    	
zGraphLowering.placeholderc                    s  t jt||` |tju rPt|d ttfrPt	 
|||W  d    S t|drv||i |W  d    S |tvr| dd }|tv rt| n\tjrt|grtnt}td|||| t| n$t|grt|||nt|||z&t| |i |}|W W  d    S  tyd } z&td t|||||W Y d }~n
d }~0 0 W d    n1 s|0    Y  d S )Nr   Z_inductor_lowering_function.z"Creating implicit fallback for:
%szError from lowering)r   r   current_originsr!   operatorgetitemr5   r   r   rK   call_functionhasattrr   r\   splitr   r   r   Zimplicit_fallbacksr   r   r   re   rf   Zoperator_strr   	exceptionr   )rB   r   r{   r   	base_nameerrorouterb   r.   r/   r   ,  s8    




zGraphLowering.call_functionc                 C   s   t | j|}t  |jdkr@t| |j|jW  d    S t|jdkr|jd dkrddl	m
} || |j|jdW  d    S W d    n1 s0    Y  | |S )Nr.   r   r      )r   )r-   rr   )r   moduler	   shaper   itemr-   rr   r@   loweringr   tolistr   )rB   r   r{   r   r   r   r.   r.   r/   get_attrS  s    
"BzGraphLowering.get_attrc                 C   s
   t  d S r%   AssertionErrorrB   r   r{   r   r.   r.   r/   call_modulea  s    zGraphLowering.call_modulec                 C   s
   t  d S r%   r   r   r.   r.   r/   call_methodd  s    zGraphLowering.call_methodc           	   	      s  t  |||}t|ttfs*J t|tdd |D sDJ |dd |D | _| j	 D ]\}}|
  t|ts|J |j}t|tjsJ |}|j}t|tr| |kr^tj|| j|  z | j|}| j| | j|< W q^ ty   Y q^0 q^|   d S )Nc              	   s   s.   | ]&}t |ttjtd tjtjtfV  qd S r%   )	r5   r   r   r   rs   r   rG   Exprintr   r.   r.   r/   	<genexpr>j  s   z'GraphLowering.output.<locals>.<genexpr>c                 S   s   g | ]}t j|qS r.   )r   r   Zrealize_inputr   r.   r.   r/   r;   x  r<   z(GraphLowering.output.<locals>.<listcomp>)rK   outputr5   r   r   rs   r   rP   rN   r   r   r   r   r   Z
StorageBoxr   get_nameZMutationLayoutZrealize_intorO   rt   
ValueErrorfinalize)	rB   r   r{   r   resultr\   r   Zvalue_storage_boxindrb   r.   r/   r   g  s.    
zGraphLowering.outputc                 C   s   | j D ]}|  qd S r%   )rS   Zdecide_layout)rB   r   r.   r.   r/   r     s    
zGraphLowering.finalize)nc           	   	      s  t j|h |jdkrf|jtv rf| |\}}t|j |g|R i |\}}| |j||}nt 	|}t
jjjjt
jjjjt
jjjjg t fdd|jD rt|jd t
jr|jd  }t
j|jd }|rt|rt j|t |}tt|j}|dkrt|tr|jD ]}|jtv r|   |jt
jjj!jt
jjj"jt
jjj#jfv rt j|t |jd  }|jdkr(t|j$j$t%t&fr(|'  q(|(t|j t|tr|) r|   W d    n1 s0    Y  |S )Nr   c                 3   s"   | ]}|j d kp|j v V  qdS )r   N)opr   r9   userZas_strided_opsr.   r/   r     s   z)GraphLowering.run_node.<locals>.<genexpr>valr   r   )*r   r   r   r   r   r   Zfetch_args_kwargs_from_envr   rK   run_noder&   opsZatenZ
as_strideddefaultZas_strided_Zas_strided_scatteranyusersr5   metaTensorr?   Z_prims_commonZis_non_overlapping_and_denser@   r   Zrequire_stride_orderZget_stride_orderrQ   r   r   Zrealize_hintZconvolutionZconvolution_backwardmmr   r   r   r   Z
mark_reuseZhas_exceeded_max_reads)	rB   r   r{   r   r   r   ZdenseZ	num_usersr   rb   r   r/   r     sR     








(zGraphLowering.run_nodec                 C   s   t jdkr| d d S )Nlinuxzplatform not linux)sysplatformr}   rj   r.   r.   r/   check_platform  s    
zGraphLowering.check_platformc                 C   s   t jr| d d S )Nzprofiler not supported)r   Zprofiler_mark_wrapper_callr}   rj   r.   r.   r/    check_profiler_mark_wrapper_call  s    z.GraphLowering.check_profiler_mark_wrapper_callc                 C   s2   t | jdkr$| j }|dkr$d S | d d S )Nr   cpuzdevice not CPU)r@   rR   popr}   )rB   rr   r.   r.   r/   check_device_for_cpp_buffer  s
    
z)GraphLowering.check_device_for_cpp_bufferc                 C   s.   | j  D ]\}}t| s
| d q
d S )Nzunsupported inputs dtype)rN   r   r0   rl   r}   )rB   rD   r   r.   r.   r/   check_input_for_cpp_buffer  s    z(GraphLowering.check_input_for_cpp_bufferc                 C   s   | j r| d d S )NZ	Constants)rT   r}   rj   r.   r.   r/   check_constant_for_cpp_buffer  s    z+GraphLowering.check_constant_for_cpp_bufferc                 C   s,   |    |   |   |   |   d S r%   )r   r   r   r   r   rj   r.   r.   r/   check_cpp_wrapper  s
    zGraphLowering.check_cpp_wrapperc                 C   s8   t jr,|   | jr,t| j| _t | _d S t	 | _d S r%   )
r   r]   r   r^   r   rA   rM   r   rV   r   rj   r.   r.   r/   init_wrapper_code  s    zGraphLowering.init_wrapper_codec                 C   sP   ddl m} |   || j| _ | j d us.J | j   | jd usFJ | j S )Nr   )	Scheduler)r`   r   r   rS   codegenrV   generate)rB   r   r.   r.   r/   r     s    
zGraphLowering.codegenc                    sn   ddl m mm} |j fdd}d}g }jD ]&}||}|||d f ||7 }q>||fS )Nr   )FusedSchedulerNodeNopKernelSchedulerNoder   c                    s   t rdS dd jjD }dd jjD }fdd t rl fdd|D }|| }|| }d}||B D ]X}|jv rj| }n|jv rxj| }nqx|tjj	t
| t|  7 }qx|S )Nr   c                 S   s   h | ]
}|j qS r.   r   r9   depr.   r.   r/   	<setcomp>  r<   zRGraphLowering.count_bytes.<locals>.get_read_write_buffers_sizes.<locals>.<setcomp>c                 S   s   h | ]
}|j qS r.   r   r   r.   r.   r/   r     r<   c                    s,   dd j |  jD }t|t j dkS )Nc                 S   s   h | ]
}|j qS r.   )r6   r   r.   r.   r/   r     r<   zkGraphLowering.count_bytes.<locals>.get_read_write_buffers_sizes.<locals>.is_materialized.<locals>.<setcomp>r   )Zname_to_noder   r@   rQ   Zsnodes)r   Zbuf_uses)r6   r`   r.   r/   is_materialized  s    zXGraphLowering.count_bytes.<locals>.get_read_write_buffers_sizes.<locals>.is_materializedc                    s   h | ]} |s|qS r.   r.   r   )r   r.   r/   r     r<   )r5   Zread_writesreadswritesrZ   rN   r$   graphrM   Z	size_hintr#   get_sizer"   rl   )r6   r   r   rU   Z
node_bytesr   r   r   r`   rB   )r   r6   r/   get_read_write_buffers_sizes  s,    






z?GraphLowering.count_bytes.<locals>.get_read_write_buffers_sizesr      )r`   r   r   r   rS   nodesrv   )rB   r   r   total_bytesZnode_countsr6   	num_bytesr.   r   r/   count_bytes  s    


zGraphLowering.count_bytesc                 C   s   ddl m} |  }tjr"t| ||}| j D ]\}}t	||| q6t
jr`td|j tj|j tjtj|jd d  |S )Nr   )PyCodeCachezOutput code: %sr   z.debug)Z	codecacher   r   r   r|   printloadrT   r   setattrr   Zoutput_codere   rf   __file__r$   renameospathsplitext)rB   r   codemodr\   r   r.   r.   r/   compile_to_module7  s    
zGraphLowering.compile_to_modulec                 C   s
   |   jS r%   )r  callrj   r.   r.   r/   compile_to_fnI  s    zGraphLowering.compile_to_fnc                 C   s   | j d usJ dd | j D S )Nc                 S   s,   g | ]$}t |tjst |tjs| qS r.   )r5   r   ZNoneAsConstantBufferZShapeAsConstantBufferr   )r9   r6   r.   r.   r/   r;   N  s   z2GraphLowering.get_output_names.<locals>.<listcomp>)rP   rj   r.   r.   r/   get_output_namesL  s    zGraphLowering.get_output_namesc                 C   s4   || j  v o2| j |  dko2| j |  jdkS )Nr   r   )rN   keysZ	get_numelZ
get_devicers   rg   r.   r.   r/   is_unspec_argU  s
    zGraphLowering.is_unspec_arg)NNN)3__name__
__module____qualname__r&   r   rE   rI   fxZGraphModulerL   rh   propertyri   r   rl   rr   rw   ry   r   rz   r}   r   ZComputedBufferr   r   r   r   r   r   r   r   r   r   r   r   Noder   r   r   r   r   r   r   r   r   r   r  r  r  r  __classcell__r.   r.   rb   r/   r1   ?   sX      &
	'%J
)
	r1   )?loggingr   r   rm   r   r[   typingr   r   r   r   rG   r&   Ztorch.fxZtorch._decompr   Ztorch._dynamo.utilsr   Z%torch.fx.experimental.symbolic_shapesr   Ztorch.utils._mode_utilsr	   Z_dynamor   r    r   Zcodegen.wrapperr   r   excr   r   r   r   r   r   r   r   r   r   r   r   r   r   r   rM   r   r   utilsr    r!   r"   r#   Zvirtualizedr$   	getLoggerr  re   r0   r  ZInterpreterr1   r.   r.   r.   r/   <module>   s2    
