a
    db%                     @   sP  d dl mZmZmZmZmZmZmZ d dlZd dl	m
Z
 d dlZd dlZd dlmZ d dlmZ g dZeeej eej f Zeejef Zeejj Zeejj Zee Zh dZedd	d
d Zedd	eeejjf ejjedddZedd	ejje dddZ!edd	G dd dZ"edd	ejj#ejj#dddZ$dS )    )ListTupleUnionDictAnySetMappingN)	dataclass)_get_qualified_name)compatibility)get_acc_ops_nameget_node_targetis_node_output_tensorFxNetAccFusionsFinderlegalize_graph>   call_moduleZcall_methodcall_functionF)Zis_backward_compatiblec                 C   sT   t | tr| S | jr*d| jv r*d| j S | jdd}|r@|nd d| j S d S )Nacc_opsacc_ops.z
torch._opsz	torch.ops .)
isinstancestr
__module____name__replace)kmodule r   e/var/www/html/stable-diffusion-webui/venv/lib/python3.9/site-packages/torch/fx/passes/tools_common.pyr      s    
r   )
submodulesnodereturnc                 C   s   |j tv s(J ddt d|j   |j dkrdt|jtsBJ | |j }t|dt|}t|S |j dkr|j}|j	durd|j	v rd	|j
 S t|S t|jtsJ |jS dS )
a,  
    Given a `node` returns its target typename.

    For "call_method" node, return node.target which is the name of that method being called.
    This could potential lead to conflict but should be okay because normally it's on a tensor.

    For "call_function" node, return typename of node.target.

    For "call_module" node, return typename of the module that node.target point to.

    If seeing "_VariableFunctionsClass" in the target name string, it will be replaced by
    "torch". e.g. _VariableFunctionsClass.relu would become torch.relu.
    zExpect op types of z, z, but found r   Z_base_class_originr   Nr   r   )opCALLABLE_NODE_OPSjoinr   targetr   getattrtyper   r   r   r
   )r    r!   ZsubmodZsubmod_typer&   r   r   r   r      s$    


r   )r!   r"   c                 C   s"   | j dd}|duo t|tjS )a  Checks if the node output produces a Tensor or not.

    NOTE: This requires to run `ShapeProp` on the containing fx graph before
    calling this function. This is because it works by checking the `type`
    metadata on the node. This metadata is produced by the `ShapeProp`.
    r(   N)metaget
issubclasstorchTensor)r!   type_r   r   r   r   C   s    r   c                   @   sh   e Zd ZdZejjedddZe	G dd dZ
deeef dd	d
Zeejjef dddZdS )r   z
    Finds groups of connected ACC nodes that pass non-tensor data between each other.
    Such groups are called fusion groups.
    )r   	acc_nodesc                 C   s   || _ t|jj| _|| _d S N)r   listgraphnodesr/   )selfr   r/   r   r   r   __init__U   s    zFxNetAccFusionsFinder.__init__c                   @   s6   e Zd ZU eed< eed< eed< eed< dd ZdS )!FxNetAccFusionsFinder.FusionGrouptop_node_idxr3   inputsnodes_need_processc                    sR   | j v rdS  j|  j |  j|  j fdd|jD  dS )z5
            Add a node to fusion group.
            Nc                    s$   h | ]}|j tv r| jvr|qS r   )r#   r$   r3   ).0nr4   r   r   	<setcomp>s   s   z=FxNetAccFusionsFinder.FusionGroup.add_node.<locals>.<setcomp>)r3   r9   addr8   discardupdateall_input_nodes)r4   r!   r   r<   r   add_nodeh   s    

z*FxNetAccFusionsFinder.FusionGroup.add_nodeN)r   r   __qualname__int__annotations__NodeSetrB   r   r   r   r   FusionGroupZ   s
   
rG   r6   )fusion_groupr8   c                 C   s\   |D ]R}|j tvrq| j||jk r(q||jv r8 dS | ||jr||  dS qdS )z
        Start from inputs and going reverse topological order. If any upstream node
        is in the fusion group, add all the nodes in this path to fusion group.
        TF)r#   r$   r3   indexr7   recursive_add_noderA   rB   )r4   rH   r8   argr   r   r   rJ   z   s    	


z(FxNetAccFusionsFinder.recursive_add_node)r"   c                 C   sr  i }t | j}|D ]X}||v r"q|jtvr.qd|jv r:q|| jvrFq| j| j||ht|j	|hd}|j
r0|j
 }| ||j d|jvr|jD ]4}|jtvrq||jv rq|| | ||j q|j	D ]V}|jtvrqd|jv rq||jv rq|| t|j| j||_| ||j qqjt|j| jksT|  j|j8  _q|jD ]}|j||< qZq|S )NZtensor_meta)r7   r3   r8   r9   )r1   r/   r#   r$   r)   rG   r3   rI   setrA   r9   poprJ   r8   usersrB   minr7   )r4   resultr/   r!   rH   userrK   r;   r   r   r   __call__   sZ    

















zFxNetAccFusionsFinder.__call__N)r   r   rC   __doc__r,   fxGraphModulerF   r5   r	   rG   r   NodeListrJ   r   NoderR   r   r   r   r   r   N   s   !
r   )gmr"   c                    s  dd | j jD tj }| j jD ] }|jD ]}|  d7  < q.q$t }| j jD ]}| dkrV|| qVi  t	|dkr|
 }|| fdd |< |jD ]*}|  d8  < | dkr|| qqvt	|jt	| j jk r
tdfdd	D  | j j|_|| _ | S )
a  
    Replace the graph of the given GraphModule with one that contains the same nodes as the
    original, but in topologically sorted order.

    This is used by the merge_matmul transformation below, which disturbs the topologically sorted
    order of its input GraphModule, so that this order is restored before further transformation.

    Arguments:
        gm: The graph module to topologically sort. It is modified in-place.

    Returns:
        The graph module in-place sorted
    c                 S   s   i | ]
}|d qS r   r   r:   r!   r   r   r   
<dictcomp>       z"legalize_graph.<locals>.<dictcomp>   r   c                    s    |  S r0   r   )x)envr   r   <lambda>   r\   z legalize_graph.<locals>.<lambda>z&Input graph has cycles, unable to add c                    s   g | ]} | d kr|qS rY   r   rZ   )indegr   r   
<listcomp>   r\   z"legalize_graph.<locals>.<listcomp>)r2   r3   r,   rT   ZGraphrN   collectionsdequeappendlenpopleftZ	node_copyRuntimeErrorZ_codegen)rX   Z	new_graphr!   rQ   queuecurr   )r_   ra   r   r      s,    



r   )%typingr   r   r   r   r   r   r   rc   Zdataclassesr	   r,   Ztorch.fxZtorch.fx.noder
   Ztorch.fx._compatibilityr   __all__r-   ZTensorsZTensorOrTensorsrT   rW   rV   rF   r   ZNamesr$   r   nnModuler   boolr   r   rU   r   r   r   r   r   <module>   s0   $

$#
 