a
    d1                     @   s   d dl mZmZmZmZmZmZmZ d dlm	Z	 d dl
mZ d dlmZmZ d dlmZ d dlZd dlZd dlmZ d dlmZ eeZeej G d	d
 d
ZG dd dZdS )    )DictListSetIterableSequenceOptionalDeque)fuse_by_partitions)GraphModule)Node_get_qualified_name)OperatorSupportBaseN)copy)dequec                   @   sT   e Zd Zdeee dddZedddZedd	d
Z	edddZ
dd ZdS )	PartitionNidnodesc                 C   s"   || _ |d urt|nt | _d S N)r   setr   )selfr   r    r   j/var/www/html/stable-diffusion-webui/venv/lib/python3.9/site-packages/torch/fx/passes/infra/partitioner.py__init__   s    zPartition.__init__returnc                 C   s
   t | jS r   )strr   r   r   r   r   __repr__   s    zPartition.__repr__nodec                 C   s   | j | d S r   )r   addr   r    r   r   r   add_node   s    zPartition.add_nodec                 C   s   | j | d S r   )r   remover"   r   r   r   remove_node   s    zPartition.remove_nodec                 C   s
   t | jS r   )lenr   r   r   r   r   size   s    zPartition.size)NN)__name__
__module____qualname__intr   r   r   r   r   r#   r%   r'   r   r   r   r   r      s
   r   c                   @   s   e Zd Zdeeeeee  eee  ddddZ	e
edddZee d	d
dZee edddZee dddZed	ddZdS )CapabilityBasedPartitionerFN)graph_moduleoperator_supportallows_single_node_partitionnon_compute_ops!allowed_single_node_partition_opsr   c                 C   s:   || _ || _|| _|d ur|ng | _|d ur0|ng | _d S r   )r-   r.   r/   r0   r1   )r   r-   r.   r/   r0   r1   r   r   r   r   $   s    z#CapabilityBasedPartitioner.__init__)r    r   c                 C   s   | j t| j |S r   )r.   Zis_node_supporteddictr-   Znamed_modulesr"   r   r   r   Z__is_node_supported5   s    z.CapabilityBasedPartitioner.__is_node_supportedr   c                    s  i  i t  }ttd fdd}ttt d fdd}td t| jj	j
D ]}i }| |r| vrt|}||| d ||< |jD ]}| v rd | | < qt| }t|dkrX|d	 }	|dd  D ]}
||	|
 qqXtd
 i }| jj	j
D ]x}d}|jD ],}|jdks0t|jdkrd} q>q|r |d }|jD ] } |d |krV|||< qVq| D ]\}}||| q| js^td ddh}|t| j}g } D ]z\}}d	}|j
D ]P}|jdkr
t|j|vr
|d7 }|jdkrt|j| jv r|d7 }q|dkr|| q|D ]}|= qPtd  D ](\}}td| dd |j
D  qpt S )N)self_idother_idc                    s   t |  j  | j t  fdd} D ](}|jD ]}| vrD||rD  dS qDq: |  _| jD ]}| |< qx|= dS )Nc                    s   t  }||  |r| }|v r&q|v r2dS | v rx |  jD ],}|jD ] }| |  jvrR|| qRqHn|jD ]}|| q~| qdS )NTF)r   appendpopr   usersr!   )Z	root_nodestackr    Zp_node	user_node)
assignmentmerged_nodespartitions_by_idvisitedr   r   dfs_iter_find_cycleM   s"    


ziCapabilityBasedPartitioner.propose_partitions.<locals>.maybe_merge_partition.<locals>.dfs_iter_find_cycleFT)r   r   updater   r7   )r3   r4   r>   r    r9   r:   r<   )r;   r=   r   maybe_merge_partitionC   s    


zLCapabilityBasedPartitioner.propose_partitions.<locals>.maybe_merge_partitionr    r   c                    sl   |  v r |    |  |d u r. |  n:|vrR| | < t|| gd|< n| | < | |  d S )Nr   )r%   r6   r   r#   rB   r@   r   r   merge_single_node~   s    zHCapabilityBasedPartitioner.propose_partitions.<locals>.merge_single_nodezProposing partitions...   r   z=Reassigning getitem nodes to its producer node's partition...Tcall_functionz_operator.getitemFz'Filtering out single node partitions...ztorch.ops.aten.viewzPartitions proposed:zpartition #c                 S   s   g | ]
}|j qS r   )name).0r    r   r   r   
<listcomp>       zACapabilityBasedPartitioner.propose_partitions.<locals>.<listcomp>)	itertoolscountr+   r   r   loggerdebugreversedr-   graphr   ._CapabilityBasedPartitioner__is_node_supportednextr7   listkeysr&   opr   targetgetitemsr/   unionr   r0   r1   r5   values)r   Znew_partition_idrA   rC   r    Zmerge_candidatesZpartition_idr9   Zmerge_candidates_listr3   r4   Znodes_reassignmentZis_tuple_outputuserr   Zdefault_non_compute_opsr0   Zpartitions_to_remove	partitionZcompute_node_countr   r@   r   propose_partitions:   sz    ;











"z-CapabilityBasedPartitioner.propose_partitions)
partitionsr   c                 C   s    t d t| jdd |D S )NzFusing partitions...c                 S   s   g | ]}t |jqS r   )rR   r   )rG   r[   r   r   r   rH      rI   z>CapabilityBasedPartitioner.fuse_partitions.<locals>.<listcomp>)rL   rM   r	   r-   )r   r]   r   r   r   fuse_partitions   s    
z*CapabilityBasedPartitioner.fuse_partitions)r]   c                    s   t | jtdfdd i i ttt tt d fddttt tt d fdd|D ]\}t  }|jD ]2} |r|||j|s||j|r||| q|t|d	krl|j| |_qld S )
Nr   c                    s   | j dkot| j v S )NrE   )rT   r   rU   r   )r0   r   r   is_non_compute_node   s    
zVCapabilityBasedPartitioner.remove_bookend_non_compute_ops.<locals>.is_non_compute_node)r    r[   removed_nodesc                    st   | j dks| |vs| |v rdS | v r.|  S  | rh| jD ]}|||s<d| <  dS q<d| < dS d| < dS NplaceholderTF)rT   Zall_input_nodes)r    r[   r`   Zinput_n)r_   is_transparent_input_nodetransparent_input_nodesr   r   rc      s    
z\CapabilityBasedPartitioner.remove_bookend_non_compute_ops.<locals>.is_transparent_input_nodec                    st   | j dks| |vs| |v rdS | v r.|  S  | rh| jD ]}|||s<d| <  dS q<d| < dS d| < dS ra   )rT   r7   )r    r[   r`   Zoutput_n)r_   is_transparent_output_nodetransparent_output_nodesr   r   re      s    
z]CapabilityBasedPartitioner.remove_bookend_non_compute_ops.<locals>.is_transparent_output_noder   )r   r0   r   r   r   r!   r&   )r   r]   r[   r%   r    r   )r_   rc   re   r0   rd   rf   r   remove_bookend_non_compute_ops   s"    
""
z9CapabilityBasedPartitioner.remove_bookend_non_compute_opsc                 C   s   |   }| |}|S r   )r\   r^   )r   r]   Zfused_gmr   r   r   partition_and_fuse  s    
z-CapabilityBasedPartitioner.partition_and_fuse)FNN)r(   r)   r*   r
   r   boolr   r   r   r   r   rP   r   r   r\   r^   rg   rh   r   r   r   r   r,   "   s"      

 7r,   )typingr   r   r   r   r   r   r   Z!torch.fx.passes.utils.fuser_utilsr	   Ztorch.fx.graph_moduler
   Ztorch.fx.noder   r   Z torch.fx.passes.operator_supportr   loggingrJ   r   collectionsr   	getLoggerr(   rL   setLevelWARNINGr   r,   r   r   r   r   <module>   s   $
