a
    dǬ                     @   sf  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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mZ ddl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 dd	lmZ e e!Z"d
d Z#G dd dZ$G dd dZ%G dd de%Z&G dd de%Z'G dd de%Z(G dd de%Z)dddZ*ej+G dd dZ,G dd dZ-dS )     N)DictListOptionalSet)dynamo_timed   )configdependenciesirmetrics)StarDepWeakDep)SimplifyIndexing)cache_on_selfcmp
has_triton)Vc                 C   sB   t | trt| td} tj| dd}d|v r>dt|d S |S )Nkey   )indent
    )
isinstancesetsortedstrpprintpformattextwrapr   )objresult r"   b/var/www/html/stable-diffusion-webui/venv/lib/python3.9/site-packages/torch/_inductor/scheduler.pyr      s    
r   c                   @   s0   e Zd Zdd Zdd Zdd Zdd ZeZd	S )

OutputNodec                 C   s   |h| _ g | _d S N)unmet_dependenciesinverse_usersselfdepr"   r"   r#   __init__$   s    zOutputNode.__init__c                 C   s   dS NFr"   r)   r"   r"   r#   is_reduction(   s    zOutputNode.is_reductionc                 C   s   dS )Nr"   r"   r-   r"   r"   r#   get_alias_names+   s    zOutputNode.get_alias_namesc                 C   s   dS )NZOUTPUTr"   r-   r"   r"   r#   get_name.   s    zOutputNode.get_nameN)__name__
__module____qualname__r+   r.   r/   r0   __repr__r"   r"   r"   r#   r$   #   s
   r$   c                   @   sB  e Zd ZdejdddZdd Zdd Zd	d
 Zdd Z	e
eef dddZdd Zed dddZdd Zdd Zdd ZejdddZee dd d!Zd"d# Zd$d% Zedd&d'Zedd(d)Zee dd*d+Zed  dd,d-Zd.d/ Zd0d1 Zd2d3 Z d4d5 Z!ej"d6d7d8Z#d9d: Z$d;d< Z%dAd>d?Z&d@S )BBaseSchedulerNode	Scheduler	schedulernodec                 C   sH   || _ || _d | _g | _| |  d | _d | _d | _d | _	d| _
d S r,   )r8   r9   usersr'   set_read_writesZget_read_writesrecursive_predecessors	min_order	max_order
last_usagewritten)r)   r8   r9   r"   r"   r#   r+   5   s    zBaseSchedulerNode.__init__c                 C   s   t | j d|  dS )Nz(name=)typer1   r0   r-   r"   r"   r#   r4   A   s    zBaseSchedulerNode.__repr__c                 C   s   |   }| dt| j dt| jj d| dt| jj | dt| j | dt| jj| j  g}z|| 	 g7 }W n  t
y   tjddd	 Y n0 d
| S )z#Longer form printout for trace logsz: (rA   z
.writes = z.unmet_dependencies = z.met_dependencies = zIgnoring error in debug_str()T)exc_infor   )r0   rC   r1   r9   r   read_writeswritesr&   readsdebug_str_extra	Exceptionlogwarningjoinrstripr)   namelinesr"   r"   r#   	debug_strD   s    "
zBaseSchedulerNode.debug_strc                 C   s   dS )N r"   r-   r"   r"   r#   rI   U   s    z!BaseSchedulerNode.debug_str_extrac                 C   s   t d| | j| jj d S )Nz(%s: unmet_dependencies = %s, writes = %s)rK   infor&   rF   rG   r-   r"   r"   r#   log_detailsX   s    zBaseSchedulerNode.log_detailsrenamesc                 C   s   |  | j| d S r%   )r;   rF   renamer)   rW   r"   r"   r#   update_mutated_names`   s    z&BaseSchedulerNode.update_mutated_namesc                 C   s   |  | j| d S r%   )r;   rF   Z	with_readr(   r"   r"   r#   add_mutation_depc   s    z"BaseSchedulerNode.add_mutation_depNodeUserr:   c                 C   sf   i }|D ]J}t |j|v rDt|j|t |j jo4|j|t |j< q||t |j< qt| | _d S r%   )idr9   r\   can_inplacelistvaluesr:   )r)   r:   r!   user"   r"   r#   	set_usersf   s    zBaseSchedulerNode.set_usersc                 C   s
   | j  S r%   )r9   r/   r-   r"   r"   r#   get_aliasesr   s    zBaseSchedulerNode.get_aliasesc                 C   s
   | j  S r%   )r9   get_mutation_namesr-   r"   r"   r#   get_mutationsu   s    zBaseSchedulerNode.get_mutationsc                 C   s   t |  p|  S r%   )boolrd   rf   r-   r"   r"   r#   has_aliasing_or_mutationx   s    z*BaseSchedulerNode.has_aliasing_or_mutation)rwc                 C   s   || _ | j j| _|   d S r%   )rF   rH   r&   
prune_deps)r)   ri   r"   r"   r#   r;   {   s    
z!BaseSchedulerNode.set_read_writesreturnc                 C   s   dd t | jj| jjD S )Nc                 S   s   h | ]
}|j qS r"   rP   .0r*   r"   r"   r#   	<setcomp>   s   z6BaseSchedulerNode.used_buffer_names.<locals>.<setcomp>)	itertoolschainrF   rH   rG   r-   r"   r"   r#   used_buffer_names   s    z#BaseSchedulerNode.used_buffer_namesc                    s    fdd j D  _ d S )Nc                    s   h | ]}|j  jjvr|qS r"   )rP   r8   available_buffer_namesrn   r-   r"   r#   rp      s   z/BaseSchedulerNode.prune_deps.<locals>.<setcomp>r&   r-   r"   r-   r#   rj      s    
zBaseSchedulerNode.prune_depsc                    s~   t   jD ](}t|ts |j    d7  < q fddfddjD }j| _j	| dS )a  
        Prunes stardeps intended for mutation ordering
        on an upstream fused node if after fusion there is another dependency
        on the fused upstream node, making the stardep redundant

        In essence this enforces an ordering on fusions. As fusions occur, prunable stardeps will
        be incrementally removed, enabling other fusions, ensuring they are fused in order.
        r   c                    s>   t | tr6 | j   dk}| j k}|p4|S dS d S )Nr   F)r   r   rP   r0   )r*   Zis_redundantZis_self_dep)name_to_dep_countname_to_fused_noder)   r"   r#   should_prune   s    
z<BaseSchedulerNode.prune_redundant_deps.<locals>.should_prunec                    s   h | ]} |r|qS r"   r"   rn   )rx   r"   r#   rp          z9BaseSchedulerNode.prune_redundant_deps.<locals>.<setcomp>N)
collectionsCounterr&   r   r   rP   r0   r;   rF   Zremove_reads)r)   rw   r*   Zdeps_to_pruner"   )rv   rw   r)   rx   r#   prune_redundant_deps   s    	

z&BaseSchedulerNode.prune_redundant_depsc                 C   s
   | j  S r%   r9   r0   r-   r"   r"   r#   r0      s    zBaseSchedulerNode.get_namec                 C   s   |   S r%   r0   r-   r"   r"   r#   get_first_name   s    z BaseSchedulerNode.get_first_namec                 C   s
   |   hS r%   r~   r-   r"   r"   r#   	get_names   s    zBaseSchedulerNode.get_namesc                 C   s   | gS r%   r"   r-   r"   r"   r#   	get_nodes   s    zBaseSchedulerNode.get_nodesc                 C   s
   | j  S r%   )r9   
get_devicer-   r"   r"   r#   r      s    zBaseSchedulerNode.get_devicec                 C   s   dS r,   r"   r-   r"   r"   r#   r.      s    zBaseSchedulerNode.is_reductionc                 C   s   dS r,   r"   r-   r"   r"   r#   is_template   s    zBaseSchedulerNode.is_templatec                 C   s   dS r,   r"   r-   r"   r"   r#   	is_extern   s    zBaseSchedulerNode.is_externread_depc                 C   s   dS r,   r"   r)   r   r"   r"   r#   r_      s    zBaseSchedulerNode.can_inplacec                    s   j  sd S t tfrB j  s. j  rBtjj	 j  d S t tfrt
jrttjtjjjjrttjdd d urddlm} t jjdd d}|D  ]} jj|j}|rtjj|r fdd|jD }t|dkr|d	 jr|d	 j  u rt|j   t!j"t!j#t!j$fs||j | j krtjj%|j  j  tjj&'|(  (  ttjtjjjjrtjj)*|(  tjj)* (   d S qtjj	 j  d S )
N	mutationsr   )buffer_reuse_keyc                 S   s   | j S r%   rm   xr"   r"   r#   <lambda>   ry   z,BaseSchedulerNode.allocate.<locals>.<lambda>r   c                    s"   g | ]}|j   jjvr|qS r"   )r9   r0   r8   rt   ro   r   r-   r"   r#   
<listcomp>   s
   z.BaseSchedulerNode.allocate.<locals>.<listcomp>r   )+r9   Zshould_allocater   SchedulerNoder/   re   r   graphwrapper_codeZcodegen_allocationr   inplace_bufferskerneltorchZ	_inductorcodegentritonZTritonKernelgetattrZcodegen.wrapperr   r   rF   rH   r8   name_to_nodegetrP   Z	can_reuser:   lenr_   Z
get_layoutr
   ZMultiOutputLayoutZMutationLayoutZAliasedLayoutZcodegen_inplace_reuseargsZmake_inplacer0   r   add)r)   r   Zordered_readsreadZ
input_nodeZremaining_usesr"   r-   r#   allocate   st    




zBaseSchedulerNode.allocatec                 C   s"   | j D ]}t|jtr dS qdS )NFT)r:   r   r9   r$   )r)   rb   r"   r"   r#   can_free  s    
zBaseSchedulerNode.can_freeTc                 C   s   t js
d S |r| jrd S | jj}g }|D ]}|jdkr8q(|d |d |d|j d|j  d|jv r(|jd  }|	dd }|d|
d	d

dd
dd  |d |d q(t|dkrd S || d| _d S )NoutputrS   z#pragma CMT ORIGIN:z#pragma CMT  stack_trace|{z{{}z}}r   \z#pragma CMT END ORIGINr   T)r   Zcomment_originr@   r9   originsopappendtargetmetasplitreplacer   
writelines)r)   bufferZ	only_oncer   Z	out_linesor   Zstack_trace_last_liner"   r"   r#   codegen_originating_info  s<    






z*BaseSchedulerNode.codegen_originating_infoN)T)'r1   r2   r3   r
   ZBufferr+   r4   rR   rI   rU   r   r   rZ   r[   r   rc   rd   rf   rh   r	   
ReadWritesr;   r   rs   rj   r|   r0   r   r   r   r   r.   r   r   	MemoryDepr_   r   r   r   r"   r"   r"   r#   r5   4   s6    @r5   c                   @   s   e Zd Zdd Zdd ZdS )ExternKernelSchedulerNodec                 C   s   |    dt| jdd  S )Nz.node.kernel = r   )r0   r   r9   r-   r"   r"   r#   rI   6  s    z)ExternKernelSchedulerNode.debug_str_extrac                 C   s   dS )NTr"   r-   r"   r"   r#   r   9  s    z#ExternKernelSchedulerNode.is_externN)r1   r2   r3   rI   r   r"   r"   r"   r#   r   5  s   r   c                   @   s   e Zd ZdS )NopKernelSchedulerNodeN)r1   r2   r3   r"   r"   r"   r#   r   =  s   r   c                       s~   e Zd Zdejd fddZdd Zdd Zd	d
 Zdd Z	dd Z
dd Zdd Zdd Zdd ZejdddZ  ZS )r   r6   r7   c                    s   t  || | \| _| _| || jf| _|  rJ| |	  n"| t
j| jg| jR ddi |  r| jjdd | jjD B | j_d S )N	normalizeTc                 S   s   h | ]}|  qS r"   )Zstrip_last_size)ro   wr"   r"   r#   rp   W  s   z)SchedulerNode.__init__.<locals>.<setcomp>)superr+   Zsimplify_and_reorder_sizes_bodyr   groupr   r;   Znormalized_read_writesr	   extract_read_writesr.   rF   rG   )r)   r8   r9   group_fn	__class__r"   r#   r+   B  s*    zSchedulerNode.__init__c                 C   s   |   }| d| jd  | d| jd  | d| j g}|  rb|| dt|    |  r|| dt|    t| jt	j
r|d| d	 |t| j d
 d|S )Nz.group.device = r   z.group.iteration = r   z	.sizes = z.aliases = z.mutations = zclass z_loop_body:r   r   )r0   r   r   rd   r   r   rf   r   r   r
   ZLoopBodyr   r   rR   rM   rO   r"   r"   r#   rI   e  s    zSchedulerNode.debug_str_extrac                 C   s   | j S r%   )r   r-   r"   r"   r#   
get_rangesu  s    zSchedulerNode.get_rangesc                 C   s   t | j S r%   )rg   r9   Zget_reduction_typer-   r"   r"   r#   r.   x  s    zSchedulerNode.is_reductionc                 C   s   t | jtjS r%   )r   r9   r
   TemplateBufferr-   r"   r"   r#   r   {  s    zSchedulerNode.is_templatec                 G   s   |    | | d S r%   )mark_runr   )r)   
index_varsr"   r"   r#   run~  s    zSchedulerNode.runc                 C   s   |    d S r%   )r   r-   r"   r"   r#   r     s    zSchedulerNode.mark_runc                 C   sH   | j }ttt|ttt|ks&J tttj|tj|}|S r%   )	r   summapr   dictziprq   rr   from_iterable)r)   r   sizes
var_rangesr"   r"   r#   ranges_from_index_vars  s     

z$SchedulerNode.ranges_from_index_varsc              	   C   s   |  |}znttt |F tj|  | j|  W d    n1 sN0    Y  W d    n1 sl0    Y  W n" ty   t	
d| j  Y n0 d S )NzError in codegen for %s)r   r   Zset_ops_handlerr   Zget_ops_handlerr   Zset_current_noder   rJ   rK   fatalr9   )r)   r   r   r"   r"   r#   r     s    

JzSchedulerNode.codegenc                    s$   j \}  fdd}t||S )zH
        Get the memory dependencies in the non-reduction axis.
        c                    s    | dd  D S )Nc                 S   s   g | ]}t d qS )r   )sympyZInteger)ro   _r"   r"   r#   r     ry   zCSchedulerNode.pointwise_read_writes.<locals>.fn.<locals>.<listcomp>)r   )indexZreduction_sizesr)   r"   r#   fn  s    z/SchedulerNode.pointwise_read_writes.<locals>.fn)r   r	   r   )r)   r   r   r"   r   r#   pointwise_read_writes  s    
z#SchedulerNode.pointwise_read_writesr   c                 C   sZ   |   s|  rdS t| jjdkrVt|drVtt| jj}|j|jkoT|j	|j	kS dS )NFr   r   )
rd   r   r   rF   rG   hasattrnextiterr   size)r)   r   	write_depr"   r"   r#   r_     s    zSchedulerNode.can_inplace)r1   r2   r3   r
   ComputedBufferr+   rI   r   r.   r   r   r   r   r   r   r	   r   r_   __classcell__r"   r"   r   r#   r   A  s   #r   c                   @   s2  e Zd ZdZeeedddZdee dddZ	e
ed	d
dZed	ddZe
ee d	ddZdd Ze
ee d	ddZee d	ddZdd Ze
dd Ze
dd Zdd Ze
dd Zeeef d d!d"Zd#d$ Zed% d&d'd(Zd)d* Zd+d, Zejd-d.d/Z d0d1 Z!d2d3 Z"d4S )5FusedSchedulerNodez
    This is a "fake" scheduler node that represents a group of scheduler nodes
    that are meant to be fused together. The way it does this is by maintaining
    its unmet dependencies as the union of its constituent nodes.
    node1node2c                 C   s(   |j |j u sJ | |j | |  S r%   )r8   r   )clsr   r   r"   r"   r#   fuse  s    zFusedSchedulerNode.fuser6   )r8   snodesc                    s   || _ || _d | _d | _g | _t|dd dj| _tt	j
dd |D | _| ttjjdd |D  t	|    fddtt	j
d	d |D D | jj | _td
d | j D | _tdd | j D | _d S )Nc                 S   s   t |  S r%   )intr.   r   r"   r"   r#   r     ry   z-FusedSchedulerNode.__init__.<locals>.<lambda>r   c                 S   s   g | ]
}|j qS r"   )r<   r   r"   r"   r#   r     ry   z/FusedSchedulerNode.__init__.<locals>.<listcomp>c                 S   s   g | ]
}|j qS r"   )rF   r   r"   r"   r#   r     ry   c                    s   h | ]}|j  vr|qS r"   rm   rn   namesr"   r#   rp     s   
z.FusedSchedulerNode.__init__.<locals>.<setcomp>c                 S   s   g | ]
}|j qS r"   ru   r   r"   r"   r#   r     ry   c                 S   s   g | ]
}|j qS r"   r=   r   r"   r"   r#   r     ry   c                 S   s   g | ]
}|j qS r"   )r>   r   r"   r"   r#   r     ry   )r   r8   r9   r:   r'   maxr   	functoolsreducer   unionr<   r;   r	   r   merger   rF   rG   r&   minr=   r>   )r)   r8   r   r"   r   r#   r+     s0    
zFusedSchedulerNode.__init__rk   c                 C   s   d dd | jD S )Nr   c                 S   s   g | ]}|  qS r"   r~   r   r"   r"   r#   r     ry   z/FusedSchedulerNode.get_name.<locals>.<listcomp>)rM   r   r-   r"   r"   r#   r0     s    zFusedSchedulerNode.get_namec                 C   s   | j d  S Nr   )r   r0   r-   r"   r"   r#   r     s    z!FusedSchedulerNode.get_first_namec                 C   s   t tjdd | jD S )Nc                 S   s   g | ]}|  qS r"   )r   r   r"   r"   r#   r     ry   z0FusedSchedulerNode.get_names.<locals>.<listcomp>r   r   r   r   r   r-   r"   r"   r#   r     s    zFusedSchedulerNode.get_namesc                 C   s"   |    dtdd | jD  S )Nz
.snodes = c                 S   s   g | ]}|  qS r"   r~   r   r"   r"   r#   r     ry   z6FusedSchedulerNode.debug_str_extra.<locals>.<listcomp>)r0   r   r   r-   r"   r"   r#   rI     s     z"FusedSchedulerNode.debug_str_extrac                 C   s   t tjdd | jD S )Nc                 S   s   g | ]}|  qS r"   )rs   r   r"   r"   r#   r     ry   z8FusedSchedulerNode.used_buffer_names.<locals>.<listcomp>r   r-   r"   r"   r#   rs     s    z$FusedSchedulerNode.used_buffer_namesc                 C   s   | j S r%   )r   r-   r"   r"   r#   r     s    zFusedSchedulerNode.get_nodesc                 C   s   t | j d|   dS )Nz(nodes=rA   rB   r-   r"   r"   r#   r4     s    zFusedSchedulerNode.__repr__c                 C   s   t dd | jD S )Nc                 s   s   | ]}|  V  qd S r%   )r.   r   r"   r"   r#   	<genexpr>  ry   z2FusedSchedulerNode.is_reduction.<locals>.<genexpr>anyr   r-   r"   r"   r#   r.     s    zFusedSchedulerNode.is_reductionc                 C   s   t dd | jD S )Nc                 s   s   | ]}|  V  qd S r%   )r   r   r"   r"   r#   r     ry   z1FusedSchedulerNode.is_template.<locals>.<genexpr>r   r-   r"   r"   r#   r     s    zFusedSchedulerNode.is_templatec                 C   s
   | j d S r   )r   r-   r"   r"   r#   r     s    zFusedSchedulerNode.get_devicec                 C   s   t dd | jD S )Nc                 s   s   | ]}|  V  qd S r%   )rh   r   r"   r"   r#   r     ry   z>FusedSchedulerNode.has_aliasing_or_mutation.<locals>.<genexpr>r   r-   r"   r"   r#   rh     s    z+FusedSchedulerNode.has_aliasing_or_mutationrV   c                 C   s   t d S r%   NotImplementedErrorrY   r"   r"   r#   rZ     s    z'FusedSchedulerNode.update_mutated_namesc                 C   s   t d S r%   r   r)   rP   r"   r"   r#   r[     s    z#FusedSchedulerNode.add_mutation_depr\   r]   c                 C   s   t d S r%   r   )r)   r:   r"   r"   r#   rc     s    zFusedSchedulerNode.set_usersc                 C   s   t d S r%   r   r-   r"   r"   r#   rd   
  s    zFusedSchedulerNode.get_aliasesc                 C   s   t d S r%   r   r-   r"   r"   r#   rf     s    z FusedSchedulerNode.get_mutationsr   c                 C   s   t d S r%   r   r   r"   r"   r#   r_     s    zFusedSchedulerNode.can_inplacec                 C   s   t d S r%   r   r-   r"   r"   r#   r     s    zFusedSchedulerNode.allocatec                 C   s   t d S r%   r   r-   r"   r"   r#   r     s    zFusedSchedulerNode.can_freeN)#r1   r2   r3   __doc__classmethodr5   r   r   r   r+   r   r   r0   r   r   r   rI   rs   r   r4   r.   r   r   rh   r   rZ   r[   rc   rd   rf   r	   r   r_   r   r   r"   r"   r"   r#   r     s:   


r   r"   c                    s`   t j fdd}ttttd }t|dkrJfdd|D tjr\|j|d |S )z
    A heuristic to decide loop iteration orders.  This has not been well
    tuned and may be something we should autotune.
    c                    s     dks dkr2t   dk dkS  fddD }fddD }tdd t||D }tdd t||D }|r|sdS |r|sdS t  S )	Nr   c                    s   g | ]}|  qS r"   r"   ro   sl)ar"   r#   r   &  ry   z6pick_loop_order.<locals>.index_cmp.<locals>.<listcomp>c                    s   g | ]}|  qS r"   r"   r   )br"   r#   r   '  ry   c                 s   s"   | ]\}}|d kp||k V  qdS r   Nr"   ro   Zsl_aZsl_br"   r"   r#   r   +  s   z5pick_loop_order.<locals>.index_cmp.<locals>.<genexpr>c                 s   s"   | ]\}}|d kp||k V  qdS r   r"   r   r"   r"   r#   r   .  s   r   )r   allr   )r   r   Zstride_len_aZstride_len_bZa_firstZb_firstr   stride_lengths)r   r   r#   	index_cmp   s    z"pick_loop_order.<locals>.index_cmpr   c                    s   g | ]} | qS r"   r"   )ro   pi)r   r"   r#   r   <  ry   z#pick_loop_order.<locals>.<listcomp>r   )	r   
cmp_to_keyr`   reversedranger   r   Zpick_loop_orderssort)r   r   Zpriority_idxr   orderr"   r   r#   pick_loop_order  s    r  c                   @   s*   e Zd ZU eed< dZeed< dd ZdS )r\   r9   Fr_   c                 C   s
   | j  S r%   r}   r-   r"   r"   r#   r0   G  s    zNodeUser.get_nameN)r1   r2   r3   r5   __annotations__r_   rg   r0   r"   r"   r"   r#   r\   B  s   
r\   c                       s  e Zd Ze fddZdd Zdd Zdd Zd	d
 Zdd Z	dd Z
dd Zdd Zdd Zdd Zdd ZeedddZdd ZeedddZd d! Zd"d# Zd$d% Zd&d' Zd(d) Zd*d+ Zd,d- Zed.d/d0Zejd1d2d3Zejd1d4d5Z ed6d7 Z!  Z"S )8r6   c                    s  t    i | _g | _h tjj tjj | _	|D ]}|j
d usNJ d| rj| jt| | q8t|tjtjfr| | j}| jt| || q8t|tjr| jt| | q8t|q8| j	tjj  | jD ]}|  qdd | jD | _d | _i | _i | _|   |    | !  | "  t# j$t%| j7  _$tj&'| j t%| j| _(dd | jD | _| )  | *  tj&+| j tj&,| j | -  d | _.t/ | _0t/ | _1d S )Nz2All nodes passed to scheduling must have an originc                 S   s   i | ]}|  |qS r"   r~   )ro   r9   r"   r"   r#   
<dictcomp>h  ry   z&Scheduler.__init__.<locals>.<dictcomp>c                 S   s   i | ]}|  |qS r"   r~   ro   nr"   r"   r#   r	  {  ry   )2r   r+   backendsnodesr   r   graph_inputskeys	constantsrt   r   Zis_no_opr   r   r   r
   r   r   get_backendr   r   r   ZExternKernelr   r   updaterj   r   rw   mutation_real_namemutation_renamescompute_dependenciestopological_sort_schedulecompute_predecessorsdead_node_eliminationr   Zir_nodes_pre_fusionr   debugZir_pre_fusionZnum_orig_nodes
fuse_nodescompute_last_usageZir_post_fusionZgraph_diagramdebug_draw_graphcurrent_devicer   buffer_names_to_freebuffer_names_no_longer_needed)r)   r  r9   r   r   r"   r#   r+   L  sX    





zScheduler.__init__c                 C   s0   t jdddkr,ddlm} || jdd dS )z,Generate an image of the graph for debuggingZINDUCTOR_WRITE_SCHEDULER_GRAPHN1r   )draw_buffersT)Zprint_graph)osenvironr   r  r!  r  )r)   r!  r"   r"   r#   r    s    zScheduler.debug_draw_graphc                 C   s0   t tjr,t d| | jD ]}|  qd S )Nz%s:)rK   isEnabledForloggingINFOrT   r  rU   )r)   labelr9   r"   r"   r#   debug_print_nodes  s    
zScheduler.debug_print_nodesc                    s|  t tjD ]}| }| D ]~}|v r|v r| }| }|| } D ]$}| |u st| |u rX||< qXq$|v r| |< q$| |< q$qfdd fdd d	fdd	}jD ]}	|	 D ]n}
|
}
||
|	 |	t	|
 |
 D ]@}| } |	 }||vr|	t
| |||	 qq|	jjD ]}||j|	|	| q\|	j |	 D ]>}
|	 j|
< |	 j|
< j|
|
j|	 < qqtj D ]}||tt	| q؈jD ]4}|tjjv r||tt	| tjj| qjD ]}	|	|	   q4jD ]"}	|	jD ]}|jj|	 q^qTdS )
zi
        Create dependency edges between nodes, handling aliasing and
        mutation properly.
        c                    s   | j v r j |  S | S r%   )r  )r  )rX   r)   r"   r#   rX     s    
z.Scheduler.compute_dependencies.<locals>.renamec                    sf   | h}j |  }t|jjd }|jjD ]8}|jj v r(|j|jkr(|j|jkr(| |j q(|S r   )	r   r`   rF   rG   rH   rP   r   r   r  )	node_nameZreachable_namesr9   r   r   )dep_closurer)   r"   r#   r*    s    



z3Scheduler.compute_dependencies.<locals>.dep_closureFc                    s    |   t|| d S r%   )r   r\   )Zused_by_nameZ	user_noder_   )name_to_usersrX   r"   r#   add_user  s    z0Scheduler.compute_dependencies.<locals>.add_userN)F)rz   defaultdictr`   r  r0   rd   r  rf   r[   r   r   rF   rH   rP   r_   rZ   r  r  r   r   r   get_output_namesr$   r  Zmutated_inputsr   rc   r:   r9   r'   r   )r)   r   Z
node1_nameZ
node2_nameZlist1Zlist2Zcombinedr   r,  r9   Zalt_nameZ
other_nodeZ
other_nameZknown_dep_node_namesr   r)  rP   userr"   )r*  r+  rX   r)   r#   r    s`    








zScheduler.compute_dependenciesc                 C   sN   g }| j D ]8}|jr || q
td|  tjj	|  q
|| _ dS )z0
        Remove any nodes without users
        zremoved dead node: %sN)
r  r:   r   rK   r  r0   r   r   removed_buffersr   )r)   Zupdated_nodesr9   r"   r"   r#   r    s    
zScheduler.dead_node_eliminationc                    sb   t  t  g  fdd| jD ]}| D ]}| |< q4q(| jD ]}| qJ| _dS )zD
        Ensure self.nodes is in topologically sorted order
        c                    sF   | vrB |  t| jdd dD ]} |j  q$|  d S )Nc                 S   s   | j S r%   rm   )dr"   r"   r#   r     ry   zDScheduler.topological_sort_schedule.<locals>.visit.<locals>.<lambda>r   )r   r   r&   rP   r   )r  r*   r   r!   seenvisitr"   r#   r4    s
    
z2Scheduler.topological_sort_schedule.<locals>.visitN)r   r   r  r   )r)   r9   rP   r"   r2  r#   r    s    


z#Scheduler.topological_sort_schedulec                 C   sr   i }| j D ]B}t }|jD ]}||j |||j O }q||| < ||_q
t| j D ]\}}||_||_	qXdS )z;
        Populate each node.recursive_predecessors
        N)
r  r   r&   r   rP   r0   r<   	enumerater=   r>   )r)   Zname_to_predecessorsr9   r<   r*   r  r"   r"   r#   r    s    

zScheduler.compute_predecessorsc                 C   s6   t dD ](}t| j}|   t| j|kr q2qdS )zO
        Mutates self.nodes to combine nodes into FusedSchedulerNodes.
        
   N)r  r   r  fuse_nodes_once)r)   r   Zold_lenr"   r"   r#   r  +  s
    
zScheduler.fuse_nodesc                    s   t | j}|  D ]\}}| j|  }| j|  }| ||r| ||st|| |	| |	| |
  | j fdd  D  qt|dd d| _|   |   dS )a  
        Mutates self.nodes to combine nodes into FusedSchedulerNodes.

        This relies on two key functions to control the logic:
            - self.can_fuses(): checks if a fusion is legal
            - self.score_fusion(): assigns priority to a given fusion
        c                    s   i | ]}|   qS r"   r~   r
  Znode3r"   r#   r	  I  ry   z-Scheduler.fuse_nodes_once.<locals>.<dictcomp>c                 S   s   | j S r%   r   r   r"   r"   r#   r   K  ry   z+Scheduler.fuse_nodes_once.<locals>.<lambda>r   N)r   r  get_possible_fusionsrw   r   can_fusewill_fusion_create_cycler   r   remover   r  r   r   r  r|   )r)   Zfused_nodesr   r   r"   r8  r#   r7  5  s"    



zScheduler.fuse_nodes_oncec                 C   s   | j D ]}|| j qd S r%   )r  r|   rw   )r)   r9   r"   r"   r#   r|   O  s    
zScheduler.prune_redundant_depsc                    s   g  t   fdd}tt}jD ] }| D ]}|| | q6q*| D ]}|| qTtj	rtt}jD ]"}t
|dd}|rx|| | qx| D ]}|| qt jddS )z^
        Helper to find all legal fusion opportunities, sorted by self.score_fusion()
        c                    s   t | D ]t\}}| |d d  D ]Z}||f}|v r6q | ||rX | q | r ||r  ||f q qd S )Nr   )r5  r   r:  r   r   )r  Znode1_indexr   r   r   Zpossible_fusionsr3  r)   r"   r#   check_all_pairsZ  s    
z7Scheduler.get_possible_fusions.<locals>.check_all_pairsr   NT)r   reverse)r   rz   r-  r`   r  rs   r   ra   r   aggressive_fusionr   r   score_fusion_key)r)   r>  Zbuffer_names_groupingr9   bufZnode_groupingZgroup_groupingr   r"   r=  r#   r9  S  s$    





zScheduler.get_possible_fusionsc                    sR    fdd t  | | B |j|jB  t fddD S )zHFinds whether there's a path from src to dst caused indirectly by fusionc                    sL   t | trH| vrH|  t| j@ pFt fdd| j D S dS )Nc                 3   s   | ]} j | V  qd S r%   rw   r
  checkr)   r"   r#   r     s   zDScheduler.will_fusion_create_cycle.<locals>.check.<locals>.<genexpr>F)r   r   r   rg   r<   r   )r9   rE  Zcombined_namesZcombined_predecessorsr)   visitedr"   r#   rE  }  s    
z1Scheduler.will_fusion_create_cycle.<locals>.checkc                 3   s   | ]} j | V  qd S r%   rC  r
  rD  r"   r#   r     ry   z5Scheduler.will_fusion_create_cycle.<locals>.<genexpr>)r   r   r<   r   )r)   r   r   r"   rF  r#   r;  z  s    	
z"Scheduler.will_fusion_create_cycler   c                 C   s2  ||u rdS t |ttfr&| s&dS t |ttfr@| s@dS | |j@ rRdS | r^dS | r| s|| s|tj	sdS |
 }||
 krdS | ||dk}|rtjr| s| rdS t| t|  tjkrdS | |j@ r| ||s
dS | |||S | |||S dS )zj
        Determine if it is possible to combine node1 and node2 into a
        single fused node.
        Fr   N)r   r   r   r   r   r<   rh   r.   r   Zepilogue_fusionr   score_fusion_memoryr@  r   r   Zmax_fusion_sizecan_fuse_verticalr  Zcan_fuse_horizontal)r)   r   r   deviceZno_shared_datar"   r"   r#   r:    sV    zScheduler.can_fusec           	      C   s   |  }t }|jD ]r}|jjD ]d}|j|jkr t|t|kr |j|jkr t|j	t|j	kr |j	dt|j	 |j	kr |
| q qdd |j| D }||@ rdS |D ]}|| j| j@ r dS qdS )a  
        Check if it is legal to fuse a consumer (node2) into a producer (node1).

        We can fuse them if all the reads of node2 either match
        corresponding writes in node1, or are written by nodes that can
        be scheduled before the fusion of node1 and node2.
        Nc                 S   s   h | ]
}|j qS r"   rm   rn   r"   r"   r#   rp     ry   z.Scheduler.can_fuse_vertical.<locals>.<setcomp>FT)r   r   r&   rF   rG   rP   rC   r   r   r   r   rw   r<   )	r)   r   r   Znode1_namesZcomputed_depsrdZcdZremaining_depsrP   r"   r"   r#   rI    s,    


zScheduler.can_fuse_verticalc                 C   sb   |  ||}tt|j|j t|j|j  }| tjkoD|dk| | koZ|dk||fS )a\  
        Assign a score (higher comes first) to the fusion of node1
        and node2.  When different fusions conflict with each other,
        this is the way we decide what order to run them in.

        Our current score is based on:
        - Estimate of the saved memory operations
        - Fusions closer together in original order
        r   )	rH  r   absr=   r>   r   r   Zepilogue_fusion_firstr.   )r)   r   r   Zmemory_scoreZproximity_scorer"   r"   r#   score_fusion  s    
zScheduler.score_fusionc                 C   s2   |j j|j jB |j j|j jB @ }tdd |D S )zf
        The first term in our fusion score that estimates number of saved memory operations.
        c                 s   s   | ]}|  V  qd S r%   )Znumbytes_hintrn   r"   r"   r#   r      ry   z0Scheduler.score_fusion_memory.<locals>.<genexpr>)rF   rH   rG   r   )r)   r   r   Zcommon_memory_depsr"   r"   r#   rH    s    zScheduler.score_fusion_memoryc                 C   s   |\}}|  ||S )z-
        Shim for list.sort(key=...)
        )rM  )r)   r  r   r   r"   r"   r#   rA    s    zScheduler.score_fusion_keyc                    sb   t  }tj D ]}|| qt jD ]2}| } fdd|D }|| |_|	| q*dS )z*
        Populate node.last_usage
        c                    s   h | ]} j ||qS r"   )r  r   )ro   kr-   r"   r#   rp     ry   z/Scheduler.compute_last_usage.<locals>.<setcomp>N)
r   r   r   r.  r   r  r  rs   r?   r  )r)   Zfuture_used_buffersr)  r9   Zused_buffersr"   r-   r#   r  	  s    
zScheduler.compute_last_usagec                 C   s   t | jtjj D ]h}|| jv rD| j| }| rztjj|j	 q|tjj
v rtjj
| j}| sjJ tjj|j q| j  dS )z*Free any buffers that are no longer neededN)r   r  r   r   r0  r   r   r   Zcodegen_freer9   r  dataZis_input_bufferclear)r)   rP   r9   Zstorager"   r"   r#   free_buffers  s    

zScheduler.free_buffersc                 C   sn   t jj| j@ D ]Z}|t jjvr|t jjjvr|| jvr|| jvr|t jjj	v r^t j
j| q| | qdS )zr
        Any buffers that are both created and have a last use in the
        same kernel can be removed.
        N)r   r   Zstore_buffer_namesr  Zmust_keep_buffersr   Zinput_buffersr  r  r   r   Zinplaced_to_remover   remove_bufferr   r"   r"   r#   remove_kernel_local_buffers&  s    
z%Scheduler.remove_kernel_local_buffersc                 C   s,   t d| dtjjj|< tjj| d S )Nzremove_buffer(%r)ZREMOVED)	rK   r  r   r   r   Zoutput_buffersr   r0  r   r   r"   r"   r#   rR  =  s    zScheduler.remove_bufferc                 C   s$   | j  D ]}|  q
|   d S r%   )r  ra   flushrQ  )r)   backendr"   r"   r#   rT  E  s    
zScheduler.flush)scheduler_nodec                 C   s6   t |tsJ |  |j}|tjj |   d S r%   )	r   r   r   r9   r   r   r   r   rQ  )r)   rV  r9   r"   r"   r#   codegen_extern_callJ  s
    zScheduler.codegen_extern_call)rJ  c                 C   s   |j dks"|jd us"J | dtjj|j  |j dkrPddlm} || S t st	j
|}|jdk rtd|j d|j d	|j ntd
ddlm} || S d S )Ncudaz( should have been normalized in loweringcpur   )CppScheduling   zFound z which is too old to be supported by the triton GPU compiler, which is used as the backend. Triton only supports devices of CUDA Capability >= 7.0, but your device is of CUDA capability .zCannot find a working triton installation. More information on installing Triton can be found at https://github.com/openai/triton)TritonScheduling)rC   r   r   r   Zdevice_typesr   Zcodegen.cpprZ  r   r   rX  Zget_device_propertiesmajorRuntimeErrorrP   minorZcodegen.tritonr]  )r)   rJ  rZ  Zdevice_propsr]  r"   r"   r#   create_backendQ  s*    

zScheduler.create_backendc                 C   s$   || j vr| || j |< | j | S r%   )r  ra  )r)   rJ  r"   r"   r#   r  i  s    
zScheduler.get_backendc                 C   s  | j D ]r}| j|j t|ts| }|| jksF| sF|	 rN| 
  || jkr|jdkr| jr| jjdkrtjj  |jd usJ dtjj|j n| jr| jjdkrtjj  || _| j|j |	 r| ^}}| ||| nT| r| | n>t|ttfr:| ||  nt|tsJJ |  tjjrj| |  | j|   q| 
  d S )NrX  zdevice should have an index)!r  r  r  r?   r   r   r   r  r   r   rT  rC   r   r   r   Zcodegen_cuda_device_guard_exitr   Zcodegen_cuda_device_guard_enterr  r   r  Zcodegen_templaterW  r   r   Zcodegen_nodesr   r   r   Zdebug_sync_kernelZcodegen_syncrt   r   )r)   r9   rJ  epiloguer"   r"   r#   r   n  sH    




zScheduler.codegen)#r1   r2   r3   r   r+   r  r(  r  r  r  r  r  r7  r|   r9  r;  r5   r:  rI  rM  rH  rA  r  rQ  rS  rR  rT  r   rW  r   rJ  ra  r  r   r   r"   r"   r   r#   r6   K  s8   :a
'1%	r6   )r"   ).rz   Zdataclassesr   rq   r%  r"  r   r   typingr   r   r   r   r   r   Ztorch._dynamo.utilsr   rS   r   r	   r
   r   r   r   Zsizevarsr   utilsr   r   r   Zvirtualizedr   	getLoggerr1   rK   r   r$   r5   r   r   r   r   r  Z	dataclassr\   r6   r"   r"   r"   r#   <module>   s<   

  nk
(