
    ^jP                   	   U d dl m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	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Zd dlZd dlmZmZ d dlmZmZ d dlmZ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& er"d d
l'm(Z(m)Z)m*Z* d dl+m,Z, d dl-m.Z. ddl/m0Z0 ddl1m2Z2 d dl3Z3d dl4Z4d dl5Z4d dl6m7c m8Z9 d dl:m;Z;m<Z< d dl=m>Z> d dl?m@Z@mAZA d dlBmCZC d dlDmEZEmFZF d dlGmHZH d dlImJZJ d dlKmLZL d dlMmNZNmOZOmPZP d dlQmRZR ddlSmTZTmUZUmVZVmWZWm$Z$mXZX ddlYmZZZ ddl[m\Z\m]Z]m^Z^ ddl_m`Z`maZa ddlWmbZbmcZcmdZdmeZe ddlfmgZgmhZh ddlimjZj dd l$mkZkmlZlmmZmmnZnmoZompZp dd!lqmrZr dd"lsmtZtmuZu dd#lvmwZwmxZx dd$lymzZzm{Z{m|Z| dd%l}m~Z~ dd&l7mZmZmZmZmZmZmZmZmZmZmZmZmZmZmZmZmZmZmZ dd'lmZ  e	j(                  e      Ze4j.                  j1                  ed(      Ze4j.                  j1                  ed)      Ze4j.                  j1                  ed*      Ze4j.                  j1                  ed+      Zed,   Zd-ed.<    ed/      Z e!d0      ZejD                   G d1 d2             ZejD                   G d3 d4             Z G d5 d6e      Z ejD                  d78       G d9 d:             Z	 	 	 	 	 	 d}d;Z G d< d=      Z G d> d?      ZejD                   G d@ dA             ZejD                   G dB dCe             Z G dD d,      ZejZ                  d~dE       ZddFZ	 	 	 	 ddGZddHZ ejD                  d78       G dI dJ             ZddKZ G dL dM      Z	 	 	 	 	 	 	 	 ddNZ G dO dPe      Z G dQ dRe      Z G dS dTe      Z	 	 	 	 ddUZ	 	 	 	 	 	 	 	 ddWZ G dX dYe      Z G dZ d[e      Z G d\ d]e      Z G d^ d_e      Z G d` dae      Z G db dce      Z	 d	 	 	 	 	 	 	 dddZ	 	 	 	 	 	 ddeZddfZ	 	 	 	 	 	 	 	 	 	 	 	 ddgZ	 	 	 	 	 	 	 	 	 	 	 	 	 	 	 	 ddhZejD                   G di dj             Z ej                         ZddkZ	 	 	 	 ddlZddmZd7dn	 	 	 	 	 ddoZ	 	 	 	 	 	 ddpZddqZddrZddsZddtZdduZejD                   G dv dw             ZejD                   G dx dy             Z G dz dV      Z G d{ d|      Zy)    )annotationsN)Counterdefaultdict)as_completedFuture)AnyGenericLiteral
NamedTupleoverloadTYPE_CHECKING	TypeAliasTypeVar)	ParamSpec
OrderedSet   )ComputedBuffer	Pointwise)CallableIteratorSequence)
ModuleType)EnterCudaStreamContextLine)PythonWrapperCodegen)CoalesceVarAnalysis)countersdynamo_timed)use_pipelined_autotuning)LambdaFuturePyCodeCache)TritonTemplateCallerBase)get_metric_tableis_metric_table_enabled)get_stream_name)free_symbols)FloorDiv)free_symbol_is_typesymbol_is_typeSymT)
has_triton)commsconfigconfig_commsdependenciesirmetrics)can_codegen_without_upcasts)BackendFeatureget_scheduling_for_deviceKernel) estimate_nccl_collective_runtime/estimate_nccl_collective_runtime_nccl_estimator)Dep	MemoryDepStarDepWeakDep)GPUTooOldForTritonTritonMissing)count_flops_fx)assign_origin_nodeget_device_typeGraphPartitionSignatureMultiOutputMultiOutputLayout
NoneLayout)LoopBody)MemoryPlanningInfoForBufferMemoryPlanningInfoForNode)DevicePropertiesReductionHint)
green_textis_power_of_2red_text)SimplifyIndexing)&_unstable_customized_partition_wrappercache_on_selfcmpdevice_need_guardget_current_backendget_device_tflopsget_dtype_sizeget_gpu_dram_gbpsget_op_namesGraphPartitionMapIndentedBufferis_collectiveis_cudagraph_unsafe_opis_gpuis_multi_outputs_template#is_output_of_multi_outputs_templateis_waitsympy_product
sympy_subs)Vfusionloop_orderingcompute_dependencies
cudagraphsBaseSchedulerNoder   PartitionType_T_Pc                  l    e Zd ZU dZded<   dZded<   dZded<   d Zedd	       Z	e	 d	 	 	 dd
       Z
y)FusionResultNzbool | Noneshould_fusezCallable[[], bool] | Nonecallable_fnLambdaFuture | Nonefuturec                L    | j                   d u| j                  d uz  sJ d       y )NzLFusion result should contain either fusion decision or callable_fn, not both)rl   rm   selfs    d/var/www/ramen.bs-engineer-server.com/venv/lib/python3.12/site-packages/torch/_inductor/scheduler.py__post_init__zFusionResult.__post_init__   s0      ,1A1A1MN 	
Z	
N    c                    t        |      S )N)rl   rk   )clsrl   s     rs   fusezFusionResult.fuse   s    44ru   c                    t        ||      S )Nrm   ro   rw   )rx   rm   ro   s      rs   from_callablezFusionResult.from_callable   s     FCCru   )rl   boolN)rm   Callable[[], bool]ro   rn   )__name__
__module____qualname__rl   __annotations__rm   ro   rt   classmethodry   r|    ru   rs   rk   rk   |   sf    #K#-1K*1"&F&

 5 5 LPD,D6ID Dru   rk   c                  B    e Zd ZU ded<   ded<   ded<   dZded<   d
d	Zy)PendingFusionr   rm   rf   node1node2Nrn   ro   c                2    | j                   | j                  fS r~   r   r   rq   s    rs   get_fusion_nodeszPendingFusion.get_fusion_nodes   s    

DJJ''ru   )return+tuple[BaseSchedulerNode, BaseSchedulerNode])r   r   r   r   ro   r   r   ru   rs   r   r      s$    ##"&F&(ru   r   c                  0    e Zd ZU dZded<   ded<   ded<   y)_LocalEntryzOne row of the post-rewrite slice the gate builds.

    `cur` is the step the node currently wants to run at, `baseline`
    is its original index (used to break ties when sorting), and
    `node` is the slice member itself.
    intcurbaselinerf   nodeN)r   r   r   __doc__r   r   ru   rs   r   r      s     
HM
ru   r   T)slotsc                  r    e Zd ZU dZded<   ded<   dZded<   dZded	<    ej                  e	
      Z
ded<   y)ComboKernelMemoryContexta  Shared state used by the memory-aware combo gate.

    Candidate windows are evaluated independently against the original
    schedule. Earlier accepted combos only contribute through `running_peak`,
    which caps the cumulative peak drift from the original graph.
    OrderedSet[str]graph_outputsdict[BaseSchedulerNode, int]node_to_idxr   r   baseline_peakrunning_peakdefault_factory	list[int]baseline_live_beforeN)r   r   r   r   r   r   r   dataclassesfieldlistr   r   ru   rs   r   r      sE     #"--M3 L# '8k&7&7&M)Mru   r   c                    | j                         r|j                         sy| j                         j                  }|dv xr t        |      dk(  S )NF)cudaxputriton)r[   
get_devicetyperR   )r   r   device_types      rs   _is_gpu_triton_backendr      sI     <<>""$))K&W+>{+Kx+Wru   c                  $   e Zd ZdZedd       Zedd       Ze	 	 	 	 	 	 dd       Zedd       Z	e	 	 	 	 	 	 dd       Z
e	 	 	 	 	 	 dd       Zedd       Ze	 	 	 	 	 	 dd	       Zedd
       Ze	 	 	 	 	 	 dd       Zedd       Zedd       Zy)MixOrderReductionz
    This class contains utility functions to decide if we should fuse reductions
    reducing across different dimensions of the same input tensor.
    c                f    | j                         xr  t        d | j                         D              S )Nc              3     K   | ]V  }t        |t              rD|j                         r4t        |j                  t              r|j                  j
                  d u X y wr~   )
isinstanceSchedulerNodeis_reductionr   r   _split_size.0subnodes     rs   	<genexpr>z7MixOrderReduction.is_split_reduction.<locals>.<genexpr>   sK      +
'=1$$&7<<8	 LL$$D0+
s   AA)r   all	get_nodesr   s    rs   is_split_reductionz$MixOrderReduction.is_split_reduction   s3      " 
s +
>>++
 (
 	
ru   c                `   | j                  |      rd }d }|j                         D ]n  }t        |t              r*|j	                         rt        |j
                  t              s?|j
                  j                  J t        j                  j                  j                  t        |j
                  j                              }|j
                  j                  J t        j                  j                  j                  t        |j
                  j                              }||}|}t        j                  j                  j                  ||      sJ | d|        t        j                  j                  j                  ||      reJ | d|         |J ||fS |j                  d   S )N v.s. r   )r   r   r   r   r   r   r   _original_rangesra   graphsizevarssimplifyr_   _original_reduction_rangesstatically_known_equalsgroup)rx   r   xnumelrnumelr   	curxnumel	currnumels          rs   get_numel_rnumelz"MixOrderReduction.get_numel_rnumel   s   !!$'FF>>+ 4w6,,."7<<@||44@@@GG,,55!',,"?"?@	 ||>>JJJGG,,55!',,"I"IJ	 >&F&F77++CC	 4 	{34  77++CC	 4 	{34 148 %%%F##::a= ru   c                    | j                  |      }| j                  |      }t        |      dk7  st        |      dk7  s||k(  ryt        |      t        t        |            k(  S )N   F)r   lentuplereversed)rx   r   r   g1g2s        rs   has_mix_reduction_ordersz*MixOrderReduction.has_mix_reduction_orders  sZ     !!%(!!%(r7a<3r7a<28RyE(2,///ru   c                R   d}|j                   j                  D ]&  }t        |t              s|j                  |k(  s$|} n |sy|j
                  }|j                   j                  }|sDt        |t              sJ t        |              |j                  d   j                   j                  }|sJ t        |      t        |j                        z
  syt        j                  j                  j                  t!        |j"                        t!        |j%                                     ryy)z@
        The access to 'buf' is not a broadcast access.
        NFr   T)read_writesreadsr   r9   nameindex
var_rangesFusedSchedulerNoder   snodesr   r&   ra   r   r   r   r_   sizevalues)rx   bufr   	found_depdepr   r   s          rs   _is_full_accessz!MixOrderReduction._is_full_access  s    
 	##)) 	C#y)chh#o		
 %%00
d$67HDJ<H7Q33>>Jz:&E4F4F)GG
 7733)..)=9J9J9L+M
 ru   c                    g }|j                         |j                         z  }|D ]9  }| j                  ||      s| j                  ||      s)|j                  |       ; |S r~   )used_buffer_namesr   append)rx   r   r   outcommon_readsr   s         rs   get_common_readz!MixOrderReduction.get_common_read0  se     ..053J3J3LL 	 C""3.33F3FsE3R

3	  
ru   c                >    t        | j                  ||            dkD  S Nr   )r   r   rx   r   r   s      rs   has_common_readz!MixOrderReduction.has_common_read<  s!     3&&ue4599ru   c                    | j                  |      }t        j                  j                  j	                  |d   |d   z  d      S )Nr   r   fallback)r   ra   r   r   optimization_hint)rx   r   r   s      rs   	get_numelzMixOrderReduction.get_numelB  s>    !!$'ww11"Q%"Q%-!1LLru   c                $    | j                  |      S r~   )r   r   s      rs   get_fusion_scorez"MixOrderReduction.get_fusion_scoreG  s    
 }}U##ru   c                   t         j                  j                  syt        j                  j
                  ryt        ||      sy|j                         r|j                         sy|j                  |j                         z  s|j                  |j                         z  ry| j                  ||      syt        j                  ||      }t        |      dk(  ry| j                  |      r||}}n| j                  |      r||}}ny| j                  |      }|\  }}t         j                  j                   sd}	t        j                  j"                  j%                  t'        j(                  ||z  |	      dd      syt        j                  j"                  j%                  t'        j(                  ||dz        dd      syt        j                  j"                  j%                  t'        j(                  |d      dd      syt+        d |j-                         D              ryt        j                  j"                  j/                  |d	      syt        j1                  |      ryt3        d
 |j-                         D              }
|
S )zP
        Check whether we can fuse two reductions with mix loop orders.
        Fr   i  P T)size_obliviousfallback_valuer   i   c              3     K   | ]T  }|j                         rB|j                  j                  j                  t        j
                  t        j                  fv V y wr~   )r   r   datareduction_hintrI   INNERDEFAULTr   s     rs   r   z-MixOrderReduction.can_fuse.<locals>.<genexpr>  sR      
 ##% LL,,##%%
s   AAi @  c              3  t   K   | ]0  }|j                         r|j                  j                         d v  2 yw)>   sumprodN)r   r   get_reduction_typer   s     rs   r   z-MixOrderReduction.can_fuse.<locals>.<genexpr>  s=      
 ##% LL++-
s   68)r-   r   mix_order_reductionra   r   cpp_wrapperr   r   	ancestorsget_operation_namesr   r   r   r   is_contiguous_noder   #mix_order_reduction_non_strict_moder   evaluate_exprsympyGeanyr   statically_known_leqr   r   )rx   r   r   r   contiguous_node
other_noder   nrowncol
size_thresr   s              rs   can_fusezMixOrderReduction.can_fuseO  sG   
 }}00 77%eU3!!#5+=+=+?OOe7799OOe7799  ++E59 )88F|!!!%(*/ZO##E**/ZO!!/2
d }}@@ #J
 77##11j1#$ 2 
 
 77##11tax(#$ 2 
 
 77##11t$#$ 2 
   
 +446
 
 
 ww44T9E//@  
 &//1
 
 
ru   c                &    | j                  ||      S r~   )r  r   s      rs   are_mix_order_reductionsz*MixOrderReduction.are_mix_order_reductions  s     ||E5))ru   c                \     t         fdj                  j                  D              syy)Nc              3  V   K   | ]   }j                  |j                         " y wr~   )is_contiguous_loadr   )r   r   rx   r   s     rs   r   z7MixOrderReduction.is_contiguous_node.<locals>.<genexpr>  s'      
7:C""388T2
   &)FT)r   r   r   )rx   r   s   ``rs   r  z$MixOrderReduction.is_contiguous_node  s,     
>B>N>N>T>T
 
 ru   c                   ddl m} |j                         D ]  }t        |t              sJ |j
                  }|j                  |j                     }|D cg c]  }|j                  |k(  s|j                    }}t        |      dk(  rr|D ]u  }	|j                  |	   }
|j                  }t        |j                               }t        j                   j"                  j%                  |
||      }|d   dk(  rk|d   dk(  rt  y  yc c}w )Nr   MemoryUsageTyper   FT)torch._inductor.loop_bodyr  r   r   r   _bodymemory_usageLOADbuffer_name
index_namer   indexing_exprsr   r   keysra   r   r   stride_vars)rx   r   parent_noder  r   	loop_bodyentrieseindex_namesr  
index_exprr   var_symbolsr"  s                 rs   r  z$MixOrderReduction.is_contiguous_load  s   =))+ 	!DdM222

I,,_-A-ABG18QAAMMS<P1<<QKQ;1$ * !
&55jA
&11
 #:??#45gg..:: $B1,B10D !	!2 + Rs   D*DNr   rf   r   r}   )r   rf   r   ztuple[sympy.Expr, sympy.Expr]r   rf   r   rf   r   r}   )r   strr   rf   r   r}   )r   rf   r   rf   r   	list[str]r   rf   r   r   r   rf   r   rf   r   r   )r   r,  r#  rf   r   r}   )r   r   r   r   staticmethodr   r   r   r   r   r   r   r   r   r  r  r  r  r   ru   rs   r   r      sv   
 
 
 #! #!J 	0%	0.?	0		0 	0  B 	%	.?			 	 :%:.?:	: :
 M M $%$.?$	$ $ n n` *%*.?*	* *
    ru   r   c                  `   e Zd ZdZdZdZe	 	 	 	 	 	 dd       Ze	 	 	 	 	 	 dd       Ze	dd       Z
 G d dej                        Z G d	 d
ej                        Z ej                   d       G d d             Ze		 	 	 	 	 	 dd       Ze		 	 	 	 	 	 	 	 dd       Ze		 	 	 	 	 	 	 	 d d       Ze		 	 	 	 	 	 d!d       Ze		 	 	 	 	 	 d"d       Ze		 	 	 	 	 	 	 	 d#d       Ze		 	 	 	 	 	 	 	 	 	 	 	 d$d       Ze		 	 	 	 	 	 d%d       Ze	dd       Ze	dd	 	 	 	 	 	 	 	 	 	 	 d&d       Ze		 	 	 	 	 	 d'd       Zy)(NestedReductiona   
    Detects when an outer reduction and a dependent grouped reduction can be
    fused into one kernel. The outer reduction reduces over a large dimension
    (e.g. D) producing per-row statistics; the grouped reduction performs a
    small local reduction over the same logical elements. The grouped reduction
    may re-read the outer reduction's large input, consume its full-resolution
    output, or both.

    This is deliberately limited to same-total-numel pairs:
    both reductions must traverse the same number of logical elements.
    General output-size-reducing nested reductions, including split reductions,
    need different grid ownership and are rejected here.

    Example:
    - layernorm + block amax: amax over groups of G after layer_norm
    i      c                    | j                         xr8 |j                         xr& t        | j                         |j                  z        S )zFCheck that the grouped reduction is a consumer of the outer reduction.)r   r}   r  r  
outer_nodegrouped_nodes     rs   _is_dependent_reduction_pairz,NestedReduction._is_dependent_reduction_pair  sI     ##% P))+PZ3358N8NNO	
ru   c                    t         j                  j                  xr) t        j                  j
                   xr t        | |      S r~   )r-   r   nested_reductionra   r   r   r   r5  s     rs   _is_enabled_forzNestedReduction._is_enabled_for  s<    
 MM** AGG'''A&z<@	
ru   c                    | j                  ||      r| j                  ||      sy|j                  \  }\  }}|j                  \  }\  }}t        j                  j
                  j                  ||       S )zECheap filter for dependent reductions with different reduction sizes.F)r;  r8  r   ra   r   r   r   )rx   r   r   _rnumel1rnumel2s         rs   is_candidatezNestedReduction.is_candidate  sl     ""5
11%?++<Aw++<Aw77##;;GWMMMru   c                  v    e Zd ZdZ ej
                         Z ej
                         Z ej
                         Zy)NestedReduction.PointwiseDomaina  
        Where a pointwise node runs in the nested pipeline.

        The local reduction stage has three meaningful domains: its reduced
        output, its input before reducing the local lane, and the outer
        reduction's parent tile after broadcast-back.
        N)	r   r   r   r   enumautoREDUCEDLOCAL_REDUCTION_INPUTPARENT_FULLr   ru   rs   PointwiseDomainrB  #  s1    	 $))+ )		diikru   rH  c                  P    e Zd Z ej                         Z ej                         Zy)NestedReduction.GroupedAxisN)r   r   r   rC  rD  RXr   ru   rs   GroupedAxisrJ  3  s    DIIKDIIKru   rM  T)frozenc                  @    e Zd ZU ded<   ded<   ded<   ded<   ded<   y	)
&NestedReduction.PointwiseDomainContextr   grouped_reduction
sympy.Exprgrouped_numelgrouped_rnumelztuple[sympy.Expr, ...]local_reduction_domainparent_full_domainN)r   r   r   r   r   ru   rs   PointwiseDomainContextrP  7  s     ((!!"" 6622ru   rW  c                ~   |j                         sy|j                         D cg c]  }|j                         s| }}t        |      dk7  ry|d   }t        |t              rt        |j
                  t              sy|j                         \  }}t        |      dvst        |      dk7  ry|j
                  j                         dvryt        j                  j                  j                  |      }t        |t        t        j                  f      rt        |      dk  ry|t        j                  |      fS c c}w )z<Validate the candidate as a single simple grouped reduction.Nr   r   )r   r   >   r  maxminr   r   xor_sum)r   r   r   r   r   r   r   
get_rangesr   ra   r   r   r   r   r  Integer)	rx   r7  rT  sn
reductions	reductioniter_rangesreduce_ranges
group_sizes	            rs   _get_grouped_reduction_and_sizez/NestedReduction._get_grouped_reduction_and_size?  s   
 ((*#/#9#9#;QRr?PbQ
Qz?aqM	)]3:NNN<
 %.%9%9%;"] {6)S-?1-D >>,,. 7
 
 WW%%..~>
*sEMM&:;s:QR?R%--
333A Rs
   D:D:c                R    | j                  |||      }|y| j                  ||      S )a\  Classify pointwise nodes by nested stage and validate their ranges.

        Each pointwise node must run either on the grouped reduction output,
        on the grouped reduction's full local-group domain, or on the outer
        reduction's full parent domain, based on its producer/consumer
        relationship to the grouped reduction.
        F) _classify_nested_pointwise_nodes!_pointwise_domains_are_compatible)rx   r6  r7  domain_contextpointwise_domainss        rs   %_pointwise_nodes_match_nested_domainsz5NestedReduction._pointwise_nodes_match_nested_domainsh  s=      @@

 $44^EVWWru   c                   g }t               }|j                         D ]&  }|j                         s||j                         z  }( |j                         D ]\  }|j                         r||j                  z  s$t        |t              s y |j                  || j                  j                  f       ^ | j                  ||j                               }|y g ||S r~   )r   r   r   r  r  r   r   r   rH  rF  !_classify_grouped_pointwise_nodes)rx   r6  r7  rh  outer_pointwise_domainsouter_reduction_namesr^  grouped_pointwise_domainss           rs   rf  z0NestedReduction._classify_nested_pointwise_nodes  s      	 
 2<&&( 	BB %)?)?)AA%	B &&( 	B $r||3!"m4'..,,BBC	 %($I$I""$%
! %,E(E+DEEru   c                L   |j                   }|j                         }t        j                  j                  j                  |j                  |j                  z        }g }|D ]?  }|j                         rt        |t              s y|j                         }t        ||j                  z        }	t        ||j                  z        }
|	r|
r y|	s|
s y|	r| j                  j                  n| j                  j                  }|j                   \  }\  }}t        j                  j                  j#                  ||j                        r| j                  j$                  }n/t        j                  j                  j#                  ||      r|}n y|j'                  ||f       B |S )aV  Classify pointwise nodes relative to the grouped reduction.

        A node must be on exactly one side of the grouped reduction: either a
        producer feeding its local-group body, or a consumer of its reduced
        output. Its numel then determines whether it runs at reduced,
        grouped-full, or parent-full resolution.
        N)rQ  r  ra   r   r   r   rS  rT  r   r   r   r}   r  rH  rF  rG  r   r   rE  r   )rx   rh  nodesrQ  reduction_names
full_numelri  r^  sn_namesis_produceris_consumerfull_domainr=  sn_numeldomains                  rs   rl  z1NestedReduction._classify_grouped_pointwise_nodes  s~    +<<+??AWW%%..((>+H+HH


  	   	3B b-0--/Hx*;*E*EEFK=>K{ {   ##99((44 
  "xxA}!ww77.66 ,,44!!99(JO$$$b&\2A 	3B ! ru   c                0     t         fd|D              S )Nc              3  J   K   | ]  \  }}j                  ||        y wr~   )_pointwise_domain_is_compatible)r   r^  ry  rx   rh  s      rs   r   zDNestedReduction._pointwise_domains_are_compatible.<locals>.<genexpr>  s,      
F //FNK
    #)r   )rx   rh  ri  s   `` rs   rg  z1NestedReduction._pointwise_domains_are_compatible  s      
/
 
 	
ru   c                   ddl m} |j                  j                         \  }}|j                  \  }\  }}|| j
                  j                  u r|j                  }t        |      }	n|| j
                  j                  u rMt        j                  j                  j                  |j                  |j                  z        }|j                  }	nf|| j
                  j                   u sJ t        j                  j                  j                  |j                  |j                  z        }|j"                  }	t        j                  j                  j%                  ||      xr  |j'                  |	|j                               S )Nr   
SIMDKernel)codegen.simdr  rQ  r\  r   rH  rE  rS  r   rF  ra   r   r   r   rT  rU  rG  rV  r   is_compatible)
rx   r^  ry  rh  r  ra  r=  rx  expected_numelexpected_groupss
             rs   r|  z/NestedReduction._pointwise_domain_is_compatible  s=    	-'99DDFQ88=HaS((000+99N49+4FOs**@@@WW--66,,~/L/LLN -CCOS00<<<<<WW--66,,~/L/LLN -??Oww77n
 I&&H	Iru   c                  ddl m} t        |t        t        f      syt
        j                  j                  r|j                         nd }t        |j                               }|j                  ||||      }	t        |	      t        d      k7  ryt        j                  j                  j!                  |	d   |      r-t        j                  j                  j!                  |	d   |      sy|| j#                  ||      kD  S )Nr   SIMDSchedulingT)xr0_r  r  )r  r  r   r   r   r-   r   coalesce_tiling_analysisget_coalesce_analysisr   r   select_tilingr   ra   r   r   r   _max_min_block_group_size)
rx   r6  outer_numelouter_rnumelgrouped_axisrc  r  coalesce_analysisnode_scheduletilings
             rs   "_min_block_unprofitable_for_kernelz2NestedReduction._min_block_unprofitable_for_kernel  s     	1*}6H&IJ }}55 ,,. 	
 Z1134--	
 fL!99GG44VC[+N  88UC99-VVVru   c                P   ddl m} |D cg c]M  }t        |t              r;|j	                         r+t        |j
                  t              r|j                  |      O }}|| j                  j                  u r |rt        d |D              r| j                  S | j                  S c c}w )Nr   SIMDKernelFeaturesc              3  @   K   | ]  }|t         j                  u   y wr~   )rI   r   )r   hints     rs   r   z<NestedReduction._max_min_block_group_size.<locals>.<genexpr>?  s     LDDM///L   )codegen.simd_kernel_featuresr  r   r   r   r   r   r   rM  rK  r   MAX_INNER_R_GROUP_SIZEMAX_NON_INNER_GROUP_SIZE)rx   r  r  r  r^  reduction_hintss         rs   r  z)NestedReduction._max_min_block_group_size-  s     	E $
"m,!277N3	 --b1
 
 COO---LOLL---+++
s   AB#c                    | j                  ||      sy| j                  ||      sJ |}|}|j                  \  }}|j                  \  }}|\  }}	|\  }
}t        j                  j
                  j                  |	|      ry| j                  ||      }|yt        j                  j
                  j                  ||	z        }t        j                  j
                  j                  |
|z        }t        j                  j
                  j                  ||      sy|\  }}| j                  |||	||      }|y|| j                  j                  u r|	n|}|j                         \  }}t        |      dk(  rW|| j                  j                  u r|d   n|d   }t        j                  j
                  j                  t        ||      |      s@yt        j                  j
                  j                  t        j                   ||      d      syt#        |      }d|k  rt%        |      sy| j'                  |||	||      ry|j                         \  }}| j)                  ||
|g ||||	f      }| j+                  |||      syy)	zACheck whether a dependent cross-axis reduction pair can be fused.Fr6  r   r   r   )r  rc  rQ  rS  rT  rU  rV  T)r;  r8  r   ra   r   r   r   rd  r   get_grouped_axisrM  rK  r\  r   r'   r  Modr   rK   r  rW  rj  )rx   r   r   r6  r7  r=  outer_groupgrouped_groupr  r  rS  rT  grouped_reduction_infoouter_totalgrouped_totalrQ  rc  r  parent_grouped_axisra  grouped_axis_groupsgroup_size_intrb  rh  s                           rs   r  zNestedReduction.can_fuseD  s   
 ""5%0//u===
 $));'--=$/!\(5%~7733L.Q "%!D!D."
 ") gg&&//l0JK((11-.2PQww77]S(>%:++! , 
 (COO,=,==L; 	 +557Q{q ".#//2C2C"CAUV   77##;;,j9;N !!99II):6
 Z^#n(E11%% 2 
 %6%A%A%C"]33/')#A[#A=#A +\: 4 
 88

 ru   Nr  c                  t         j                  j                  }|j                         \  }}t	        |      dk7  st	        |      dk7  rt	        |      dk(  rt	        |      dk(  r|j                  |d   |      sy|j                  t        ||      |d         r(|j                  |d      r| j                  j                  S |j                  t        ||      d      r+|j                  |d   |      r| j                  j                  S y|j                  |d   |      sy|j                  |d   |      r5|j                  t        ||      |d         r| j                  j                  S |j                  |d   |      r5|j                  t        ||      |d         r| j                  j                  S || j                  ||      S y)zAReturn which parent axis is split by the grouped local reduction.r   r   r   N)ra   r   r   r\  r   r   r'   rM  rK  rL   _get_grouped_axis_from_loop_body)	rx   rQ  r  r  rc  r6  r   ra  rb  s	            rs   r  z NestedReduction.get_grouped_axis  s    77##%6%A%A%C"]{q C$6!$;;1$]);q)@77a8H*U 33\:6A66{AF??,,, 33[*5q66{1~|T??,,,//a0@*M ++NK
..\:.A
 ??$$$++NL
..[*-{1~
 ??$$$!77
DUVVru   c                x   ddl m |j                         D cg c]  }|j                         s| }}t	        |      dk7  ryt        j                  t        |d         }t        |dd      }t        |dd      }||y|j                         \  }}	|j                         \  }
}t	        |	      dk7  st	        |      dk7  ryd	fd}t	        |j                        dk7  st	        |j                        dk7  ryt	        |j                        t	        |      k7  s!t	        |j                        t	        |
      k7  ry ||      }d}|j                  d   } ||      j                         D ]  \  }}|j                  |      }|s|D ]  }|j                  |      dk(  r|D ]  |j                  d   }j                  |      }|k(  }t        fd|j                  D              }||k(  rM|r| j                   j"                  n| j                   j$                  }|	||k7  r   y|}   |S c c}w )
z>Use LoopBody iter/reduce vars to disambiguate equal-size axes.r   r  r   Nr  c                    t        t              }| j                  j                  j                  d      D ]D  }|j
                  ||j
                     j                  | j                  |j                            F |S Nr   )	r   r   r  getr  r  r   r   r  )bodyresultentryr  s      rs   load_exprs_by_namezLNestedReduction._get_grouped_axis_from_loop_body.<locals>.load_exprs_by_name  su    2=d2CF**../C/CRH $$05,,-44++E,<,<=
 Mru   r  c              3  F   K   | ]  }j                  |      k(    y wr~   )coeff)r   vargrouped_coeffouter_read_exprs     rs   r   zCNestedReduction._get_grouped_axis_from_loop_body.<locals>.<genexpr>  s)      ' &)>)>s)CC's   !)r  rE   r   zdict[str, list[sympy.Expr]])r  r  r   r   r   typingcastr   getattrr\  reduce_vars	iter_varsitemsr  r  r  rM  rK  rL  )rx   r6  rQ  r^  outer_reductionsouter_reduction
outer_bodygrouped_bodyouter_iter_rangesouter_reduce_rangesgrouped_iter_rangesgrouped_reduce_rangesr  outer_reads_by_namer  grouped_reduce_varr   grouped_read_exprsouter_read_exprsgrouped_read_exprouter_reduce_varouter_reduce_coeffmatches_reductionmatches_iter	candidater  r  r  s                            @@@rs   r  z0NestedReduction._get_grouped_axis_from_loop_body  sj   
 	>)3)=)=)?U22??CTBUU A% ++m5Ea5HI_gt<
0'4@!51@1K1K1M..5F5Q5Q5S22"#q(C0E,F!,K	 z%%&!+s<3K3K/LPQ/Qz##$,=(>>#""C
$%C& 0<59)55b9(:<(H(N(N(P 	'$D$266t<#%7 '! 1 7 78J K A%'7 'O'1'='=b'A$)8)>)>?O)P&(59K(K%#& '#-#7#7' $L )L8 ->))COODUDU  )f	.A#&F!'	'		'2 u Vs
   H7H7)r6  rf   r7  rf   r   r}   r+  )r7  rf   rT  rR  r   z*tuple[SchedulerNode, sympy.Integer] | None)r6  rf   r7  rf   rh  rW  r   r}   )r6  rf   r7  rf   rh  rW  r   2list[tuple[SchedulerNode, PointwiseDomain]] | None)rh  rW  rq  Sequence[BaseSchedulerNode]r   r  )rh  rW  ri  z/Sequence[tuple[SchedulerNode, PointwiseDomain]]r   r}   )r^  r   ry  rH  rh  rW  r   r}   )r6  rf   r  rR  r  rR  r  rM  rc  r   r   r}   )r  r  r  rM  r   r   )rQ  r   r  rR  r  rR  rc  rR  r6  BaseSchedulerNode | Noner   GroupedAxis | None)r6  rf   rQ  r   r   r  )r   r   r   r   r  r  r0  r8  r;  r   r@  rC  EnumrH  rM  r   	dataclassrW  rd  rj  rf  rl  rg  r|  r  r  r  r  r  r   ru   rs   r2  r2    s:   " !"
%
5F
	
 
 
%
5F
	
 
 N N"$)) " dii  [$'3 3 (3 &4,&4>H&4	3&4 &4P X%X (X /	X
 
X X. "F%"F ("F /	"F
 
<"F "FH 5!.5! +5! 
<	5! 5!n 
.
 K
 
	
 
 II  I /	I
 
I I8 %W%%W  %W !	%W "%W %W 
%W %WN ,2, ", 
	, ,, _ _B  041(1  1 !	1
 1 -1 
1 1f @*@?L@	@ @ru   r2  c                      e Zd ZU ded<   ded<   ded<    ej
                  e      Zded	<    ej
                  e      Z	d
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ddZddZy)SchedulerBuffer	Scheduler	schedulerz	ir.Bufferr   r  defining_opr   list[NodeUser]usersrF   
mpi_bufferc                B    | j                   }|J |j                         S r~   )r  get_name)rr   ops     rs   defining_op_namez SchedulerBuffer.defining_op_name(  s#    ~~{{}ru   c                @    t        | j                  j                        S r~   )hashr   r   rq   s    rs   __hash__zSchedulerBuffer.__hash__-  s    DIINN##ru   c                v   t               }| j                         }|j                  | dt        | j                        j
                          |j                  | d| j                  j                          | j                         r-|j                  | dt        | j                                       | j                         r-|j                  | dt        | j                                       t        | j                        dk  r0|j                  | d| j                          |j                         S |j                  | d       |j                  d      5  | j                  D ]  }|j                  | d        	 d d d        |j                  d	       |j                         S # 1 sw Y   *xY w)
N: z
.layout = z.aliases = z.mutations = r   z	.users = z
.users = [,])rX   r  	writeliner   r   r   layoutget_aliasespformatget_mutationsr   r  indentgetrawvalue)rr   r  r   users       rs   	debug_strzSchedulerBuffer.debug_str0  s   !}}D6DO$<$<#=>?D6DII,<,<+=>?v[9I9I9K1L0MNOv]74;M;M;O3P2QRStzz?avYtzzl;< !!## vZ01q! 1 JJ 1D$$vQZ011 S!!!##	1 1s   &F//F8c                6    | j                   j                         S r~   r   r  rq   s    rs   r  zSchedulerBuffer.get_nameD      yy!!##ru   c                   | j                   J | j                   j                         sy | j                   j                         sL| j                   j                         s2t	        | j                   j                         t        j                        r4t        j                  j                  j                  | j                          y t        t        j                  d      r| j                         t        j                  j                  v rt        j                  j                  | j                            }|| j                   j"                  v r$| j                   j"                  |   j                   }n#| j                   j$                  |   j                   }t        j                  j                  j'                  || j                          y t        j                  j                  j                  | j                          y )Nargs)r   should_allocateget_inputs_that_alias_outputget_mutation_namesr   get_output_specr0   CommBufferLayoutra   r   wrapper_codecodegen_allocationhasattrkernelr  inplace_update_buffersr  name_to_donated_buffername_to_bufcodegen_inplace_reuse)rr   input_buffer_nameinput_buffers      rs   allocatezSchedulerBuffer.allocateG  sV   yy$$$yy((* II224yy++-$))335r7J7JKGG  33DII> AHHf%188#B#BB !" ? ? P DNN$I$II#~~DD% $   $~~99:KLQQGG  66		
 GG  33DII>ru   c                   | j                   J t        | j                   j                  t        j                        st        | j                         ry| j                  D ]  }t        |j                   t              s y yNFT)r   r   r  r0   rD   r\   r  
OutputNode)rr   uses     rs   can_freezSchedulerBuffer.can_freeh  sg    yy$$$dii&&6:SII;
 :: 	C#((J/	 ru   c                ,   i }|D ]o  }t        |j                        |v r>|j                  |t        |j                                 |t        |j                        <   X||t        |j                        <   q t        |j	                               | _        y r~   )idr   merger   r   r  )rr   r  r  r  s       rs   	set_userszSchedulerBuffer.set_userst  st    &( 	+C#((|v%'*yy3881E'Fr#((|$'*r#((|$		+
 &--/*
ru   c                R    | j                   J | j                   j                         S r~   )r   r  rq   s    rs   r  zSchedulerBuffer.get_aliases~  s%    yy$$$yy5577ru   c                R    | j                   J | j                   j                         S r~   )r   r  rq   s    rs   r  zSchedulerBuffer.get_mutations  %    yy$$$yy++--ru   c                R    | j                   j                         j                         S r~   )r   r  r   rq   s    rs   r   zSchedulerBuffer.get_device  s    yy((*5577ru   Nr   r,  r   r   r   Noner   r}   )r  r  r   r  r   zSequence[str]r   torch.device | None)r   r   r   r   r   r   r   r  rF   r  r  r  r  r  r	  r  r  r  r  r   r   ru   rs   r  r    sz    
O))-K--dCE>C.?k.?.?3/J+ 
$$($?B
+8.8ru   r  c                      e Zd ZU dZded<   y)SchedulerDonatedBufferNr  r  )r   r   r   r  r   r   ru   rs   r   r     s    ,0K)0ru   r   c                  .   e Zd ZU ded<   ded<   ded<   ded<   ded<   ded	<   ded
<   ded<   ded<   dZded<   ded<   ded<   dZded<   ded<   ded<   dZded<   dVdZdWd ZdXd!Z	dXd"Z
dXd#ZdYd$ZdXd%ZdZd&Z	 	 	 	 	 	 d[d'Zd\d(Zd]d)Zd^d*Zd_d+ZdZd,Zed`d-       Z	 	 	 	 	 	 dad.ZdZd/Zdbd0Zdbd1ZdZd2ZdZd3Z	 	 	 	 dcd4ZdXd5ZdXd6Zedbd7       Z edbd8       Z!ed^d9       Z"ed^d:       Z#ddd;Z$ded<Z%dfd=Z&dgd>Z'd^d?Z(d^d@Z)d^dAZ*d^dBZ+d^dCZ,d^dDZ-d^dEZ.d^dFZ/dhdGZ0d^dHZ1dZdIZ2	 di	 	 	 	 	 djdJZ3edkdK       Z4edkdL       Z5edkdM       Z6	 	 	 	 	 	 dldNZ7	 	 	 	 	 	 dmdOZ8edndP       Z9dodQZ:edodR       Z;dpdSZ<dqdTZ=e>	 	 	 	 drdU       Z?y)srf   r   r  z7tuple[torch.device, tuple[tuple[sympy.Expr, ...], ...]]r   
last_usager   min_input_distancemax_input_distance	min_order	max_orderrG   mpi_nodedict[str, str]mutation_renamesNir.Operation | Noner   list[SchedulerBuffer]outputsdict[str, SchedulerBuffer]outputs_by_namefloat | Noneoverride_estimated_runtimedependencies.ReadWritesr   OrderedSet[Dep]unmet_dependenciesFr}   writtenc                "    || _         d | _        y )Nc                     g S r~   r   )r  kwargss     rs   <lambda>z,BaseSchedulerNode.__init__.<locals>.<lambda>  s    B ru   )r  debug_device_str)rr   r  s     rs   __init__zBaseSchedulerNode.__init__  s    $-& 	ru   c                v   || _         t               | _        d| _        d| _        t        t
                  | _        d| _        |j                         D cg c]  }t        | j                  ||        c}| _        | j                  D ci c]  }|j                         | c}| _        i | _        y c c}w c c}w )Nr   F)r  r   r  )r   r   r  r#  r$  r,  r"  r4  get_outputsr  r  r,  r  r.  r)  )rr   r   outputr   s       rs   _init_from_nodez!BaseSchedulerNode._init_from_node  s    	#"#"#$
   **,
  .. 
 @D||L 3L !#
  Ms   B1	B6c                T    t        |       j                   d| j                         dS )Nz(name=)r   r   r  rq   s    rs   __repr__zBaseSchedulerNode.__repr__  s'    t*%%&fT]]_,?qAAru   c                   | j                         }t               }|j                  | dt        |       j                   dt        t        | dd            j                   d| dt        | j                  j                         d| dt        | j                         d| d	t        | j                  j                  | j                  z
         d| d
| j                   d| d| j                   d| d       |j                         5  | j                         D ]!  }|j                  |j                                # 	 ddd       |j!                  d       	 |j                  | j#                                |j+                         j-                         S # 1 sw Y   XxY w# t$        $ r t&        j)                  dd       Y Lw xY w)#Longer form printout for trace logsr  (r   N)

.writes = 
.unmet_dependencies = .met_dependencies = .min_input_distance = .max_input_distance = z.outputs = [
        r  Ignoring error in debug_str()Texc_info)r  rX   splicer   r   r  r  r   writesr3  r   r#  r$  r  r<  r  r  debug_str_extra	Exceptionlogwarningr  rstrip)rr   r   r   r   s       rs   r  zBaseSchedulerNode.debug_str  s   }}

bd		QtGD&$$?@IIJ Kj))0012 3WT%<%<=> ?74#3#3#9#9D<S<S#STU VT445 6T445 6 	
	
 ZZ\ 	,'') ,

3==?+,	, 	c	HJJt++-.  ''))	, 	,  	HKK7$KG	Hs   5FF F G Gc                     y)N r   rq   s    rs   rR  z!BaseSchedulerNode.debug_str_extra      ru   c                $    | j                  |       S r~   )r9  rq   s    rs   _debug_str_for_devicez'BaseSchedulerNode._debug_str_for_device  s    $$T**ru   c                   t        | j                  dd       }d}t        |t        j                  j
                  j                        r'd|j                  |j                         gdd      z   }nct        |t        j                  j
                  j                        r5d|j                  |j                         |j                         gdd      z   }|  | S )Nr   rX  , F)shorten	multiline)r  r   r   torch	_inductorr0   r   
str_helperget_size	Reductionget_reduction_sizer   )rr   
maybe_datadata_strs      rs   debug_str_shortz!BaseSchedulerNode.debug_str_short  s    TYY5
j%//"4"4">">?j33$$&'% 4  H 
EOO$6$6$@$@Aj33..0*2O2O2QR 4  H
 z""ru   c                p    t         j                  d| | j                  | j                  j                         y )Nz(%s: unmet_dependencies = %s, writes = %s)rT  infor3  r   rQ  rq   s    rs   log_detailszBaseSchedulerNode.log_details  s,    6####		
ru   c                     yNFr   )rr   self_dep	other_deps      rs   reorder_loops_by_dep_pairz+BaseSchedulerNode.reorder_loops_by_dep_pair       ru   c                    d | j                   j                         D        D ci c]  }||v r|||    c}| _        | j                  | j                   j	                  | j                               y c c}w )Nc              3  4   K   | ]  }|j                     y wr~   r   r   r   s     rs   r   z9BaseSchedulerNode.update_mutated_names.<locals>.<genexpr>	  s     QcQ   )r   reads_and_writesr)  set_read_writesrename)rr   renamesr   s      rs   update_mutated_namesz&BaseSchedulerNode.update_mutated_names  sp     RT-=-=-N-N-PQ!
w '$-!

 	T--44T5J5JKL!
s   A2c                X    | j                  | j                  j                  |             y r~   )rx  r   	with_readrr   r   s     rs   add_fake_depzBaseSchedulerNode.add_fake_dep  s!    T--77<=ru   c                B    t        d | j                         D              S )Nc              3  `   K   | ]&  }|j                         xs |j                          ( y wr~   )r  r  r   r   s     rs   r   z=BaseSchedulerNode.has_aliasing_or_mutation.<locals>.<genexpr>  s-      
9<COO4!2!2!44
s   ,.)r  r<  rq   s    rs   has_aliasing_or_mutationz*BaseSchedulerNode.has_aliasing_or_mutation  s%     
@D@P@P@R
 
 	
ru   c                    || _         | j                   j                  | _        | j                          | j	                          y r~   )r   r   r3  "clear_read_writes_dependent_caches
prune_deps)rr   rws     rs   rx  z!BaseSchedulerNode.set_read_writes  s4    "&"2"2"8"8//1ru   c                :    | j                   j                  |        y r~   )r  clear_cacherq   s    rs   r  z4BaseSchedulerNode.clear_read_writes_dependent_caches  s    ""..t4ru   c                L    ddl m} t        | t        t        f      sy  ||       S )Nr   )_analyze_memory_coalescing)tiling_utilsr  r   r   r   )rr   r  s     rs   r  z'BaseSchedulerNode.get_coalesce_analysis  s#    <$0B CD)$//ru   c                b    | j                         }t        fd|D              }||z
  | _        y )Nc              3  B   K   | ]  }j                  ||        y wr~   )r  )r   kmutation_real_names     rs   r   z3BaseSchedulerNode.set_last_usage.<locals>.<genexpr>+  s     !U1"4"8"8A">!U   )used_or_aliased_buffer_namesr   r"  )rr   future_used_buffersr  used_bufferss     ` rs   set_last_usagez BaseSchedulerNode.set_last_usage'  s0     88:!!U!UU&)<<ru   c                F    | j                   D ]  }|j                           y r~   )r,  r	  )rr   r   s     rs   mark_runzBaseSchedulerNode.mark_run.  s    << 	CLLN	ru   c                    t        d t        j                  | j                  j                  | j                  j
                        D              S )Nc              3  4   K   | ]  }|j                     y wr~   rt  ru  s     rs   r   z6BaseSchedulerNode.used_buffer_names.<locals>.<genexpr>3  s      
 HH
rv  )r   	itertoolschainr   r   rQ  rq   s    rs   r   z#BaseSchedulerNode.used_buffer_names2  s?     
 t'7'7'='=t?O?O?V?VW
 
 	
ru   c                \   t               t        j                  | j                  j                  | j                  j
                        D cg c]*  }t        |t              r|j                  s|j                  , }}t        |      dkD  r|j                         }j                  |       t        j                  j                  j!                  |      rC|j#                  fdt        j                  j                  |   j%                         D               t        |      dkD  rS c c}w )z
        Returns buffer names used by this node, including aliases.

        Note: is_fake WeakDeps are excluded since they are purely for ordering
        and should not affect buffer lifetime.
        r   c              3  *   K   | ]
  }|vr|  y wr~   r   )r   alias
used_namess     rs   r   zABaseSchedulerNode.used_or_aliased_buffer_names.<locals>.<genexpr>J  s#       J.	 s   )r   r  r  r   r   rQ  r   r;   is_faker   r   popaddra   r   name_to_bufferr  extendr  )rr   r   depsr  s      @rs   r  z.BaseSchedulerNode.used_or_aliased_buffer_names8  s     '1l
 !t'7'7'='=t?O?O?V?VW
sG, HH
 

 $i!m((*CNN3ww%%))#. !"!7!7"224	 	 $i!m !
s   /D)c                L     t         fd j                  D               _        y )Nc              3  f   K   | ](  }|j                   j                  j                  vr| * y wr~   )r   r  available_buffer_namesr   r   rr   s     rs   r   z/BaseSchedulerNode.prune_deps.<locals>.<genexpr>T  s/      -
xxt~~DDD -
s   .1r   r3  rq   s   `rs   r  zBaseSchedulerNode.prune_depsS  s#    ", -
..-
 #
ru   c                     d fdt        fd j                  j                  D              } j                   j                  j	                  |             y )Nc                   t        | t              sy| j                  j                  j                  vryj                  j                  | j                     j                         }|t        j                  j                  v S rm  )	r   r;   r   r  r  r  ra   r   removed_operations)r   op_namerr   s     rs   should_prunez7BaseSchedulerNode.prune_weak_deps.<locals>.should_prune\  s_    c7+xxt~~999nn00:KKMGagg8888ru   c              3  4   K   | ]  } |      s|  y wr~   r   r   r   r  s     rs   r   z4BaseSchedulerNode.prune_weak_deps.<locals>.<genexpr>e  s      
\#5FC
   r   r8   r   r}   )r   r   r   rx  remove_reads)rr   	to_remover  s   ` @rs   prune_weak_depsz!BaseSchedulerNode.prune_weak_depsZ  sN    	9  
++11
 
	 	T--::9EFru   c                F    t        | || j                  j                         y r~   )_prune_redundant_depsr  r  )rr   name_to_fused_nodes     rs   prune_redundant_depsz&BaseSchedulerNode.prune_redundant_depsj  s     	d$68R8RSru   c                R    | j                   J | j                   j                         S r~   )r   get_operation_namerq   s    rs   r  zBaseSchedulerNode.get_nameo  r  ru   c                "    | j                         S r~   r  rq   s    rs   get_first_namez BaseSchedulerNode.get_first_names  s    }}ru   c                B    t        d | j                         D              S )Nc              3  <   K   | ]  }|j                           y wr~   r  r   r   s     rs   r   z8BaseSchedulerNode.get_operation_names.<locals>.<genexpr>x  s     Gd$--/G   )r   r   rq   s    rs   r  z%BaseSchedulerNode.get_operation_namesv  s    Gdnn6FGGGru   c                :    t        d | j                  D              S )Nc              3  <   K   | ]  }|j                           y wr~   r  r   r   s     rs   r   z5BaseSchedulerNode.get_buffer_names.<locals>.<genexpr>|  s     AS#,,.Ar  )r   r,  rq   s    rs   get_buffer_namesz"BaseSchedulerNode.get_buffer_namesz  s    ADLLAAAru   c                B    t        d | j                         D              S )Nc              3  Z   K   | ]#  }t        |t              xr t        |d        % yw)T)disallow_fp32_opsNr   r   r2   r   ns     rs   r   zABaseSchedulerNode.can_codegen_in_low_precision.<locals>.<genexpr>  s7      
  q-( G+AFG
   )+r   r   rq   s    rs   can_codegen_in_low_precisionz.BaseSchedulerNode.can_codegen_in_low_precision~  s%     
 ^^%
 
 	
ru   c                B    t        d | j                         D              S )Nc              3  V   K   | ]!  }t        |t              xr t        |       # y wr~   r  r  s     rs   r   z@BaseSchedulerNode.can_codegen_without_upcasts.<locals>.<genexpr>  s-      
 q-(K-H-KK
s   ')r  rq   s    rs   r2   z-BaseSchedulerNode.can_codegen_without_upcasts  s#     
^^%
 
 	
ru   c                    | gS r~   r   rq   s    rs   r   zBaseSchedulerNode.get_nodes  s	    vru   c                    | j                   S r~   )r,  rq   s    rs   r<  zBaseSchedulerNode.get_outputs  s    ||ru   c                     | j                   |   S r~   )r.  )rr   buf_names     rs   
get_outputzBaseSchedulerNode.get_output  s    ##H--ru   c                R    | j                   J | j                   j                         S r~   )r   r   rq   s    rs   r   zBaseSchedulerNode.get_device  s%    yy$$$yy##%%ru   c                L    | j                         }|d uxr |j                  dk(  S Ncpu)r   r   rr   devices     rs   is_cpuzBaseSchedulerNode.is_cpu  s'    "T!:fkkU&::ru   c                X    | j                         }|d uxr t        |j                        S r~   )r   r[   r   r  s     rs   r[   zBaseSchedulerNode.is_gpu  s'    "T!9fV[[&99ru   c                     yrm  r   rq   s    rs   r   zBaseSchedulerNode.is_reduction      ru   c                     yrm  r   rq   s    rs   is_native_matmulz"BaseSchedulerNode.is_native_matmul  r  ru   c                     yrm  r   rq   s    rs   is_split_scanzBaseSchedulerNode.is_split_scan  r  ru   c                     yrm  r   rq   s    rs   is_templatezBaseSchedulerNode.is_template  r  ru   c                     yrm  r   rq   s    rs   	is_externzBaseSchedulerNode.is_extern  r  ru   c                     yrm  r   rq   s    rs   
is_foreachzBaseSchedulerNode.is_foreach  r  ru   c                     yrm  r   rr   read_deps     rs   can_inplacezBaseSchedulerNode.can_inplace  r  ru   c                     yrm  r   rq   s    rs   has_side_effectsz"BaseSchedulerNode.has_side_effects  r  ru   c           	         ddl m} t         t              rt        j
                  rt        j                  j                   j                         t        j                        r{t        t        j                  t        j                  j                  j                   j"                        rt%        t        j                  dd      t'        t        j                  d      sy j(                  t        j                  j*                  z   j,                  j.                  z  }d fd} j1                         D ]Q  }|j2                  }|J |j5                         rr|j7                         sb|j9                         sR|j;                         t        j                  j<                  v s(t        |j?                         t@        jB                        r jD                  jF                  D ]  }|jH                   j,                  jJ                  v r$ j,                  jJ                  |jH                     }n/ j,                  jL                  jO                  |jH                        }|s|t        j                  jP                  jS                  |       st        |jT                  tV              r|jX                  J |jX                  D cg c]   }|j2                  j;                         |vr|" }	} j,                  j[                  |jH                         }
|
r/t]        |	      dk(  s?|	d   j^                  sP|	d   j2                   u sc|j2                  qt        |j2                  j?                         t@        j`                  t@        jb                  t@        jd                  t@        jB                  f      r|jT                  rft        |jT                  j2                  t@        jf                  t@        jh                  f      r(t]        |j2                  j7                               dkD  rE ||j2                  |j2                        sd ||      snt        j                  jj                  jm                  |j;                         |j;                                t        t        j                  t        j                  j                  j                   j"                        rnt        j                  jn                  jq                  |j;                                t        j                  jn                  jq                  |j;                                |j;                         t        j                  jr                  |j;                         <    Q T yc c}w )	z~
        Decide if there should be inplace updates for the node
        and record the decision in the active kernel.
        r   )can_match_buffer_size	mutationsNr  c                   | j                   j                        }| j                         t               }| j                  D ]  }|j
                  }t        |t              s |j                         | j                   j                  vs| j                   j                  |      |urd|fd|j                  j                         D        z  }t        |      dkD  s y y)Nc              3  @   K   | ]  }|j                   k(  r|  y wr~   rt  )r   or  s     rs   r   z^BaseSchedulerNode.decide_inplace_update.<locals>.single_index_in_fused_node.<locals>.<genexpr>  s%      vv)    r   FT)r  get_fused_noder  r   r  r   r   rf   r  r  r   rw  r   )buf_to_be_inplaced
fused_noder  r  	user_noder  rr   s        @rs   single_index_in_fused_nodezKBaseSchedulerNode.decide_inplace_update.<locals>.single_index_in_fused_node  s    
 ,55DDTJJ)224H %/LD*00 ! II	!)->? ,,.-77JJK)33BB9M%&  &22CCE 
 t9q= '!* ru   r   )r  r  r   r}   ):codegen.wrapperr  r   r   r-   inplace_buffersra   r   has_featurer   r3   INPLACE_BUFFERSr  r`  ra  codegensimdr  r  r  r  r  r  completed_operationsr<  r   r  r  r  r  removed_buffersr  r0   r  r   r   r   r  r  r  r  	can_reuser  NopKernelSchedulerNoder  has_cross_stream_hazardr   r  rD   rC   MutationLayoutSHOULDREMOVEFallbackKernelrB   r  make_inplacer  r  r  )rr   r  inconsequential_nodesr  r   buf_noderead	input_bufr  remaining_usesr  s   `          rs   decide_inplace_updatez'BaseSchedulerNode.decide_inplace_update  s   
 	; t]+&&##DOO$5~7U7UVqxx)@)@)E)E)P)PQ188[$7C &) NNgg(()nn112 	 	D ##% L	CxxH''',,.88:..0<<>QWW%<%<< h668":M:MN((.. >99 E EE $ E Edii PI $ : : > >tyy II ,,66y$G&y'<'<>TU$??666 "+&66??,4II &N &
 /3nn.T.T		4/+
 4/14*1-99*1-22d:%NN6 *%NN::< " " 4 4 " = = " 3 3	! &11 * ) 5 5 : :!#!2!2BNN C! !$INN$O$O$Q RUV V1)..#((K6yA
 2293E3E3GX%HHeoo&=&=&B&B&M&M HH..2293E3E3GHHH..223<<>B &..0 77G }>L	6&s   ?%V
c                R   t         j                  sy |r| j                  ry | j                  J | j                  j	                         }g }|D ]0  }|j
                  dk(  r|j                  d       |j                  d       d|j
                   d|j                   }d|j                  v r|d|j                  d    z   }|j                  |       d|j                  v s|j                  d    }|j                  d	d
      d   }|j                  d|j                  dd      j                  dd      j                  dd      j                  dd      z          |j                  d       |j                  d       3 t        |      dk(  ry |j                  |       d| _        y )Nr=  rX  z#pragma CMT ORIGIN:z#pragma CMT  seq_nrz seq_nr:stack_trace|r   )maxsplitr  {z{{}z}}rH  \z\\z#pragma CMT END ORIGINr   T)r-   comment_originr4  r   get_originsr  r   targetmetarsplitreplacer   
writelines)	rr   buffer	only_onceorigins	out_linesr  op_info_strr  stack_trace_last_lines	            rs   codegen_originating_infoz*BaseSchedulerNode.codegen_originating_infoE  s    $$yy$$$))'')	 	%AttxR 23(az:K166!)hqvvh7G6H,II[)&!"!6 7(3(:(:3(:(KB(O%  "+33C>WS$'WT4(Wf	   !9:  $3	%6 y>Q 	)$ru   c                (    | j                  dd      S )NTinclude_readsinclude_writes!get_read_write_buffers_sizes_implrq   s    rs   get_read_write_buffers_sizesz.BaseSchedulerNode.get_read_write_buffers_sizest  s    55t 6 
 	
ru   c                (    | j                  dd      S )NTFr*  r-  rq   s    rs   get_read_buffer_sizesz'BaseSchedulerNode.get_read_buffer_sizesz  s    55u 6 
 	
ru   c                (    | j                  dd      S )NFTr*  r-  rq   s    rs   get_write_buffer_sizesz(BaseSchedulerNode.get_write_buffer_sizes  s    55 6 
 	
ru   c                Z    t        | j                  ||      j                         d      S )Nr*  r   )start)r   get_read_write_buffer_accessesr   )rr   r+  r,  s      rs   r.  z3BaseSchedulerNode.get_read_write_buffers_sizes_impl  s3     //+N 0 fh	
 	
ru   c                    t         t              ri S t         t              rt         j                  t              ri S t         t              r`t         j                  t
        j                        r< j                  j                  t        j                  j                  j                  u ri S ddt         t              r@ t         j                         d         t         j                         d         z        nt        d      t!        j"                  t$              }|r9 j&                  j(                  D ]   }||j*                     j-                  |       " |r9 j&                  j.                  D ]   }||j*                     j-                  |       " |r&t1        d  j&                  j(                  D              n	t1               }|r&t1        d  j&                  j.                  D              n	t1               }d fdt         t2              rt1         fd|D              }||z
  }||z
  }i }||z  D ]  }	t5        fd	||	   D              |	t6        j8                  j:                  v rt6        j8                  j:                  |	   }
n;|	t6        j8                  j<                  v rt6        j8                  j<                  |	   }
n	 	 	 	 d fd
 |
      }|	|vr|||	<   ||	xx   |z  cc<    |S )az  
        Counting the number of bytes accessed for a kernel is
        surprisingly tricky. In particular, there is a differentiation
        between 'theoretical' memory accesses and practical memory
        accesses. For example, a layernorm kernel may actually access an
        input 3 times, but in theory, it only needs to access its input
        once (and may be optimized to do so through say, persistent
        reductions)

        Another example is that even though a buffer is passed in, we may
        not access the entire buffer. This may occur if we are accessing
        a slice of the buffer. Another tricky case is for indirect
        indexing, where the amount of bytes accessed depends on the
        values of the input.

        What this function aims to compute is the memory accesses for
        worst-case inputs, best-case optimization. What this means is
        that for each buffer we compute the amount of potential accesses in two ways and take the minimum.

        1. Numel in ranges multiplied by number of deps the buffer has
        2. The buffer size

        Returns memory accesses per buffer.
        c                X    t         j                  j                  j                  | d      S )Nr   r   )ra   r   r   r   )ss    rs   try_size_hintzGBaseSchedulerNode.get_read_write_buffer_accesses.<locals>.try_size_hint  s"    77##55a!5DDru   r   r       eAc              3  4   K   | ]  }|j                     y wr~   rt  ru  s     rs   r   zCBaseSchedulerNode.get_read_write_buffer_accesses.<locals>.<genexpr>  s     BCsxxBrv  c              3  4   K   | ]  }|j                     y wr~   rt  ru  s     rs   r   zCBaseSchedulerNode.get_read_write_buffer_accesses.<locals>.<genexpr>       CCsxxCrv  c                    j                   j                  |    j                  }t        d |D              }t	        |t        |      z
        dkD  S )Nc              3  4   K   | ]  }|j                     y wr~   r   )r   r  s     rs   r   z\BaseSchedulerNode.get_read_write_buffer_accesses.<locals>.is_materialized.<locals>.<genexpr>  s     !>$))!>rv  r   )r  r  r  r   r   )r   r   r  buf_usesrr   s       rs   is_materializedzIBaseSchedulerNode.get_read_write_buffer_accesses.<locals>.is_materialized  sG    NN..s399E!!>!>>Hx*V"44599ru   c              3  J   K   | ]  } |j                         r|  y wr~   r   )r   r   rB  rr   s     rs   r   zCBaseSchedulerNode.get_read_write_buffer_accesses.<locals>.<genexpr>  s#      )_S$++-N)s   ##c              3  "   K   | ]  }  y wr~   r   )r   r   
node_numels     rs   r   zCBaseSchedulerNode.get_read_write_buffer_accesses.<locals>.<genexpr>  s     $RCZ$Rs   c                B   | syt        | t        j                        r| j                         S t        | j                  t
              r͉j                  j                  | j                            j                  }d}|D ]  }t        |j                  t              rt        |j                  t              sJ t        |j                  j                  t              r5|j                  j                         D ]  }| |j                        z  }  y |S t        | j                  t        j                        r"t!        fd| j#                         D              S  	t%        | j'                                     }t)        | j+                               t-        |      z  S )Nr   c              3  h   K   | ])  } t         j                  j                  |             + y wr~   )ra   r   
get_buffer)r   mut_nameget_buf_bytess     rs   r   zZBaseSchedulerNode.get_read_write_buffer_accesses.<locals>.get_buf_bytes.<locals>.<genexpr>  s-      $ &agg&8&8&BCs   /2)r   r0   TorchBindObjectrK  r  rC   r  r  r  r  r   r  rf   rB   r<  rD   r   r  r_   rc  rT   	get_dtyperZ  )
r   r  totr  	sched_buf	buf_elemsbuf_accessed_elemsrK  rr   r:  s
         rs   rK  zGBaseSchedulerNode.get_read_write_buffer_accesses.<locals>.get_buf_bytes  sT    c2#5#56,,..

,=> !NN66s||~FLLEC % %%dii<$)$))5FGGG%diinnkB-1YY-B-B-D E	 #}Y^^'D DE $%% J

BMM: (+(>(>(@  
 !.mCLLN.K LI)#--/:S*I>  ru   )r9  rR  r   r   )r   r,  r   r  r   r}   )r   z4ir.Buffer | ir.TensorBox | ir.TorchBindObject | Noner   r   )r   r  ExternKernelSchedulerNoder   rB   r0   r
  op_overloadr`  _prims	rng_primsgraphsafe_run_with_rng_stater   r_   r\  r   collectionsr   r   r   r   r   r   rQ  r   r   r   ra   r   r  graph_inputs)rr   r+  r,  buf_accessesr   r   rQ  r  buf_byte_accessesr  r   	buf_bytesrQ  rK  rB  rF  r:  s   `           @@@@@rs   r6  z0BaseSchedulerNode.get_read_write_buffer_accesses  s   6 d23Id56:II{<
 It67499b&7&78		%%||%%BBC I	E dM*&doo/23 1! 456J
 SJ"..t4''-- 3SXX&--c23 ''.. 3SXX&--c23
  B4+;+;+A+ABB 	  C4+;+;+B+BCC 		:
 d./( )%) O o-FO+E,. 3	9H!$$R<;Q$R!R177111gg,,X6QWW111gg**84#I## #J &c*I00.7!(+!(+y8+g3	9j ! ru   c                T   | j                   y | j                   j                         }|y t        |      }|y t        |t        j
                        r|j                   j                  }t        j                  j                  j                  |d      }t        d   dxx   |z  cc<   |S )Nr   r   inductor
flop_count)r   get_origin_noder>   r   r`  SymIntexprra   r   r   r   r   )rr   fx_nodeflopsresolved_flopss       rs   estimate_flopsz BaseSchedulerNode.estimate_flops  s    99))++-?w'=eU\\*JJOOE));;EA;N\*n<*ru   c                R    | j                   | j                   S | j                         S r~   )r0  _get_estimated_runtimerq   s    rs   get_estimated_runtimez'BaseSchedulerNode.get_estimated_runtime1  s)    **6222**,,ru   c                   | j                         d   j                         d   }|j                  j                         }t	        t        |            syt        | j                        rt        | j                  t        j                        sJ 	 t        j                  rst        |       }t               }|j                  |      }|t        |t              sJ |S t!        |       }|t#        | j                        }|j%                  ||       |S t#        | j                        S t/        | j                        ryt1        |       }||S |j                  j3                         }		 t5               }
t7        |	      dz  }|
dk  rt9        d|
       |dk  rt9        d|       	 | j=                         }|dk(  s|| j?                         |
z  }|dz  }|S d}| j?                         }|dn|}||z  |z  d	z  }||
z  }tA        ||      }|dz  }|S # t&        $ r}t(        j+                  |       Y d}~yd}~wt,        $ r}t(        j+                  |       Y d}~yd}~ww xY w# t:        $ r Y yw xY w)
zC
        Returns estimated op runtime in milliseconds (ms)
        r   Nvaluel    J)z-gpu_memory_bandwidth cannot be <= 0, but got z"gpu_flops cannot be <= 0, but got g    .Ag      ?r;  )!r   r<  r   r  r[   r@   rY   r   r0   IRNoder.   ,runtime_estimations_use_nccl_lib_estimations)get_estimate_runtime_cache_key_from_snodeget_estimate_runtime_cachelookupfloatr7   r6   	set_value
ValueErrorrT  rj  	TypeErrorr^    maybe_estimate_runtime_benchmarkmaybe_get_dtyperU   rS   AssertionErrorrS  re  r/  rY  )rr   r   r  	cache_keycache	cache_valmsr&  retdtypegpu_memory_bandwidth	gpu_flops	flops_estnsfactorcounted_bytescompute_timetransfer_times                     rs   rg  z(BaseSchedulerNode._get_estimated_runtime7  sw   
 nnq!--/2))+of-. #dii333LL I$ OI68E %Y 7I ,))U;;;((HNBz=diiHOOIRO8I7		BB TYY
 .t4?J((*	#4#6 )%069I $q($CDXCYZ  A~$'I)%UVV 
 '')	>Y.2247KKBcBI 99;*2*Y6#=%(<< }-#X	o    :  		sC   AH 6H H (>I$ 	I!H66I!II!$	I0/I0c                     y r~   r   rq   s    rs   get_template_nodez#BaseSchedulerNode.get_template_node      ru   c                .    | j                         }|J |S r~   r  )rr   templates     rs   get_template_node_or_throwz,BaseSchedulerNode.get_template_node_or_throw  s!    ))+###ru   c                f    t        d t        |       D              }| d| }| |   }| |dz   d }|||fS )zQ
        For the list of nodes, get the prologue, template, and epilogue
        c              3  H   K   | ]  \  }}|j                         s|  y wr~   r  )r   ir  s      rs   r   zCBaseSchedulerNode.get_prologue_template_epilogue.<locals>.<genexpr>  s     PDAqaPs   ""Nr   )next	enumerate)rq  template_indexprologuetemplate_nodeepilogues        rs   get_prologue_template_epiloguez0BaseSchedulerNode.get_prologue_template_epilogue  sN     PIe,<PP.)n-!+-.00ru   )r  r  r   r  )r   ir.Operationr   r  r  )r   r-  r  rn  r9   ro  r9   r   r}   rz  r(  r   r  )r   r8   r   r  r  )r  r1  r   r  )r   zCoalesceVarAnalysis | Noner  r   r  r(  r   r  r   r   r  dict[str, BaseSchedulerNode]r   r  r   r  )r   zSequence[SchedulerBuffer])r  r,  r   r  r  r  zdependencies.Depr   r}   T)r"  rX   r#  r}   r   r  r  )r+  r}   r,  r}   r   r   )r+  r}   r,  r}   r   zdict[str, int]r   
int | Noner   rq  r   zir.TemplateBuffer | None)r   zir.TemplateBuffer)rq  list[BaseSchedulerNode]r   zJtuple[list[BaseSchedulerNode], BaseSchedulerNode, list[BaseSchedulerNode]])@r   r   r   r   r   r0  r4  r:  r>  rB  r  rR  r[  rh  rk  rp  r{  r  r  rx  r  rO   r  r  r  r   r  r  r  r  r  r  r  r  r  r2   r   r<  r  r   r  r[   r   r  r  r  r  r  r  r  r  r(  r/  r1  r3  r.  r6  re  rh  rg  r  r  r0  r  r   ru   rs   rf   rf     s   BB NN''$$ $D
$""///33((''GT
#4B*6+#
!.7	
M>

5 0 0=#2=HV=	=
6
G T">T	T
. H H B B 
 
 
 
.&;:IX 9=-$-15-	-^ 
 

 
 

 
 


!
37
	
L!!L!37L!	L!\  $- U Un
 1&1	S1 1ru   c                 R    t         j                  j                  j                         S r~   )r`  ra  	codecache
LocalCacher   ru   rs   ro  ro    s    ??$$//11ru   c                   t        | j                  dd      }| j                  j                  }| j                  j                  g || j                  j                  | j                  j
                        }| j                  j
                  }t        j                  ||f      \  }}ddt        |ft        fd|D              z         }|S )Npython_kernel_namerX  c                    t        | t        j                        xr+ t        | t        j                  t        j                  f       S r~   )r   r0   rl  GeneratorStateOpaqueObjectStater  s    rs   _is_tensor_irz@get_estimate_runtime_cache_key_from_snode.<locals>._is_tensor_ir  s<    !RYY' 

!!2#7#781
 -
 	
ru   c              3  d   K   | ]'  } |      rt        |j                               nd  ) y wr~   )r   rc  )r   ar  s     rs   r   z<get_estimate_runtime_cache_key_from_snode.<locals>.<genexpr>  s(     Ua}Q'7ajjl#TAUs   -0r  )
r  r   inputsfill_non_provided_argsconstant_argsr7  pytreetree_flattenr,  r   )snoder  r  r7  	flat_argsflat_args_pytree_specrx  r  s          @rs   rn  rn    s     -A2F::D::,,*$*))*

D ZZF'-':':D&>'J$I$

 	
U9U
U	VI ru   c                   t        | t              sy t        | j                  t        j                        sy | j                  j
                  }t        |t        j                  j                        sy |j                  }ddl
m} ||vry |S )Nr   )flop_registry)r   rR  r   r0   ExternKernelrS  r`  _ops
OpOverloadoverloadpackettorch.utils.flop_counterr  )r  rS  r  r  s       rs   _get_benchmarkable_extern_fnr    sk     e67ejj"//2**((Kk5::#8#89		#	#B6	Iru   c                Z    d }d }t         j                  rt               }|y |} fd}ny t               }t	               }|j                  |      }|t        |t              sJ |S ddlm	  |       \  }}ddl
m}	 |	j                  |||ddd      }
|j                  ||
	       |
S )
Nc                             S r~   r   )r  snode_args_kwargss   rs   r8  z2maybe_estimate_runtime_benchmark.<locals>.<lambda>  s    !25!9 ru   r   )r  r   )benchmarker   
   )memory_warmup_itersbenchmark_itersmax_benchmark_durationrj  )r-   !runtime_estimations_mms_benchmarkr  rn  ro  rp  r   rq  utilsr  $torch._inductor.runtime.benchmarkingr  	benchmarkrr  )r  bench_fnargs_kwargs_fnmm_fnrx  ry  rz  r  r7  r  r{  r  s   `          @rs   ru  ru    s    HN//,U3=99%@I&(EY'I)U+++(!#LD&@			! 
 
B 
OOIRO(Iru   c                  N    e Zd ZU ded<   ded<   ded<   ded<   ddZddZdd	Zy
)	WhyNoFuser,  name1name2reasontuple[Any, ...]r  c                X    |j                         | _        |j                         | _        y r~   )r  r  r  rr   r   r   s      rs   r:  zWhyNoFuse.__init__  s    ^^%
^^%
ru   c                J    || _         || _        t        j                  |        y r~   )r  r  
fusion_logdebug)rr   r  r  s      rs   __call__zWhyNoFuse.__call__  s    	ru   c                p    d| j                    d| j                   d| j                  | j                  z  z   S )Nzcannot fuse z with r  )r  r  r  r  rq   s    rs   __str__zWhyNoFuse.__str__  s6    djj\

|2>KK$))#
 	
ru   Nr   rf   r   rf   r   r  )r  r,  r  r   r   r  r  )r   r   r   r   r:  r  r  r   ru   rs   r  r    s&    JJK
&

ru   r  c                    t        | t        t        f      rt        | t              } t        j                  | d      }d|v rdt        j                  |d       S |S )Nkey   )r  rH      )	r   r   setsortedr,  pprintr  textwrapr  )objr  s     rs   r  r    sR    #
C()Sc"^^C*Fv~HOOFG4566Mru   c                  0    e Zd ZddZddZddZd	dZeZy)
r  c                &    t        |g      | _        y r~   r  r~  s     rs   r:  zOutputNode.__init__  s    ",cU"3ru   c                     yrm  r   rq   s    rs   r   zOutputNode.is_reduction!  r  ru   c                     yr  r   rq   s    rs   r  z'OutputNode.get_inputs_that_alias_output$  rY  ru   c                     y)NOUTPUTr   rq   s    rs   r  zOutputNode.get_name'  s    ru   N)r   r:   r   r  r  r  r  )r   r   r   r:  r   r  r  rB  r   ru   rs   r  r    s    4 Hru   r  c                    t        j                          j                  D ]N  }t        |t              r|j
                     j                         }|   j                         xx   dz  cc<   P d fdt        fd j                  D              }|r? j                  |z
   _         j                   j                  j                  |             yy)am  
    Prunes weakdeps intended for mutation ordering
    on an upstream fused node if after fusion there is another dependency
    on the fused upstream node, making the weakdep redundant

    In essence this enforces an ordering on fusions. As fusions occur, weakdeps will
    be incrementally removed, enabling other fusions, ensuring they are fused in order.
    r   c                    t        | t              rf| j                     j                         }|   j	                            dkD  xr  j
                  j                  | |         }|   k(  }|xs |S y)Nr   F)r   r;   r   r  r  r  fusable_weak_dep)r   r  is_redundantis_self_depr  name_to_dep_countr  r   s       rs   r  z+_prune_redundant_deps.<locals>.should_pruneA  s    c7#!#((+<<>G,"7+446 nn55'0$  -W5=K.;.ru   c              3  4   K   | ]  } |      s|  y wr~   r   r  s     rs   r   z(_prune_redundant_deps.<locals>.<genexpr>Q  s      ,s2Cr  Nr  )rW  r   r3  r   r;   r   r  r  r   rx  r   r  )r   r  r  r   r  deps_to_pruner  r  s   ```   @@rs   r  r  -  s     '2&9&9&;&& K#w'!#((+<<>G09BBDEJEK
    .. M "&"9"9M"IT--::=IJ ru   c                  H     e Zd Zd fdZddZd	dZd	dZd
dZddZ xZ	S )rR  c                n   t         |   |       | j                  |       | j                  |j	                                t        |t        j                        r[|j                         rJt        j                  |j                  d   j                        }d}|j                         }|||ff| _        y y y Nr   r   )superr:  r>  rx  get_read_writesr   r0   UserDefinedTritonKernelcan_fuse_epiloguemathr   mutable_argsshapeget_device_or_errorr   )rr   r  r   numelr   r  	__class__s         rs   r:  z"ExternKernelSchedulerNode.__init__[  s    #T"T1134dB667D<R<R<TIId//2889EF--/F 5&/2DJ =U7ru   c                V    | j                          dt        | j                  dd        S )Nz.node.kernel = r  )r  r  r   rq   s    rs   rR  z)ExternKernelSchedulerNode.debug_str_extrag  s*    --/"/'$))EY[_2`1abbru   c                     yNTr   rq   s    rs   r  z#ExternKernelSchedulerNode.is_externj  r  ru   c                    | j                   J t        | j                   d      xr | j                   j                         S )Nr  )r   r  r  rq   s    rs   r  z*ExternKernelSchedulerNode.has_side_effectsm  s6    yy$$$tyy"45V$)):T:T:VVru   c                    t        | j                  t        j                        rU| j                  j	                         r;t        j                  | j                  j                  d   j                        }|gg fS g g fS r   )	r   r   r0   r  r   r  r   r  r  )rr   r  s     rs   r\  z$ExternKernelSchedulerNode.get_rangesq  s^    tyy""<"<=		++-IIdii44Q7==>EGR= Bxru   c                    t        | j                  t        j                        sJ | j                  j	                  |      S r~   )r   r   r0   r  r  )rr   wrappers     rs   r  z!ExternKernelSchedulerNode.codegenz  s/    $))R__555yy  ))ru   r  r  r   r  r   r  r  r  r   Sequence[Sequence[sympy.Expr]]r  r   r   r  )
r   r   r   r:  rR  r  r  r\  r  __classcell__r  s   @rs   rR  rR  Z  s#    
3cW*ru   rR  c                        e Zd Zd fdZ xZS )r  c                    t         |   |       | j                  |       | j                  |j	                                y r~   )r  r:  r>  rx  r  rr   r  r   r  s      rs   r:  zNopKernelSchedulerNode.__init__  s5    #T"T1134ru   r  )r   r   r   r:  r  r  s   @rs   r  r    s    5 5ru   r  c                      e Zd ZU dZded<   ded<   	 	 	 	 	 	 d& fdZ	 	 d'	 	 	 	 	 d(dZ	 	 d'	 	 	 	 	 d)dZ	 	 	 	 	 	 d*d	Zd+d
Z	d,dZ
d-dZd.dZd/dZd0dZd.dZd1dZd.dZ	 	 	 	 	 	 d2dZd.dZ	 	 	 	 	 	 d3dZd4dZd5dZd6dZd6dZd6dZd6dZd7dZd8dZ	 	 	 	 d9dZd:dZ	 d;	 	 	 d<d Ze d=d!       Z!e d=d"       Z"d>d#Z#e d?d$       Z$e d6 fd%       Z% xZ&S )@r   zu
    A SchedulerNode is a node for scheduling that encapsulates either
    a ComputedBuffer or a TemplateBuffer.
    z tuple[Sequence[sympy.Expr], ...]_sizesrE   r  c                t    t         |   |       d | _        | j                  |       | j	                          y r~   )r  r:  _loop_mutation_listenerr>  _compute_attrsr  s      rs   r:  zSchedulerNode.__init__  s4    
 	#OS$T"ru   c                   t        | j                  t        j                  t        j                  f      sJ | j                  j                  ||      \  | _        }|| _        | j                  j                         }| j                  j                  |      j                  }| || j                        f| _        t        j                   xs t        |j                          }t        | j                  t        j                        r,| j#                  | j                  j%                  |             y | j#                  t'        j$                  | j                  g| j                  d|i       y )Nextra_indexing_constraintsrecompute_sizes_body_func)	normalizer   )r   r   r0   r   TemplateBuffersimplify_and_reorderr  r  r  r  get_backendgroup_fnr   r-   loop_ordering_after_fusionr[   r   rx  extract_read_writesr/   )rr   r  r  r  r  r$  should_normalizes          rs   r  zSchedulerNode._compute_attrs  s7   
 $))b&7&79J9J%KLLL II::'A&? ; 
T 
..0>>--f5>>ht{{34
  &@@@ 
KKI
 E
 dii!2!23  		--8H-I   00JJ!%8Hru   c                   t        d | j                  j                  D              }| j                  ||       |rD| j	                  | j                  j                  |      j                  | j                               y y )Nc              3  N   K   | ]  }t        |t        t        f      s|  y wr~   r   r;   r:   ru  s     rs   r   z8SchedulerNode.recompute_size_and_body.<locals>.<genexpr>  #      0
ZgwEW5XC0
   %%r  )r   r   r   r  rx  r}  ry  r)  )rr   r  r  	fake_depss       rs   recompute_size_and_bodyz%SchedulerNode.recompute_size_and_body  s    
 &0 0
++110
 &
	 	'A&? 	 	
     **95<<T=R=RS ru   c                :   t        d | j                  j                  D              }| j                  t	        j
                  | j                  g| j                  d|ij                  |      j                  | j                               | j                  |       y )Nc              3  N   K   | ]  }t        |t        t        f      s|  y wr~   r*  ru  s     rs   r   z5SchedulerNode.refresh_dependencies.<locals>.<genexpr>  r+  r,  r   )r   r   r   rx  r/   r&  r  r  r}  ry  r)   clear_loop_body_dependent_caches)rr   r   need_clear_tiling_cacher-  s       rs   refresh_dependenciesz"SchedulerNode.refresh_dependencies  s    
 &0 0
++110
 &
	 	,,

![[4= Yy!VD))*	
 	--.EFru   c                    | j                          | j                  j                  |        |r!ddlm} |j
                  j                          y y )Nr   r  )r  pointwise_read_writesr  r  r  candidate_tilingscache_clear)rr   r2  r  s      rs   r1  z.SchedulerNode.clear_loop_body_dependent_caches  sA    //1""..t4"4 ,,88: #ru   c                t    | j                   | j                  | j                  | j                  | j                  fS )zSnapshot mutable state modified by loop transformations
        (apply_new_loop_order, apply_loop_reindexing). Must be kept
        in sync with those methods and restore_loop_state.)r  r  r   r   r3  rq   s    rs   snapshot_loop_statez!SchedulerNode.snapshot_loop_state  s5    
 JJKKJJ##
 	
ru   c                j    |\  | _         | _        | _        | _        | _        | j                  d       y)z'Restore state from snapshot_loop_state.T)r2  N)r  r  r   r   r3  r1  )rr   states     rs   restore_loop_statez SchedulerNode.restore_loop_state  s8     	
JKJ#--d-Kru   c                @    | j                   | j                  |        y y r~   )r  rq   s    rs   _before_loop_state_mutationz)SchedulerNode._before_loop_state_mutation	  s!    ''3((. 4ru   c                    | j                          | j                  j                  |      | _        | j                  j                  | _        | j                  dd       y NFTr   r2  )r>  r  reorder_iter_loopssizesr  r3  )rr   	new_orders     rs   apply_new_loop_orderz"SchedulerNode.apply_new_loop_order	  sM    ((*ZZ22

 jj&&!!E4!Pru   c                   t        | j                  t        j                  t        j                  f      sJ | j                          | j                  j                  |      | _        | j                  j                  | _	        | j                  j                         }| j                  j                  |      j                  }| || j                        f| _        | j                  dd       y r@  )r   r   r0   r   r!  r>  r  reindex_iter_loopsrC  r  r  r  r#  r$  r   r3  )rr   new_iter_sizesr  r$  s       rs   apply_loop_reindexingz#SchedulerNode.apply_loop_reindexing	  s    $))b&7&79J9J%KLLL((*ZZ22>B
jj&&..0>>--f5>>ht{{34
!!E4!Pru   c                   | j                   j                         }t        | j                   j                        |z
  }t	        t        |            }t	        t        |||z               }| j                  ||z          t        | j                  d         dk(  sJ | j                  d   | j                  d   d   | j                  d   d   ff| _        y )Nr   r   r   )r  get_original_num_rdimsr   r  r   rangerE  r   )rr   	num_rdims
num_pwdimspwdimsrdimss        rs   swap_pw_red_dimensionz#SchedulerNode.swap_pw_red_dimension	  s    JJ557	--.:
uZ()eJ
Y(>?@!!%&.14::a=!Q&&&ZZ]TZZ]1%5tzz!}Q7G$HH
ru   c                D    | j                   j                         | _         | S r~   )r  extract_pw_from_reductionrq   s    rs   rS  z'SchedulerNode.extract_pw_from_reduction(	  s    ZZ99;
ru   c                    t         j                  |       sy t        | j                  t        j
                        sJ | j                  j                         5  | j                          d d d        y # 1 sw Y   y xY wr~   )r   r   r   r   r0   r   with_original_inner_fnr  rq   s    rs   cancel_reduction_splitz$SchedulerNode.cancel_reduction_split,	  s^     33D9$))R%6%6777YY--/ 	"!	" 	" 	"s   A11A:c                   t        | j                  t        j                  t        j                  f      sJ | j
                  j                  ||      | _        | j
                  j                  | _        | j                  j                         }| j                  j                  |      j                  }| || j                        f| _        | j                  dd       y )NTrA  )r   r   r0   r   r!  r  #expand_dimension_for_pointwise_noderC  r  r  r  r#  r$  r   r3  )rr   	dimension	new_ranger  r$  s        rs   rX  z1SchedulerNode.expand_dimension_for_pointwise_node3	  s     $))b&7&79J9J%KLLLZZCCy

 jj&&..0>>--f5>>ht{{34
 	!!D$!Oru   c                    | j                   j                         | _         | j                   j                  | _        | j	                  dd       y )NTFrA  )r  merge_loopsrC  r  r3  rq   s    rs   r\  zSchedulerNode.merge_loopsD	  s<    ZZ++-
jj&& 	!!D%!Pru   c                   d }| j                   d   }t        |      |j                  cxk(  r|j                  k(  rn n|j                  |      }|rPt        xj
                  dz  c_        t        j                  d| j                         |       | j                  |       yt        j                  d| j                                y)Nr   r   z"Reorder loops for %s with order %sTzEDon't reordering %s because we can not decide the suitable loop orderF)
r  r   num_varsdecide_loop_order_to_matchr1   num_loop_reorderingloop_ordering_logr  r  rE  )rr   rn  ro  rD  
self_sizess        rs   rp  z'SchedulerNode.reorder_loops_by_dep_pairP	  s     	[[^
z?h//E93E3EE ;;IFI''1,'##4dmmoy %%i0##W ru   c                $   | j                         }| d| j                  d    | d| j                  d    | d| j                   g}| j                  j	                         D ]  }t        |t              r|j                  }t        j                  j                  |      }t        |t        j                        rZ|j                  | dt        |j                                 t        | j                   t"              rR|j                  d| d       |j                  t%        j&                  | j                   j)                         d	             | j*                  J |j-                  | j/                                d
j1                  |      S )Nz.group.device = r   z.group.iteration = r   z	.sizes = z
_layout = zclass z_loop_body:r  rH  )r  r   r  r   rw  r   r;   r   ra   r   rI  r0   rL  r   r  r  r  rE   r  r  r  r   r  r[  join)rr   r   linesr   r  r   s         rs   rR  zSchedulerNode.debug_str_extrag	  sK   }}f$TZZ]O4f'

17fIdkk]+

 ##446 	OCc7+88gg((2!#r'9'9:LLH:Z

8K7L!MN	O djj(+LL6${34LL)=)=)?HIyy$$$T//12yyru   c                    | j                   S r~   )r  rq   s    rs   r\  zSchedulerNode.get_ranges}	      {{ru   c                <   t        | j                  t        j                  t        j                  f      sJ dt        | j                               t        | j                  j                               xr' | j                  d u xs | j                  j                   S Ntype(self.node)=)
r   r   r0   r   r!  r   r}   r   r  has_partial_accumulaterq   s    rs   r   zSchedulerNode.is_reduction	  s    $))b&7&79J9J%KL 	
tDII !	
L DII0023 
JJ$Gdjj&G&G"G	
ru   c                    t        | j                  t        j                        sJ dt	        | j                               | j                  j                         dk(  S )Nrj  dot)r   r   r0   r   r   r   rq   s    rs   r  zSchedulerNode.is_native_matmul	  sJ    $))R%6%67N<LDO;M9NN7yy++-66ru   c                L   t        | j                  t        j                  t        j                  f      sJ dt        | j                               t        | j                  t        j                        xr. t        | j                  j                  t        j                        S ri  )r   r   r0   r   r!  r   r   	SplitScanrq   s    rs   r  zSchedulerNode.is_split_scan	  sy    $))b&7&79J9J%KL 	
tDII !	
L $))R%6%67 
JIINNBLL=
 	
ru   c                J    t        | j                  t        j                        S r~   r   r   r0   r!  rq   s    rs   r  zSchedulerNode.is_template	  s    $))R%6%677ru   c                f    t        | j                  t        j                        r| j                  S d S r~   rq  rq   s    rs   r  zSchedulerNode.get_template_node	  s$    &tyy"2C2CDtyyN$Nru   c                f    | j                          | j                          | j                  |       y r~   )r  r  r  )rr   
index_varss     rs   runzSchedulerNode.run	  s#    ""$Z ru   c                &   | j                   }t        t        t        |            t        t        t        |            k(  sJ t	        t        t        j                  j                  |      t        j                  j                  |                  }|S r~   )	r  r   mapr   dictzipr  r  from_iterable)rr   rt  rC  r   s       rs   ranges_from_index_varsz$SchedulerNode.ranges_from_index_vars	  sp     3sE?#s3sJ+?'@@@@--j9--e4

 ru   c                   | j                  |      }	 t        j                  t        t        j                         |            5  t        j
                  j                  |       5   | j                  |  ddd       ddd       y# 1 sw Y   xY w# 1 sw Y   yxY w# t        $ r" t        j                  d| j                          w xY w)a  
        Generate code for this node using the provided index variables.

        This method sets up the appropriate context for code generation, including
        simplifying indexing expressions based on the variable ranges, and then
        calls the node's body function with the index variables.

        Args:
            index_vars: A sequence of sequences of sympy expressions representing
                        the index variables for each dimension of the computation.
        NzError in codegen for %s)r{  ra   set_ops_handlerrM   get_ops_handlerr  set_current_noder  rS  rT  fatalr   )rr   rt  r   s      rs   r  zSchedulerNode.codegen	  s     00<
	!!"213D3D3F
"ST())$/( 

J'	( ( ( ( ( (
  	II/;	sA   1B  B$B4B<B B	
BBB B +Cc                    |r| j                   nt        | j                         \  }}t        j                  | j                  |t
        j                  j                  gt        |      z  g      S )z\
        Get the memory dependencies in either the pointwise or the reduction axes.
        )hidden_args)	r  r   r/   r&  r  r  SZeror   )rr   	pointwise
keep_sizesignore_sizess       rs   "pointwise_or_reduction_read_writesz0SchedulerNode.pointwise_or_reduction_read_writes	  sT     3<4;;$++AV 
L//JJ
%'',,#lBS1S0T
 	
ru   c                &    | j                  d      S )zH
        Get the memory dependencies in the non-reduction axes.
        Tr  r  rq   s    rs   r5  z#SchedulerNode.pointwise_read_writes	  s    
 666FFru   c                &    | j                  d      S )zD
        Get the memory dependencies in the reduction axes.
        Fr  r  rq   s    rs   reduction_read_writesz#SchedulerNode.reduction_read_writes	  s    
 666GGru   c                   | j                         ryt        d | j                         D              ryt        | j                  j
                        dk(  rt        |t        j                        rt        t        | j                  j
                              }t        |t        j                        sJ dt        |             |j                  |j                  k(  xr |j                  |j                  k(  S y)NFc              3  <   K   | ]  }|j                           y wr~   )r  r  s     rs   r   z,SchedulerNode.can_inplace.<locals>.<genexpr>	  s     ?Ss ?r  r   ztype(write_dep)=)r  r  r<  r   r   rQ  r   r/   r9   r  iterr   r   r   )rr   r  	write_deps      rs   r  zSchedulerNode.can_inplace	  s    ?D,<,<,>??t&&'1,l,,2
 T$"2"2"9"9:;Ii)?)?@WEUT)_DVBWW@>>Y__4X)..9XXru   c                   t               }t        | j                  t              r| j                  j	                         D ]  }|j
                  dk(  s|j                  dk(  s#d|j                  v r|j                  d   dk(  s,t        |j                        dk(  s\|j                  d   dk(  so|j                  d|j                  v r|j                  d   n(t        |j                        dk\  r|j                  d	   nd
        |S )Ncall_methodstoremode
atomic_addr  r  r   r   r   rX  )r   r   r  rE   r   r  r  r7  r   r  r  )rr   buffers_store_as_atomic_addr   s      rs   _get_atomic_add_buffersz%SchedulerNode._get_atomic_add_buffers	  s    7A|#djj(+

,,. GG},w.4;;.4;;v3F,3V		Na/DIIaLL4P 033!T[[0 F+.1$))n.Adiilr +*ru   c                p    | j                   | j                   j                  d      ryt        |          S )Ndevice_assert_asyncT)r  has_opr  r  rr   r  s    rs   r  zSchedulerNode.has_side_effects
  s2     ::!djj&7&78M&Nw'))ru   )r  r  r   z%ir.ComputedBuffer | ir.TemplateBufferr   r  NN)r  'tuple[dict[Any, Any], list[Any]] | Noner  zCallable[_P, _T] | Noner   r  )r  r  r  zCallable[..., Any] | Noner   r  )r   r}   r2  r}   r   r  )r2  r}   r   r  )r   r  )r;  r  r   r  r  )rD  Sequence[int]r   r  )rH  Sequence[sympy.Expr]r   r  r   rf   )rY  r   rZ  r   r   r  r  r  r  r  r  )rt  r  r   r  )rt  r  r   zdict[sympy.Expr, sympy.Expr])rt  r  r   r  r  )r  r}   r   r1  )r   r1  r  r  )'r   r   r   r   r   r:  r  r.  r3  r1  r9  r<  r>  rE  rI  rQ  rS  rV  rX  r\  rp  rR  r\  r   r  r  r  r  ru  r{  r  r  rO   r5  r  r  r  r  r  r  s   @rs   r   r     s   
 -,O 4 
	 OS=A$K $; 
	F OS?C$K $= 
	"GG8<G	G*;

	L/QQI"PP),P	P"
Q!.7	. ,
7
8O!
8	%0 !%	
	
	 	
 G G H H + +& * *ru   r   c           	     n     j                   } j                  t        j                  j	                  |D cg c]  }|j
                   c}             t         fdt        j                  |D cg c]  }|j                   c} D               j
                  j                  z
   _        y c c}w c c}w )Nc              3  Z   K   | ]"  }|j                   j                         vr| $ y wr~   r   r  )r   r   group_snodes     rs   r   z2refresh_group_node_dependencies.<locals>.<genexpr>
  s.      
xx{;;== 
   (+)
r   rx  r/   
ReadWrites
merge_listr   r   unionr3  rQ  )r  r   r  s   `  rs   refresh_group_node_dependenciesr  

  s     F**6+JaAMM+JK
 	 
!'')O1!*>*>)OP
 	

 
!
!
(
(	) " ,K *Ps   B-0B2r  c                l   t        | t        t        f      sJ || _        || _        d | _        t        j                  |D cg c]  }|j                  |j                   c} | _        t        |        t        d | j                  D              | _        t        d | j                  D              | _        t        d | j                  D              | _        t        d | j                  D              | _        | j!                         D ci c]  }|j#                         | c}| _        y c c}w c c}w )Nc              3  4   K   | ]  }|j                     y wr~   r%  r   r  s     rs   r   z"init_group_node.<locals>.<genexpr>+
       HHrv  c              3  4   K   | ]  }|j                     y wr~   )r&  r  s     rs   r   z"init_group_node.<locals>.<genexpr>,
  r  rv  c              3  4   K   | ]  }|j                     y wr~   )r#  r  s     rs   r   z"init_group_node.<locals>.<genexpr>-
        )!")rv  c              3  4   K   | ]  }|j                     y wr~   )r$  r  s     rs   r   z"init_group_node.<locals>.<genexpr>0
  r  rv  )r   r   GroupedSchedulerNoder   r  r   r   r  r  r  rZ  r%  rY  r&  r#  r$  r<  r  r.  )r  r  r   r  r   s        rs   init_group_noder  
  s   
 k$68L#MNNNK%KK&,,%	A!)@!++	AK $K0H[5G5GHHKH[5G5GHHK%( )&1&8&8) &K" &) )&1&8&8) &K" (3'>'>'@# ##K 
B#s   D,D,D1c                      e Zd ZU dZded<   e	 	 	 	 	 	 d!d       Zd"dZd#dZe	d$d       Z
	 	 	 	 	 	 d%dZd& fd	Ze	d'd
       Zd'dZe	d(d       Zd)dZd'dZd'dZ	 	 	 	 	 	 d* fdZe	d(d       Ze	d(d       Zd+dZd'dZe	d,d       Ze	d,d       Ze	d,d       Ze	d,d       Ze	d-d       Zd.dZe	d,d       Zd/dZd0dZ d1dZ!d'dZ"e	d, fd        Z# xZ$S )2r   z
    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.
    r  r   c                   |j                   |j                   u sJ t        |t        t        f      sJ |j	                         r6t        |t
              r&t        |j                  t        j                        sJ t        |t        t        f      sJ t        t        j                  |j                         |j                                     } | |j                   |      S r~   )r  r   r   r   r  rR  r   r0   rB   r   r  r  r   )rx   r   r   rq  s       rs   ry   zFusedSchedulerNode.fuseA
  s     %//111%-1C!DEEE:e5N#Oejj"..999em5G%HIIIY__U__%68IJK5??E**ru   c                    | j                   D ]6  }t        |t              sJ |j                         sJ |j	                          8 | S r~   )r   r   r   r   rS  rr   r   s     rs   rS  z,FusedSchedulerNode.extract_pw_from_reductionN
  sJ    {{ 	0Gg}555'')))--/	0 ru   c                j    | j                   D ]$  }t        |t              sJ |j                          & y r~   )r   r   r   rQ  r  s     rs   rQ  z(FusedSchedulerNode.swap_pw_red_dimensionU
  s1    {{ 	,Gg}555))+	,ru   c                    t        t        d d | j                         D                    }t        |      dk(  ry t	        |      }|S )Nc              3  |   K   | ]4  }|j                         s|j                         r|j                          6 y wr~   r  r  re  r  s     rs   r   z4FusedSchedulerNode.estimate_flops.<locals>.<genexpr>`
  6      '')T^^-= '')   :<r   r   filterr   r   r   rr   fpsr|  s      rs   re  z!FusedSchedulerNode.estimate_flopsZ
  K      $ 0	
 s8q=#h
ru   c                   | j                         ryd}| j                  D ]`  }t        |t              s y|;t	        |      t	        |j
                  d         k7  rt        j                  d        y|j
                  d   }b d}|J t        |      |j                  cxk(  r|j                  k(  rn n|j                  |      }|s%t        j                  d| j                                yt        xj                  dz  c_        t        j                  d| j                         |       | j                  D ]%  }t        |t              sJ |j                  |       ' t        |        y)	z@
        Return true if a loop reordering is performed.
        FNr   z1Can not reorder fused node due to different sizeszODont reordering fused node %s because we can not decide the suitable loop orderr   z-Reorder loops for fused node %s with order %sT)r  r   r   r   r   r  ra  r  r   r^  r_  r  r1   r`  rE  r  )rr   rn  ro  rb  r  rD  s         rs   rp  z,FusedSchedulerNode.reorder_loops_by_dep_pairl
  sI    
[[ 	)Ee]3%%
*;uU\\RS_?U*U!''G aJ	) 	%%%z?h//E93E3EE ;;IFI##a ##q(#;T]]_i	
 [[ 	2Ee]333&&y1	2 	(-ru   c                    t         |   |       t        | ||       g | _        t	        |d       j
                  | _        y )Nc                4    t        | j                               S r~   )r   r   r  s    rs   r8  z-FusedSchedulerNode.__init__.<locals>.<lambda>
  s    s1>>3C/D ru   r  )r  r:  r  r  rY  r   )rr   r  r   r  s      rs   r:  zFusedSchedulerNode.__init__
  s8    #i0%'
%DEKK
ru   c                z    dj                  | j                  D cg c]  }|j                          c}      S c c}w Nr=  rd  r   r  rr   r  s     rs   r  zFusedSchedulerNode.get_name
  )    xxt{{;!;<<;   8c                <    | j                   d   j                         S r   r   r  rq   s    rs   r  z!FusedSchedulerNode.get_first_name
      {{1~&&((ru   c                |    t        j                  | j                  D cg c]  }|j                          c} S c c}w r~   r   r  r   r  r  s     rs   r  z#FusedSchedulerNode.get_buffer_names
  .    !L1!"4"4"6!LMM!L   9c                j    g }| j                   D ]!  }|j                  |j                                # |S r~   r   r  r<  rr   r  r   s      rs   r<  zFusedSchedulerNode.get_outputs
  4    (*KK 	.DMM$**,-	.ru   c           
     ~   t        | j                        D cg c]+  \  }}| j                          d| d|j                          - }}}| j                  d   j                  }||j                  | j                                t        j                  dj                  |      j                         d      S c c}}w )Nz.snodes[z] =
r   rH  r  )r  r   r  r  r   r  r[  r  r  rd  rV  )rr   r  r   re  s       rs   rR  z"FusedSchedulerNode.debug_str_extra
  s     %T[[1
4 }}xs%0@/AB
 
 {{1~""LL3356tyy/668&AA
s   0B9c                h    | j                   D cg c]  }|j                          }}|  d| S c c}w )Nz
, snodes: )r   rh  )rr   r   
snodes_strs      rs   rh  z"FusedSchedulerNode.debug_str_short
  s9    9=Ed**,E
Ez*.. Fs   /c                    t         |   ||       t               }t        | j                        D ]/  }|j                  ||       |j                  |j                         1 y r~   )r  r  r   r   r   updater"  )rr   r  r  r   r  s       rs   r  z!FusedSchedulerNode.set_last_usage
  s\    
 	24FG 0:|T[[) 	8D 35GH&&t7	8ru   c                |    t        j                  | j                  D cg c]  }|j                          c} S c c}w r~   )r   r  r   r   r  s     rs   r   z$FusedSchedulerNode.used_buffer_names
  s.    !MA!"5"5"7!MNN!Mr  c                |    t        j                  | j                  D cg c]  }|j                          c} S c c}w r~   )r   r  r   r  r  s     rs   r  z/FusedSchedulerNode.used_or_aliased_buffer_names
  s3    8<D1a,,.D
 	
Dr  c                    | j                   S r~   rD  rq   s    rs   r   zFusedSchedulerNode.get_nodes
  rg  ru   c                T    t        |       j                   d| j                          dS )Nz(nodes=r@  rA  rq   s    rs   rB  zFusedSchedulerNode.__repr__
  s'    t*%%&gdmmo->a@@ru   c                :    t        d | j                  D              S )Nc              3  <   K   | ]  }|j                           y wr~   )r   r  s     rs   r   z2FusedSchedulerNode.is_reduction.<locals>.<genexpr>
  s     91>>#9r  r  r   rq   s    rs   r   zFusedSchedulerNode.is_reduction
  s    9T[[999ru   c                :    t        d | j                  D              S )Nc              3  <   K   | ]  }|j                           y wr~   )r  r  s     rs   r   z6FusedSchedulerNode.is_native_matmul.<locals>.<genexpr>
  s     =A1%%'=r  r  rq   s    rs   r  z#FusedSchedulerNode.is_native_matmul
  s    ====ru   c                :    t        d | j                  D              S )Nc              3  <   K   | ]  }|j                           y wr~   )r  r  s     rs   r   z3FusedSchedulerNode.is_split_scan.<locals>.<genexpr>
  s     :1??$:r  r  rq   s    rs   r  z FusedSchedulerNode.is_split_scan
  s    :dkk:::ru   c                :    t        d | j                  D              S )Nc              3  <   K   | ]  }|j                           y wr~   r  r  s     rs   r   z1FusedSchedulerNode.is_template.<locals>.<genexpr>
  s     8q1==?8r  r  rq   s    rs   r  zFusedSchedulerNode.is_template
  s    8DKK888ru   c                j    | j                   D ]$  }|j                         s|j                         c S  y r~   )r   r  r  rr   r   s     rs   r  z$FusedSchedulerNode.get_template_node
  s5    KK 	0D!--//	0 ru   c                     | j                   d   S r   )r   rq   s    rs   r   zFusedSchedulerNode.get_device
  s    zz!}ru   c                :    t        d | j                  D              S )Nc              3  <   K   | ]  }|j                           y wr~   )r  r  s     rs   r   z>FusedSchedulerNode.has_aliasing_or_mutation.<locals>.<genexpr>
  s     EA1--/Er  r  rq   s    rs   r  z+FusedSchedulerNode.has_aliasing_or_mutation
  s    EEEEru   c                    t         r~   NotImplementedError)rr   rz  s     rs   r{  z'FusedSchedulerNode.update_mutated_names
      !!ru   c                    t         r~   r  )rr   r   s     rs   r  zFusedSchedulerNode.add_fake_dep
  r  ru   c                    t         r~   r  r  s     rs   r  zFusedSchedulerNode.can_inplace   r  ru   c                   | j                         }dj                  d | j                  D              }t               }|j	                  | dt        |       j                   d| d| dt        | j                  j                         d| dt        | j                         d| d	t        | j                  j                  | j                  z
         d| d
| j                   d| d| j                   d| d       |j                         5  | j                         D ]!  }|j	                  |j!                                # 	 ddd       |j#                  d       	 |j	                  | j%                                |j-                         j/                         S # 1 sw Y   XxY w# t&        $ r t(        j+                  dd       Y Lw xY w)rD  r  c              3  F   K   | ]  }t        |      j                    y wr~   )r   r   r  s     rs   r   z/FusedSchedulerNode.debug_str.<locals>.<genexpr>  s     FQQ 0 0Fs   !r  rE  rF  rG  rH  rI  rJ  rK  rL  z.outputs = [
            Nr  rM  TrN  )r  rd  r   rX   rP  r   r   r  r   rQ  r3  r   r#  r$  r  r<  r  r  rR  rS  rT  rU  r  rV  )rr   r   node_typestrr   r   s        rs   r  zFusedSchedulerNode.debug_str  s   }}xxF$++FF

bd		Q|n -j))0012 3WT%<%<=> ?74#3#3#9#9D<S<S#STU VT445 6T445 6 
	
 ZZ\ 	,'') ,

3==?+,	, 	c	HJJt++-.  ''))	, 	,  	HKK7$KG	Hs   	5FF" F" GGc                p    | j                   t        d | j                   D              S t        |          S )Nc              3  <   K   | ]  }|j                           y wr~   )r  r  s     rs   r   z6FusedSchedulerNode.has_side_effects.<locals>.<genexpr>"  s     G4t,,.Gr  )r   r  r  r  r  s    rs   r  z#FusedSchedulerNode.has_side_effects  s0    ;;"G4;;GGGw'))ru   r   rf   r   rf   r   r   r  r  r  r  )r  r  r   r  r   r  r  r  r   r+  r  r  r  r  )r   torch.devicer  )r   r8   r   r  r  )%r   r   r   r   r   r   ry   rS  rQ  rO   re  rp  r:  r  r  r  r<  rR  rh  r  r   r  r   rB  r   r  r  r  r  r   r  r{  r  r  r  r  r  r  s   @rs   r   r   8
  s    $#
+%
+.?
+	
+ 
+,
  ")!).7)	)VL = =) N N	B/8#28HV8	8 O O 
 

A : : > > ; ; 9 9   F F
"""*8 * *ru   r   c                  H     e Zd ZdZd fdZ	 	 	 	 	 	 ddZddZddZ xZS )	FusedMixOrderReductionszNFused node for two reductions with different iteration orders (inner + outer).c                `   t         j                  |      st         j                  |      sJ ||}}|| _        || _        t        |   |j                  t        |j                               t        |j                               z          t         j                  | j                        | _
        y r~   )r   r  r   r   r  r:  r  r   r   r   r  )rr   r   r   r  s      rs   r:  z FusedMixOrderReductions.__init__)  s     33E:$77>>> %5E

OOT%//"34tEOO<M7NN	
 '00<
ru   c                   t        |t              rJ t        |t              rJ | j                  j                  ||d      syt        j                  |      rt        j                  |      sydd}	 	 	 	 dd}|r' |||f       ||      z  s ||       |||f      z  ry|j                          xs+ | j                  j                  ||d      | j                  k\  S )a  
        node1 is from the current mix order reduction; node2 is another node we want to fuse in.

        other_nodes are passed in to check if fusion will introduce producer/consumer relationship
        between the inner and outer reduction. If yes, we don't fuse.
        Fallow_mix_order_reductionc                B    t               } |j                  d | D         S )Nc              3  4   K   | ]  }|j                     y wr~   )r  r  s     rs   r   zTFusedMixOrderReductions.sub_node_can_fuse.<locals>._get_ancestors.<locals>.<genexpr>S  s     :qq{{:rv  r   r  rq  r   s     rs   _get_ancestorszAFusedMixOrderReductions.sub_node_can_fuse.<locals>._get_ancestorsQ  s     ,C399:E:;;ru   c                B    t               } |j                  d | D         S )Nc              3  <   K   | ]  }|j                           y wr~   )r  r  s     rs   r   zZFusedMixOrderReductions.sub_node_can_fuse.<locals>._get_operation_names.<locals>.<genexpr>Y  s     F1q446Fr  r  r  s     rs   _get_operation_nameszGFusedMixOrderReductions.sub_node_can_fuse.<locals>._get_operation_namesU  s"     ,C399FFGGru   )count_bytes)rq  tuple[BaseSchedulerNode, ...]r   r   )	r   r  r  r  r   r  r   score_fusion_memoryr  )rr   r   r   other_nodesr  r  s         rs   sub_node_can_fusez)FusedMixOrderReductions.sub_node_can_fuse5  s     e%<===e%<===
 ~~&&ueu&U //
#66u=	<	H0	H	H u~.1Ek1RR{+.BE5>.RR ""$$ ~~11%E1Rzz	
ru   c                @   t         j                  j                  }|dkD  rt               }t	        j
                  | j                         |j                               D ]I  }|j                  j                  D ].  }t        |t              s|j                  |j                         0 K t        |      |kD  rt        xj                  dz  c_        yt        |t               sR| j#                  | j$                  || j&                  f      xs( | j#                  | j&                  || j$                  f      S | j#                  | j$                  |j$                  | j&                  |j&                  f      xr/ | j#                  | j&                  |j&                  t)                     S )Nr   r   F)r-   r   mix_order_reduction_max_readsr   r  r  r   r   r   r   r9   r  r   r   r1   #rejected_mix_order_reduction_fusionr  r  r   r   r   )rr   other	max_reads	all_readsr^  r   s         rs   can_fuse_withz%FusedMixOrderReductions.can_fuse_withg  sD    MM??	q=)3Ioodnn&68IJ 0>>// 0C!#y1!chh/00 9~	);;q@;%!89))

EDJJ= J''

EDJJ=IJ ))

EKK$**ekk)B K((U[[%'JKru   c                T   | j                   j                         }| j                  j                  |      }t	        |t
              rX|j                  | j                   |j                         }|j                  | j                  |j                        }t        ||      S | j                  | j                   || j                  f      r2|j                  | j                   |      }t        || j                        S |j                  | j                  |      }t        | j                   |      S r~   )	r   r   r  r#  r   r  ry   r   r  )rr   r  r  backendfused_node1fused_node2r  s          rs   	fuse_withz!FusedMixOrderReductions.fuse_with  s    &&(..,,V4e45!,,tzz5;;?K!,,tzz5;;?K*;DD%%djj%$**G$\\$**e<
.z4::FF$\\$**e<
.tzz:FFru   r  )r   rf   r   rf   r  r  )r  rf   )	r   r   r   r   r:  r  r  r  r  r  s   @rs   r  r  &  s9    X
=0
 0
 !0
 3	0
dK2Gru   r  c                  4     e Zd ZdZd fdZddZddZ xZS )FusedNestedReductionsz
    Fused node for two dependent reductions over the same logical elements.
    The outer reduction owns the codegen grid; the grouped reduction is staged
    inside that grid.
    c                   || _         || _        t        |   |j                  t        |j                               t        |j                               z          | xj                  | j                         z  c_        |}|j                  \  }\  }}t        j                  ||      }|J |\  }}	|| _        |	| _        |j                  \  }\  }
}t        j                  ||
||	|      }|J || _        | j                  t        j                   j"                  u | _        |j'                         \  }}t        j)                  |||g |||
|f      | _        y )Nr  r  )r   r   r  r:  r  r   r   r  r  r   r2  rd  rQ  rc  r  r  rM  rK  group_size_in_rr\  rW  rh  )rr   r   r   r7  r=  rS  rT  r  rQ  exact_group_sizer  r  r  ra  rb  r  s                  rs   r:  zFusedNestedReductions.__init__  sn   

OOT%//"34tEOO<M7NN	
 	$2244-9-?-?**M>!0!P!P."
 &111.D++0A)9).&&K&77 8 
 '''9E%)%6%6/:U:U:W:W%W%6%A%A%C"]22"3+-'E'E}'E$/#> 3  	ru   c                  |j                         ry| j                  j                         |j                  z  syt        j                  | j                  |j                               }|yt        j                  | j                  |      syt        d |D              ry| j                  j                  | j                  |||      S )zAllow downstream pointwise of the grouped reduction to fuse in.

        Consumers fused directly into the grouped reduction must run at either
        reduced-output resolution or full parent-tile resolution.
        Fc              3  Z   K   | ]#  \  }}|t         j                  j                  u  % y wr~   )r2  rH  rF  )r   r=  ry  s      rs   r   z6FusedNestedReductions.can_fuse_with.<locals>.<genexpr>  s-      
6 o55KKK
r  can_reorder)r   r   r  r  r2  rl  rh  r   rg  r  r  !_can_fuse_nested_reduction_append)rr   r  r%  ri  s       rs   r  z#FusedNestedReductions.can_fuse_with  s      

..05??B ==OO
 	 $@@!2
   
.
 
 ~~??JJ#	 @ 
 	
ru   c                    | j                   j                         }| j                  j                  |      }|j	                  | j                   |      }t        | j                  |      S r~   )r   r   r  r#  ry   r  r   )rr   r  r  r  	new_node2s        rs   r  zFusedNestedReductions.fuse_with  sM    &&(..,,V4LLU3	$TZZ;;ru   r  )r  rf   r%  r}   r   r}   )r  rf   r   r  )r   r   r   r   r:  r  r  r  r  s   @rs   r  r    s    (
T%
N<ru   r  c                  f     e Zd Z	 	 	 	 	 	 	 	 d fdZe	 	 	 	 	 	 dd       ZddZd	dZd
dZ xZ	S )$FusedExternTritonKernelSchedulerNodec                ,   t        |j                  t        j                        sJ t	        j
                  t        t           ||g      }t        | %  ||       || _
        || _        | j                  j                  | _        |j                  | _        y r~   )r   r   r0   r  r  r  r   rf   r  r:  kernel_nodefused_epiloguer%  r,  )rr   r  r,  r-  r   r  s        rs   r:  z-FusedExternTritonKernelSchedulerNode.__init__  s{     +**B,F,FGGGT"34{N6STF+&,))33%--ru   c                   t        |j                  t        j                        sJ |j                  }t        |j                  j                        dk(  sJ |j                  j                  d   j                  }|j                  j                  ||      }|j                  |   j                  j                  t        |              | |||      S )Nr   r   )r   r   r0   r  r  r   mutation_outputsr   r  r  r  r  removeNodeUser)rx   r   r   r  mutated_name	real_names         rs   epilogue_fusez2FusedExternTritonKernelSchedulerNode.epilogue_fuse  s     %**b&@&@AAAOO	5::../1444!JJ77:?? 0044\<P	i(..55huoF9eU++ru   c                   t        | j                  j                  t        j                        sJ t        | j
                  j                  t        j                        sJ | j
                  j                  j                         sJ t        j                  | j
                  j                  j                  d   j                        }ddlm} |j                  | j                  g|      \  }}ddlm}  || j                  g|      }ddlm}  ||||       }	|	j'                         }
| j
                  j                  j)                  || j                  j                  |
f      S )Nr   r  r  )FusedUserDefinedTritonKernel)r   r-  r   r0   r   r,  r  r   r  r   r  r  torch._inductor.codegen.simdr  get_tiling_and_scores,torch._inductor.codegen.simd_kernel_featuresr  torch._inductor.codegen.tritonr6  r  codegen_with_epilogue_fusion)rr   r  r  r  r  r=  r  kernel_featuresr6  fused_user_kernelnew_kernel_srcs              rs   r  z,FusedExternTritonKernelSchedulerNode.codegen  s    $--22B4E4EFFF$**//1K1KLLL$$66888		$**//<<Q?EEF?"88$:M:M9NPUV	S,d.A.A-BEJO8RVW*224$$AAd))..?
 	
ru   c                     yr	  r   rq   s    rs   r  z.FusedExternTritonKernelSchedulerNode.is_extern(  r  ru   c                6    | j                   j                         S r~   )r,  r\  rq   s    rs   r\  z/FusedExternTritonKernelSchedulerNode.get_ranges+  s    **,,ru   )r  r  r,  rR  r-  r   r   r  )r   rR  r   r   r   r   r  r  r  )
r   r   r   r:  r   r4  r  r  r\  r  r  s   @rs   r*  r*    sm    .. /. &	.
 
. ,(, , 
	, ,$
,-ru   r*  c                  T    e Zd ZU dZ	 	 	 	 ddZ	 	 	 	 ddZedd       Ze	 	 	 	 	 	 dd       Z	 	 	 	 d	 	 	 	 	 	 	 	 	 	 	 	 	 	 	 d fdZ	e	 	 	 	 dd       Z
e	 	 	 	 dd       ZeZd	ed
<   e	 	 	 	 dd       Ze	 	 	 	 dd       ZddZddZddZddZd dZd!dZ	 	 	 	 d"dZ xZS )#ForeachKernelSchedulerNodez
    This is a schedular node that consists of a set of scheduler nodes that
    has no data dependencies among them and can be executed in parallel.
    c                    |j                         D ]=  }|j                         | j                  v s | j                  |j                            c S  y r~   )r<  r  read_to_node)rr   producerr   s      rs   get_consumer_subnode_forz3ForeachKernelSchedulerNode.get_consumer_subnode_for5  sL     '') 	9C||~!2!22((88	9 ru   c                   t        t                  }|j                  j                  D ]  }|j                  | j
                  j                  vr&| j
                  j                  |j                     j                         }|| j                  v sf|j                  | j                  |           t        |      dk(  rt        t        |            S y Nr   )r   rf   r   r   r   r  r  r  name_to_noder  r   r  r  )rr   consumer	producersrd	node_names        rs   get_producer_subnode_forz3ForeachKernelSchedulerNode.get_producer_subnode_for>  s     013	&&,, 	<Bwwdnn88822277;LLNID---d//	:;	< y>QY((ru   c                   t        |      }j                         r|j                         rt        j                  t              t        j                  t        |      }t        j                        t        |j                        k(  }|s |d       |xr2 t        fdt        j                  |j                        D              S |j                         rkj                         r	 |d       yt        j                  t        |      }|j                        }||j                  j                  |      S  |d       yj                         rk|j                         r	 |d       yt        j                  t              j                  |      }|j                  j                  ||      S  |d       yt        d      )	Nzforeach do not have same lengthc              3  \   K   | ]#  \  }}j                   j                  ||       % y wr~   )r  r  )r   lrrE  s      rs   r   z6ForeachKernelSchedulerNode.can_fuse.<locals>.<genexpr>Z  s0      )Aq ""++Aq1)s   ),zXcandidate producer is a reduction, foreach ops cannot be fused with reductions currentlyFz5candidate producer is not dep of any foreach consumerzXcandidate consumer is a reduction, foreach ops cannot be fused with reductions currentlyz5candidate consumer has no dep in any foreach producerzXAt least one node passed to ForeachKernelSchedulerNode.can_fuse should be a foreach node)r  r  r  r  rB  r   r   r   ry  r   rF  r  r  rN  rw  )rx   rE  rJ  whyforeach_matchconsumer_subnodeproducer_subnodes    `     rs   r  z#ForeachKernelSchedulerNode.can_fuseQ  s   (+ X%8%8%:{{#=xHH{{#=xHH0C4HHM 56  S )A) &    "$$&n {{#=xHH'@@J+))228=MNNGH  "$$&n {{#=xHH'@@J+))223CXNNGHf
 	
ru   c                
   |j                         s|j                         sJ |j                         r3t        j                  t        |      }|j                  }|j
                  }n2t        j                  t        |      }|j                  }|j
                  }d }d }|j                         r|j                         r|t        j                  t        |      }t        j                  t        |      }t        |j                  |j                        D cg c]  \  }}t        j                  ||       }	}}n/|j                         rt        j                  t        |      }|j                  |      }
g }	|}d }|j                  D ]A  }||
u r*t        j                  ||      }|}|	j                  |       1|	j                  |       C n|j                         rt        j                  t        |      }|j                  |      }g }	|}d }|j                  D ]A  }||u r*t        j                  ||      }|}|	j                  |       1|	j                  |       C nt        d       | |j                  |	||||      S c c}}w )NzTAt least one node passed to ForeachKernelSchedulerNode.fuse should be a foreach node)use_custom_partition_algoprev_node_1prev_node_2enable_autotune)r  r  r  rB  rX  r[  ry  r   r   ry   rN  r   rF  rw  r  )rx   rE  rJ  rX  r[  rY  rZ  rQ  rR  fused_nodesrV  r   new_noderU  s                 rs   ry   zForeachKernelSchedulerNode.fuse  sZ    ""$(;(;(=== {{#=xHH(0(J(J%&66O{{#=xHH(0(J(J%&66O X%8%8%:{{#=xHH{{#=xHH  AAq #''1-K    "{{#=xHH'@@JK"KK  -++166tXFH"*K&&x0&&t,-   "{{#=xHH'@@JK"KK  -++166xFH"*K&&x0&&t,- !f  &?##+
 	
Ks    I?c                |    i  _         i  _        ||qt           ||       |D ]Z  }|j                  j
                  D ]  }	| j                   |	j                  <    |j                         D ]  }
| j                  |
<    \ n<| _        | _	        d  _
        g  _         j                  t        j                  j                  |j                  |j                  g             t!         fdt!        j"                  |j$                  |j$                        D               j                  j&                  z
   _        t)        |j*                  |j*                  g       _        t-        |j.                  |j.                  g       _        t)        |j0                  |j0                         _        t-        |j2                  |j2                         _        |j5                         rt7        |t8              sJ ||}}nt7        |t8              sJ ||}}|j:                   _         j:                  j=                  |j:                         |j                   _        |j                         D ]  }
| j                  |
<     j                  D ci c]'  }|j>                  jA                         D ]  \  }}||
 ) c}}} _        | _!        |d   jE                         }|sJ |tG        jH                  d      fff _%        t!        tL        jN                  jP                             _)        | _*        | _+        y c c}}}w )Nc              3  Z   K   | ]"  }|j                   j                         vr| $ y wr~   r  r  s     rs   r   z6ForeachKernelSchedulerNode.__init__.<locals>.<genexpr>  s0       xxt'<'<'>>	 r  r   combo_kernel),rD  rI  r  r:  r   r   r   r  r  r   r   r  rx  r/   r  r  r   r  r3  rQ  rZ  r%  rY  r&  r#  r$  r  r   rB  r  r  r.  r  rX  r   r  Exprr   r`  fxNoder$  r[  per_subkernel_blocks)rr   r  r   rX  rY  rZ  r[  rd  r   r  r   foreach_noder  r  r  vr  r  s   `                rs   r:  z#ForeachKernelSchedulerNode.__init__  s    +"5GY/ 3 ,,22 8D37D%%dii08 !446 3D.2D%%d+3	3 'DN DKDI)+DJ  ''22 ,,k.E.EF  )//#668V8V   ""))* # !+"7"79N9N!OPDN +"7"79N9N!OPDN&)..0N0N'D# '*..0N0N'D# %%'!+/IJJJ+6j!+/IJJJ+6j)33DNNN!!*"6"67 , 9 9D"668 5*4!!$'5 #'++@ @%:O:O:U:U:W@26!Q1@@D  *C&%%'v

> :<>?
!%((--02.$8!@s   ,L7c                   |D cg c]  }t        |t              s| }}|rSt        j                  dt	        |      |D cg c])  }|j
                  |j
                  j                         + c}       |D cg c]  }t        |t              s| }}|rt        j                  dt	        |             |D cg c]  }t        |t              s| }}|rt        j                  dt	        |             |D cg c]  }t        |t              s| }}|rt        j                  dt	        |             |D cg c])  }t        |t        t        t        t        t        f      s|+ }}|D cg c]  }t        |t              s| }	}|	rt        j                  dt	        |	             |D cg c]  }t        |t              r| }}|D cg c]  }|j                         s| }
}|
r t        j                  dt	        |
      |
       |D cg c]	  }||
vs| }}t        j                  ra|D cg c]  }|j                         s| }}|rt        j                  dt	        |             |D cg c]  }|j                         r| }}|S c c}w c c}w c c}w c c}w c c}w c c}w c c}w c c}w c c}w c c}w c c}w c c}w )Nz/ComboKernels: %d external nodes are filtered %sz+ComboKernels: %d grouped nodes are filteredz;ComboKernels: %d FusedMixOrderReductions nodes are filteredz9ComboKernels: %d FusedNestedReductions nodes are filteredz+ComboKernels: %d foreach nodes are filteredz0ComboKernels: %d template nodes are filtered: %szCComboKernels: %d reduction nodes are filtered (pointwise_only mode))r   rR  rT  r  r   r   r  r  r  r  r  rB  r  r-   combo_kernels_pointwise_onlyr   )rx   rq  r  externr   grouped	mix_ordernested_reductionsfiltered_nodesforeach_nodestemplate_nodesreduction_nodess               rs   combinable_nodesz+ForeachKernelSchedulerNode.combinable_nodes  s    #Oj4M&N!OOIIAF5;UTtyy?T&&(U
 $Kz!5I'J1KKII=G !&P1A7N)OQP	PIIMI ).V1A?T1UQVVIIK%& 
*-(+)	 
 
 &
A7Q)RA
 
 IICSEWX%
Z;U-VA
 
 &4Gq}}!GGIIBN#
 &4Oq7N!OO ..*8MQANN<LqMOM		Y( *8PAq~~?OaPNPG P
 VK Q W



 H P N Qs   JJJJ:J!J!;J&J&<J+J+=.J01J5J52J:J:J?(J?	KK8K	K	9KKc                2   | j                         }g }t        j                  }t        |D cg c]0  }|D ])  }t	        |t
              r|j                         D ]  }| + 2 c}}}      }|D ]  }t        t              }	|D ][  }|j                         }
|
r|
j                  dk(  s|
j                  dk(  r4|j                         |z  rH|	|
   j                  |       ] |	j                         D ]  }t        t              }|D ]0  }|| j                  j                  |d         j                  |       2 |j                         D ];  }|j!                  t#        dt%        |      |      D cg c]
  }||||z     c}       =   |S c c}}}w c c}w )zS
        Returns a list of lists of nodes that are to be grouped together.
        mpsr  r   )_topological_sort_nodesr-   combo_kernel_max_num_nodesr   r   r  r  r   r   r   r   r   r   r   node_to_streamr  r  rL  r   )r  sorted_nodesgrouped_nodesmax_num_nodesr   r   r  excluded_buffer_namesrq  device_groupsr  device_nodesstream_groupsstream_nodesr  s                  rs   &_default_group_nodes_for_combo_kernelszAForeachKernelSchedulerNode._default_group_nodes_for_combo_kernelsZ  s    !88:991; * ! d$;< $ 5 5 7
 	 2
 " 	E D!   3*v{{e3v{{e7K ))+.CCf%,,T23 !. 4 4 6 
DOPTDU( VD!)":":">">tQ"GHOOPTUV$1$8$8$: L!(( &+1c,.?%O ! )Q->?	
%	: K@s   5F.F4Callable[[Scheduler], list[list[BaseSchedulerNode]]]!group_algorithm_for_combo_kernelsc                    | t         _        y r~   rB  r  )custom_group_algorithms    rs   %set_group_algorithm_for_combo_kernelsz@ForeachKernelSchedulerNode.set_group_algorithm_for_combo_kernels  s    
 # 	#Dru   c                ,    t         j                  |       S r~   r  r  s    rs   group_nodes_for_combo_kernelsz8ForeachKernelSchedulerNode.group_nodes_for_combo_kernels  s     *KKIVVru   c                    t         r~   r  rq   s    rs   r  z#ForeachKernelSchedulerNode.mark_run  r  ru   c                    t         r~   r  rq   s    rs   r  z"ForeachKernelSchedulerNode.codegen  r  ru   c                     yr	  r   rq   s    rs   r  z%ForeachKernelSchedulerNode.is_foreach  r  ru   c                ,    t        | j                        S )zeReturns a list of nodes which comprise the combo kernel.
        These nodes may be vertically fused.)r   r   rq   s    rs   get_subkernel_nodesz.ForeachKernelSchedulerNode.get_subkernel_nodes  s     DKK  ru   c                t    t        t        j                  j                  d | j                  D                    S )zqReturns all nodes contained in this kernel, unpacking fused nodes
        into their constituent scheduler nodes.c              3  <   K   | ]  }|j                           y wr~   )r   r  s     rs   r   z7ForeachKernelSchedulerNode.get_nodes.<locals>.<genexpr>  s     1UA!++-1Ur  )r   r  r  rz  r   rq   s    rs   r   z$ForeachKernelSchedulerNode.get_nodes  s(     IOO111U1UUVVru   c                <    | j                   d   j                         S r   )r   r  rq   s    rs   r  z)ForeachKernelSchedulerNode.get_first_name  s    {{1~,,..ru   c                    t        | || j                  j                         | j                  D ]  }|j	                  |        y r~   )r  r  r  r   r  )rr   r  r   s      rs   r  z/ForeachKernelSchedulerNode.prune_redundant_deps  s=     	d$68R8RSKK 	:D%%&89	:ru   )rE  rf   r   r  )rJ  rf   r   r  rE  rf   rJ  rf   r   r}   )rE  rf   rJ  rf   r   rB  )NNFF)r  r  r   r  rX  r}   rY  r  rZ  r  r[  r}   rd  r}   r   r  rq  r  r   r  )r  r  r   list[list[BaseSchedulerNode]])r  r  r   r  r  r  r   r  r  r  r  )r   r   r   r   rF  rN  r   r  ry   r:  rq  r0  r  r  r   r  r  r  r  r  r  r   r  r  r  r  s   @rs   rB  rB  /  s   
)	!)	!& ,
 ,
\ >
(>
4E>
	#>
 >
J 1504 %%*N9N9 (N9 $(	N9
 .N9 .N9 N9 #N9 
N9` F+F	 F FP 00	&0 0h 	/ & ( / 
 T
	
 
 WW	&W W
""!
W
/:">:	:ru   rB  c                       e Zd ZU dZded<   edd       Z	 d	 	 	 	 	 	 	 d fdZddZddZ	e
dd       Zdd	Ze
dd
       ZddZe
dd       ZddZddZedd       Z xZS )r  aC  
    This is a "fake" scheduler node that represents a group of scheduler nodes
    that are meant to be *grouped* together (it does not allow another node to be scheduled
    in between its constituent nodes, nor does it allow another node to fuse into any of its constituent nodes).
    The way it does this is by maintaining its unmet dependencies as the union of its constituent nodes.
    Fusion will still happen among the nodes within each GroupedSchedulerNode.
    At codegen time, this scheduler node will be unpacked and codegen is called on each constituent node.
    r  r   c                    |d   j                   t        fd|D              sJ  | |      }|D ]  }|j                  |j                         <   ! |j                  |j                         <   |S )Nr   c              3  :   K   | ]  }|j                   u   y wr~   r  )r   r   r  s     rs   r   z.GroupedSchedulerNode.create.<locals>.<genexpr>  s     B44>>Y.B   )r  r   r  r  )rx   r   grouped_snoder  r  s       @rs   createzGroupedSchedulerNode.create  sy    1I''	B6BBBBIv. 	KE=JI(()9:	KAN	$$]%;%;%=>ru   c                L    t         |   |       t        | ||       || _        y r~   )r  r:  r  temp_grouping)rr   r  r   r  r  s       rs   r:  zGroupedSchedulerNode.__init__  s(     	#i0 +ru   c                6   | j                   r| j                  S | j                  D ])  }|| j                  j                  |j	                         <   + | j                  j                  | j	                         = | j                  j                  | j                        S )z
        Do fusion among nodes within this GroupedSchedulerNode,
        and then unpack this GroupedSchedulerNode into regular nodes.
        )r  r   r  r  r  
fuse_nodes)rr   r  s     rs   unpackzGroupedSchedulerNode.unpack  sx    
 ;;[[ 	HEBGDNN--enn.>?	HNN--dmmo>~~((55ru   c                    | j                  | j                  j                  |             | j                  j	                  |       y r~   )rx  r   r}  r3  r  )rr   fake_deps     rs   r  z!GroupedSchedulerNode.add_fake_dep  s5    T--77AB##H-ru   c                z    dj                  | j                  D cg c]  }|j                          c}      S c c}w r  r  r  s     rs   r  zGroupedSchedulerNode.get_name  r  r  c                <    | j                   d   j                         S r   r  rq   s    rs   r  z#GroupedSchedulerNode.get_first_name  r  ru   c                |    t        j                  | j                  D cg c]  }|j                          c} S c c}w r~   r  r  s     rs   r  z%GroupedSchedulerNode.get_buffer_names  r  r  c                j    g }| j                   D ]!  }|j                  |j                                # |S r~   r  r  s      rs   r<  z GroupedSchedulerNode.get_outputs  r  ru   c                    t        t        d d | j                         D                    }t        |      dk(  ry t	        |      }|S )Nc              3  |   K   | ]4  }|j                         s|j                         r|j                          6 y wr~   r  r  s     rs   r   z6GroupedSchedulerNode.estimate_flops.<locals>.<genexpr>  r  r  r   r  r  s      rs   re  z#GroupedSchedulerNode.estimate_flops  r  ru   c                    | j                   S r~   rD  rq   s    rs   r   zGroupedSchedulerNode.get_nodes  rg  ru   c                X    | j                   r| j                   d   j                         S d S r   )r   r   rq   s    rs   r   zGroupedSchedulerNode.get_device  s$    .2kkt{{1~((*CtCru   c                     yrm  r   )rx   rE  rJ  s      rs   r  zGroupedSchedulerNode.can_fuse  rq  ru   )r   r  r   r  )F)r  r  r   r  r  r}   r   r  r  )r  r8   r   r  r  r  r  r  r  r  r  )r   r   r   r   r   r   r  r:  r  r  rO   r  r  r  r<  re  r   r   r  r  r  s   @rs   r  r    s     $#  $	++ (+ 	+
 
+6. = =) N N  "D  ru   r  c           
          t         j                  d fd       }t        t        t	        t         d                           }t        |      dkD  r|D cg c]  } |   	 c} t        j                  r|j                  |       |S c c}w )z
    A heuristic to decide loop iteration orders.  This has not been well
    tuned and may be something we should autotune.
    c                t   |    dk(  s|   dk(  rt        |    dk(  |   dk(        S D cg c]  }t        ||           }}D cg c]  }t        ||          }}t        d t        ||      D              }t        d t        ||      D              }||kD  ry||kD  ryt        ||       S c c}w c c}w )Nr   c              3  :   K   | ]  \  }}|d k(  xs ||k    ywr   Nr   r   sl_asl_bs      rs   r   z5pick_loop_order.<locals>.index_cmp.<locals>.<genexpr>:  )      
)3tDAI$$
   c              3  :   K   | ]  \  }}|d k(  xs ||k    ywr  r   r  s      rs   r   z5pick_loop_order.<locals>.index_cmp.<locals>.<genexpr>=  r  r  r  )rP   absr   ry  )	r  bslstride_len_astride_len_ba_firstb_firstrC  stride_lengthss	          rs   	index_cmpz"pick_loop_order.<locals>.index_cmp-  s    8q=E!HMuQx1}eAh!m44 .<<rBqE
<<-;<rBqE
<<  
7:<7V
 
  
7:<7V
 
 WW 1ay# =<s   B0	B5r   r  )r  r   r  r   r   r   )		functools
cmp_to_keyr   r   rL  r   r-   pick_loop_orderssort)r  rC  priority_idxr  orderpis   ``    rs   pick_loop_orderr  #  s      4 %N1$5 6789E
<17CD.,D

y
!L Es   Bc                   |j                         }| j                         }t        |t              rt        |t              sJ |j                         }| j                         }t        |t              rt        |t              sJ t        j
                  j                  |= ||_        t        j
                  j                  |= ||_	        t        j
                  j                  j                  |       }t        j
                  j                  j                  |       |t        j
                  j                  |<   |t        j
                  j                  |<   t        j
                  j                  j                  |       }t        j
                  j                  j                  |       |t        j
                  j                  |<   |t        j
                  j                  |<   y r~   )r  r   r,  r  ra   r   r  r   
name_to_opoperation_namebuffersr   r0  
operations)	orig_noder]  replaced_buf_nameorig_buf_namereplaced_op_nameorig_op_nameorigs          rs   _replace_operation_bufferr  Q  sU    !))+&&(MmS)j9JC.PPP224//1LlC(Z8H#-NNN	01!HM	+,*H77??  +DGGOO8$$AGGOOD,4AGG=)77##I.DGGh''AGGt'/AGG|$ru   c                p    | j                         }|j                         }||z
  }||z  }|d|z   z  }||z  S rH  )r3  r1  )r   r   epilogue_runtimetemplate_write_bytesepilogue_read_bytesextra_bytesextra_bytes_ratioextra_memory_ratios           rs    _estimate_fused_epilogue_runtimer  m  sX     779557%(<<K $&:: +a2C.CD 000ru   c                    |dk\  ry|j                   }|y|sJ ||j                  z  }| |z  }||z  }||z  }	||z  }
|	|
fS )N   )r   r  )r   r   )regs_per_multiprocessorwarp_size_or_default)unfused_n_regsfused_n_regsfused_n_spills	num_warpsdevice_propsregs_per_smthreads_per_blockregs_per_block_unfusedregs_per_block_fusedblocks_unfusedblocks_fuseds              rs   "_occupancy_before_and_after_fusionr  z  sw      66K9!L$E$EE+.??'*;; $::N"66L<''ru   c                    d}d}t        |||||      \  }	}
|d| z  kD  xr |
dkD  }|
dk7  xr |
|k\  xs |
|	z  |kD  xs |S )zE
    Determine whether to fuse an epilogue into a GEMM template.
    r  g      ?r   r   r  )r  )ms1ms2r  r  r  r  r  MIN_ACCEPTED_OCCUPANCYREGRESSED_OCCUPANCY_RATIOr  r  ,epilogue_dominated_with_sufficient_occupancys               rs   _fuse_epiloguer    s      # $Fni$ NL 47S=3U\TUEU0
 2 .. 	8.(+DD	87ru   c                  T    e Zd ZU ded<   dZded<   dZded<   ddZddZdd	Zdd
Z	y)r1  BaseSchedulerNode | OutputNoder   Fr}   r  is_weakc                v    t        | j                  j                         | j                  | j                  f      S r~   )r  r   r  r  r  rq   s    rs   r  zNodeUser.__hash__  s+    TYY'')4+;+;T\\JKKru   c                    t        |t              xrW | j                         |j                         k(  xr4 | j                  |j                  k(  xr | j                  |j                  k(  S r~   )r   r1  r  r  r  rr   r  s     rs   __eq__zNodeUser.__eq__  s[    uh' .5>>#33.  E$5$55. -		
ru   c                6    | j                   j                         S r~   r  rq   s    rs   r  zNodeUser.get_name  r  ru   c                    | j                   |j                   u sJ t        | j                   | j                  xr |j                  | j                  xr |j                        S r~   )r   r1  r  r  r  s     rs   r  zNodeUser.merge  sP    yyEJJ&&&II2!2!2LL*U]]
 	
ru   Nr  )r  objectr   r}   r  )r  r1  r   r1  )
r   r   r   r   r  r  r  r  r  r  r   ru   rs   r1  r1    s3    
((K GTL
$
ru   r1  c                   t               }| j                         }t        |t        j                        r|j                  t        |j                        t        |j                        z  t        |j                        z         t        |t        j                        r$|j                  t        |j                               |S |
J d|        |S )z=Get free symbols from a node's layout (size, stride, offset).z*Expect layout to be None but found layout=)r   maybe_get_layoutr   r0   Layoutr  r&   r   strideoffsetr	  get_layout_symintsr  )r   free_symbol_usesr  s      rs   r  r    s    1;""$F&"))$%6==)*6==)*	

 fb;;<##$6v}}$EF  ~T!KF8TT~ru   c                "   t        | t              r( t               j                  d | j                  D         S | j
                  J | j
                  j                         } |j                  d | j
                  j                         D          |S )z
    Gets symbols used in a scheduler node, including free symbols from
    the node's operations and layout symints from outputs.
    c              3  2   K   | ]  }t        |        y wr~   get_scheduler_node_symbol_uses)r   r  s     rs   r   z1get_scheduler_node_symbol_uses.<locals>.<genexpr>  s     M,U3M   c              3  2   K   | ]  }t        |        y wr~   )r  )r   ir_nodes     rs   r   z1get_scheduler_node_symbol_uses.<locals>.<genexpr>  s     	M'
W
%	Mr  )	r   r   r   r  r   r   get_free_symbol_usesr  r<  )r   r  s     rs   r  r    s     $*+!z|!!MM
 	
 99   yy557	MTYY5J5J5L	M ru   c                v    | j                         }||j                  |j                  S t        j                  S z4Check per-template flag, fall back to global config.)r  allow_epilogue_fusionr-   epilogue_fusionr  tbs     rs   _is_epilogue_fusion_enabledr    8    		(	(	*B	~"22>'''!!!ru   check_configc               H   |rt         j                  syt        | t              syt        | j                  t
        j                        syt        | j                  j                  t
        j                        sy| j                  j                  j                  dk7  ry| j                         }t        |      dk7  ry|d   j                         st        |d   j                               dk7  ry| j                  j                  }t!        |      xr t#        d |D              S )NFr  r   r   c              3  ^   K   | ]%  }t        |t              xr |j                  d k(   ' yw)r  N)r   r9   r  )r   writes     rs   r   z3_is_atomic_add_mutation_epilogue.<locals>.<genexpr>  s-       HM
5)$C|)CC s   +-)r-   epilogue_fusion_with_atomic_addr   r   r   r0   r   r   Scatterscatter_moder<  r   r  r  r   rQ  r}   r   )r   r  r,  rQ  s       rs    _is_atomic_add_mutation_epiloguer    s     FBBdM*dii!2!23diinnbjj1yy~~""l2 G
7|qqz3wqz'?'?'A#Ba#G$$F< C  QW   ru   c                p    | j                         }t        |t        j                        xr t	        |      S r~   )r  r   r0   TritonTemplateBufferr  )r  epilogue_nodetemplate_bufs      rs   &_can_fuse_atomic_add_template_epiloguer     s7     !224Lb-- :
*=
9:ru   c                v    | j                         }||j                  |j                  S t        j                  S r  )r  allow_prologue_fusionr-   prologue_fusionr  s     rs   _is_prologue_fusion_enabledr$  (  r  ru   c                b    | j                         xr |j                          xr t        |       S r~   )r  r  r   s     rs   is_epilogue_fusionr&  0  4     	/!!##	/'.ru   c                b    |j                         xr | j                          xr t        |      S r~   )r  r$  r   s     rs   is_prologue_fusionr)  8  r'  ru   c                6    t        | |      xs t        | |      S r~   )r&  r)  r   s     rs   is_template_fusionr+  @  s    eU+O/A%/OOru   c                "    t        | |      r|S | S r~   )r&  r   s     rs   template_fusion_pw_noder-  D  s    &ue45?%?ru   c                      e Zd ZU dZ ej
                  e      Zded<    ej
                  e      Z	ded<   e
dd       ZddZdd	Zdd
ZddZy)_LoopStateSnapshota  Captured loop state for a set of scheduler nodes, restorable on rollback.

    Stores both SchedulerNode loop state (body, sizes, deps) and
    FusedSchedulerNode group assignments, since the latter is reassigned
    directly after child reindexing and has no mutation listener.
    r   z$dict[SchedulerNode, tuple[Any, ...]]scheduler_node_stateszdict[FusedSchedulerNode, Any]fused_node_groupsc                D     |        }|D ]  }|j                  |        |S )zDCapture scheduler-node boundaries and their mutable leaf loop state.)snapshot_node)rx   rq  snapshotr   s       rs   r  z_LoopStateSnapshot.createX  s-     5 	)D""4(	)ru   c                ^    || j                   vsJ |j                         | j                   |<   y)z?Capture one leaf scheduler node before its first loop mutation.N)r0  r9  rr   r^  s     rs   _snapshot_scheduler_nodez+_LoopStateSnapshot._snapshot_scheduler_node`  s/    33333)+)?)?)A""2&ru   c                V    || j                   vsJ |j                  | j                   |<   y)zACapture fused-node group metadata changed outside leaf listeners.N)r1  r   r  s     rs   _snapshot_fused_nodez'_LoopStateSnapshot._snapshot_fused_nodee  s*    411111'+zzt$ru   c                    t        |t              r| j                  |       |j                         D ]$  }t        |t              s| j                  |       & y)zBCapture a scheduler node boundary and all mutable leaf loop state.N)r   r   r9  r   r   r7  )rr   r   r^  s      rs   r3  z _LoopStateSnapshot.snapshot_nodej  sI    d./%%d+.." 	2B"m,--b1	2ru   c                    | j                   j                         D ]  \  }}|j                  |        | j                  j                         D ]  \  }}||_        t        |        y)z>Restore all captured loop state and fused-node group metadata.N)r0  r  r<  r1  r   r  )rr   r^  r;  r   r   s        rs   restorez_LoopStateSnapshot.restorer  sc    3399; 	)IB!!%(	)11779 	2KD%DJ+D1	2ru   N)rq  r  r   r/  r^  r   r   r  )r   r   r   r  r   rf   r   r  r  )r   r   r   r   r   r   rx  r0  r   r1  r   r  r7  r9  r3  r<  r   ru   rs   r/  r/  H  st     CT+BSBSC?  8I{7H7H84   B
2
22ru   r/  c                      e Zd ZU dZded<    ej                  e      Zded<   dZ	ded	<   e
dd
       ZddZddZddZy)_LoopMutationTrackera+  Rollback scope for speculative loop mutations during can_fuse().

    can_fuse() may speculatively reorder or reindex loops while evaluating
    whether a fusion is legal. If the final decision rejects the fusion,
    this tracker restores the original loop structure so later fusion
    candidates do not inherit a speculative layout chosen for a fusion
    that did not happen.

    The first active tracker for a SchedulerNode leaf owns that leaf's
    listener. Recursive can_fuse() calls reuse the outer listener instead of
    installing nested listeners, so the captured state is the original state at
    the outermost decision boundary.

    Usage: call finish(commit=True) to keep mutations, or finish(commit=False)
    to restore the original state. If no mutation occurred, finish() is a no-op.
    r  rq  r   zOrderedSet[SchedulerNode]watched_nodesNz_LoopStateSnapshot | Noner;  c                    t        |      } | t        |            }|D ]9  }|j                         D ]$  }t        |t              s|j                  |       & ; |S )z?Create a rollback scope and watch mutable leaf scheduler nodes.)rq  )r   r   r   r   r   watch)rx   rq  seentrackerr   r^  s         rs   r  z_LoopMutationTracker.create  s^     % E$K( 	&Dnn& &b-0MM"%&	& ru   c                v    |j                   y| j                  j                  |       | j                  |_         y)z<Install this scope as the mutation listener for a leaf node.N)r  rA  r  trackr6  s     rs   rC  z_LoopMutationTracker.watch  s1    %%1r"%)ZZ"ru   c                    || j                   v sJ | j                  yt        j                  | j                        | _        y)z?Lazily snapshot candidate roots when the first mutation occurs.N)rA  r;  r/  r  rq  r6  s     rs   rG  z_LoopMutationTracker.track  s;    T'''''::!
 (..tzz:
ru   c                   | j                   D ]	  }d|_         |r| j                  y| j                  j                          y)z<Detach listeners and restore captured state if rolling back.N)rA  r  r;  r<  )rr   rollbackr^  s      rs   finishz_LoopMutationTracker.finish  s>    $$ 	.B)-B&	.4::-

ru   )rq  r  r   r@  r=  )rJ  r}   r   r  )r   r   r   r   r   r   r   r   rA  r;  r   r  rC  rG  rK  r   ru   rs   r@  r@  {  sZ    " )(/@{/@/@"0M,  (,E$+ 0
;ru   r@  c                  .    e Zd ZdZddZd fdZddZddZddZddZ	ddZ
edd	       Zej                  dd
       ZddZddZddZddZddZddZddZddZ	 	 	 	 ddZddZddZddZddZddZddZddZddZ	 	 	 	 ddZ	 d	 	 	 	 	 	 	 ddZ 	 	 	 	 	 	 ddZ!	 	 	 	 dd Z"dd!Z#	 	 	 	 	 	 	 	 	 	 dd"Z$dd#Z%	 d	 	 	 	 	 dd$Z&	 	 	 	 	 	 dd%Z'dd&Z(	 	 	 	 	 	 	 	 dd'Z)	 	 	 	 	 	 	 	 dd(Z*	 	 	 	 	 	 dd)Z+	 	 	 	 	 	 	 	 	 	 dd*Z,	 	 	 	 dd+Z-	 	 	 	 dd,Z.	 	 	 	 	 	 dd-Z/e0	 	 	 	 	 	 	 	 dd.       Z1ddd/Z2dd0Z3	 	 	 	 	 	 	 	 dd1Z4	 	 	 	 	 	 	 	 	 	 	 	 dd2Z5dd3Z6	 	 	 	 	 	 dd4Z7	 	 	 	 	 	 dd5Z8	 	 	 	 	 	 dd6Z9	 	 	 	 	 	 	 	 dd7Z:	 	 	 	 	 	 dd8Z;	 	 	 	 	 	 	 	 dd9Z<	 	 	 	 	 	 dd:Z=	 	 	 	 	 	 dd;Z>	 	 	 	 	 	 dd<Z?	 	 	 	 	 	 dd=Z@dd>ZA	 	 	 	 	 	 	 	 dd?ZB	 	 	 	 	 	 dd@ZC	 	 	 	 	 	 ddAZD	 	 	 	 	 	 ddBZE	 	 	 	 	 	 	 	 	 	 ddCZF	 	 d	 	 	 	 	 	 	 	 	 ddEZG	 	 d	 	 	 	 	 	 	 	 	 ddFZH	 	 	 d	 	 	 	 	 	 	 	 	 	 	 ddGZIddH	 	 	 	 	 	 	 ddIZJ	 	 	 	 	 	 	 	 ddJZKe0ddK       ZLdDdL	 	 	 	 	 	 	 ddMZM	 	 	 	 	 	 ddNZNe0ddO       ZO	 	 	 	 	 	 	 	 ddPZPe0ddQ       ZQdddRZR	 	 d	 	 	 	 	 	 	 	 	 ddSZS	 	 d	 	 	 	 	 	 	 	 	 ddTZTeU	 	 	 d	 	 	 	 	 	 	 	 	 	 	 ddU       ZVeU	 	 	 d	 	 	 	 	 	 	 	 	 	 	 ddV       ZV	 	 	 d	 	 	 	 	 	 	 	 	 	 	 ddWZV	 	 	 	 	 	 ddXZW	 	 	 	 	 	 ddYZX	 	 	 	 ddZZY	 	 	 	 dd[ZZdd\Z[dd]Z\dd^Z]	 	 	 	 dd_Z^dd`Z_ddaZ`ddbZa	 	 	 	 	 	 ddcZbdddZceddde       Ze	 	 	 	 ddfZf	 	 ddgZg	 	 	 	 ddhZh	 	 	 	 	 	 ddiZi	 	 	 	 	 	 ddjZj	 	 	 	 	 	 ddkZk	 	 	 	 ddlZl	 	 	 	 ddmZm	 	 	 	 ddnZn	 	 ddoZo	 	 	 	 	 	 ddpZpddqZqddrZr	 	 	 	 	 	 ddsZs	 	 	 	 	 	 ddtZt	 	 	 	 	 	 dduZuddvZvddwZw	 	 	 	 ddxZxddyZyddzZzdd{Z{edd|       Z|edd}       Z}dd~Z~ddZddZ xZS )r  z
    A Scheduler is a graph of BaseSchedulerNodes. It is responsible for
    optimizations such as fusion, reorder, and graph partition.
    c                f    t        d      5  | j                  |       d d d        y # 1 sw Y   y xY w)NzScheduler.__init__)r   _initrr   rq  s     rs   r:  zScheduler.__init__  s,    ./ 	JJu	 	 	s   '0c           
         t                     t        j                  _        i  _        t        t               _        t        j                          _        t                _        t        g t        j                  j                  j                         t        j                  j                   j                         t        j                  j"                  j                                _        |D cg c]  } j'                  |       c} _        d  _        d  _         j/                           j$                  j1                  t        j                  j                   j                                 j(                  D ]  }|j3                           d  _         j7                          _         j(                  D ci c]  }|j;                         | c} _         j(                  D ci c](  }|j?                         D ]  }|j;                         | * c}} _          j<                  jC                          _"        i  _#        i  _$        t                _%        tM        jN                   j(                   j@                   jD                         _         jQ                           jS                   j(                         _         jU                           j(                  D ci c]  }|j;                         | c} _"         jW                           jY                          tZ        xj\                  t_         j(                        z  c_.        ddl0m1}m2}  | j(                         t_         j(                         _3         ji                           jS                   j(                         _        t        tj        tl        tl        f              _7        tp        jr                  $tq        jr                   j(                         _        tp        jt                  r'ddl;m<} |j{                           jW                          i  _>        i  _?        d _@        i  _A         j                           j                   j(                         _        tp        j                  $tq        j                   j(                         _        t        d  j(                  D              r jU                           j                           j                          tp        j                  stp        j                  r<t               r2t        j                  j                  j                  j                          tp        j                  r)t        ddd	      5   j                  d 
       d d d         j                          tp        j                  rddlUmT}  | j(                   j@                   jD                  t        t        j                  j                  j                               t        t        j                  j                                      _        tp        j                  stp        j                  rtp        j                  s#ddlUmY}	  |	 j(                   j@                         t        j                  rFd}
 j(                  D ]  }t        |j                        sd}
 n |
rddl&m^}  | j(                         t        j                  rddl`ma}  |dd  fd       tM        j                   j(                         _         j                          tp        j                  rttp        j                  j                  rZtp        j                  j                  r@ j                   j(                         _         j                   j(                         _         j                          t        j                  jp                  j                  j                  r j                           | j(                         t        j                  j                   j(                          j                          t                _q        i  _r        d  _s        t        d      j                   fd       t                _v        y c c}w c c}w c c}}w c c}w # 1 sw Y   1xY w)Nr   )log_ir_post_fusionlog_ir_pre_fusionr   )distributed_autotuneFc              3  <   K   | ]  }t        |t                y wr~   )r   r*  r  s     rs   r   z"Scheduler._init.<locals>.<genexpr>/  s       
 tAB
r  z#Scheduler.create_combo_kernel_nodesTlog_pt2_compile_eventlog_waitcounter)num_ck_nodes)reorder_for_peak_memory)1assign_memory_planning_info_for_scheduler_buffers)6align_runtime_estimations_across_all_distributed_ranks)trace_structuredartifactc                     dddS )N#scheduler_nodes_before_comm_overlapstring)r   encodingr   r   ru   rs   r8  z!Scheduler._init.<locals>.<lambda>u  s     E$,) ru   c            
         dj                  t        j                        D  cg c]0  \  } }d|  d|j                         z   d|j	                          z   2 c}}       S c c}} w )Nz

zsnode[r  z buffer_names:)rd  r  rq  r  r  )r  r  rr   s     rs   r8  z!Scheduler._init.<locals>.<lambda>y  sl    v{{
 )2$**(=	 !%1 %QCqMkkm, .q/A/A/C.DEF( s   5A"
)metadata_fn
payload_fngraph_statsc                 ^     j                    j                  t         j                        dS )N)graph_idnum_nodes_before_fusionnum_nodes_after_fusion)post_grad_graph_idnum_orig_nodesr   rq  rq   s   rs   r8  z!Scheduler._init.<locals>.<lambda>  s'     33+/+>+>*-djj/ ru   )wr  r:  ra   r   r  backendsr  _post_grad_graph_counterrj  r  count_graph_partition_counterr   r  rX  r!  	constantstorchbind_constantsr  create_scheduler_noderq  previous_nodecurrent_nodeupdate_zero_dim_cpu_tensorr  r  default_device_contextget_donated_buffersr  r  rI  r<  r  copyr  r  r)  seen_template_fusionsr,   decide_global_ordering_of_commsrd   topological_sort_scheduledead_node_eliminationcompute_ancestorscompute_input_distancesr1   ir_nodes_pre_fusionr   torch._inductor.debugrQ  rR  rk  create_foreach_nodesr   r,  logged_slow_fusionr-   _pre_fusion_custom_passdistributed_max_autotune_gemmrX  rS  schedulerv  buff_to_stream_multi_stream_nodesstream_idx_to_user_obj_idx_populate_stream_assignmentsr  _post_fusion_custom_passr  r\  finalize_multi_template_buffersmax_autotune_gemmmax_autotuner   r`  ra  select_algorithmPrecompileThreadPoolshutdown_instancecombo_kernelsr   create_combo_kernel_nodes_enforce_conditional_orderingrY  memoryget_output_namesdeterministic reorder_for_compute_comm_overlaprZ  r.   6runtime_estimations_align_across_all_distributed_ranksrY   r   r[  reorder_sink_verbose_loggingtorch._loggingr\  $reorder_compute_and_comm_for_overlapprocess_grouped_nodesgraph_partitionr   re   %reorder_for_reducing_graph_partitions&maybe_reorder_for_minimizing_partition,reorder_for_partition_with_simple_dependencycompute_last_usagetest_configstrack_memory_lifecycleinsert_memory_check_nodesr  graph_diagramdebug_draw_graphbuffer_names_to_freeorigin_to_index_current_stream_ctxr#   add_rowremoved_ops)rr   rq  r  r   r   rQ  rR  rS  rY  rZ  has_collectivesr[  r\  r  s   `            rs   rN  zScheduler._init  s    <>"&'?"@(1(9%5?\!&0%%**,""'') ,,113'
# >CCd003C
7;6:'')##**177+<+<+A+A+CDJJ 	DOO	 <@# $$& 	# &*ZZ;
 !AJJL!O;

 -1JJ8
$($BRBRBT8
;>CLLNC8
8
 AE@Q@Q@V@V@X 35 13 L 	"
 ::JJ##

 	!!#33DJJ?
""$<@JJ"Gq1::<?"G $$& 	##s4::6#O$**%!$**o!!#33DJJ?
",U38_"="?))577

CDJ//. ))$/""$ =?.0). :<'))+__TZZ0
**688DDJ 



 
 &&(,,.$$(;(;&(OO,,AASSU5&* $ B
 ..D.AB 	**, ))70

  ''177//44671773356DJ ##(O(O11UAJJ 0 0 RR"' JJ D$TYY/*. # K4::V 88; !  CCDJJODJ""$ ""((CCDDTZZPDJJJ4::VDJ!??!!..EE**,4::&	djj) 6@\! :< GK '//	
 -7Ls D;
8
H #HhB Bs$   5b9?b>.-c(c	ccc                   i }t         j                  j                  D ]d  }t        t         j                  j                  |   t        j
                        s9t        | t         j                  j                  |   d       ||<   f |S )N)r  )ra   r   graph_inputs_originalr   r0   DonatedBufferr   )rr   name_to_donated_bufr   s      rs   rw  zScheduler.get_donated_buffers  sp     GG11 	D!''77=r?O?OP,BGG11$7 $-#D)	 #"ru   c                v  
 ddl m
 i }t        j                  d      }| j                  D ]  }
}|j
                  D|j
                  j                         }|(||vrt        |      }|||<   || j                  |<   ||   }|| j                  |<   |j                         D ]  }|| j                  |<     t        
fd| j                  j                         D              rt        d | j                  D        d      }|| j                  D ]z  }|j
                  }	|j                          t        |	t         j"                        s;t        |	j$                  t         j&                        s`t!        j&                  |      |	_        | t        
fd| j                  j                         D              | _        y)a=  Populate node_to_stream and buff_to_stream from IR node stream_idx.

        Reads the stream_idx field set on IR nodes during lowering to determine
        which stream each scheduler node should run on. This field is propagated
        from 'custom.stream' FX node metadata via IRNode.current_stream_idx().
        r   )DEFAULT_STREAM_IDXNc              3  (   K   | ]	  }|k7    y wr~   r   )r   r9  r  s     rs   r   z9Scheduler._populate_stream_assignments.<locals>.<genexpr>  s     M1q&&M   c              3  ^   K   | ]%  }|j                         |j                          ' y wr~   r   r  s     rs   r   z9Scheduler._populate_stream_assignments.<locals>.<genexpr>  s!     RAq||~7QRs   --r  c              3  (   K   | ]	  }|k7    y wr~   r   )r   
stream_idxr  s     rs   r   z9Scheduler._populate_stream_assignments.<locals>.<genexpr>  s      '
 ,,'
r  )stream_constantsr  r  rn  rq  r   get_stream_idxr  r  rv  r  r  r  r   r   r   r0   Bufferr  rD   r  )rr   user_obj_to_stream_idxstream_idx_counterr   r  user_obj_idxnew_stream_idxr   r  r	  r  s             @rs   r  z&Scheduler._populate_stream_assignments  s    	9 24&__Q/JJ 	6D+Jyy$#yy779+#+AA)-.@)A?M.|<JV77G!7!EJ(2D% ,,. 6+5##C(6!	6, M0C0C0J0J0LMMRRTXF ! JJ FD"iiG)1&w		:&w~~r}}E *,f)EF $' '
"1188:'
 $
 ru   c                    | j                   S )z7Check if any nodes are assigned to non-default streams.)r  rq   s    rs   _has_multi_stream_nodesz!Scheduler._has_multi_stream_nodes  s    '''ru   c                    | j                   j                  ||      }| j                  j                  || j                  j                  |d            S )zAReturn the stream index for a buffer, resolving mutation renames.r   )r)  r  r  )rr   r  reals      rs   get_buf_streamzScheduler.get_buf_stream  sF    $$((8<""&&tT-@-@-D-DXq-QRRru   c                    | j                         sy| j                  |      | j                  j                  |d      k7  S )zTrue if buf_name was produced on a different stream than node.

        Resolves mutation renames so that mutated buffers inherit the
        stream of their original definition.
        Fr   )r  r  rv  r  )rr   r  r   s      rs   r  z!Scheduler.has_cross_stream_hazard  s<     ++-""8,0C0C0G0Ga0PPPru   c                6    t         j                  j                  S r~   ra   r   current_devicerq   s    rs   r  zScheduler.current_device  s    ww%%%ru   c                .    |t         j                  _        y r~   r  r  s     rs   r  zScheduler.current_device
  s    !'ru   c                    t         j                  j                  dd      dk(  rddlm}  || j
                  d       yy)z,Generate an image of the graph for debuggingINDUCTOR_WRITE_SCHEDULER_GRAPHN1r   )draw_buffersT)print_graph)osenvironr  r  r  rq  )rr   r  s     rs   r  zScheduler.debug_draw_graph  s1    ::>>:DASH+6 Iru   c                    t         j                  t        j                        r8t         j	                  d|       | j
                  D ]  }|j                           y y )Nz%s:)rT  isEnabledForloggingINFOrj  rq  rk  )rr   labelr   s      rs   debug_print_nodeszScheduler.debug_print_nodes  sF    GLL)HHUE"

 #  "# *ru   c                6   |j                         J d       |j                         rt        | |      S t        |t        j
                  t        j                  f      rt        | |      S t        |t        j                        rt        | |      S t        |      )Nz2All nodes passed to scheduling must have an origin)r  is_no_opr  r   r0   r   r!  r   r  rR  r  r  s     rs   rr  zScheduler.create_scheduler_node  s    !- 	
@	
- ==?)$55r00"2C2CDE t,,boo.,T488%d++ru   c                   t               }g }| j                  j                         }t        j                  j
                  j                         D ]  }|D cg c]%  }||v rt        | j                  |   t              s|' }}|s6|j                  |       |D cg c]  }| j                  |    }}t        j                  dkD  }t        | |d|      }|j                  |       |D ]  }|| j                  |<     | j                  D 	cg c]  }	|	j!                         |vs|	 c}	t#        |      z   | _        y c c}w c c}w c c}	w )Nr   F)rX  r[  )r   r  r!  ra   r   listsr   r   rI  r  r  r-   combo_kernels_autotunerB  r   rq  r  r   )
rr   removed_node_namesfe_nodeskept_node_namesnamesr   r   r[  fe_noder   s
             rs   r  zScheduler.create_foreach_nodes(  sN   .8l11668WW]]))+ 	8E "?*"4#4#4T#:<RS E  %%e,:?@$d''-@F@$;;a?O0*/ /	G OOG$ 807''-81	88 "ZZ
4==?BT+TD
N
5 A
s   *D<EE#Ec                *   '()  G 'fddt         t                 't        j                  '      ( j                  D ]  }|j                         D ]  }|j                         }t        |j                  j                  t        j                        rt        |j                               dkD  r^|j                         D ]J  }|(v r/|(v r+(|   }(|   }||z   }(D ]  }(|   |u s(|   |u s|(|<    6|(v r	(|   (|<   C(|   (|<   L   d) fd)	 	 d	 	 	 	 	 	 	 	 	 d()fd}	i }
t        j                  j                   j#                         D ]  }t        |t$        j&                        r|j(                  D ]  }d|
|<   	 4t        |t        j*                        sO|j-                         D cg c]  }t        |t$        j&                        s|! }}|D ]  }|j(                  D ]  }d|
|<   	   d} j                  D ]s  }|j                  J t/        |j                  j1                         d 	      }|D ]8  }t        |t$        j2                        sJ d
}||
vs&|j                         |
|<   : u  j                  D ]@  }t4        j7                  d|j                         |r|j                  J t/        |j                  j9                  d
      d 	      }|D ]d  }||
v sJ | d|
        |
|   x} j:                  |   j                         D ]*  }|j=                  t?        |j                                      , f t        |j@                  jB                        dk(  rGtE        tG        |j@                  jB                              x}rt        |tH              r|jJ                  }nd}|j                         D ]7  }t        |jM                               dk  sJ |jM                         D ]  } )|      } |	||       |j=                  t?        ||             (|   jN                  D ]  }|j                         |j                         k(  r%t        |j                  tP              sJ |j                  j                         D ]c  }|j                         } )|      }||j                         v }|j=                  tS        ||j                         |               |	||d
       e   : t        j                  jT                  |j                            D ]8  } |	||d
       |j=                  tS        ||j                         d
             : t        j                  jV                  |j                            D ]'  } |	||d       |j=                  t?        |             ) |j@                  jX                  D ]6  }t        |tR              r |	|jZ                  ||j]                  |             8 |j_                   j`                         |j                         D ]  }|jM                         D ]y  }|j                          j`                   )|      <   |j                          j`                  |<    jb                  je                  ||       jb                  |j                         <   {  C t        j                  jg                         D ]3  }t4        j7                  d|        |	|ti        t?        |                   5 |rt        j                  jj                  D ]  }|j9                  d
      D ]|  }||
v sJ | d|
jm                                 |
|   x}s) j:                  |   jo                         D ]4  }t4        j7                  d||        |	|ti        t?        |                   6 ~   j`                  D ]  }|t        j                  j                   v rE |	|ti        t?        |                   t        j                  jp                  js                  |       d|t        j                  jt                  v s |	|ti        t?        |                    tw        t        j                  j                   jm                               D  ci c]  \  } }|| 
 }!} }t        j                  jp                  D cg c]  }|!|   	 c}t        j                  _<         j                  D ]C  }|j                         D ].  }|j{                  (|j                            jN                         0 E  j|                  D ]-  } j|                  |   j{                  (|   jN                         / t               }"|"j                  d       (jO                         D ]]  \  }}#|"j                         5  |#jN                  D $cg c]  }$|$j                          }%}$|"j                  d| d|% d       ddd       _ |"j                  d       |"j                         j                         }&t        j7                  d       t        j7                  d|&       yc c}w c c}} w c c}w c c}$w # 1 sw Y   xY w)zi
        Create dependency edges between nodes, handling aliasing and
        mutation properly.
        c                  >    e Zd ZdZ	 	 d	 	 	 	 	 ddZddZd	 fdZy)
1Scheduler.compute_dependencies.<locals>.DedupListan  
            This data structure behaves like a list except it makes sure the
            elements remain unique.
            Normally one could use a OrderedSet/dict for this purpose however
            the list in question gets elements appended as it is being
            iterated over which means that we need to keep the list
            semantics.
            Nc                @    |xs g | _         |xs
 t               | _        y r~   )r  r   
membership)rr   r  r  s      rs   r:  z:Scheduler.compute_dependencies.<locals>.DedupList.__init__\  s    
 #[b
","<
ru   c                    || j                   v ry | j                  j                  |       | j                   j                  |       y r~   )r  r  r   r  )rr   	node_users     rs   r   z8Scheduler.compute_dependencies.<locals>.DedupList.appendd  s5    /

!!),##I.ru   c                    t        j                  | j                  |j                        }| j                  |j                  D cg c]  }|| j                  vs| c}z   } ||      S c c}w r~   )r   r  r  r  )rr   r  new_membershipr  	new_items	DedupLists        rs   __add__z9Scheduler.compute_dependencies.<locals>.DedupList.__add__j  sc    !+!1!1$//5CSCS!T JJ${{*at.FA* 	 !N;;*s   A+A+r  )r  zlist[_T] | Noner  zOrderedSet[_T] | Noner   r  )r  rh   r   r  )r  DedupList[_T]r   r  )r   r   r   r   r:  r   r  )r  s   rs   r  r  R  s;     *.48=&= 2= 	=/<ru   r  r   c                N    | j                   v r j                   |          S | S r~   )r)  )r  ry  rr   s    rs   ry  z.Scheduler.compute_dependencies.<locals>.rename  s,    D)))d33A677Hru   Fc                P     |          j                  t        |||             y r~   )r   r1  )used_by_namer  r  r  name_to_usersry  s       rs   add_userz0Scheduler.compute_dependencies.<locals>.add_user  s)     &./66K9ru   Nc                    | j                   S r~   rt  r  s    rs   r8  z0Scheduler.compute_dependencies.<locals>.<lambda>  s
    AFF ru   r  Tzscheduling %s)unbacked_onlyc                    | j                   S r~   rt  r  s    rs   r8  z0Scheduler.compute_dependencies.<locals>.<lambda>  s
    !&& ru   z not in )r  mutating_bufr  )r  )r  zscheduling output %sz+scheduling output %s for unbacked symint %sr  'z': r  r  zBUFFER USER LIST
z===== AFTER SCHEDULING =====
%s)r  r,  r   r,  )FF)
r  r,  r  r  r  r}   r  r}   r   r  )Er	   rh   rW  r   rq  r<  r  r   r   r  r0   rD   r   r  ra   r   rX  r   r  ra  r&   	TensorBoxrc  r  get_unbacked_symbol_defsSymbolrT  r  r
  rI  r  r:   r   rQ  r  r  r9   r  r  r  rf   r;   additional_buffer_depsadditional_star_depsr   r   r  r{  r)  r  r  r  r  r   r!  r  mutated_inputsr  rp  r  mutated_input_idxsr  r  rX   rP  r  r  rV  compute_dependencies_log)*rr   r   buf1	buf1_name	buf2_namelist1list2combinedr  r  unbacked_symbol_to_origin_nodevalfsr9  sym_sizehas_non_input_unbacked_defsunbacked_symbol_defsunbacked_symbol_usesrR  r   r   	node_modealt_namer  out_buf
other_nameis_aliasadd_depr  r  r   r   r   	inp_nameslogbufrk  rf  r  r,  r  r  ry  s*   `                                      @@@rs   rd   zScheduler.compute_dependenciesL  sB
   	< 	<@ @K?V?V@
 JJ 	LD((* L MMO	 tyy//?D,,./!3!%!1!1!3 LI M1i=6P -i 8 -i 8#(5=#0 >C -c 2e ;#0#5#>5=c 2> #m33@3Ki03@3Ki0LL	L<	 !&!				 6		 			
 		 		 JL&
 77''..0 
	BC#uzz*** >B9=226>C. (+||~S!Auzz9RASS! BAnn B=A6r:BB
	B ',#JJ 	HD99((( $*		224:J$  * H!!U\\222 /3+::8<215H	H  JJ V	DIIotyy1*yy,,,'-II222F(($
 . GA >> #X&D%EF> <A>>K#'#4#4Q#7#C#C#E GC --gclln.EFGG D$$++,1 d&6&6&=&=!>??S?sI.HH	 	 '') E3,,./1444 # 1 1 3 EH%h/HXt,%%ghY&GH -h 7 = = E==?dmmo=$)$))5FGGG'+yy'<'<'> EG)0)9)9);J)/
);J (073F3F3H'HH -- '$.1408L!" %ZtD%EEEEB 7799$--/J S$5 !!''4==?D"QR	S 7777H 4$6!!''"234
 ((.. F!$0TYYd.>.>t.DEF %%d&;&;< '')  # 1 1 3 H>AllnD))&*:;69llnD))(3//33HhG ++CLLN;aV	r 002 	>HII,h7Xz'(*;<=	>
 'ww,, N111E NA >> #X&D&I&I&K%LM> ;1==q=(,(9(9!(<(M(M(O NHII M ( !
 %Xz'(:K/LMNNN )) 	:Dqww+++z'$-89&&**40***z'$-89	: ,5QWW5I5I5N5N5P+Q
'E4D%K
	 
 )*(>(>&
 $IdO&
"
 JJ 	CD'') CmCLLN;AABC	C // 	SD''-77d8K8Q8QR	S  !c'--/ 	4JC 4/4{{;!;;#c%234 4	4 	c  "))+ &&';< &&'I3OK TX
&
" <4 4s6   5i4i42i9i?j	j5j	j		j	c           
         ddl m}m}m}m} t        t        j                  j                  j                               } | j                  |      }t        j                  j                  j                  s | j                   j                         t        t        j                  j!                               } | j                  ||      \  }}	}	t#        t%         j                              D 	cg c]  }	g g f c}	|D ]}  }
|
j&                  dk(  r|
j(                  dk(  r"|
j*                  j-                         }|
j.                     d   j1                  |       |
j2                     d   j1                  |        ddlm}  |        	 	 	 	 	 	 d fd}g }t9         j                        D ]H  \  }}|j1                  |       |j1                   |||t%         j                        dz
  k(               J | _
        y c c}	w )Nr   )rZ  compute_memory_timelineFreeableInputBufferget_freeable_input_bufr   )register_check_mem_opc                X   |    d   }|    d   }|||g}t        j                  t        t        j                  d            t        j
                  j                  j                  j                  g |d       }dj                  |    j                          |_        t        |      S )Nr   r   r  r  c                $    | |d   |d   |d   dfS )Nr   r   r   )alivedeadis_final_stepr   )tensor_argsr  s     rs   r8  zWScheduler.insert_memory_check_nodes.<locals>.construct_mem_check_node.<locals>.<lambda>  s*    !.q!1 -a 0)6q)9C ru   )r  r  r!  nontensor_argsunflatten_args
mem_check_)r0   MemoryCheckKernelrD   r`  r  ops_inductor_debugcheck_memory_stepdefaultrq  r  r  rR  )step_idxr   expected_newly_aliveexpected_newly_deadr"  r   rr   step_allocs_deallocss         rs   construct_mem_check_nodezEScheduler.insert_memory_check_nodes.<locals>.construct_mem_check_node  s     $8#A!#D "6x"@"C24GWN''!e)<=yy00BBJJ- D %/tzz(/C/L/L/N.O"PD,T488ru   )r   )r*  r   r   r}   r   rR  )r  rZ  r  r  r  r   ra   r   rX  r!  rq  r`  ra  r-   rY  r  r  rL  r   
size_alloc	size_freer"  r  
start_stepr   end_step#torch._inductor.runtime.debug_utilsr  r  )rr   rZ  r  r  r  rX  name_to_freeable_input_bufr   buf_info_listr=  buf_infor  r  r.  	new_nodesr  r   r-  s   `                @rs   r  z#Scheduler.insert_memory_check_nodes[  s   	
 	
 )31773G3G3L3L3N(O"4::|< 	# %%===

D,, *4AGG4L4L4N)O5JJ&
q! $C

O4C
RHC
 & 	HH""a'H,>,>!,C//1H !4!45a8??I !2!23A6==hG	H 	N	9	9*.	9&	92 	 , 	GAtT"(1DJJRS@S;SU	 
eC
s   3Hc                |  	 t         j                  syg }t        | j                        D ]  }dd	d}|j	                         D ]  }t        	fd|j                  D              }|r\t        j                  d|j                                t        j                  j                  j                  |j                                d} |j                          xr | }|s|j                  |       t        j                  d|j                                t        j                  j                   j                  |j                                |j"                  j$                  D ]  }|j&                  | j(                  v s| j(                  |j&                     j                  }|D cg c]0  }|j*                  j                         |j                         k7  s/|2 c}| j(                  |j&                     _          t-        t        |            | _        | j                  D ]  }|j/                           yc c}w )	z0
        Remove any nodes without users
        Nc                r    | j                   xs* | j                         t        j                  j                  v S r~   )r  r  ra   r   r  )r  s    rs   can_eliminate_userz;Scheduler.dead_node_elimination.<locals>.can_eliminate_user  s&    ||Tt}}!'':T:T'TTru   Fc              3  .   K   | ]  } |        y wr~   r   )r   ur:  s     rs   r   z2Scheduler.dead_node_elimination.<locals>.<genexpr>  s     #Ma$6q$9#M   zremoved dead buffer: %sTzremoved dead operation: %s)r  r1  r   r}   )r-   use_dcer   rq  r<  r   r  rT  r  r  ra   r   r  r  r  r   r  r   r   r   r  r   r   r  )
rr   updated_nodesr   active_buffersr   can_eliminater  r  r<  r:  s
            @rs   r|  zScheduler.dead_node_elimination  s    ~~
 TZZ( 	DU #N'') * ##M399#M M II7HGG++//?%)N* !% 5 5 77N<NM $$T* 		6H**..t}}? ,,22 DyyD$4$44 $ 0 0 ; A A',="#0AT]]_0TA=((39-	8 (=12
 JJ 	#D  "	#=s   %0H9H9c                
    |duS )z:Check if store mode requires cross-thread synchronization.Nr   )rr   r  s     rs   mode_requires_synchronizationz'Scheduler.mode_requires_synchronization  s    4ru   c                    t        t                  t               g dfd|D ]  }|j                         D ]  }||<   	  |D ]
  } |        S )z?
        Ensure nodes is in topologically sorted order
        c                    | vrdj                  |        t        | j                  d       D ]&  }|j                  vr |j                            ( j	                  |        y y )Nc                    | j                   S r~   rt  )ds    rs   r8  zDScheduler.topological_sort_schedule.<locals>.visit.<locals>.<lambda>  s
    aff ru   r  )r  r  r3  r   r   )r  r   rI  r  rD  visits     rs   rH  z2Scheduler.topological_sort_schedule.<locals>.visit  se    }!!"6"6<LM 2Cxx|3 ,sxx01	2
 a  ru   )r  rf   r   r  )r   rf   rx  r  )rr   rq  r   r   rI  r  rD  rH  s       @@@@rs   r{  z#Scheduler.topological_sort_schedule  sy     +,.59V*,	! 	!  	*D--/ *%)T"*	*  	D$K	ru   c                   | j                   D cg c])  }t        |j                  t        j                        s(|+ }}t        dt        |            D ]o  }t        t        ||   j                                     }t        t        ||dz
     j                                     }||   j                  t        ||d             q y c c}w )Nr   Tr  )rq  r   r   r0   ConditionalrL  r   r  r  r  r  r;   )rr   r  conditional_nodesr  r  prev_bufs         rs   r  z'Scheduler._enforce_conditional_ordering  s    zz
Z%GA
 
 q#/01 	A%6q%9%J%J%L MNLD!21q5!9!J!J!LMNHa --|TJ	
s
   )CCc                <    t               }t        |t        t        t        t
        t        f      r-|j                  D ]  }|j                  |j                          nt        dt        |       d       fd|D        }t        t         fd|D                    S )Nz+get_unmet_dep_nodes is not implemented for .c              3  X   K   | ]!  }j                   |   j                          # y wr~   )r  r  r  s     rs   r   z1Scheduler._get_unmet_dep_nodes.<locals>.<genexpr>  s%     Xc))#.??AXs   '*c              3  <   K   | ]  }j                   |     y wr~   r  )r   r  rr   s     rs   r   z1Scheduler._get_unmet_dep_nodes.<locals>.<genexpr>  s     Qat66q9Q   )r   r   r   rR  r  r   r  r3  r  r   RuntimeErrorr   r   )rr   r  
unmet_depsr   unmet_dep_opss   `    rs   _get_unmet_dep_nodeszScheduler._get_unmet_dep_nodes  s    &0l
)&"$	
 // )sxx() =d5k]!L  YZXJQ=QQRRru   c                z   g }t         j                  | j                  d      }i }| j                  D ]P  }| j                  |      }t	        |      ||<   |D ]*  }|j                  |g       }|j                  |       |||<   , R |j                         D 	cg c]  \  }}	|	dk(  s| }
}}	|
rx|j                  |
       |
D ]7  }|j                  |g       D ]  }||xx   dz  cc<    |j                  |       9 |j                         D 	cg c]  \  }}	|	dk(  s| }
}}	|
rx|rJ d       |S c c}	}w c c}	}w )zU
        Sort nodes by their topological order, return a list of node lists.
        r   r   zTopological sort failed!)	rx  fromkeysrq  rV  r   r  r   r  r  )rr   r  rq  childrenr   r  r   cr  rf  zero_deg_nodesr  s               rs   rt  z!Scheduler._topological_sort_nodes  sF    djj!,#%JJ 	"D,,T2Dd)E$K "LLb) !"	" ).@1a!@@LL(# $LLB/ %D$K1$K%		! -2KKMDDAqQ!VaDND  444y A Es   D1%D1D7D7c                j   i }| j                   D ]w  }t               }|j                  D ]B  }| j                  |j                     j                         }|j                  |       |||   z  }D |||j                         <   ||_        y t        | j                         D ]  \  }}||_
        ||_         y)z.
        Populate each node.ancestors
        N)rq  r   r3  r  r   r  r  r  r  r  r%  r&  )rr   name_to_ancestorsr   r  r   dep_node_namer  s          rs   r}  zScheduler.compute_ancestors.  s    
 9;JJ 	'D)3I.. > $ 0 0 : K K Mm,.}==	> 2;dmmo.&DN	' %TZZ0 	#KE4"DN"DN	#ru   c                   i }i }| j                   D ]  }|j                  sd}d}n|j                  D cg c]/  }|| j                  |j                     j	                            dz   1 }}|j                  D cg c]/  }|| j                  |j                     j	                            dz   1 }}t        |      }t        |      }|||j                         <   |||j                         <   ||_        ||_	         yc c}w c c}w )z
        Populate each node's min/max_input_distance with the depth from graph
        inputs, measured as dependency hops before fusion. Nodes whose
        dependencies are all satisfied by graph inputs/constants have depth 0.
        r   r   N)
rq  r3  r  r   r  rZ  rY  r  r#  r$  )	rr   name_to_min_distancename_to_max_distancer   min_distmax_distr   dep_min_distsdep_max_distss	            rs   r~  z!Scheduler.compute_input_distancesA  s     02/1JJ 	/D**
  $66!  ))9)9#(()C)T)T)VW! !  $66!  ))9)9#(()C)T)T)VW! !
 }-}-4< 14< 1&.D#&.D#)	/
!
!s   4C:74C?c                H   t         j                  sy | j                  D ]  }t        |t        t
        f      r#|j                         st         j                  dk7  r=|j                         D ]3  }t        |t              r|j                         r$|j                          5  y )Nhalide)r-   r%  rq  r   r   r   r[   cpu_backendr   r  r\  )rr   r   r  s      rs   r\  zScheduler.merge_loops_  s    00JJ 	$D d]4F$GHKKMf&8&8H&D) $!%75;L;L;N!!#$	$ru   c                   t        ddd      5  t        d      D ]  }t        |      }t        j	                  d|dz   |       | j                  |d      }t        |      }t        j	                  d	|dz   ||       ||k(  s|dk(  slt        j	                  d
|dz           n t        j                  st        j                  r| j                  |d      }|cddd       S # 1 sw Y   yxY w)zB
        Combine eligible nodes into FusedSchedulerNodes.
        zScheduler.fused_nodesTrU  r  z/===== attempting fusion (%d/10): %d nodes =====r   F)is_reorder_roundz=completed fusion round (%d/10): fused %d nodes into %d nodes
z+===== fusion complete (%d iterations) =====N)	r   rL  r   r  r  fuse_nodes_oncer-   r%  loop_index_inversion_in_fusion)rr   rq  r  old_lennew_lens        rs   r  zScheduler.fuse_nodesz  s     #4QU
 	 2Y e*  EE
 ,,UU,Ke*  TE	 g%A$$Eq1u ', 1188,,UT,J;	 	 	s   A7C!AC!!C*c                    g }| j                   D ]4  }|j                  t        |t              r|j	                         n|g       6 || _         y)zA
        Unpack GroupedSchedulerNode into regular nodes.
        N)rq  r  r   r  r  )rr   r7  r   s      rs   r  zScheduler.process_grouped_nodes  sJ     .0	JJ 	D!+D2F!GdV	 
ru   c                    t        |      dkD  sJ |d   j                         }|| _        | j                  |      }t	        ddd      5  |j                  |      cddd       S # 1 sw Y   yxY w)
        Benchmark fused list of nodes and return the execution time
        in milliseconds on randomly generated inputs.
        r   benchmark_fused_nodesTcompile_time_autotune_time_us)rV  dynamo_compile_column_usN)r   r   r  r#  r   rr  )rr   rq  r  r  s       rs   rr  zScheduler.benchmark_fused_nodes  st     5zA~~q$$&$""6*#"&%D
 	8
 007	8 	8 	8s   
A%%A.Nc                    t        |      dkD  sJ |d   j                         }|| _        | j                  |      }t	        d      5  |j                  |||      cddd       S # 1 sw Y   yxY w)D
        Generate a kernel given a list of pre-fused nodes.
        r   generate_kernel_code_from_nodeshint_overrideN)r   r   r  r#  r   rw  )rr   rq  benchmark_kernelry  r  r  s         rs   rw  z)Scheduler.generate_kernel_code_from_nodes  sw     5zA~~q$$&$""6*;< 	::'} ; 	 	 	s   A%%A.c                    || _         | j                  |      }t        d      5  |j                  |      cddd       S # 1 sw Y   yxY w)
        Benchmark a compiled module and return the execution time
        in milliseconds on randomly generated inputs.
        benchmark_codegened_moduleN)r  r#  r   r}  )rr   moduler  r  s       rs   r}  z$Scheduler.benchmark_codegened_module  sH     %""6*67 	>55f=	> 	> 	>s	   ?Ac                   t         j                  j                  }|syt        j	                  d||       |j
                  D ]  }|j                         }t        j                  j                  j                  |      s=|j                  }||   }t        |t        j                        r'|j                  |j                          |j                  }t        |t        j"                        s||k7  st        j%                  d|||        y y)z
        Check if selecting a Triton template would cause layout conflicts.
        Returns True if there's a conflict and we should fall back to ATen.
        FzNode %s has constraints %szOLayout conflict detected for %s: template expects %s but layout is frozen to %sT)ra   r   buffer_layout_constraintsrT  r  r  r  r`  ra  r  should_use_layout_constraintsr  r   r0   FlexibleLayout freeze_layout_with_exact_stridesr  FixedLayoutrU  )rr   
multi_nodeconstraintsinpinp_namer  expected_layouts          rs   !_has_layout_conflict_for_templatez+Scheduler._has_layout_conflict_for_template  s     gg77		.
KH$$ 	C||~H??33QQRUV ZZF)(3O&""3"34 44_5K5KL&"..1o6Oe#	 3	6 ru   c           
     &   t        | j                        D ]  \  }}t        |t              st        |j                  t
        j                        s=|j                  }t        j                  j                  s%| j                  |      s|j                         \  }}n|j                  D cg c]2  }t        |t        j                  j                  j                         r|4 }}|sJ d       t#        |      dkD  r!|j%                         t'        |fd      }n|d   }t        |t        j                  j
                  j(                        rt        j*                  ri }||d<   t        j*                  D ]k  }	|j%                  |	      j-                         D 
ci c]  \  }
}t        |
t(              r|
| }}
}t'        |j-                         d       d   }|||	<   m |j                  j/                  |       n|j                  j1                  |       t
        j2                  j5                  |j6                        5  |j9                         }ddd       j:                  }t        |t
        j<                        sJ |j:                  }t        |t
        j>                        sJ |j@                  rtC        ||j@                         |jD                  |_"        | jG                  ||||        yc c}w c c}}
w # 1 sw Y   xY w)	a  
        Finalize a backing choice for MultiTemplateBuffers which did not already have a
        choice finalized through fusion. In the case of an extern choice, this will result
        in replacing the SchedulerNode.

        If a MultiTemplateBuffer did not have any fusion opportunities, finalizing a choice
        will force completion of compilation and benchmarking.
        z&No extern kernel detected for fallbackr   c                    |    S r~   r   )rZ  timingss    rs   r8  z;Scheduler.finalize_multi_template_buffers.<locals>.<lambda>"  s    WUVZ ru   r  r   Nrx  c                    | d   S rH  r   r  s    rs   r8  z;Scheduler.finalize_multi_template_buffers.<locals>.<lambda>7  s    qQRt ru   )$r  rq  r   r   r   r0   MultiTemplateBufferr-   r  %force_extern_kernel_in_multi_templater  get_min_choicechoicesr`  ra  r  ExternKernelCallerr   choice_timingsrZ  r"   multi_kernel_hintsr  finalize_as_triton_callersfinalize_as_triton_callerrl  current_originsr$  output_noder   
StorageBoxOperationBufferorigin_noder?   r  _replace_node)rr   r  r   r  min_node_unfusedr=  rZ  extern_choicescallersr  r  rf  triton_timingschoiceout_tensorboxout_storage
out_bufferr  s                    @rs   r  z)Scheduler.finalize_multi_template_buffers  s    !, ?	DGAt$.:		2114 "YY
++QQ BB:N +5*C*C*E'$a ",!3!3&%!OO<<OO &N & *S+SS>>*Q.",";";"=+.~CW+X(+9!+<($OO&&??
 00NP(8 %+$=$= 3D&0&?&?d&?&SG -4MMO.$(Aq#-a1I#J !"1.N .
 &))=)=)?^%TUV%WF,2GDM3 		<<WE		;;<LMYY..z/A/AB C$4$@$@$BMC+00!+r}}===(--
!*b.@.@AAA))&}j6L6LM$.$5$5
!"":z1dC?	D&6.C Cs   -7K< L
LL	c                   t        ||       | j                  |      }|| j                  |<   || j                  |j	                         <   || j
                  |j	                         <   i t        j                  |j                  j                  |j                        D ]:  }| j                  j                  |j                  d       x}s,|j                  |<   < dfd} ||j                        |_
         ||j                  j                        |j                  _	        t        |j                         |j                               D ]3  \  }	}
|	| j                   |
j	                         <   |
j"                  |	_        5 |j$                  |_        |j&                  |_        |j(                  |_        |j*                  |_        y )Nc                ,    t        fd| D              S )Nc              3  @   K   | ]  }|j                          y wr~   )ry  )r   r   r)  s     rs   r   z?Scheduler._replace_node.<locals>.rename_deps.<locals>.<genexpr>a  s     Kscjj)9:Kr  r   )r  r)  s    rs   rename_depsz,Scheduler._replace_node.<locals>.rename_deps`  s    KdKKKru   )r  r2  r   r2  )r  rr  rq  rI  r  r  r  r  r   r   r3  r  r  r   ry  r<  r  r  r%  r&  r  r"  )rr   r  r  r  r   new_scheduler_noder   r3  r  new_outold_outr)  s              @rs   r  zScheduler._replace_nodeL  s    	"*j9!77
C*

1-?$--/*3E0 ??4#3#3#9#94;R;RS 	7C 3377$GGyG.1hh +	7	L 1<111
- 0;**000
&&, !$**,d.>.>.@!
 	*GW 4;DW--/0#MMGM		* (,~~$'+~~$'+~~$(,%ru   c                &    t        d |D              S )Nc              3     K   | ]q  }t        |j                  d       xrU |j                  duxrE t        |j                  j                  d      xr# |j                  j                  j                  dk(   s yw)r   Nr  r  )r  r   r   r  r  s     rs   r   z,Scheduler._any_atomic_add.<locals>.<genexpr>v  so      

 	 AFFF# 9d"9^49 ((L89
s   A7A9r  )rr   	node_lists     rs   _any_atomic_addzScheduler._any_atomic_addu  s     

 
 
 	
ru   c                "   | j                  |d|      }t        j                  |      }t        j                  j
                  j                         }|j                         sd }||fS |j                  d|      }t        |t              sJ ||fS )NT)rz  ry  triton_)kernel_namesource_code)rw  r!   loadr`  ra  async_compileAsyncCompileuse_process_poolr   r   r    )rr   rq  ry  src_codemodr  futs          rs   compile_kernelzScheduler.compile_kernel~  s     77D 8 
 x(55BBD--/C
 Sz  &&9(&SCc<000Szru   c                    !"#$%&'()*+,-./0 t        d fD              }t        j                         t        j                        xr t        d      }|r%t        j                  st        j                  d      S t        j                  s|st        j                  d      S j                         r(t        j                         t        j                        r j                         sj                         rt        j                  d      S j                         }|d   j                         !!sJ !j                  dk(  r(t        j                   dk7  rt        j                  d      S j                         }t#        t%        j&                  ||            / j)                  /      }|r|st        j                  d      S ddlm t/              0/d   j                         !!J dfd
)|rSt        d fD              r>j                         d	u""rj                         nj                         .t        .t        j0                        sJ  j3                  .      rt        j                  d      S i (|sWg &t        j4                  D ]A  }.j7                  |       t9         j;                         d       D ]~  \  }	}
t        |	t<        j>                  j@                  jB                        s5.jE                  |	      5  &jG                  |	g jI                  /|	jJ                               d	d	d	        tM        d      }d	}i }&D ]V  \  }	}}	 ||jO                          .jE                  |	      5   j]                  |!      \  }}|||	<   ||k  r|}|	}d	d	d	       X |.j^                  |<   t        |t`              sJ |(|<   D t        jb                  te        d .jf                  D              }ti               xr  xr |t        jj                  k  'tM        d      tM        d      c+,d	*'sR.j7                          .jm                         \  *+t9         j;                         to        jp                  d            }n.jf                  D cg c]  }|df }}ddl9m: d".fd}	 	 	 	 d"./ fd}|r/"st        j                  d      S g t        j4                  d	D ]  }.j7                  |       d	}t9         j;                         to        jp                  d            D ]R  \  }	}
 ||	      stw        jx                  t<        j>                  j@                  jB                  |	      } ||      sP|} n |t        j                  d      c S |(|<    t        j4                  r.j{                  (       n.j}                  (d	          t        j                  d      S r("r j                  |      n j                  |      \  ,}n4"st        j                  d      S j                         ,t        ,      -g &d}|D ]  \  }	} ||	      stw        jx                  t`        |	      }	r
|+,z   k\  r n\|dz  }|t        jj                  kD  r nB.jE                  |	      5  	 &jG                  |	g jI                  /             	 d	d	d	        t        &      dk(  rt        j                  d      S d !"&'()*+,-. fd}t        j                  |&d   d         S  jI                  |      # jI                  |      % jI                  /      $d!#$%) 0fd}t        j                  |$d         S # 1 sw Y   xY w# tP        $ rI}tR        jU                  tV        jX                        rtR        j[                  d"sdnd|       Y d	}~d	}~ww xY w# 1 sw Y   xY wc c}w # $ r Y d	d	d	       w xY w# 1 sw Y   xY w)
        If config.benchmark_fusion is False, always return True.
        Otherwise, return True if fusion can brings speedup.
        c              3     K   | ]>  }|j                         xr( t        |j                         t        j                         @ y wr~   )r  r   r  r0   r  r  s     rs   r   z.Scheduler.speedup_by_fusion.<locals>.<genexpr>  sE       
  MMO J1..0"2H2HIJ 
s   AAFr  Tr   r  r   CompilationErrorNc           
     t   t         j                  t        j                        r| ||z   k  rFt         j	                  dj                         j                         t        ||z   | z  d             y t         j	                  dj                         j                         t        | ||z   z  d             y y )Nz9can fuse (benchmark): fusing %s with %s cause %sx speedup.3fz=cannot fuse (benchmark): fusing %s with %s cause %sx slowdown)r  r  r  DEBUGr  r  rJ   rL   )ms_fusedr  r  r   r   s      rs   
log_fusionz/Scheduler.speedup_by_fusion.<locals>.log_fusion  s    &&w}}5cCi'$$S..0..0"sSyH&<S%AC	 $$W..0..0 Hc	$:3#?A	 6ru   c              3  @   K   | ]  }|j                         d u  y wr~   r  r  s     rs   r   z.Scheduler.speedup_by_fusion.<locals>.<genexpr>  s#      %
23A!-%
r  c                    | d   S rH  r   r  s    rs   r8  z-Scheduler.speedup_by_fusion.<locals>.<lambda>  s    RSTURV ru   r  rx  infException in compiling %s: %sr  r  c              3  <   K   | ]  }t        |t                y wr~   )r   r"   )r   rZ  s     rs   r   z.Scheduler.speedup_by_fusion.<locals>.<genexpr>  s      %<=
167%r  r   )	CantSplitc                    t        | t        j                  j                  j                        sy xr' t        | d      xr | j                  j                  k7   S )NFallowed_prologue_inps)r   r`  ra  r  TritonTemplateCallerr  r  )r  r  r  s    rs   choice_supports_fusionz;Scheduler.speedup_by_fusion.<locals>.choice_supports_fusion5  sb    !EOO<<QQ ! (' Y(?@Y44
8X8XX ru   c                   j                  |       5  	 j                  | j                        \  }}||j                          n|j                  j                          d d d        y# $ r Y d d d        yt        $ rP}t        j                  t        j                        rt        j                  dsdnd|       Y d }~d d d        yd }~ww xY w# 1 sw Y   yxY w)Nrx  Fr  r  r  T)swap_as_triton_callerr  ry  r  r  
precompilerS  r  r  r  r  r  )	r  ro   	mod_fusedr&  r  r  r  node_list_fusedrr   s	       rs   compile_without_benchmarkingzAScheduler.speedup_by_fusion.<locals>.compile_without_benchmarkingE  s      55f= %%,0,?,?+6;O;O -@ -)	 "-"MMO%--88:%&  % %$% % % %%227==A&,, ?2A
z !
  %%% %%%& s;   CAA++C0C:C>C CCCCc                    t        d      } d }i }rhrt        t        j                        sJ j	                         j                         \  D cg c]  }|d   v r| c}t        fd      D ]0  \  }}}	 ||j                         }n!s|j                  }|j                          nd }r>j!                  |      5  j#                  |      \  }	}
|	||<   |	| k  r|	} |}d d d        }|k(  xs z   |   z   kD  }|s|s|j                          |j$                  r|j&                  sJ |j$                  d   }|j&                  }|j(                  }t+        |j&                  |||j,                  j.                  t1        j2                              }|s/|} n r
 |        r| z   k  rL|Jt4        j6                  r|d <   j9                         nj;                  |       r|j<                  d <   yy	c c}w # t        $ rI}t        j                  t        j                        rt        j                  dsdnd|       Y d }~d }~ww xY w# 1 sw Y   xY w)
Nr  r   c                    | d      S r   r   )r  r  s    rs   r8  zKScheduler.speedup_by_fusion.<locals>.benchmark_when_ready.<locals>.<lambda>  s    nQqT&: ru   r  r  r  r  TF)rq  r   r0   r  r  r  r  r  r  r  rS  r  r  r  r  r  r  r}  	launchersn_regsn_spillsr  bmreqr  rH   r  r-   r  r  r  _choice_timings)min_ms_fusedms_fused_choicenew_timings
fut_choicer  ro   r  resr&  r  pathfusible_choicecompiled_kernelr  r  should_fuse_epiloguebench_epiloguer  r  r  future_choicesget_choice_timings_async hint_override_best_fusion_choicer  
min_choicer  r  	ms2_fusedr  rr   s                   rs   benchmark_when_readyz9Scheduler.speedup_by_fusion.<locals>.benchmark_when_ready  s   $U|"& +%*ZAWAW*XXX%/%>%>%@N&0&?&?&AOJ
 +9&&%a=N: #&N &,&:&N
 2@ <&-FFI!!-"(--/C!/"+"3"3CNN,"&C &'==fE 	9-1-L-L ) &.NHd
 3;K/',6/728	9 	9 '&0 N"Sy>&+AI+MM '
 >"--/#&==V]]B#B.1mmA.>O+:+A+AL-<-E-EN3A # # & , . & 6 6 0 7 7 ?40  428 %y<&| "|S#6 ',#)*D%100AP8>"==<
 #<<_M%;F
2248 }&. % !%227==A&,, ?2A
z !
 !!	9 	9s*   H4H$I.	I+">I&&I+.I8	c                    ddl m}  	 d   d   d   fD ]  }||j                           j                  d   
      \  t	        j
                        r	 d       yj                  d   
      \  t	        j
                        r	 d       yj                  d   
      \  t	        j
                        r	 d       y        t        d      rWz   k\  rOfj                  vr?j                  j                  f       t        d      j                  fd	       z   k  S # | $ r Y y	$ r}d
t        |      v rY d }~y d }~ww xY w)Nr   )NoTritonConfigsErrorr   z%register spilling of the first kernelFz&register spilling of the second kernelz%register spilling of the fused kernelslow_fusionc            	     $      z   z  dS )N)kernel1_pathkernel1_latencykernel2_pathkernel2_latencyfused_kernel_pathfused_kernel_latencyslow_down_ratior   )r  r  r  path1path2
path_fuseds   rs   r8  zKScheduler.speedup_by_fusion.<locals>.benchmark_when_ready.<locals>.<lambda>Q  s(    053605365?8@3;sSy3I% ru   Loop-carried variableT))torch._inductor.runtime.triton_heuristicsr  r  r}  r  isinfr$   r  r  r#   r  r,  )r  r  r&  r  r  r  r  r  r  r  r  future_and_mod_l1future_and_mod_l1_fusedfuture_and_mod_l2r  rr   rS  s      @@@@@@rs   r  z9Scheduler.speedup_by_fusion.<locals>.benchmark_when_ready  s   A *!,)!,/2  )
 ?JJL) "&!@!@)!,"JC
 zz#CD$!%!@!@)!,"JC
 zz#DE$+/+J+J/2,(Hj
 zz(+CD$xc2 0>$c	1"EN$2I2II//33UENC(7?? 
 $cCi//+ ! ' .#a&8#s<   E AE +5E !5E A3E E.E.E)(E))E.r{   )r  rq  r  rq  r  rq  r   r  )r  zir.ChoiceCallerr   r}   )r  z5torch._inductor.select_algorithm.TritonTemplateCallerr   r}   r  )Dr  r   r  r0   r  r  r-   r  rk   ry   benchmark_fusionr  r  r   r   r   rh  r   r  r  r  triton.compiler.errorsr  r  r  r  r  r  r  r  r`  ra  r  r  r  r   r  ry  rq  r  rS  r  r  r  r  r  r}  r  r"   benchmark_epilogue_fusionr   r  r    max_epilogue_benchmarked_choicesr  operator
itemgetterr7  r  r  r  r  r  rr  rg  r  r   r|   )1rr   r   r   is_multi_templateatomic_add_template_epiloguenode_list_1node_list_2has_atomic_addry  r  r=  r  r  r  ro   r  r&  r  r  num_triton_callerschoice_timings_iterrZ  r  r  triton_choicer  triton_choicesunfused_timer  r  r  r  r  r  r  r   r  r  r  r  r  r  r  r  r  r  r  r  rS  s1   ```                          @@@@@@@@@@@@@@@@@@@@rs   speedup_by_fusionzScheduler.speedup_by_fusion  se       
 U^ 
 

 (2##%r'>'>(
 (J.u5I 	% (0V0V$$U++&&/@$$T** u668":Q:QR!!  $$T**oo'Q**,v ;;%F$6$6($B$$T**oo'y{KHI--o> "3$$T**;u% #..0!!!	"  %
8=u~%
 "
 $557tCO # ''),,. 
 j"*@*@AAA55jA#((//  - "TV%+%>%> ,VM%/%>%>}%MN%+N,@,@,B%W 	)"!OO<<QQ  %'==fE 	*11$*!"%)%8%8(76<6J6J &9 &&!"	 	" $)<LGKO"$K5C 91	
%%1 & (==fE 9-1-L-L )6.NHd 3;K/',6/7289 99( ALJ..}=%o7OPPPFU4]CY,V\ $==N!$ %AKASAS% "
 )* R&&R&&*Q*QQ % U|U5\HC15J+!+!:!:!<",";";"=
C&,"((*0C0CA0F'# 8B7I7I&J!1v&J#&J> M 0 &',,U33%Gv'@'@%G$%G VM%/%>%>}%MNGKO%+&,,.H4G4G4J& "	  6f=$(.!OO<<QQ") 8F.;O!" '.+0077FU4]C#V& ,,998 888> $((.. ' ..{;33K@ U '',,U33224<UE3O	 QSNN(; !$-f5%=vF!lcCi&?!#!F$K$KK55f= !!&--#Kd&9&9/&JK! !!, >"a'#((//j! j! j!X  --$nQ&7&:  !% 3 3K @ $ 3 3K @&*&9&9/&J#F FP  --09PQR9S .  e	 	$  ) %)66w}}E * 0 0$C6EJ:$%!"
 %%9 9F 'K^ % ! !! !
!! !s`   !1^55_$`:`$`=$`)5^?	`>```!)`:.`=9`::`==a	c                <    | j                   |j                            S )z0Look up the node in Scheduler name_to_fused_node)r  r  r  s     rs   r  zScheduler.get_fused_nodej  s    &&t':':'<==ru   c                L   t         j                  d|j                         |j                                |j                         }|j                         |k(  sJ | j	                  |      j                  ||      }|j                  |       |j                  |       |j                  |       | j                  j                  |j                         D ci c]  }|j                         | c}       | j                  j                  |      }||| j                  |<   |S c c}w )Nzfusing %s with %s)r  r  r  r   r#  ry   r0  r  r  r  r   rv  r  )rr   r   r   r\  r  node3r  stream1s           rs   fuse_two_nodeszScheduler.fuse_two_nodesn  s     	,enn.>@PQ!!#!V+++  (--eU;5!5!&&U__EV'W

e(;'WX %%))%0)0D& (Xs   D!c                    | j                  ||      r-| j                  ||      s |       r| j                  |||       yyNTF)r  will_fusion_create_cycler  )rr   r   r   
speedup_fnr\  s        rs   fuse_if_speedupzScheduler.fuse_if_speedup  s?     MM%'11%?uk:ru   c                   |rg }i }t               }|D ]  }||v rt        ||         dk\  sJ ||   j                  d      }t        ||         dk(  r|j                  |       |j	                         \  }}	|	|k(  rt        ||	      sJ |}
n||k(  sJ t        ||	      sJ |	}
| j                  |
      |
ur|j                  r3|j                  j                  }|J |j                  |       ||f||<   | j                  ||	|j                  |      s|j                  |        t        |      D ]l  }||   \  }}| j                  | j                  |j                        | j                  |j                        |j                  |      s\|j                  |       n |D ]  }|j                  |        |ryy)z
        Evaluate pending template fusions for a set of fusion candidate nodes.
        The fusion candidate nodes are pointwise nodes as potential epilogue
        or prologue fusions
        r   r   N)r   r   r  r  r   r&  r)  r  ro   r   r  rm   r   r   r   )rr   template_fusion_candidatesr\  template_futuresfuture_to_pending_fusionfusions_to_remover  pending_fusionr   r   r  fcands                rs   "_evaluate_pending_template_fusionsz,Scheduler._evaluate_pending_template_fusions  s    )-/  % @J|7 %9	!;;6yABaGH "<I!F!J!J1!M1)<=B%)))4->>@uI%-eU;;;$)M I----eU;;;$)M &&}5]J!((&--44A=(=$++A.3A92M,Q/ ++un&@&@+ *--i8K%9P ""23 0'?'B$''''(<(<=''(<(<="..	 &))$/0 ' 2*..q12q )ru   c                    	 	 	 	 	 	 d fd}|D ]H  \  }} |||        j                  |      } j                  |      }t        ||      r||f j                  v rO j                  |||      sc j	                  ||      rv j                  ||      }	|	j                  t        |	j                  |||	j                        }
t        ||      rY||f j                  vsJ  j                  j                  ||f       t        ||      }||vrg ||<   ||   j                  |
       n
|
|<   |
|<   (|	j                  s6 j                  ||       K y )Nc                b   j                  |       v sj                  |      v rj                  j                  |       j                  j                  |                  }|J |j                         \  }}|j                  }j	                  |d        j	                  |d        j                  |      |u sJ j                  |      |u sJ  |       rj                  | |      rj                  ||       j                  |       v rj                  |      v ry y r~   )r  r  r   rm   r  r  r  )	r   r   r#  	node_key1	node_key2
is_speedupr\  pending_fusionsrr   s	         rs   resolve_pending_fusionsz<Scheduler._try_fusion_pairs.<locals>.resolve_pending_fusions  s1   
 ##E*o=&&u-@!0!4!4''.#''(;(;E(BC" &111'5'F'F'H$	9+77
##It4##It4**95BBB**95BBB!|t'D'DUE'R##Iy+F+ ##E*o=&&u-@ru   )rm   r   r   ro   r  )r  r+  ry  r  r  r  rm   r   ro   r  r-  r   rl   r  )rr   possible_fusion_pairsr,  template_fusion_nodesr\  rj  r-  r   r   
fusion_resr#  template_pw_nodes   ` ` `       rs   _try_fusion_pairszScheduler._try_fusion_pairs  s   	G$	G$	G 	G8 2 +	?LE5 $E51''.E''.E #5%0ENd&@&@@}}u.33E5A!33E5A
))5%2$.$:$:##)00	&N *%7 %u~T5O5OOOO2266u~F+B5%+P(+3HHFH12BC-.>?FF~V1?.1?.!--##E5+>W+	?ru   c                @   t               }|j                         D ]  }|j                         \  }}|j                  }||v st	        ||      r3|j                  |       | j                  |      |u sJ | j                  |      |u sJ | j                  ||||        y r~   )r   r   r   rm   r+  r  r  r  )rr   r\  r,  seen_pair_speedup_fnr#  r)  r*  is_speedup_fns           rs   _finish_pending_fusionsz!Scheduler._finish_pending_fusions.  s    
 @J| .446 	SN#1#B#B#D Iy*66M 448J99  $$]3&&y1Y>>>&&y1Y>>>  I}kR	Sru   c           
         t        |D cg c]  \  }}t        ||      s| c}}      }g }|D ]<  \  }}t        ||      r||v r|j                  ||f       *|j                  ||f       > |S c c}}w r~   )r   r&  r)  r   )rr   possible_fusionsdeferred_prologue_fusionsn1n2epilogue_template_nodesnew_possible_fusionss          rs   _handle_template_overlapz"Scheduler._handle_template_overlapF  s     #-.MFB2DR2LRM#
  "& 	6FB!"b)b4K.K)00"b:$++RH5		6 $# Ns
   A1
A1
c                   | j                  |       t        |      }t        j                  t        j
                        r@t        j                  d       |D ]&  }t        j                  d|j                                ( i }i }g }| j                  ||      }t        j                  st        j                  r| j                  ||      }| j                  |||||       | j                  ||       | j                  ||       |j!                          |r'| j                  |||||       | j                  ||       t#        |d       }| j%                  |      }|S )a  
        Combine eligible nodes into FusedSchedulerNodes.

        This relies on two key functions to control the logic:
            - self.can_fuse(): checks if a fusion is legal
            - self.score_fusion(): assigns priority to a given fusion
        zfuse_nodes_once, candidates:z  %sc                    | j                   S r~   r  r  s    rs   r8  z+Scheduler.fuse_nodes_once.<locals>.<lambda>  s
    !++ ru   r  )r  r   r  r  r  r  r  rh  get_possible_fusionsr-   r  r  r>  r2  r6  r&  clearr  r{  )	rr   rq  rj  r\  r   r,  r/  r9  r8  s	            rs   rk  zScheduler.fuse_nodes_onceZ  sk    	!!%( '""7==1;<# A  )=)=)?@A  	
 OQ  	"  44

 ##v':':#<< ";  	!	
 	$$[/B//0E{S##%$"")%  334I;W{(=>..u5ru   c              #     K   t        | fd      }|dk  r|r| yg }d}|D ],  }|   }|r||z
  |kD  r| g }|s|}|j                  |       . |r| yyw)u  Sort `nodes` by baseline index, then yield groups whose span
        is at most `max_distance`. Start a new window whenever the next
        node would push the span past the limit. Negative `max_distance`
        means "no limit" — yield everything as one window.
        c                    |    S r~   r   )r  r   s    rs   r8  z-Scheduler._distance_windows.<locals>.<lambda>  s    k!n ru   r  r   Nr  )r  r   )rq  r   max_distanceorderedwindow
window_minr  idxs    `      rs   _distance_windowszScheduler._distance_windows  s      $<=!*,
 	Aa.C#
*\9 
MM!	 L s   AAc           	     <    t         j                        dt         j                        }t        j	                  d|       t
        j                  dkD  }t
        j                  }t
        j                  }|duxs |du}t
        j                  }d}d}	|r j                         }	|	j                  }
n(t         j                        D ci c]  \  }}||
 }
}}	 	 	 	 	 	 	 	 d fd}t        t        j                               D ]  \  }}||kD  r nt        j                  |      }t        |      dk  r3t         j#                  ||
|      D ]  }||kD  r Vt        |      dk  s j%                  |      s,|rI|	J t'        j(                         } j+                  |||	||       |t'        j(                         |z
  z  }wt        |d   j,                  |d|t
        j.                  	      } ||||         t1        d
        _         j3                   j                         _        t        j5                  d|t         j                               |rt        j5                  d|        j7                   j                         yc c}}w )a"  Group parallel nodes into combo kernels.

        Each parallel group is split into windows whose baseline-index
        span is at most `combo_kernel_max_distance` (set the config to a
        negative value to disable splitting). If a peak-memory threshold is set,
        each window goes through the gate, which simulates the
        post-fusion peak and accepts the combo only if the peak stays
        under the threshold; rejected windows are halved and retried.
        Without a threshold, every window becomes a combo directly.
        r   z2ComboKernels: Generating with num_ck_nodes = %s...N        c                   dz  t         j                  dt        |      |       |D ]  }j                  |        j	                  |        j
                  j                  | j                         D ci c]  }|j                         |  c}       j                  j                  |d         }||j                  | <   y y c c}w )Nr   z0ComboKernels: Combining %d nodes for %d-th groupr   )rT  rj  r   r0  r  r  r  r   r  rv  r  )	
combo_nodeacceptednumr   r  streamrn  r\  rr   s	         rs   _register_acceptz=Scheduler.create_combo_kernel_nodes.<locals>._register_accept  s     QJEHHBH
 ! )""4()OOJ'##**3=3G3G3IJaz)J ((,,Xa[9F!28##J/ " Ks   7Cr   )r[  	on_acceptTrX  r[  rd  c                    | j                   S r~   r  r  s    rs   r8  z5Scheduler.create_combo_kernel_nodes.<locals>.<lambda>  s
    q{{ ru   r  zDGenerated ComboKernel nodes: %d ComboKernels, totally %d -> %d nodesz9ComboKernels memory-aware: %.3fs spent in peak simulation)rN  rB  rO  r  rP  r   r   r  )r   rq  r   rT  r  r-   r  $combo_kernel_peak_memory_increase_gb&combo_kernel_peak_memory_pct_thresholdcombo_kernel_max_distance_init_peak_memory_contextr   r  rB  r  rq  r  rJ  speedup_by_combo_kerneltimeperf_counter_try_combo_with_halvingr  !combo_kernel_per_subkernel_blocksr  r{  rj  r  )rr   rX  num_nodes_origr[  
abs_thr_gbpct_thrmemory_checkrE  memory_sim_timemem_ctxr   r  r  rR  rP  r  membersrG  	sim_startrN  rn  r\  s   `                   @@rs   r  z#Scheduler.create_combo_kernel_nodes  s    !,TZZ		FU 77!;@@
??!-D1D7737446G!--K,5djj,ABDAq1a4BKB	92	9-	9 	9 		9, (&DDTJ
 #	>NC 'EL,@0AA)LG7|a#55l >  +0Dv;?$*F*Fv*N"... $ 1 1 3I00(7"2 1  $t'8'8':Y'FFO!;q	++26(7-3-U-U"J %Z=5>#	>J K-BC
33DJJ?
R

O		
 HHK 	!!$**-U Cs   Jc           	        ddl m}m}m}m}m}m} t        t        j                  j                  j                               }t        t        j                  j                               } || j                  |      }	 || j                  | j                          || j                  | j                  | j                  |	        || j                  |	|      \  }
}} ||
t!        | j                              \  }} ||
t!        | j                              }t#        |t%        | j                        D ci c]  \  }}||
 c}}|||      S c c}}w )zBuild the immutable baseline state the gate compares against:
        original buffer lifetimes, original peak, and step indices.
        r   )rZ  /assign_memory_planning_info_for_scheduler_nodesr  r  +live_memory_before_steps_from_buf_info_listpeak_memory_from_buf_info_list)r   r   r   r   r   )r  rZ  rh  r  r  ri  rj  r   ra   r   rX  r!  r  rq  r  r  r   r   r  )rr   rZ  rh  r  r  ri  rj  rX  r   name_to_freeabler5  r=  r   r   rI  r   s                   rs   rY  z#Scheduler._init_peak_memory_context#  s    	
 	
 "!''"6"6";";"=>"177#;#;#=>1$**lK9$**dFVFVW7JJ##		
 6JJ(-
q! :3tzz?
q  K3tzz? 
 ('4=djj4IJysDsJ'&!5
 	
Js   3Ec                :    ddl m} t        |d   j                  |d|t        j
                        }t               }|D ]'  }|j                  |j                  j                         ) t        |      |_        |j                   t               }t        j                  }	d}
|D ]'  }|j                  |        |   }||	k  r|}	||
kD  s&|}
) g d}dfd	}t        |	|
dz         D ].  }| j                   |   }||v r|s |||	|       d}% ||||       0 t#        d
       D cg c]  }|j$                   }}| j'                  |      }t)        |      D ci c]  \  }}||	|z    c}}|   }|D ]  }||<   	 d fd}|j*                  |	   } |||	|
||j,                  |      }|j.                  }t1        |j2                  |      }||z
  }t        j4                  }t        j6                  }|t9        |      dz  gng }||j;                  ||z         |sJ |t=        |      k  }|dkD  rd|z  |z  nd}|s"t>        jA                  dtC        |      ||       yt>        jE                  dtC        |      ||       ||_        ||fS c c}w c c}}w )a4  The gate: does fusing `group_nodes` into one combo keep peak
        memory under the threshold?

        Returns `(combo_node, combo_step)` if accepted, or `(None, 0)`
        if rejected. The running peak lives on `mem_ctx.running_peak`.

        The pretend rewrite only changes node order inside the window
        `[region_start, region_end]` (the smallest range containing all
        members). Inside that window:
          - the members collapse into one combo node, all at `combo_step`,
          - everything else is evaluated from the original schedule.
        Outside the window, nothing moves. Earlier accepts are accounted
        for only by `mem_ctx.running_peak`. The threshold is checked
        against `mem_ctx.baseline_peak` (the original graph peak) so total
        peak drift is capped.
        r   )estimate_region_peak_memoryr   TrT  )pred_buffersr  Fc                >    j                  t        |||              y r~   )r   r   )r   r   r   local_entriess      rs   	add_localz9Scheduler._try_combo_with_memory_check.<locals>.add_local  s      S(D!ABru   c                2    | j                   | j                  fS r~   )r   r   )r&  s    rs   r8  z8Scheduler._try_combo_with_memory_check.<locals>.<lambda>  s    

@S ru   r  c                     | v r|    S |    S r~   r   )r   new_stepr   s    rs   step_ofz7Scheduler._try_combo_with_memory_check.<locals>.step_of  s     x~%t$$ru   )region_start
region_endru  r   
cur_memoryi   @g      Y@rL  zLComboKernels memory-aware: rejected %d nodes (peak delta %+d bytes = %.3f%%)r   zLComboKernels memory-aware: accepted %d nodes (peak delta %+d bytes = %.3f%%))r   rf   r   r   r   r   r   r  r.  )#r  rm  rB  r  r-   r^  r   r  r'  rn  rG   r   sysmaxsizer  rL  rq  r  r   r{  r  r   r   r   rY  r   rV  rW  rq  r   rZ  rT  r  r   rj  )!rr   group_nodesrd  r[  rm  rN  combo_pred_buffersm	group_setrv  rw  r  rI  inserted_comborq  r  r&  local_nodes
combo_stepru  rx  region_peakoriginal_peaknew_peakdeltar`  ra  limitsacceptpctrp  rt  r   s!                                 @@@rs   _try_combo_with_memory_checkz&Scheduler._try_combo_with_memory_checkM  s   , 	8/N$$&*+!'!I!I

 (\ 	?A%%ajj&=&=>	?7+

 )) 4><	{{
 	!AMM!a.C\!"Z 
	! ,.	C |Z!^4 	A

1AI~%j,:%)NaA	 #=6ST
AFF
 
 44[A4=k4JKDAqA|a''Kj)
 	%A$HQK	%	% 11,?
1%!!//!
  --w++[9=(@@
??4>4J%
#w/0PRMM'M12v#f+%1>1Buu}},II2K  .	
  (:%%y

 Ls   *JJc                  |j                   }|g}|r%|j                         }t        |      dk  s| j                  |      s3t        j                  | |||      \  }	}
|	 ||	||       [|D cg c]  }||   	 }}t        |      t        |      }}||k(  r||z   dz  }|D cg c]  }||   |k  s| }}|D cg c]  }||   |kD  s| }}dt        |      cxk  rt        |      k  rn n|j                  |       dt        |      cxk  rt        |      k  rn n|j                  |       |r$yyc c}w c c}w c c}w )zkTry the full candidate; on reject, halve at the
        baseline-index midpoint and try each half.
        r   N)	r   r  r   rZ  r  r  rZ  rY  r   )rr   r  rP  rd  r[  rS  n2istacksubsetrN  r  r  idxslohimidearlylates                     rs   r]  z!Scheduler._try_combo_with_halving  sA    !!09{YY[F6{Qd&B&B6&J%.%K%Kfg&"J
 %*fc2$*+qCF+D+YD	BRx7q.C &81#a&C-Q8E8%6!Q#A6D6CI+F+T"CJ,V,U#-  ,
 96s   /D9%D>3D>=EEc                H    |D ]  }|j                  | j                          y r~   )r  r  )rr   rq  r   s      rs   r  zScheduler.prune_redundant_deps  s%     	?D%%d&=&=>	?ru   c                   
 g 
t        t        t        t        f             d
 fd}t        j                  t
              }|D ]=  } j                  |      r|j                         D ]  }||   j                  |        ? |j                         D ]
  } ||        t        j                  rat        j                  t
              }|D ]&  }t        |dd      }	|	s||	   j                  |       ( |j                         D ]
  } ||         j                  
      

j                   j                  d       t         j#                  dt%        
             
S )z^
        Helper to find all legal fusion opportunities, sorted by self.score_fusion()
        c                |   t        |       D ]  \  }}| |dz   |dz   t        j                  z    D ]  }||f}|v rj                  |       j	                  ||      rj                  |       B|j                         s|j                         scj	                  ||      swj                  ||f         y rH  )r  r-   )max_fusion_buffer_group_pairwise_attemptsr  r  r   r  r  )	rq  node1_indexr   r   r  rj  r8  rD  rr   s	        rs   check_all_pairsz7Scheduler.get_possible_fusions.<locals>.check_all_pairs
  s    &/&6 @"U"!Ok'FF'G @E
 !%.Cd{ HHSM}}UE3CD(//4++-1A1A1Cu&6J )//?!@@ru   r   NT)r  reversezfound %d possible fusionsrq  r  r   r  )r   r   rf   rW  r   r   unfusable_noder   r   r   r-   aggressive_fusionr  *get_possible_fusions_with_highest_priorityr  score_fusion_keyr  r  r   )rr   rq  rj  r  buffer_names_groupingr   r   node_groupinggroup_groupingr   r8  rD  s   ` `       @@rs   rA  zScheduler.get_possible_fusions  sn    % 13D DEFH	@ 	@( !, 7 7 = 	8D""4(--/ 8%c*11$78	8
 399; 	+MM*	+ ##(44T:N 7gt4"5)0067 "0!6!6!8 /./  JJ
 	$"7"7F4c:J6KLru   c                    t        t                  d fd|j                         j                  j	                         |j                         j                  j	                         z  |j
                  j                  j	                         |j
                  j                  j	                         z  z
  t         fdD              }|r t        ||      d       |S )z~
        Finds whether there's a path from node1 to node2 (or vice-versa)
        caused indirectly by other fusions.
        c                   t        | t              rq| vrmj                  |        | j                         j	                        ryt        | j                  z        xs" t        fd| j                  z
  D              S y)NFc              3  H   K   | ]  } j                   |           y wr~   rQ  r   r  
found_pathrr   s     rs   r   zIScheduler.will_fusion_create_cycle.<locals>.found_path.<locals>.<genexpr>R  s+      H #4#:#:1#=>H   ")r   r   r  r  issubsetr}   r  r  )r   combined_ancestorscombined_namesr  rr   visiteds    rs   r  z6Scheduler.will_fusion_create_cycle.<locals>.found_pathA  s    $ 23G8KD!++-667IJ !   ?@ C H!%2D!DH E  ru   c              3  H   K   | ]  } j                   |           y wr~   rQ  r  s     rs   r   z5Scheduler.will_fusion_create_cycle.<locals>.<genexpr>`  s!     WqJt66q9:Wr  zwill create cycler*  )r   r   r  _dictr!  r  r  r  )rr   r   r   cycler  r  r  r  s   `   @@@@rs   r  z"Scheduler.will_fusion_create_cycle7  s     /02	 	2 %%'--224'')//4467 	
 OO!!&&(5??+@+@+E+E+GG WDVWW#IeU#$78ru   c                    ddl m 	 	 	 	 d fd} ||      } ||      }t        fd|D              }t        fd|D              }|j                  |      }d}	|D ]  }
	 |	t	        |
d         z  }	  j                  ||      }t        j                  j                  j                  |	d	|z        ry
y# t
        $ r Y  yw xY w)a  
        Return true if fusing the two nodes can potentially increasing peak memory.

        The implementation is more like a heuristic since we don't really know if we are at peak
        or not when trying to fuse these two nodes. The order of nodes may change later which makes the
        peak memory estimation hard.

        Here is how we decide the LOWER BOUND of extra memory allocation if we fuse these 2 nodes:
        1. find all buffers read by each node with a single user. These buffers are supposed to
           be reused if we don't fuses these 2 nodes
        2. find the intersection of these buffers for the two node and sum the total buffer size.
           If we don't fuse these two nodes, we can at lease avoid this much memory allocation.
           Note that the extra memory allocation is not necessarily causing peak memory increase.
           This is just a heuristic.

        We return true only if the saving for fusion can not trade off the extra memory allocation.
        r   )buffer_reuse_keyc                0   g }| j                   j                  D ]y  }j                  j                  |j                        }|s+t        |j                        dk(  sD|j                  j                         s_|j                  |j                         { |S rH  )
r   r   r  r  r   r   r  r   has_tensor_outputr   )r   r=  rL  r   rr   s       rs   _find_single_user_inputszKScheduler.can_fusion_increase_peak_memory.<locals>._find_single_user_inputs|  sw     F&&,, ,&&**27733syy>Q.3883M3M3OMM#((+, Mru   c              3  .   K   | ]  } |        y wr~   r   r   r   r  s     rs   r   z<Scheduler.can_fusion_increase_peak_memory.<locals>.<genexpr>       #Sc$4S$9#Sr=  c              3  .   K   | ]  } |        y wr~   r   r  s     rs   r   z<Scheduler.can_fusion_increase_peak_memory.<locals>.<genexpr>  r  r=  r   r   F    T)r   rf   r   zlist[ir.Buffer])r  r  r   intersectionr   rs  r  ra   r   r   statically_known_gt)rr   r   r   r  lhs_dep_nodesrhs_dep_nodeslhs_reuse_keysrhs_reuse_keyscommon_reuse_keysmemory_overheadr  	bw_savingr  s   `           @rs   can_fusion_increase_peak_memoryz)Scheduler.can_fusion_increase_peak_memorye  s    * 	6	#		 1707##S]#SS##S]#SS*77G$ 	C3s1v;.	 ,,UE:	 77//iP  s   $B88	CCc                   t        |j                         D cg c]  }|j                          c}|j                         D cg c]  }|j                          c}z         }t        d |j                  j                  D              }t        d |j                  j
                  D              }||z  }t               }	|j                  j                  D ]:  }
| j                  |
j                  |      s |	j                  |
j                         < t        d |j                  j
                  D              t        d |j                  j
                  D              z  }t        d |j                  j                  D              t        d |j                  j                  D              z  }||z
  }||	z
  }||z  }t        |      |kD  S c c}w c c}w )Nc              3  4   K   | ]  }|j                     y wr~   rt  ru  s     rs   r   zEScheduler.fusion_prevent_too_many_reads_and_writes.<locals>.<genexpr>  s     &TCsxx&Trv  c              3  4   K   | ]  }|j                     y wr~   rt  ru  s     rs   r   zEScheduler.fusion_prevent_too_many_reads_and_writes.<locals>.<genexpr>       %R3chh%Rrv  c              3  4   K   | ]  }|j                     y wr~   rt  ru  s     rs   r   zEScheduler.fusion_prevent_too_many_reads_and_writes.<locals>.<genexpr>  s      $
CHH$
rv  c              3  4   K   | ]  }|j                     y wr~   rt  ru  s     rs   r   zEScheduler.fusion_prevent_too_many_reads_and_writes.<locals>.<genexpr>  r>  rv  c              3  4   K   | ]  }|j                     y wr~   rt  ru  s     rs   r   zEScheduler.fusion_prevent_too_many_reads_and_writes.<locals>.<genexpr>  s      %
CHH%
rv  c              3  4   K   | ]  }|j                     y wr~   rt  ru  s     rs   r   zEScheduler.fusion_prevent_too_many_reads_and_writes.<locals>.<genexpr>  s     DCsxxDrv  )
r   r   r  r   rQ  r   $can_buffer_be_removed_through_fusionr   r  r   )rr   r   r   	thresholdr   fused_node_namesnode1_write_namesnode2_read_namesreads_removed_through_fusionwrites_removed_through_fusionr  all_read_namesall_write_namesunique_readsunique_writesunique_io_bufferss                   rs   (fusion_prevent_too_many_reads_and_writesz2Scheduler.fusion_prevent_too_many_reads_and_writes  s    &).):;T]]_;+0??+<=4t}}=>
 '&T5;L;L;S;S&TT%%R%:K:K:Q:Q%RR'7:K'K$ :D%**11 	BI88 0 .11)..A		B $ $
 % 1 1 7 7$
 
C5+<+<+B+BCCD
 % %
 % 1 1 8 8%
 
D5+<+<+C+CDDE
 &(DD (*GG )=8$%	11M <=s   GG
c                    t        t        |j                  |j                  z
        t        |j                  |j                  z
              }|dkD  S )aA  
        This function prevents fusion for nodes that can increase memory
        footprint. This problem is more common in horizontal fusion, where nodes
        that are far apart in the original order get fused, lengthening the live
        intervals of tensors. This is very evident in models with activation
        checkpointing, where the recomputed nodes from different checkpointed
        regions get fused and significantly increase the memory footprint.

        The current attempt is a quick, possibly hacky, heuristic to prevent the
        fusion of nodes that are far away in the original order.

        A better but difficult to implement heuristic would be to use live
        intervals of the buffers, find region of peak pressure in the original
        program and prevent fusion that crosses that peak region. We might need
        special care or good approximation in this implementation, as fusion of
        node changes live intervals, and re-computing live intervals and peak
        memory after each fusion can introduce large compilation overhead.
        @   )rY  r  r%  r&  )rr   r   r   proximity_scores       rs   are_long_distant_nodesz Scheduler.are_long_distant_nodes  sE    * %//12%//12
 ##ru   c                   i }|j                   j                         D ci c]  }|j                  | }}|j                   j                         D ci c]  }|j                  | }}|D ]}  }t        j                  j                  |      }	||   }
||   }t        |
t              rt        |t              sdt        |
       dt        |       ||<   k|
j                         |j                         k7  r(d|
j                          d|j                          ||<   t        |
j                        t        |j                        k7  rd||<   |
j                         }|j                         }||k7  rd| d| ||<   |
j                         |j                         k(  rd|
 d| ||<   Ed}t        |	t        j                        sd|	j                    }d	|
 d| d
| ||<    t#        |      S c c}w c c}w )z}
        Try to decide reasons why fusion fail due to no shared memory even though
        there are common buffers.
        znot MemoryDep: r   zdifferent numel: 	broadcastzdifferent offset: zMismatch loop orders: rX  zLayout: zUnknown reason: z. )r   rw  r   ra   r   rI  r   r9   r   r   r_   r   
get_offsetnormalize_with_stride_orderr0   rL  r  r,  )rr   r   r   common_buf_namesreasonsr   node1_name2depnode2_name2depr  r   lhs_deprhs_deplhs_offrhs_off
layout_strs                  rs   decide_fusion_fail_reasonz#Scheduler.decide_fusion_fail_reason  s    383D3D3U3U3WXC#((C-XX383D3D3U3U3WXC#((C-XX( ,	H''$$X.C$X.G$X.Ggy1GY9W%d7m_F4=/J !   "g&7&7&99'(9(9(;'<F7CTCTCVBWX !  W\\*mGLL.II$/!((*G((*G'! '9	y$Q! 3356689 '=WIVG9$U! Jc2#5#56'

|4
"7)6'"ZLI HU,	\ 7|c YXs   G5G:c                   t         j                  syt        d ||fD              ry|j                  j	                         }|j                  j	                         }||z  }|syt        d |j                  D              }||z
  ryt        |      dkD  ryt        |j                  j                        dkD  s"t        |j                  j                        dkD  ryt        t        |j                  j                              }t        t        |j                  j                              }t        |t              rt        |t              sy|j                  j                  D 	ci c]  }	|	j                  |	 }
}	|j                  |
vry|
|j                     }t        |t              sy|j                         }|j                   |j                   k7  r|j"                  |j"                  k7  ry|j"                  |j"                  k7  st        |j$                        dk7  ryt        |j&                  j(                        dk7  ry|j&                  j*                  ryd|j&                  j(                  v rd|j&                  j(                  v sJ t        d |j&                  j-                         D              }t        |      dk7  ryt        t        |            }||j&                  j(                  d   k(  rd}d}n"||j&                  j(                  d   k(  sJ d}d}d	d
lm} |j&                  j2                  d	   }t        |      dk7  ryg }t4        j6                  j9                  |      D ]:  }|j;                  t<        j>                  j@                  jC                  |             < tE        |      } |||d	         }|y|j&                  j(                  |   |j&                  j(                  |<   ||j&                  j(                  |<   |jG                  dd       | jI                  ||      }t        |tJ              sJ tL        jO                  d|       |S c c}	w )aW  
        Attempts to enable fusion between two nodes by inverting indexing patterns.

        This optimization targets cases where node1 has a contiguous write and
        node2 has a contiguous write but discontiguous read. By inverting the
        indexing in node2's read and write operations, we can make them compatible
        with node1 for potential fusion.

        Args:
            node1: First scheduler node (source)
            node2: Second scheduler node (target for inversion)

        Returns:
            int: Fusion score if successful, 0 if optimization not applicable
        r  c              3  <   K   | ]  }|j                           y wr~   r  r  s     rs   r   zAScheduler.shared_data_after_inverting_indexing.<locals>.<genexpr>=       2aqxxz2r  c              3  4   K   | ]  }|j                     y wr~   rt  ru  s     rs   r   zAScheduler.shared_data_after_inverting_indexing.<locals>.<genexpr>I  s      .
CHH.
rv  r   r   index0index1c              3      K   | ]  }|  y wr~   r   )r   ra  s     rs   r   zAScheduler.shared_data_after_inverting_indexing.<locals>.<genexpr>  s     %Ttd%Ts   r   )generate_inverse_formulaTFz!Shared memory after inversion: %d)(r-   rl  r  r   buffer_namesr   r3  r   r   rQ  r  r  r   r9   r   r   r   r   	var_namesr  r   	subblocksget_read_exprs$torch._inductor.invert_expr_analysisr  varsr  Add	make_argsr   ra   r   r   combine_modular_indexing_pairsr   r3  r  r   r  rj  )rr   r   r   node1_buffer_namesnode2_buffer_namescommon_buffer_namesnode2_unmet_dependencies
node2_readnode2_writer   node1_writesnode1_writenode2_read_exprs	read_exprread_expr_indexwrite_expr_indexr  rt  simplified_termstermsimplified_read_exprinverse_formulascores                          rs   $shared_data_after_inverting_indexingz.Scheduler.shared_data_after_inverting_indexing'  s   & 442E5>22 #..;;="..;;=03EE" $. .
 % 8 8.
 $
  $&88'(1, u  &&'!+s53D3D3K3K/Lq/P$u006678
4 1 1 8 89:*i0
9
 161B1B1I1IJ##JJ??,.":??3+y1 "++- !2!22  K$4$44??k...#j6J6J2Kq2P u{{))*a/ ;;   222EKK666	
7
 &%Tu{{7Q7Q7S%TT A%./0	 228<<&O' : :8 DDDD&O'Q[[%%a(
z?aII''	2 	D##  ??E	  ##3423GTUW "
 7<kk6P6P7
""?3 8G""#34 	""4/((6%%%%;UCm Ks   "Qc                   t        d ||fD              ry|j                         s|j                         ry|j                  j                         |j                  j                         z  }|syt        j
                  r| j                  ||      }|dk\  r|S t        j                  r| j                  ||      syt        j
                  r| j                  ||      }|dk\  r|S | j                  ||      S )a  
        Right now just greedily reorder the loop of node1 to be compatible with node2,
        but ideally we should have some heuristics to reorder the loop for node2
        to be compatible with node1 if that's more efficient.

        Return the amount of shared data re-computed in this method.
        If no such recomputation happens, return -1 (not return 0 since 0 is a valid
        amount of shared data).

        c              3  <   K   | ]  }|j                           y wr~   r  r  s     rs   r   z>Scheduler.shared_data_after_reordering_loop.<locals>.<genexpr>  r  r  r  r   )
r  r  r   r  r-   r%  !_try_reorder_loops_for_candidatesloop_reindexing_after_fusion$_try_reindex_pointwise_for_reductionr  )rr   r   r   r  r	  s        rs   !shared_data_after_reordering_loopz+Scheduler.shared_data_after_reordering_loop  s      2E5>22
 %"3"3"5 **,u/@/@/M/M/OO 	 #,,::5%HEz 33<<UEJ,,::5%HEz''u55ru   c                   |j                   j                         |j                   j                         z  }|j                   j                  D ci c]  }|j                  | }}|j                   j                  D ci c]  }|j                  | }}|j                   j                  D ci c]  }|j                  | }}|j                   j                  D ci c]  }|j                  | }}g }	|D ]  }
|j                  |
      xs ||
   }|j                  |
      xs ||
   }|
|v xr |
|v xs
 |
|v xr |
|v }|j                         |j                         k(  rM|	j                  |t        j                  j                  j                  |j                         d      ||f       |s|j                  |
      xs |j                  |
      }|j                  |
      xs |j                  |
      }t        |t              st        |t              s't        j                  j                  |j                         j                   }|j                         j                   t#        fd|D              r y t%        |	      dk(  ryt'        |	t)        j*                  dd            \  }}}}t        |t              rt        |t              sy|j,                  |j,                  k7  r3|j                         |j                         k(  r| j/                  |      S yd}|j1                         s|j3                  ||      }nV|j1                         s|j3                  ||      }n3t4        j7                  d|j9                         |j9                                |r| j;                  ||      S dS c c}w c c}w c c}w c c}w )	z
        Find common buffers with matching normalized stride order but different
        loop orders, and try to reorder loops to align them.
        r   r   c              3  J   K   | ]  t        fd D                yw)c              3  B   K   | ]  }j                  |        y wr~   )r   )r   rssvwss     rs   r   zHScheduler._try_reorder_loops_for_candidates.<locals>.<genexpr>.<genexpr>%  s     Q2B66r2>Qr  Nr  )r   r	  r_sizesr	  s    @rs   r   z>Scheduler._try_reorder_loops_for_candidates.<locals>.<genexpr>$  s$       QQQs   #r  r   r  Fz?Don't reorder loops since both nodes are reductions: %s v.s. %s)r   r  r   r   rQ  r  r  r   ra   r   r   r   r   r   r9   r   r   r   r   rY  r  r  r^  dep_size_hintr   rp  ra  r  r  r  )rr   r   r   r  r   node1_readsr	  node2_readsnode2_writes
candidatesr  r  r  is_write_readwrR  w_sizes_is_wr_numel	reorderedr	  r	  s                       @@rs   r	  z+Scheduler._try_reorder_loops_for_candidates  s{    **,u/@/@/M/M/OO 	 160A0A0G0GHsxx}HH161B1B1I1IJ##JJ050A0A0G0GHsxx}HH161B1B1I1IJ##JJ
. '	"K"&&{3O{;7OG"&&{3O{;7OG |+J{0JN-L+2L 
 3356689 !!%((::#--/! ;   	  !$$[1R\5E5Ek5ROOK0PKOOK4Pa+
1i0H))Bkkm00Gkkm00G ")   "O'	"R z?a ,/H//15,
( '9-Z5Sw///
   "g&7&7&99))'22	!!#77II##%77II##Q   :Ct''u5JJc IJHJs   N6>N;,O Oc                    ddl m |j                         s|j                         ry|j                         r|j                         s||}}n&|j                         r|j                         s||}}ny|j                  \  }}t        j                  t        j                  |d         t        j                  t        j                  |d         z  t        d |j                         D              syt        j                  t        t           |j                               }t        fd|D              syt        fd|D              syft        fd|D              ryt        j                  |f      }|D ]  }	|	j                  g        t!        |t"              r|d   j                  |_        t%        |       |j&                  j)                         |j&                  j)                         z  }
|j&                  j+                         D ci c]  }|j,                  | c}|j&                  j+                         D ci c]  }|j,                  | c}t/         fd	|
D              }|s|j1                          yt2        j4                  s5|D ]  }	|	j7                  d
d        t!        |t"              rt%        |       y
c c}w c c}w )z
        Reindex a pointwise's iteration loops to match a reduction's
        groups. After reindexing, the shared reads have identical index
        expressions, enabling the codegen to CSE loads.

        Returns True if reindexing was applied.
        r   r  Fr   c              3  <   K   | ]  }t        |t                y wr~   r   r   )r   r^  s     rs   r   zAScheduler._try_reindex_pointwise_for_reduction.<locals>.<genexpr>m  s     OR:b-0Or  c              3     K   | ]D  }t         j                  j                  j                  t	        |j
                  d                 F ywr  )ra   r   r   r   r_   r  )r   r^  target_numels     rs   r   zAScheduler._try_reindex_pointwise_for_reduction.<locals>.<genexpr>s  sB      
  GG44biil+\
s   A
Ac              3  b   K   | ]&  }j                  f|j                                ( y wr~   )r  r\  )r   r^  r  	red_numel
red_rnumels     rs   r   zAScheduler._try_reindex_pointwise_for_reduction.<locals>.<genexpr>{  s1      
 $$i%<bmmoN
s   ,/c              3  T   K   | ]  }t        |j                  d          k(   ! ywr  )r   r  )r   r^  target_iter_sizess     rs   r   zAScheduler._try_reindex_pointwise_for_reduction.<locals>.<genexpr>  s$     IBuRYYq\"&77Is   %(c              3  N   K   | ]  }j                  |   |           y wr~   )deps_match_normalized)r   r   n1_depsn2_depsrr   s     rs   r   zAScheduler._try_reindex_pointwise_for_reduction.<locals>.<genexpr>  s.      
 &&wt}gdmD
s   "%TrA  )r  r  r  r   r   r  r  r  ra  r   r   r   r   r/  r  rI  r   r   r  r   r  rw  r   r  r<  r-   r%  r3  )rr   r   r   reduction_nodepw_noder=  groupsr   rollback_snapshotr^  common_namesr   has_benefitr  r3	  r4	  r-	  r.	  r0	  r+	  s   `            @@@@@@@rs   r	  z.Scheduler._try_reindex_pointwise_for_reductionN  s    	- <<>U\\^(:(:(<&+UGN!%*<*<*>&+UGN"((	6KK

F1I6	[[VAY7
 :-O7;L;L;NOOT-0'2C2C2EF  
 	
 
  

 
  '
3I&II /55wjA 	>B$$i%<=	> g12"1IOOGM+G4 **,u/@/@/M/M/OO 	 -2,=,=,N,N,PQS388S=Q,1,=,=,N,N,PQS388S=Q 
$
 
 %%' 00 W''$PU'VW'#56/8+ RQs   "KK c                f   t        |t              r)|j                          xr t        |j                         S t        |t
              rht        |j                  t        j                        r|j                  j                          S |j                          xr t        |j                         S y)z>
        Is this node unfusable under any conditions.
        F)	r   r  r  r]   r   rR  r0   r  r   r  s     rs   r  zScheduler.unfusable_node  s     d23'')) 2U		3 /  d56$))R%?%?@9966888'')) 2U		3 /  ru   c                   |j                         t        j                  j                  k  ry|j	                         }|j                         }d}|||z  kD  r	 |d       yt        d |j                         D              }|t        j                  j                  j                  j                  fk(  r	 |d       yd	d}|j                         }	|	j                         s+ ||	j                        r|j!                         s	 |d       yy)
zT
        Heuristics to avoid benchmarking predictably slow prologue fusions
        Tg?z@prologue fusion will not increase amount of bytes read in kernelFc              3     K   | ]J  }|j                   <|j                   j                         D ]  }|j                  dk(  r|j                   ! L y w)Ncall_function)r   r  r  r  )r   r  r&  s      rs   r   zEScheduler.check_prologue_fusion_heuristics_fusable.<locals>.<genexpr>  sT      
vv!VV'')	
 tt&	 HH

s   AAz\prologue fusion will not increase attempt to fuse in padding bc it increases unaligned readsc                <    | j                   dk  xr | j                  S )Nr   )itemsizeis_floating_point)r}  s    rs   low_prec_fpzGScheduler.check_prologue_fusion_heuristics_fusable.<locals>.low_prec_fp  s    >>Q&B5+B+BBru   zVprologue fusion that must be upcast to fp32 not profitable for low precision templates)r}  ztorch.dtyper   r}   )r  ra   r   invoke_quant_opsr1  r3  r   r   r`  r&  atenconstant_pad_ndr)  r  r\   r}  r  )
rr   prologue_noder  rS  
read_byteswrite_bytesBYTES_THRESHOLD_MULTIPLIERr$  rB	  r  s
             rs   (check_prologue_fusion_heuristics_fusablez2Scheduler.check_prologue_fusion_heuristics_fusable  s    ,,.!''2J2JJ"88:
#::< &)"'AABRS  
",,.
 
 uyy~~55==??n 	C %??A668L../!>>@h ru   c                <    t        |t              rt        |t              syt        |j                  t        j                        r$t        |j                  t        j                        sy|j                         s|j                         ryt        j                  dk(  ry|j                  |j                  }}|\  }}|\  }}|j                         s,|j                         s||k7  st        |      t        |      k7  ryt        |j                  j                        dkD  s"t        |j                  j                        dkD  ry j                  t        t        |j                  j                                    }	 j                  t        t        |j                  j                                    }
t!        |	|
      t        j"                  kD  ryd fd} ||      s ||      ryg }t%        t'        ||            D ]  \  }\  }}||k7  s|j)                  |       ! t        |      dk7  ry|d   }||   ||   }}t*        j,                  j.                  j1                  ||      r|||fS t*        j,                  j.                  j1                  ||      r|||fS y)ao  
        Fusing two small pointwise nodes significantly reduces kernel overhead
        and launch overhead. However, slightly different sizes would prevent fusion.
        Here, we decide if expanding sizes of one node is profitible by allowing
        fusion, and returns the dimension to expand, node with smaller sizes,
        and new size after expand.
        Nrg  r   c                ~   | j                   j                  D ]  }|j                  j                  v rj                  |j                     }n%j                  j                  |j                        }|s]t        j                  j                  j                  ||       st        |j                  t              r y yr  )r   r   r   r  r  r  ra   r   r  r  r   r  r  )r   r  r  rr   s      rs   has_reusable_bufferzIScheduler.get_expand_dim_for_pointwise_nodes.<locals>.has_reusable_buffer5  s    ((..  99 ; ;; $ ; ;DII FI $ 0 0 4 4TYY ?I ,,66y$G&y'<'<>TU  ru   r   r*  )r   r   r   r0   r   r  r-   rh  r  r   r   r   rQ  r	  r  r  rY  small_memory_access_thresholdr  ry  r   ra   r   r   statically_known_lt)rr   r   r   n1_sizesn2_sizesn1_iter_sizesn1_reduce_sizesn2_iter_sizesn2_reduce_sizesnode1_write_memorynode2_write_memoryrM	  mismatch_dimensionsrI  n1_sizen2_sizemismatch_dimmismatch_size1mismatch_size2s   `                  rs   "get_expand_dim_for_pointwise_nodesz,Scheduler.get_expand_dim_for_pointwise_nodes  s]    %/z%7W uzz2#4#455::r'8'89 ))+u/M/M/O ) #\\5<<()1&)1& !!#/1=!S%77 u  ''(1,E4E4E4L4L0MPQ0Q "//T%:K:K:R:R5S0TU!//T%:K:K:R:R5S0TU"$67223 	  u%)<U)C !'0]M1R'S 	0#C#'7'!#**3/	0 "#q(*1-,',' ' 77//O66WW11..Q66ru   c                     t         fd|j                         D              t         fd|j                  j                  D              S )Nc              3  V   K   | ]   }j                   j                  ||       " y wr~   r)  r  r   r   rr   s     rs   r   zDScheduler._producer_output_names_read_by_consumer.<locals>.<genexpr>`  s,      (
 !!%%dD1(
r  c              3     K   | ]:  }j                   j                  |j                  |j                        xv r < y wr~   r)  r  r   )r   r   r   producer_buf_namesrr   s     rs   r   zDScheduler._producer_output_names_read_by_consumer.<locals>.<genexpr>d  sD      
--11#((CHHEE!" 
s   A A)r   r  r   r   )rr   rE  rJ  r   re	  s   `  @@rs   '_producer_output_names_read_by_consumerz1Scheduler._producer_output_names_read_by_consumer]  sN     ( (
 113(
 
  
++11
 
 	
ru   c                    t         j                  ||      sy t         j                  ||      sy | j                  ||      S r~   )r2  r@  r  rf	  r  s      rs   "_nested_index_equivalent_dep_namesz,Scheduler._nested_index_equivalent_dep_namesk  s?    
 ++E59 ''u5;;E5IIru   c               H   	 t               }t         fd|j                         D              	|D ][  \  }}|t        j                  j                  ur# fd|j
                  j                  D        }|j                  	fd|D               ]  j                  ||||      S )Nc              3  V   K   | ]   }j                   j                  ||       " y wr~   ra	  rb	  s     rs   r   z>Scheduler._can_fuse_nested_reduction_append.<locals>.<genexpr>  s,      '
 !!%%dD1'
r  c              3  ~   K   | ]4  }j                   j                  |j                  |j                         6 y wr~   rd	  r  s     rs   r   z>Scheduler._can_fuse_nested_reduction_append.<locals>.<genexpr>  s4       %%))#((CHH=s   :=c              3  ,   K   | ]  }|v s|  y wr~   r   )r   dep_namegrouped_buf_namess     rs   r   z>Scheduler._can_fuse_nested_reduction_append.<locals>.<genexpr>  s      .%(>O2O.s   	)r%  index_equivalent_dep_names)	r   r  r2  rH  rG  r   r   r  	_can_fuse)
rr   r7  r  ri  r%  ro	  r^  ry  r   rn	  s
   `        @rs   r&  z+Scheduler._can_fuse_nested_reduction_append|  s     7Al"& '
$557'
 
 , 		JB_<<HHH>>//E '-- .).. 		 ~~#'A	  
 	
ru   Fc                    t         j                  ||f      }| j                  ||||      }|j                  |        |S )zDetermine if node1 and node2 can be combined into a single fused node.

        Speculative loop mutations (reordering, reindexing) are automatically
        rolled back if the fusion decision ultimately fails.
        r%  r  )rJ  )r@  r  _can_fuse_implrK  )rr   r   r   r%  r  rE  r  s          rs   r  zScheduler.can_fuse  sP     '--uen=&&#&?	 ' 
 	H-ru   c                ,    | j                  ||||      S )zj
        Determine if it is possible to combine node1 and node2 into a
        single fused node.
        rr	  )rp	  )rr   r   r   r%  r  s        rs   rs	  zScheduler._can_fuse_impl  s&     ~~#&?	  
 	
ru   c                   !"# u ry|t         fd|D              } j                         r@ j                  j                        } j                  j                        }||||k7  ryt	        t
              rj                  |      S t	        t
              ryt	        t              rj                        S t	        t              ryt              }j                         r0 j                  j                               j                        ryt	        t              st	        t              r	 |d       yt	        t              rj                         s	 |d       yt	        t              rt	        j                   t"        j$                        s	 |d       yj                   j'                         s	 |d       yt	        t(              s	 |d	       yt	        j                   t*              s	 |d
       yt	        j                   j,                  t.              s	 |d       yt1        j                   j2                        dk(  sJ j                   j2                  d   j4                  #t7        #fdj8                  j:                  D              r	 |d       yj                   j,                  j=                         }	|	D ];  }
j                   j,                  j?                  |
      }t7        d |D              s; y t1        j                   j@                        dk(  sJ j                   j@                  d   jB                  }j                   jB                  }t	        |t"        jD                        rt	        |t"        jD                        sJ |jF                  |jF                  k7  s2|jH                  |jH                  k7  s|jJ                  |jJ                  k7  r	 |d       y	 	 d'#fd!t7        !fd jL                  D              ryt	        t        t        f      rj                         s	 |d       yjO                         jP                  z  r	 |d       yj                         rtS              s	 |d       yjU                         sj                         r	 |d       yjW                         }|jY                         }|s	 |d       yt        d |jZ                  D              |z
  }j]                         |z  r	 |d       yj_                         s|ja                        r	 |d       yjc                         ""d d D ]B  }|je                         }|D ]+  }tg        "fd|jh                  D              r" |d         y D t	        tj              sgn*jl                  D cg c]  }|j                         s| c}}t1        |      dk(  sJ |d   }t1        "d   jn                        dk(  rSt1        "d   jn                  d   jh                        dk(  r+"d   jn                  d   jh                  d   j                   |u s	 |d       y jq                  |      syj                         rts              }j_                         r|rjU                         stu              s	 |d        yjw                         }|J |jy                         r-t	        j                   t"        j*                        s	 |d!       yj]                         tz        j|                  j~                  z  s+j]                         tz        j|                  j~                  z  r	 |d"       yj                         }j                         }||k7  r |d#||       y~| j                        } j                  ||$      }|rL|t        j                  k  r9t        j                  st        j                  r j                        }|dk\  r|}t        j                  r@ j                        x}r,|\  }}} |j                  ||         j                  |%      }t        j                  r,|t        j                  k  r j                        }|dk\  r|}t        j                  t        j                        r4t        j                  d&j                         j                         |       tz        j                  j                   |      syjO                         jP                  z  r։ j                  |%      rDtz        j                  j                   |      r" j                  |      j                        ryt        j                  rm j                        r[ j                  |%      xrE tz        j                  j                   |      xr!  j                  |      j                        S ytz        j                  j                   |      xr!  j                  |      j                        S c c}w )(NFc              3  V   K   | ]   }j                   j                  ||       " y wr~   ra	  rb	  s     rs   r   z&Scheduler._can_fuse.<locals>.<genexpr>  s,      4 %%))$54r  r$  Tz/grouped node must not be fused with other nodesznode1 is nopz'node1 is extern but not a triton kernelz5node1's triton kernel doesn't support epilogue fusionz.node1 is extern but node2 is not SchedulerNodez3node1 is extern but node2.node is not SchedulerNodez4node1 is extern but node2.node.data is not Pointwiser   r   c              3  <   K   | ]  }|j                   k7    y wr~   rt  )r   r   written_buffer_names     rs   r   z&Scheduler._can_fuse.<locals>.<genexpr>  s     Vs38822VrR  z9epilogue reads from buffers other than the mutated outputc              3  &   K   | ]	  }|d k7    yw)r  Nr   )r   usages     rs   r   z&Scheduler._can_fuse.<locals>.<genexpr>  s     ;5u;s   z*node1 and node2 uses different buf layoutsc                @    | uxr | uxr | j                         v S r~   )r   )r  r   r   rx	  s    rs   ._is_other_node_that_references_mutation_bufferzKScheduler._can_fuse.<locals>._is_other_node_that_references_mutation_buffer-  s7      u, N#50N+z/K/K/MMru   c              3  .   K   | ]  } |        y wr~   r   )r   r   r|	  s     rs   r   z&Scheduler._can_fuse.<locals>.<genexpr>6  s       ?tDr=  znode2 is extern or nopznode1 must go before node2zprologue fusion turned offz2prologue fusion only supported for pointwise nodesz'template has no allowed prologue inputsc              3  <   K   | ]  }|j                           y wr~   r  )r   r  s     rs   r   z&Scheduler._can_fuse.<locals>.<genexpr>W  s     Ec3<<>Er  z;prologue fusion not implemented for kernel for these inputsz:template prologue can only fuse functional pointwise nodesr  c              3  :   K   | ]  }|j                   v   y wr~   r   )r   r  prologue_nodess     rs   r   z&Scheduler._can_fuse.<locals>.<genexpr>i  s     QttyyN:Qr  z7template prologue can only fuse nodes with a single usezEtemplate prologue can only fuse nodes with a single use into templateztemplate epilogue not satisfiedz6multi-output template epilogue requires ComputedBufferz#fusion for buffer explicit disabledzdevice mismatch (%s vs %s))r  ro	  ro	  z%s and %s has %s shared data)r  rf   )Wr   r  rv  r  r   r  r  r  r  r  r#  r   can_fuse_multi_outputs_templater  r  rR  r   r0   r  r   r   r   r   r   r   r/  r   r  r   r   inner_fn_free_symbolscollect_inner_fn_symbol_usager  r  r  r   r  r  rq  r  r  r$  r   r  get_allowed_prologue_inpsr  r  r  ,has_aliasing_or_mutation_for_prologue_fusionr   r<  r   r  r   r   r,  rJ	  r   r  r  r\   ra   r   no_fuse_buffer_namesrh	  !_score_fusion_memory_for_can_fuser-   score_fusion_memory_thresholdr%  r	  r	  $expand_dimension_for_pointwise_nodesr^	  rX  rl  r	  ra  r  r  r  r  r  r  r  can_fuse_verticalr	  can_fuse_horizontal)$rr   r   r   r%  r  ro	  r  stream2rS  node2_inner_fn_free_symbolssymbolusageslayout1layout2r  r  unsupported_prologue_argsr   	node_outsr   r  template_snodestemplate_snodeatomic_add_mutation_epiloguer  r  device2shared_data_scorenew_shared_data_scoreexpand_analysis
expand_dimsmaller_nodeexpand_sizer|	  r	  rx	  s$   ```                              @@@rs   rp	  zScheduler._can_fuse  s	    E>%1)3 464 *& '')))--e4G))--e4G"w':w'?Qe23&&u+&FFe23e45&&u--e45 u%4#3#3$

)
)%
7$8 e12j'7
 ABe34U=N=N=Pe67ejj"*D*DE=>:://1KLe]3DEejj.9IJejjooy9JKuzz223q888"'**"="=a"@"E"E Ve>O>O>U>UVVOP +0**//*O*O*Q'5 !FFvN;F;; ! uzz../1444 jj--a077Gjj''Ggryy1j"))6TTT,>>W^^3>>W^^3@A-   JJ   u8:PQR%%'()$$&8,-.u501!!#u'8'8':HI779H$,$F$F$H!(=> EX__EE'( &
 %%'*CCQR--/EEeLPQ"__.N&s+ % ,,.	$ %CQsyyQQUV$%% "%);< !&AAaA 
 '1,,,,Q/N N2&../14r*2215;;<A"2&..q177:??>Q[ @@sS+Qu,( //1:V%%'25956 224L+++557


B--A LM""$qww'C'CC""$qww'C'CC56!!#""$W,fg>%-)-)P)Pu*& !BB&?'A	 C 
 !F$H$HH11V5X5X %)$J$J5RW$X!$)$9!66#FFueTTOT6E3Z{<<ZU $ F F+E !G ! 11!F$H$HH$($M$Mu%! %)$9!))'--8##.  !	 yy!!$u6GH$$&8 &&/I ' 
 II//eUDUV$$V,>>ueL 33==eUK **3M +  Q
 		33eU,=Q ((0BB5%P
 9900eU$5 M""6*>>ueLMU Bs   i?0i?r	  c                  |j                         }t        ||      }t        t              }|j                  D ]j  }| j
                  j                  |j                  |j                        }t        |t              r| j                  |||      rW||   j                  |       l |j                  j                  D ]  }	t        |	t              st        |	t              s%| j
                  j                  |	j                  |	j                        }
|j                  |
      }|si|D ]  }t        |	t              rG| j!                  |j#                  | j
                        |	|duxr |
|v       r|j%                  |       Zt        |	t              sk| j'                  ||	|j(                        s|j%                  |        	 t+        d t,        j.                  j1                  |j3                               D              }||z  r	 |d       y|j5                         }|D ]E  }| j6                  |   j9                         }|| j:                  |   j<                  z  s= |d        y y)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.

        ``index_equivalent_dep_names`` relaxes write/read matching only for
        named producer outputs; the remaining intermediate-dependency checks
        still run normally.
        Nallow_index_equivalencec              3  4   K   | ]  }|j                     y wr~   rt  ru  s     rs   r   z.Scheduler.can_fuse_vertical.<locals>.<genexpr>1  s      $
 HH$
rv  zmemory deps did not matchFz(intermediate nodes between node1 & node2T)r  r  r   r   r3  r)  r  r   r   r;   r  r   r   rQ  r9   r:   fusable_read_and_writery  r0  .fusable_stardep_write_and_read_on_empty_tensorr   r   r  r  rz  r   r  r  r  r  r  )rr   r   r   ro	  node1_buf_namesrS  remaining_deps_by_namer   r   cd
write_name	remainingrL  remaining_depsnode1_op_namesr  s                   rs   r	  zScheduler.can_fuse_vertical  s:   $  002u%7B47H++ 	5C((,,SXXsxx@D#w'D,A,A#ue,T"4(//4		5 ##** 	-Bb),ZG5L..22277BGGDJ.22:>I# -B!"i0T5P5P		$"7"786dB I *.H H 6Q 6 "((,#B0KKEJJ "((,-	-. $ $
 445K5R5R5TU$
 

 O+
 +,224" 	D&&t,==?G 7 7 @ J JJ>?		 ru   c                   |j                   |j                         vry|j                  j                  D cg c]  }|j                   |j                  k(  r| }}t        |      dk7  ry|d   t        t              ryt        t              sJ t        j                  t        j                        ryt        j                        j                  j                  k  sy| j                   |j                     }|g}t        |t"              r|j$                  }d}|D ]R  }	|	j                  j&                  D 
cg c]  }
|
j                   |k(  r|
 }}
|s8|dz  }t)        fd|D              rR y |dk  S c c}w c c}
w )NFr   r   c              3     K   | ]q  }t        |t              xr[ t        |j                  t        j
                         xr4 |j                  j                  k(  xr |j                  j                  k(   s y wr~   )r   r9   r(   r   r*   TMPr   )r   r  r  s     rs   r   z-Scheduler.fusable_weak_dep.<locals>.<genexpr>r  sm      
 	 4+ ,+DJJAA,JJ%++-, II+,s   A7A:)r   r  r   rQ  r  r   r   r:   r9   r(   r   r*   r	  r   r  r&   r  rB  r   r   r   )rr   weak_depr   r   r  mutating_writesr3  relevant_reading_nodesnum_concurrent_readsreading_noder  relevant_readss       `       rs   r  zScheduler.fusable_weak_depG  s    == 6 6 88 **11
zzX222 
 

 1$"eW%%+++u{{DHH5
 %//*ekk.F.FF++H,A,AB	"'e78%*\\" 2 	L )44::99	) N 
 " A%  
 +  !	" $q((W
6s   "E:;E?c                    | j                   |j                   k(  xr\ t        | j                        t        |j                        k\  xr/ | j                  d t        |j                         |j                  k(  S r~   )r   r   r   )r  r  s     rs   _same_index_with_prefix_sizez&Scheduler._same_index_with_prefix_size|  s[     JJ%++% ;DII#ejj/1;		+C

O,

:	
ru   r	  c                  t        |t              r|j                  |j                  k7  sHt        |j                  t
        j                        s$t        |j                  t
        j                        ry|}|}| j                  |j                        ry| j                  ||      ryt        j                  r9|j                  |j                  k7  r |j                         }|j                         }| j                  ||      ry|sy| j                  ||      S t        |t              r?|j                  |j                  k(  r&|j                  |j                  |j                  k(  ryy)a  Return whether a producer write can satisfy a consumer read.

        The default path accepts exact matches, plus the existing
        loop-ordering-normalized exact match when that config is enabled.
        ``allow_index_equivalence`` only runs after those checks fail. It keeps
        the producer write dense and injective, then accepts conservative
        consumer-side equivalent reads such as broadcasts or normalized
        loop-order changes.
        FT)r   r9   r   r(   r   r*   r	  rC  r  r	  r-   r%  r^  r   %_fusable_read_after_index_equivalencer:   )rr   r  r  r	  original_readoriginal_writes         rs   r	  z Scheduler.fusable_read_and_write  s    dI&		UZZ'&tzz488<&u{{DHH= M"N 11.2E2EF 00O00T]]enn5T ~~')
 00u=*==~  g&		UZZ'JJ*II+ru   c                "   t        |j                        |j                  j                  k  sy|j	                         j                         s|j                         j                         sy| j                  ||      xs | j                  ||      S rm  )	r   r  r   r&   r   is_contiguousr  r2	  _fusable_read_after_broadcast)rr   r  r  s      rs   r	  z/Scheduler._fusable_read_after_index_equivalence  s|     %//*ekk.F.FF
 OO++-002@@B))%
 =//e<	=ru   c           
         t         fd j                  D              }t        |       j                  k7  r|D ci c]  }| j                  |    }}t         j                   j                  t        |      t        |j                                j                          j                         |j                         k(  ry j                  |j                  k7  ryt        j                  j                  }t        d t         j                        D              }i }i }t         j                   j                   ||j                         D ]  \  }	}
}}|j#                  |
|      r|||	<    |j#                  t%        j&                  |
|      d      s y|j)                  t+        |
|            }|j-                  |d      s yt%        j.                  dt        |       dd      }||z  |z   ||	<   |||<    |sy|j1                  t3         j                  |      i t5        t        ||j                               |      }t3        |j                  t5        t        |j                  |                  }|j#                  ||      S c c}w )	aQ  Match conservative broadcasted read forms.

        This handles two nested-reduction dependency shapes:

        - Pure broadcast dims absent from the read index:

              read:  d1, {d0: 1024, d1: 16}
              write: d0, {d0: 16}

        - Same-rank expanded dims used through a quotient:

              read:  32*d0 + FloorDiv(d1, 128), {d0: 128, d1: 4096}
              write: 32*d0 + d1,                {d0: 128, d1: 32}

        Producer-side broadcast and non-dense writes are rejected before this
        helper, so these cases only relax consumer-side broadcasts.
        c              3  T   K   | ]  }|j                   j                  v s| ! y wr~   )r   r&   )r   r  r  s     rs   r   z:Scheduler._fusable_read_after_broadcast.<locals>.<genexpr>  s'      
SDJJ4K4K-KC
s   ((TFc              3  R   K   | ]  }t        j                  d | dd       ! yw)_fusable_broadcast_TintegernonnegativeN)r  r  )r   r  s     rs   r   z:Scheduler._fusable_read_after_broadcast.<locals>.<genexpr>  s0      
 LL.qc2DdSS
s   %'r   r   _fusable_broadcast_tail_r	  )r   r  r   r^  rangesr9   r   r   r   r  r   ra   r   r   rL  ry  r   r   r  r  r   r'   r  r  simplify_with_rangesr`   rx  )r  r  	read_varsr  read_rangesr   
write_varsreplacementstail_rangesread_var	read_size	write_var
write_sizer  tail_var
read_indexwrite_indexs   `                rs   r	  z'Scheduler._fusable_read_after_broadcast  sW   *  
>>
 
	 y>T]]*<EFS3C 00FKF		

k"k((*+		D ~~5??#44==ENN*
 77## 
4==)
 

 6868:=NNDIIz5::;
 	+6HiJ //	:F)2X&33		)Z0! &&x	:'FGF//:||*3{+;*<= H
 &/%7(%BL"$*K!+	+. 22tzz<0@tC
EJJ/0@K@

 !d3u
3S.TU//
KHHo Gs   J	c                H   t        |t        j                        sy|j                         sy| j                  j                  |j                  |j                        }| j                  j                  |j                  |j                        }t        |t              r||k(  ryyr  )r   r0   r  r   r)  r  r   r:   )rr   r  r  writing_node	read_namer	  s         rs   r	  z8Scheduler.fusable_stardep_write_and_read_on_empty_tensor.  s}     ,(B(BC--/))--diiC	**..uzz5::F
eW%)z*Aru   c                   t        | t              rt        |t              sy| |k(  ry| j                  |j                  k(  r!| j                         |j                         k(  S | j	                         |j	                         k(  S )a  Check if two deps refer to the same access pattern after normalization.

        Handles the case where FusedSchedulerNodes have more loop vars
        than a single SchedulerNode (e.g., 3 vars vs 2) by falling back
        to normalize() which merges loops before comparing.
        FT)r   r9   r^  r  r   dep1dep2s     rs   r2	  zScheduler.deps_match_normalized;  sn     $	**T92M4<==DMM)002d6V6V6XX ~~4>>#333ru   c                B    t         j                  j                  ||      S r~   )ra   r   get_dep_size_hint)rr   r   r  s      rs   r	  zScheduler.dep_size_hintM  s    ww((k::ru   c           
     n   |j                         |j                  z  syt        |j                  j                        }t               }d}|j                  j                  D ]  }t        |t              s| j                  j                  |j                  |j                        }	t        |      D ]  \  }
}|
|v r| j                  |j                  | j                        ||duxr |	|v       sA|j                  |
       |t!        | j#                  ||      | j#                  ||            z  }   |S )a0  Score vertical producer-output deps missed by exact dep scoring.

        Exact scoring is set-intersection based, but vertical legality can also
        accept normalized equivalent read/write deps. Give those pairs a memory
        score so heuristics do not discard them before legality runs.
        r   Nr	  )r  r  r   r   r   r   rQ  r   r9   r)  r  r   r  r	  ry  r  rY  r	  )rr   rE  rJ  r  ro	  r   matched_readsr	  r  r	  r  r  s               rs   *_score_fusion_memory_by_fusable_read_writez4Scheduler._score_fusion_memory_by_fusable_read_writeP  s7    ,,.1C1CCX))//0)3))00 	EeY/..225::uzzJJ$U+ 4%..KK 5 562$> E&*DD /  "%%a(S**4=**5+> E !		* ru   c                    | j                  |||      }t        |t              sJ |dk(  r|r| j                  |||      }|S )Nr  r   r	  )r  r   r   r	  )rr   r   r   r  ro	  r	  s         rs   r	  z+Scheduler._score_fusion_memory_for_can_fusez  sd     ((&? ) 

 %%%%A:4CC+E D E
 ru   c                     y r~   r   rr   r   r   r  return_is_mix_order_reductionr  s         rs   r  zScheduler.score_fusion_memory  s     ru   c                     y r~   r   r	  s         rs   r  zScheduler.score_fusion_memory  s     !$ru   c           	     d    fd}|r6t         j                  ||      r t         j                  ||      } ||dd      S t        |j                  t
        j                        r|j                  j                         s |j                         s|j                         r|j                  j                  |j                  j                  z  }|j                  j                  |j                  j                  z  }	dd}
d}|D ]@  }|	D ]9  } |
||      s|t         j                  |       j                  |            z  }; B  ||dd      S |j                  j                  |j                  j                  z  }|j                  j                  |j                  j                  z  }	t        |      t        |	      kD  r|	|}	}|D cg c]	  }||	v s| }}t         fd|D              }|r
 ||dd      S d}|dk(  r$ j!                  ||      r j#                  ||      } |||d      S c c}w )a2  
        The first term in our fusion score that estimates number of saved
        memory operations.

        This function scores fusion candidates based on shared memory access patterns.
        Higher scores indicate better fusion candidates.

        Scoring strategy:
        1. If nodes share exact memory deps (same buffer + same indexing), return
           the sum of shared dep sizes (original behavior).
        2. If no dependency score remains, check for same-buffer reads with
           different indexing (e.g., split operations reading different slices).
           - Give bonus if nodes read from exactly the same set of buffers
           - Score based on overlap ratio: common_buffer_size / total_read_size
           - High overlap (>50%) suggests good cache locality benefit from fusion
        c                    r| ||fS | |z   S r~   r   )r	  buffer_overlap_scoreis_mix_order_reductionr	  s      rs   _construct_return_valuez>Scheduler.score_fusion_memory.<locals>._construct_return_value  s#     -35KLL///ru   r   Tc                    | |k(  ryt        | t        t        f      r/t        |t        t        f      r| j                  |j                  k(  S yr  )r   r:   r9   r   r	  s     rs   _matchz-Scheduler.score_fusion_memory.<locals>._match  sD    4<dWi$89j7I.?  99		11ru   Fc              3  B   K   | ]  }j                  |        y wr~   r	  )r   r   r  rr   s     rs   r   z0Scheduler.score_fusion_memory.<locals>.<genexpr>  s     WSD&&sK8Wr  )r	  r8   r	  r8   )r   r  r   r   r   r0   r  r   r  r   r   rQ  rY  r	  r   r   _can_use_buffer_overlap_scoring&_score_fusion_memory_by_buffer_overlap)rr   r   r   r  r	  r  r	  r	  
node1_deps
node2_depsr	  	node1_dep	node2_depr   common_memory_depsr	  s   `  ``           rs   r  zScheduler.score_fusion_memory  s   2	0 %):)C)CE5)Q
 &66ueDE*5!T:: 5::r'A'ABJJ002  "  "**0053D3D3K3KKJ**0053D3D3K3KKJ E' 	!+ Ii3 ..y94;M;Mi;X"  +5!U;;&&,,u/@/@/G/GG
&&,,u/@/@/G/GG
z?S_,%/
J-7Mc3*;LcMMWDVWW*5!U;;  !A:$>>ueL#'#N#Nu$  'u.BEJJ Ns   	H-H-c                t   |j                         s|j                         ry|j                         s|j                         ryt        j                  st        j                  rU|j                         }|j                         }|r|syt        d |D              }t        d |D              }t               }|D ]  }|j                  D ]  }	t        |	j                  t              s|	j                  j                         s9t        |	j                        sO|	j                  j                         }
|
Lt        |
t        j                        r2|
j                         }||z  s|j!                  |	j                         |j!                  |	j                           |r|D ]  }|j                  D ]  }	t        |	j                  t              s|	j                  j                         s9|	j                  |v sH|	j                  j                         }
|
3t        |
t        j                        r|
j                         }||z  s  y  y  ||fD ]d  }|j"                  j$                  D ]I  }| j&                  j)                  |j*                        }|+|j                         s<t-        |      sH  y f y)a@  
        Check if buffer overlap scoring should be used for this node pair.

        Buffer overlap scoring handles split/cat patterns where nodes read from
        the same buffer at different indices. We skip it when:
        - Either node is a reduction (different memory access patterns)
        - Either node is a template
        - Both nodes are prologue/epilogue candidates for the same template,
          because horizontal fusion would prevent them from being absorbed
          into the template kernel. For example, in:
            q = a[:64, :]; k = a[64:, :]
            return mm(q + 2, k - 2)
          "q + 2" and "k - 2" both read from `a` and would get a high overlap
          score, but fusing them horizontally prevents prologue fusion into mm
          (resulting in 2 kernels instead of 1).

        We allow buffer overlap scoring when:
        - The node outputs are not actually in the template's allowed_prologue_inps,
          meaning they can't be prologue-fused anyway, so horizontal fusion doesn't
          prevent any optimization opportunity.
        FTc              3  <   K   | ]  }|j                           y wr~   r  r  s     rs   r   z<Scheduler._can_use_buffer_overlap_scoring.<locals>.<genexpr>*        +TsCLLN+Tr  c              3  <   K   | ]  }|j                           y wr~   r  r  s     rs   r   z<Scheduler._can_use_buffer_overlap_scoring.<locals>.<genexpr>+   r	  r  )r   r  r-   r  r  r<  r   r  r   r   rf   r$  r  r0   r  r	  r  r   r   r  r  r   r  )rr   r   r   node1_outputsnode2_outputsnode1_output_namesnode2_output_names&node1_prologue_eligible_template_usersr   r  r  allowed_inpsr   r   rE  s                  rs   r	  z)Scheduler._can_use_buffer_overlap_scoring   sf   4 5#5#5#7%"3"3"5&":":!--/M!--/M !!++Tm+T!T!++Tm+T!T
  3 % RII RD"499.?@ II1137		B
 )-		(C(C(E(4)2+B+B: ,9+R+R+TL1L@ F J J499 U CFFtyyQ%RR. 6( -C #		 -&tyy2CD $		 5 5 7 $		-S S -1II,G,G,IM,8Z -r/F/F> 0=/V/V/X#5#D+0 (-!--&  %++11 %C#66::388DH ,$0027A$%% ru   c                |    dd
 fdt        d |j                  j                  D              }t        d |j                  j                  D              }||z  syt        fd|j                  j                  D              }t        fd|j                  j                  D              }t	        ||      }|dk(  ryt        fd|j                  j                  D              }t        fd	|j                  j                  D              }	t	        ||	      }
|
|z  }|t
        j                  k\  r|
S dS )a8  
        Score fusion based on buffer name overlap when exact dep matching fails.

        This handles the split/cat fusion case where nodes read from the same buffer
        but at different indices (e.g., different slices from a split operation).

        Scoring logic:
        - If nodes read from exactly the same buffers: high bonus (encourages fusion)
        - For common buffers: score based on overlap ratio
          - overlap_ratio = common_buffer_size /
            max(node1_total_reads, node2_total_reads)
          - If overlap_ratio > threshold (e.g., 0.5): give proportional score
          - If overlap_ratio < threshold: minimal/no score (not worth fusing)

        Note on dynamic shapes:
        - When deps have unbacked symbols (dynamic shapes), dep_size_hint returns 0
        - In this case, we use count * 10 as a proxy for size
        - This ensures fusion still works for models with dynamic batch sizes

        Note on multiple deps from same buffer:
        - A node may have multiple MemoryDep entries for the same buffer name
          (e.g., 4 split reads from arg0_1 at different indices)
        - We sum ALL dep sizes for each buffer, not just take max
        - This ensures overlap ratio is calculated correctly when nodes read
          multiple slices from the same underlying buffer
        r  c                8    j                  |       }|dkD  r|S S r   r	  )r   r   FALLBACK_DEP_SIZErr   s     rs   get_dep_sizezFScheduler._score_fusion_memory_by_buffer_overlap.<locals>.get_dep_size   s%    %%c*D!84:)::ru   c              3  4   K   | ]  }|j                     y wr~   rt  ru  s     rs   r   zCScheduler._score_fusion_memory_by_buffer_overlap.<locals>.<genexpr>   r  rv  c              3  4   K   | ]  }|j                     y wr~   rt  ru  s     rs   r   zCScheduler._score_fusion_memory_by_buffer_overlap.<locals>.<genexpr>   r  rv  r   c              3  .   K   | ]  } |        y wr~   r   r   r   r
  s     rs   r   zCScheduler._score_fusion_memory_by_buffer_overlap.<locals>.<genexpr>         $
"%L$
r=  c              3  .   K   | ]  } |        y wr~   r   r
  s     rs   r   zCScheduler._score_fusion_memory_by_buffer_overlap.<locals>.<genexpr>   r
  r=  c              3  J   K   | ]  }|j                   v r
 |        y wr~   rt  r   r   r9	  r
  s     rs   r   zCScheduler._score_fusion_memory_by_buffer_overlap.<locals>.<genexpr>   *      %
xx<' %
r}  c              3  J   K   | ]  }|j                   v r
 |        y wr~   rt  r

  s     rs   r   zCScheduler._score_fusion_memory_by_buffer_overlap.<locals>.<genexpr>   r
  r}  )r   r8   r   r   )r   r   r   r   rY  r-   min_overlap_ratio)rr   r   r   node1_read_namesr  node1_total_read_sizenode2_total_read_sizemax_total_read_sizenode1_common_read_sizenode2_common_read_sizecommon_read_buffer_sizeoverlap_ratior
  r9	  r
  s   `           @@@rs   r	  z0Scheduler._score_fusion_memory_by_buffer_overlapi   sQ   < 	; &%R%:K:K:Q:Q%RR%%R%:K:K:Q:Q%RR (*:: !$ $
).):):)@)@$
 !
 !$ $
).):):)@)@$
 !
 ""79NO!#
 "% %
((..%
 "

 "% %
((..%
 "
 #&&<>T"U 02EE
 (58P8P'P#	
VW	
ru   c                   t        |      dk(  r|S i }|D ]  \  }}|j                         |j                         k(  sJ |j                         }t        | j                  |      j	                  ||            }||vr	||fg||<   p||   j                  ||f        t        |j                         t        j                  d            d   }t        |      dkD  sJ |S )Nr   r  r   )
r   r   r   r#  get_fusion_pair_priorityr   rZ  r  r  r  )rr   r8  "possible_fusions_group_by_priorityr   r   r  fusion_pair_priority&possible_fusions_with_highest_prioritys           rs   r  z4Scheduler.get_possible_fusions_with_highest_priority   s   
  A%##  	+ - 	LE5##%)9)9);;;;%%'F#&  (AA%O$  $+MMENL23GH 33GHOOEN	 25.446H<O<OPQ<R2

2. 9:Q>>>55ru   c                B    t        j                  j                  | g| S )z-
        Shim for list.sort(key=...)
        )ra   r  score_fusionrO  s     rs   r  zScheduler.score_fusion_key   s     yy%%d3U33ru   c                    t        t        j                  j                               }t	        | j
                        D ]9  }|j                  || j                         |j                  |j                         ; y)zg
        Populate node.last_usage recursively (also for the nodes within a FusedSchedulerNode)
        N)
r   ra   r   r  r   rq  r  r  r  r"  )rr   r  r   s      rs   r  zScheduler.compute_last_usage   s]    
 ))A)A)CDTZZ( 	8D 3T5L5LM&&t7	8ru   c                   t        | j                  t        j                  j                  z
  t        j                  j
                  j                  z
        D ]z  }|| j                  v rT| j                  |   }|j                         s2t        j                  j
                  j                  |j                         f|t        j                  j                  v st        j                  j                  |   }t        |t        j                        r*t        j                  j
                  j                  |       t        |t        j                  t        j                   f      r|j"                  }t        |t        j$                        r|j'                         sJ t        j                  j
                  j                  |j"                         } | j                  j)                          y)z*Free any buffers that are no longer neededN)r  r  ra   r   r  r  freedr  r  codegen_freer   rX  r   r0   rL  r  r  r   r  is_input_bufferrB  )rr   r   r   r  storages        rs   free_bufferszScheduler.free_buffers   sV   %%gg%%&gg""(()
 	DD
 t'''&&t,<<>GG((55chh?---gg**40c2#5#56GG((55c:b&7&79M9M%NO!hhG"7BMM:w?V?V?XXGG((55gllC)	D, 	!!'')ru   c                    | j                   j                         D ]  }|j                           | j                          y r~   )rl  r   flushr#
  )rr   r  s     rs   r%
  zScheduler.flush
!  s3    }}++- 	GMMO	ru   c                v   t        |t        t        f      sJ t        d   dxx   dz  cc<   t	        j
                  t        d            5  |j                          |j                          d d d        |j                  t        j                  j                         | j                          y # 1 sw Y   CxY w)Nr]  extern_callsr   F)increase_kernel_count)r   rR  r*  r   ra   set_kernel_handlerr5   r  r  r  r   r  r#
  )rr   scheduler_nodes     rs   codegen_extern_callzScheduler.codegen_extern_call!  s     &(LM
 	
 
 	^,1,!!&u"EF 	&002##%	& 	qww334		& 	&s   !B//B8c                P   t        |j                        r|j                  
J | d       t        j                  j                  |       t        |j                        }|t        d|j                         t               s|j                  dk(  rLt        j                  j                  |      x}j                  dk  rt        |t        j                               t        |j                        r,|j                  dk(  st!        t        j                                ||       S )Nz( should have been normalized in loweringzUnsupported device type: r      rs  )r[   r   r   ra   r   add_device_infor4   rS  r+   r`  r   get_device_propertiesmajorr<   inspectcurrentframer=   )rr   r  device_schedulingr  s       rs   create_backendzScheduler.create_backend"!  s    &++&&,,*B 	
h>?	
B 	
'5fkkB$!:6;;-HII|v%%*ZZ%E%Ef%MM\TTWXX(w7K7K7MNN$V[[E-A#G$8$8$:;; &&ru   c                    |J || j                   vr| j                  |      | j                   |<   | j                   |   S r~   )rl  r4
  r  s     rs   r#  zScheduler.get_backend7!  sB    !!!&$($7$7$?DMM&!}}V$$ru   c                    d fd}|j                         D ci c]8  }|j                  *|j                  j                         D ]  } ||      |fd  : }}}t        |j	                               }|rMt        |t        j                  d            \  }}t        j                  j                  j                  |       y y c c}}w )Nc                    | j                   vrLj                   j                  t        | j                  j                        D  ci c]  \  }} | |
 c} }       j                       S c c} }w r~   )r  r  r  r   rq  )r  r  rr   s     rs   	get_orderz*Scheduler.enter_context.<locals>.get_order>!  s\    ,,,$$++i>V,WdaQT,WX''** -Xs   A+
r   r  )r  ztorch.fx.Noder   r   )r   r   r  r   r!  rY  r  r  ra   r   r  enter_context)rr   r   r8
  r  r&  r$  r=  lasts   `       rs   r9
  zScheduler.enter_context=!  s    	+ ^^%
vv!VV'')	
  q\1t#

 
 w||~&'x':':1'=>GAtGG  ..t4 
s   =Cc                    	 | j                   |   j                  }t        fd|D              xr || j                  vxr || j
                  vS # t        $ r Y yw xY w)NFc              3  ^   K   | ]$  }|j                   xs |j                         v  & y wr~   )r  r  )r   r  r  s     rs   r   zAScheduler.can_buffer_be_removed_through_fusion.<locals>.<genexpr>W!  s)     VC3C CCVs   *-)r  r  KeyErrorr   r)  r  )rr   r   r  r  s     ` rs   r  z.Scheduler.can_buffer_be_removed_through_fusionO!  sn    	$$T*00E VPUVV 4D1114D333	
  		s   A 	AAc                   |j                   }t        |t        j                  j                  j
                        rk|j                  x}r]t        |      \  }}|t        j                  v s|t        j                  v r+t        |t        j                  j                        sJ d| S t        j                  j                  j                  j                  st        j                  yt        |t               r)|j"                  D ]  }| j%                  |      }|s|c S  y|j                   J |j'                         s|j)                          dS t        |j                   t        j*                        ryt        |j                   t        j,                        ryt/        |j                   dd      ryt1        |j                         ry	| j3                  |      x}r|S t        j                  j4                  rt7        |      ry
y)z
        Return the reason why we should partition the inductor graph on this node,
        or None if the node is cudagraphable.
        zcustom partition op: Nz6partition includes all ops when cudagraphs is disabledz opszDeviceCopy opszConditional opsunbacked_bindingszunbacked binding opszCUDAGraph-unsafe custom opszdynamic shape ops)r   r   r`  ra  r0   r
  rS  rV   r-   custom_should_partition_opsr  r  r   re   rN   r  r   r   should_partitionr[   r   
DeviceCopyrJ  r  rZ   &_uses_cudagraph_unsafe_unbacked_symintcudagraph_skip_dynamic_graphsr  )rr   r   r	  r  op_overload_packet_nameop_overload_namer  r  s           rs   rA
  zScheduler.should_partition\!  s    ))gu11@@A%%%B%8DR8H5#%5'6+M+MM#v'I'II!"ejj&;&;<<<./?.@AA &&--886>>FKd./ "..u5!M" yy$$${{}oo'(--dii/#dii0$499148)!$)),0@@FF6FM ==66-d3*ru   c                T   t               }t        j                  s|S | j                  D ]  }|j                  }|t        |t        j                  j                  j                        sA|j                  }|Pt        |      \  }}|t        j                  vr|t        j                  vr|j                         D ]g  }t        j                  j                  j!                  |      }t#        |t$        j&                  t$        j(                  f      sW|j+                  |       i  |S )zc
        Collect output unbacked symints from ops in config.cudagraph_unsafe_unbacked_ops.
        )r   r-   cudagraph_unsafe_unbacked_opsrq  r   r   r`  ra  r0   r
  rS  rV   r  ra   r   r   r   r)   r*   UNBACKED_INTUNBACKED_FLOATr  )rr   unsafe_symintsr   r	  r  rE
  rF
  syms           rs   &_get_cudagraph_unsafe_unbacked_symintsz0Scheduler._get_cudagraph_unsafe_unbacked_symints!  s   
 4><33!!JJ 	,DiiGgu'9'9'H'HI$$Bz8DR8H5#%5'v/S/SS$F,P,PP779 ,gg&&//4!#(9(94;N;N'OP"&&s+,'	,0 ru   c                    | j                         }|sy t        |      }|D ]I  }t        j                  j                  j                  |      }|j                  D ]  }||v sd| c c S  K y )Nz'uses cudagraph-unsafe unbacked symint: )rM
  r  ra   r   r   r   r&   )rr   r   rK
  node_symbolsrL
  simplified_symfree_syms          rs   rC
  z0Scheduler._uses_cudagraph_unsafe_unbacked_symint!  s~     DDF5d; 	PCWW--66s;N*77 P~-DXJOOP	P ru   c                    i }|j                  t        j                  j                         | j                  D ]3  }|j
                  j                         D ]  \  }}|j                  ||<    5 |S )z~
        Return a mapping from name strings to the corresponding graph inputs or
        base scheduler node outputs.
        )r  ra   r   rX  rq  r.  r  r   )rr   rI  r   r   scheduler_buffers        rs   get_name_to_nodeszScheduler.get_name_to_nodes!  sr     PRAGG001JJ 	;D*.*>*>*D*D*F ;&&%5%:%:T";	; ru   c           	        t        t        j                  j                        D ci c]  \  }}||
 }}}t        t        j                  j	                               D ci c]  \  }}||
 }}}g t        j                  _        t        |      D ]  \  }}|j                  rg }|j                  D ]"  }|j                  |j                  |             $ g }	|j                  D ]0  }
|	j                  |j                  |
j                                      2 t        j                  j
                  j                  t        |||	|j                                yc c}}w c c}}w )z
        computes a mapping from partition input/output indices to graph input/output
        indices for each partition.
        N)r  ra   r   rX  r  partition_mapsskip_cudagraphinput_nodesr   r  output_nodesr  rW   constant_names)rr   
signaturesrI  r   name_to_graph_input_indexname_to_graph_output_indexpartition_id	signatureinput_mappingoutput_mappingr   s              rs   compute_graph_partition_mapsz&Scheduler.compute_graph_partition_maps!  sT    (11E1E'F%
##tD#I%
! %
 (11I1I1K'L&
##tD#I&
" &
 "$'0'< 	#L)''
 M!-- J$$%>%B%B4%HIJ  N!.. W%%&@&D&DT]]_&UVW GG""))! !",,	!	%
&
s   E!E c                   	 	 	 	 dd	 	 	 	 dd} t               j                  d |D         } |j                  fd|j                         D           ||      }t               }|D ]F  }t        j
                  j                  j                  |      }|j                  |j                         H t        t        |t        j                  d                  S )	ai  
        Returns all symbol inputs which are required to be in scope to successfully
        perform codegen for this graph partition, including:
        - free symbols used in partition nodes
        - free symbols in partition input/node shapes, strides, and offsets. This is needed
          for recording cudagraphs for tensors with dynamic shapes.
        c                    t        | t        j                        r
t               S t        | t        j                        rt        |       S t        dt        |              )zW
            Gets symbols used in input node shapes, strides, and offsets.
            zUnsupported input node type: )r   r0   rL  r   rl  r  r  r   r   s    rs   get_input_node_symbolszKScheduler.get_graph_partition_symbol_inputs.<locals>.get_input_node_symbols"  sN     $ 2 23!|#D")),)$// *,I$t**VWWru   c                &    t        d | D              S )z
            Filters a set of symbols that are required for codegen. Skip symbols
            that are always internal to kernels, such as SymT.TMP, SymT.INDEX,
            and SymT.R0_INDEX.
            c              3     K   | ]N  }t        |t        j                  t        j                  t        j                  t        j
                  f      r| P y wr~   )r)   r*   SIZEFLOATrI
  rJ
  r   r9  s     rs   r   zVScheduler.get_graph_partition_symbol_inputs.<locals>.filter_symbols.<locals>.<genexpr>."  sH      !		

))++	 s   AAr   )symbolss    rs   filter_symbolszCScheduler.get_graph_partition_symbol_inputs.<locals>.filter_symbols&"  s         ru   c              3  2   K   | ]  }t        |        y wr~   r  r  s     rs   r   z>Scheduler.get_graph_partition_symbol_inputs.<locals>.<genexpr>="  s     It,T2Ir  c              3  .   K   | ]  } |        y wr~   r   )r   r   re
  s     rs   r   z>Scheduler.get_graph_partition_symbol_inputs.<locals>.<genexpr>@"  s     Lt$T*Lr=  r   r  )r   z+ir.IRNode | sympy.Expr | ir.TorchBindObjectr   OrderedSet[sympy.Symbol])rk
  ro
  r   ro
  )r   r  r  r   ra   r   r   r   r&   r  r  
attrgetter)	rr   	partitionrX
  rl
  candidate_symbolsr  r9  symplified_sre
  s	           @rs   !get_graph_partition_symbol_inputsz+Scheduler.get_graph_partition_symbol_inputs	"  s    	X=	X%	X 	-	%	, 7Ijl6H6HIyI7
 	!  L{7I7I7KL	
 ++<=(2" 	2A77++44Q7LJJ|001	2
 &(*=*=f*EFGGru   c           
         g }t        t        j                  j                               } j	                         }d fdt        t        |      t        |            D ]2  \  }}t               }|D ]+  }	|j                  |	j                  j                                - |j                  |      }
t        j                  j                  |D 	cg c]  }	|	j                   c}	      }t        |j                  |j                   z  D cg c]  }t#        |t$              s|j&                    c}      |z
  }t         fd|D              }t               }|D ]  }	|j                  |	j(                          ||z
  D cg c]  }||v r|
 }}|j                  |       |D ci c]  }||v r|||    }}|D ci c]  }||v r|||v  }}|D cg c]  }||v r||vr| }}|
j                  |       t         fd|
D              }
|
D cg c]  } |      s||    }}|D cg c]!  }|t        j                  j*                  v s |# }} j-                  ||      }t/        ||||||      }|j1                  |       |j3                  ||
z
        }5 |ddd   S c c}	w c c}w c c}w c c}w c c}w c c}w c c}w c c}w )z
        Gets signature for each graph partition, including input nodes, output nodes, and
        whether deallocating an input within graph partition.
        c                    j                   j                  | d      }|yt        |j                  j                  t
              r'j                  j                  | d      x}r |      S yy)z
            Checks if buf_name resolves to a NoneLayout buffer (following mutation_real_name).
            Buffers with NoneLayout are not allocated so graph partition should not
            take them as inputs or outputs.
            NFT)r  r  r   r   r  rD   r  )r  r   r3  is_unallocated_bufferrr   s      rs   rw
  zFScheduler.get_graph_partition_signature.<locals>.is_unallocated_bufferY"  sh     ""&&x6C{#((//:6 !% 7 7 ; ;Hd KK9K0;;ru   c              3  V   K   | ]   }j                   j                  ||       " y wr~   r  r  rb	  s     rs   r   z:Scheduler.get_graph_partition_signature.<locals>.<genexpr>"  ,      / ''++D$7/r  c              3  V   K   | ]   }j                   j                  ||       " y wr~   ry
  rb	  s     rs   r   z:Scheduler.get_graph_partition_signature.<locals>.<genexpr>"  rz
  r  Nr  )r  r,  r   r}   )r   ra   r   r  rT
  ry  r   r  r.  r!  r  r/   r  r  r   r   rQ  r   r;   r   r"  rp  rt
  rA   r   r  )rr   
partitionsskip_cudagraphsr[
  unmet_output_namesrI  rq
  rW
  output_namesr   returned_output_namesr   r  partition_input_namesr  r   extra_input_namesrX
  input_deallocationextra_output_namesrY
  rZ
  symbol_inputspartition_signaturerw
  s   `                       @rs   get_graph_partition_signaturez'Scheduler.get_graph_partition_signatureM"  s^    
'(@(@(BC--/	, *-Z (?";*
 g	%I~ -7LL! A##D$8$8$=$=$?@A %1$=$=>P$Q! '11<<.78d!!8K  "-!2!2[5G5G!G)!W5   " %/ /1/ %!
 5?L ! =$++DOO<= 2L@!<' ! !
 "(():; 2<' l4((K  2"<' d222" " 2"<'D8L,L " " "(();<$. /1/ %! 2,T2 T"L  "7$!''BSBS:SN  !BB;M #:"# 12!6!<!<"%::"Kg	R $B$y 9*!
""s6   J

#J$
>J)"J.9J3J8J=%!KKc                   |j                   j                         D ci c]  \  }}||vr|| }}}|j                  j                         D ci c]  \  }}||vr|| }}}|j                  D cg c]  }|j	                         |vr| }	}|j
                  D cg c]	  }||vs| }
}t        |j                  ||	||j                  |
      S c c}}w c c}}w c c}w c c}w )z
        Updates the partition signature by removing buffers specified in
        removed_buffers. See [Note: Removed Graph Partition Arguments]
        )	rX
  r  r
  rY
  maybe_get_namerZ
  rA   r
  rW
  )rr   r_
  r  r   r"  rX
  r	  r
  r   rY
  rZ
  s              rs   .clean_removed_buffer_from_partition_signaturesz8Scheduler.clean_removed_buffer_from_partition_signatures"  s    !* 5 5 ; ; =
f?* &L
 
 '99??A
c?* #I
 
 "..
""$O; 
 
 '55
_9TD
 
 '##$$
 	
%






s   CC1C	C!&C!c                p   	
 ddl 	t               g g t        |      D ci c]  \  }}||
 c}}d	 fd
d
fd}|D ]5  }t        |j                  j
                        |<   |   dk(  s. 
|       7 g }d}|t        |      k  rsr}r0	j                        \  }}|j                  |        ||       r0r0	j                        \  }}|j                  |        ||       r0|dz  }|t        |      k  rrzr}|t        |      kD  rt        d      |S c c}}w )a  
        Reorder nodes to minimize the number of partitions via a bfs
        topological sort. This is the optimal reordering such that the
        number of partitions cannot be reduced further. This may be
        sub-optimal for other metrics such as peak memory. This does not
        change relative orders of two cudagraphable nodes, nor the
        relative order of two non_cudagraphable nodes.
        r   Nc                    |    | f}j                  |       rj                  |       y j                  |       y r~   )rA
  heappush)r   node_with_indexcudagraphable_nodesheapqnode_to_indexnon_cudagraphable_nodesrr   s     rs   insert_pending_nodeszHScheduler.reorder_for_minimizing_partition.<locals>.insert_pending_nodes#  s>    ,T2D9O$$T*6H2ODru   c                    | j                   j                  D ]*  }|   dkD  sJ |xx   dz  cc<   |   dk(  s# |       , y r  )r'  
succ_nodes)r   	succ_noder
  node_to_indegrees     rs   update_indegreezCScheduler.reorder_for_minimizing_partition.<locals>.update_indegree#  sT    !]]55 4	'	2Q666 +q0+#I.!3(3	4ru   r   z
                Failed to schedule, while loop ran too long when
                reordering for minimizing the num of partitions
                r>  )	r
  rx  r  r   r'  
pred_nodesheappopr   rS  )rr   rq  rI  r   r
  r  	num_itersr=  r
  r
  r
  r
  r
  r
  s   `       @@@@@@rs    reorder_for_minimizing_partitionz*Scheduler.reorder_for_minimizing_partition"  sU    	9=CEGI4=e4DEysDsE	E 	E	4  	+D%()A)A%BT"%*$T*	+
 -/	#e*$#':)--(?@4%% *
 &--(;<4%% &
 NI #e*$#': s5z!  ] Fs   D2c           	        ddl m}m} t        t        j
                  j                               } ||| j                  | j                  t        t        j
                  j                  j                               |      \  }}| j                  |      } ||||      \  }}	||t        j                  j                  z  k  r|S |S )zx
        Reorder nodes to minimize the number of partitions if this only slightly
        increase peak memory.
        r   )estimate_peak_memoryprepare_planning_info)r  r
  r
  r   ra   r   r  r  r  rX  r!  r
  r-   r   !cudagraph_partition_memory_budget)
rr   rq  r
  r
  r   default_peak_memoryr4  reordered_nodesreorder_peak_memoryr=  s
             rs   r  z0Scheduler.maybe_reorder_for_minimizing_partition?#  s     	H"177#;#;#=>:O##qww++0023;
77 ??F!57"
Q
  !FMM$S$SST #"ru   c                   g }g }g }dd}|D ]n  }| j                  |      du}|r*t        |j                        dk(  r|j                  |       B|r ||      r|j                  |       ^|j                  |       p ||z   |z   S )a  
        Reorder a node if it should be partitioned and has simple dependency:
        1. move a partitioned node to the front if it has no dependency
        2. move a partitioned node to the back if it is only used by OutputNode
        3. otherwise do not reorder
        c                    | j                         D ]0  }|j                  D ]  }t        |j                  t              r  y 2 yr  )r<  r  r   r   r  )r   r   r  s      rs   only_output_userzPScheduler.reorder_for_partition_with_simple_dependency.<locals>.only_output_usern#  sC    '') %99 %C%chh
;$%% ru   Nr   r*  )rA
  r   r3  r   )rr   rq  frontmiddlebackr
  r   rA
  s           rs   r  z6Scheduler.reorder_for_partition_with_simple_dependency`#  s     *,*,(*	  	$D#44T:$FC(?(?$@A$ET"!&6t&<D!d#	$ v~$$ru   c                   g }d}g }g }| j                   D ]S  }| j                  |      du}|r)||k7  r$|j                  |       |j                  |       g }|}|j                  |       U |r"|j                  |       |j                  |       t        j                  j
                  }|dkD  rXt        t        ||            D ]@  \  }\  }	}
|
rt        d |	D              }||k  s$d||<   t        j                  d|||       B | j                  ||      }| j                  |       | j                  ||       ||fS )z
        Given a list of BaseSchedulerNodes, split into a list of
        graph partitions and compute partition input/output signatures.
        TNr   c              3  @   K   | ]  }t        |t              sd   ywr   N)r   r  r  s     rs   r   z,Scheduler.graph_partition.<locals>.<genexpr>#  s#      ')!-CD 'r  zFPartition %d has %d kernels, below minimum size %d, skipping cudagraph)r|
  r}
  )rq  rA
  r   r-   r   cudagraph_min_partition_sizer  ry  r   cudagraphs_logr  r
  rb
  _log_graph_partitions)rr   r|
  rW
  cur_partitionr}
  r   node_should_partitionmin_sizer  rq
  skipkernel_countr[
  s                rs   r  zScheduler.graph_partition#  sm    +-
')JJ 	'D$($9$9$$?t$K!3H!H!!-0&&~6 "2N  &	' m,"">2 ====a<(1#j/2R(S $$It#& '!*' $L
 $h.-1*&,,d($	" 77!? 8 

 	))*5"":z::%%ru   c                `   t         j                  t        j                        sy t	        d t
        j                  j                  D              }|sy t        d |D              }t        |      |z
  }t         j                  dt        |      ||       t        t        ||            D ]  \  }\  }}t         j                  d|t        |      |j                  rdndt        |j                        t        |j                               |j                  sm|D ]  }	| j!                  |	         y )Nc              3  2   K   | ]  }t        |        y wr~   )r[   )r   r  s     rs   r   z2Scheduler._log_graph_partitions.<locals>.<genexpr>#  s     OVF^Or  c              3  :   K   | ]  }|j                   rd   ywr
  )rW
  rj
  s     rs   r   z2Scheduler._log_graph_partitions.<locals>.<genexpr>#  s     !Pq?O?O!!Ps   zCCreated %d graph partitions: %d cudagraphable, %d non-cudagraphablez3  Partition %d: %d nodes, %s, inputs=%d, outputs=%dznon-cudagraphablecudagraphable)r
  r  r  r  r  ra   r   device_typesr   r   r  r  ry  rW
  rX
  rY
  _log_non_cudagraphable_node)
rr   r|
  r[
  has_gpu_devicecudagraphable_countnon_cudagraphable_countr  rq
  r_
  r   s
             rs   r
  zScheduler._log_graph_partitions#  s   
 **7==9 O!'':N:NOO!!PZ!PP"%j/4G"GQ
O#		
 *33z:3N)O 	;%A%	9  EI'0'?'?#_I))*I**+ ''% ;D44T:;	;ru   c                   | j                  |      }|sy|j                         }|j                  |j                  j                         nd}d| g}t	        |j                        j
                  }|j                  d|        |F|j                   ddj                  d |j                  D               d}|j                  d|        t        j                  d	|dj                  |             |Z|j                  j                  d
d      }|r;|j                         j                  d      D ]  }	t        j                  d|	        yyy)z)Log details for a non-cudagraphable node.Nzreason=zir=rE  r]  c              3  2   K   | ]  }t        |        y wr~   )r,  )r   r  s     rs   r   z8Scheduler._log_non_cudagraphable_node.<locals>.<genexpr>#  s     2Pa3q62Pr  r@  zfx=z
    %s: %sr  rH  z         %s)rA
  r  r   r_  r   r   r   r  rd  r  r
  r  r  r  stripsplit)
rr   r   r  rM  rb  partsir_typefx_strr  lines
             rs   r
  z%Scheduler._log_non_cudagraphable_node#  s1   &&t,MMO	151F$))++-D6(#$tyy/**s7)_%'q2P7<<2P)P(QQRSFLL3vh(\9dii6FG !,,**=$?K'--/55d; >D"((=>  ru   c                    t        d      5  t        j                  j                  j                  r| j                         n| j                  | j                        	 cd d d        S # 1 sw Y   y xY w)NzScheduler.codegen)r   r`  ra  r-   r  _codegen_partitions_codegenrq  rq   s    rs   r  zScheduler.codegen#  sX    -. 	 ??))99 ((*]]4::.	 	 	s   AA&&A/c                   ddl m} t        j                  j                  }t        | j                        }t        j                  j                         5  t        j                  j                  dd| ||       t        j                  j                  j                         }| j                  |       t        t        j                  j                  |      sJ t        j                  j                  |z
  }| j                  ||      }|t        j                  j                  _        t        j                  j                  j                          t        j                  j                   }t        j                  j                  j#                  t        j                  j$                        \  }	}
ddd       t        j                  j                  j'                  	       t        j                  j                  j)                  ||       t        j                  j                  j*                  j-                  |j.                  D cg c]  }|j1                          c}       y# 1 sw Y   xY wc c}w )z,Codegen a partition given its inputs/outputsr   )SubgraphPythonWrapperCodegenT
partition_)is_subgraphsubgraph_nameparent_wrapper_codepartition_signaturesN)r  r
  ra   r   r  r  ro  set_current_wrapper_codeinit_wrapper_coder  rx  r
  r   r
  r
  write_prefixr   generateis_inferencedefine_subgraph_launcher_fncodegen_partition_call	allocatedr  rY
  r  )rr   rq
  r_
  r
  r
  graph_partition_idremoved_buffers_before_codegenremoved_buffers_during_codegen
graph_namepartition_coder=  r   s               rs   _codegen_partition_wrapperz$Scheduler._codegen_partition_wrapper#  s    	Bgg22!$"?"?@WW--/ 	TGG%%  *+=*>?$7%.	 &  ./WW-D-D-I-I-K*MM)$ agg224PQQQ''*HH + KK9I 9BAGG  5GG  --/J ! 4 4 = =agg>R>R SNA?	TB 	
88^T	334F	R	&&--)2)?)?@T]]_@	
I	T 	TJ As   EI$I0$I-c                L     t         j                  d fd       } |       S )Nc               3    K   j                          j                  ryt        j                  j                        rZj                  j                  J d       t
        j                  j                  j                  j                  j                         	 d  j                  rGt        j                  j                        r(t
        j                  j                  j                          d _        y # j                  rGt        j                  j                        r(t
        j                  j                  j                          d _        w xY ww)Ndevice should have an index)
%update_graph_partition_default_devicerv  rQ   r   r   ra   r   r  codegen_device_guard_entercodegen_device_guard_exit)r|
  rr   r[
  s   rs   ctxz1Scheduler.use_default_device_context.<locals>.ctx1$  s    66z:N**/@++000 2288D 1D $$??//553..3D//444 GG((BBD.2+	 ..3D//444 GG((BBD.2+s    BEC;  AE;AEE)r   zIterator[None])
contextlibcontextmanager)rr   r|
  r[
  r
  s   ``` rs   use_default_device_contextz$Scheduler.use_default_device_context.$  s&     
	"	"	3 
#	3* uru   c                    t        |      dk(  r|d   j                  sy dd}	 	 	 	 	 	 dd}d }t        ||      D ]  \  }}|j                  r ||      } n |y t        ||      D ]  \  }}|j                  s |||      r y  || _        y )Nr   r   c                4    | d   j                         }|J |S r   r  )rq
  partition_devices     rs   get_cudagraph_partition_devicezWScheduler.update_graph_partition_default_device.<locals>.get_cudagraph_partition_deviceX$  s'    (|668#///##ru   c                @    | D ]  }|j                         }||k7  s y yr  r  )rq
  target_devicer   r  s       rs   all_on_target_devicezMScheduler.update_graph_partition_default_device.<locals>.all_on_target_device]$  s/     " !*]* ! ru   )rq
  rg   r   r  )rq
  rg   r
  r  r   r}   )r   rW
  ry  rv  )rr   r|
  r[
  r
  r
  cudagraph_partition_devicerq
  r_
  s           rs   r
  z/Scheduler.update_graph_partition_default_deviceI$  s     z?a
1(D(D 	$
	$	5A		 &*"$'
J$? 	 Iy++-KI-V*	 &-$'
J$? 	 Iy''0D51 		 'A#ru   c                   | j                         \  }}t        |      dkD  rt        d   dxx   t        |      z  cc<   | j                  ||      5  t	        ||      D ]V  \  }}t        |      dk\  sJ dt        |              |j
                  r| j                  |       E| j                  ||       X 	 ddd       t        | j                        }t        j                  j                  j                  |       |dkD  rqt        j                  j                  J |t        t        j                  j                        k(  s.J d| dt        t        j                  j                                yy# 1 sw Y   xY w)	z
        Split nodes into partitions and codegen each partition into separate functions.
        This allows further applying different optimizations (e.g., cudagraph) to
        each function.
        r   r]  cudagraph_partitionsz5Each partition must have at least one node but found Nr   zExpect z partition maps but got )r  r   r   r
  ry  rW
  r
  r
  r  ro  ra   r   r  set_all_partition_namesrV
  )rr   r|
  r[
  rq
  r_
  num_partitionss         rs   r
  zScheduler._codegen_partitionsx$  sh    "&!5!5!7
Jz?QZ !78C
OK8,,ZD 		J(+J
(C J$	99~* KCPYNK[\* ++MM),33IyIJ		J d;;<	44^D A77))555!S)?)?%@@ .))A#aggF\F\B]A^_@ 		J 		Js   A&E55E>c                   t         j                  rdd l}t        j                         }t               }t        |      D ]  }|j                  dk(  r/|j                  |j                  j                  j                  k(  r nQ|j                  |j                  f}||vs"J d|j                   d|j                   d       |j                  |        | j                  | _        | j                   J | j                  rBt         j"                  j$                  r(t&        j(                  j*                  j-                          t&        j(                  j*                  j/                          |D ]  }t0        j3                  t4        j6                        r4	 t0        j9                  d|j;                         |j=                                | jA                  |       t         jB                  rDt&        j(                  j*                  jE                  d |jF                  jH                  D               |jK                         x}rs|| j                  k7  s |jM                         s|jO                         r| jQ                          || j                  k7  r$| j                  rctS        | j                  jT                        rD| jV                  | jY                          t&        j(                  j*                  j[                          || _        tS        |jT                        r|j\                  J d	       d
}	| j_                         r5t        | j`                  jc                               }
|
rte        |
      d
z   nd
}	t&        j(                  j*                  jg                  |j\                  |	| jh                         | j_                         r| j                  | jk                  |       t&        j(                  j*                  jm                  d |jF                  jH                  D               || _7        | jp                  js                  |jt                         |jO                         rP|jw                  ty        |j{                                     \  }}}| j}                  |      j                  |||       nf|jM                         r| j                  |       nC|j                         rxt        j                  t        |      }| j}                  |      }d
dlEmF} d
dlGmH} d
dlImJ} t        ||||f      r|}nt        dtU        |             |j                  |       nt        |t              r!| j}                  |      j                  |       nt        |t              r!| j}                  |      j                  |       nYt        |t        t        f      r!| j}                  |      j                  |       n"t        |t              sJ |j                          t         j"                  j                  r| j}                  |      j                          | j                  js                  |j                                | j                  js                  |j                                t        |t              sP|jK                         }|>|jT                  dk7  r/| j}                  |      j                         r| jQ                          t        d |j{                         D              r	|| _        d | _         | j                  | j                  k7  rU| j                  J tS        | j                  jT                        r(t&        j(                  j*                  j[                          d | _        | jQ                          y # t>        $ r( t0        j9                  d|j;                                Y "w xY w)Nr   _compile_innerzDuplicate stack frame :zs; did you add a decorator to one of the functions in this stack trace?  If so, try using a context manager instead.z5Generating code for node %s with estimated runtime %fz6Generating code for node %s with estimated runtime 0.0c              3  4   K   | ]  }|j                     y wr~   rt  ru  s     rs   r   z%Scheduler._codegen.<locals>.<genexpr>$  s      D!$CHHDrv  r
  r   c              3  4   K   | ]  }|j                     y wr~   rt  ru  s     rs   r   z%Scheduler._codegen.<locals>.<genexpr>$  s      C Crv  )CUDACombinedSchedulingr  )XPUCombinedSchedulingztype(self)=r  c              3  <   K   | ]  }t        |t                y wr~   r)	  r  s     rs   r   z%Scheduler._codegen.<locals>.<genexpr>:%  s     JA:a/Jr  )_r-   "check_stack_no_cycles_TESTING_ONLYtorch._dynamo.convert_frame	tracebackextract_stackr   r   r   filename_dynamoconvert_frame__file__linenor  rv  r  rs  r   autotune_at_compile_timera   r   r  write_get_raw_stream_headerregister_alignment_check_inputsrT  r  r  r  r  r  rh  rS  r9
  size_assertscodegen_deferred_input_assertsr   r   r   r  r  r%
  rQ   r   current_stream_idxgenerate_stream_ctx_exitr
  r   r  rv  r   rY  r
  r  generate_stream_ctx_switching!codegen_deferred_alignment_copiesrt  r  r  r"  r  r   r   r#  codegen_templater+
  r  r  r  rB   codegen.cuda_combined_schedulingr
  r  r  #codegen.xpu.xpu_combined_schedulingr
  r   rw  codegen_combo_kernelr  codegen_nested_reductionr  codegen_mix_order_reductionr   r   codegen_noder  r  debug_sync_kernelcodegen_syncr  r  r  r  ready_to_flushr   )rr   rq  r`  r  rD  framer  r   r  num_streamsunique_streamsr  r  r  backend_r
  r  r
  r  s                      rs   r
  zScheduler._codegen$  s>   44.++-E7A|D!%  JJ"22%--*E*E*N*NN~~u||4$ ,U^^,<Aell^ LJ J
  #99!!))) &&6==+Q+QGG  <<> 	
<<> E	*D.
IIO224 t$ ""$$CC D(,(8(8(>(>D  **v*d111~~''')JJLT000**/@++000  22> 99;,,FFH*0D'(5%||7V9VV7&'779-78K8K8R8R8T-UN;IN 3a 7q ( ,,GG"LL' ;; ++-$2E2E2Q2248 GG  BB C$($4$4$:$:C  !%D%%,,T__=!484W4W)*51-   (99!8X !((."{{#=tD++F3T8V#%;=RS 'G(KDJ=)9::,,T2D"78  (AA$GD"9:  (DDTJD#5}"EF  (55d;!$(>??? }}..  (557''..t/D/D/FG%%,,T-E-E-GHd$:;*&v-((0??AJJLJ9IJJ%)"%)"KE	*N $"="== &&222 !4!4!9!9: $$>>@!

U ! IIPs   3^44-_%$_%c                    |d   j                         }| t        j                  _        || _        |J | j                  |      }|j                  ||      S )rq  r   )r   ra   r   r  r  r#  benchmark_combo_kernel)rr   r  node_benchmark_resultsr  r  s        rs   r  z Scheduler.benchmark_combo_kernelL%  sZ     1((* $!!!""6*--i9OPPru   c                   |}|d   j                         t        fd|D              sJ d       t        j                  syddlm} dg }}i }t        |      D ]  \  }}|j                         }	| j                  |	      rt        j                  d       	 | j                  |	      \  }
}|
|f||<   t        j                  |
      rt        j                  d|        y		 ||
z  }|j                  |        	 | j                  ||      \  }}}||z
  dk  xs |dk  }t        j!                  t"        j$                        rP||kD  s|r%t        j                  dt'        ||z  d             n$t        j                  dt)        ||z  d             ||z
  |k  xs |S # |$ r.}d
t        |      v rt        j                  d       Y d}~ y d}~ww xY w# |$ r-}d
t        |      v rt        j                  d       Y d}~y d}~ww xY w)r  r   c              3  D   K   | ]  }|j                         k(    y wr~   r  )r   r   r  s     rs   r   z4Scheduler.speedup_by_combo_kernel.<locals>.<genexpr>c%  s     K44??$.Ks    z<All nodes in a combo kernel group must be on the same deviceTr  rL  z<ComboKernel: benchmarking may not accurate due to atomic_addz;ComboKernel benchmark: register spilling of %d-th subkernelFr  zCComboKernel benchmark: return True because of loop-carried variableNg333333?z/can fuse (benchmark): fusing causes %sx speedupr  z3cannot fuse (benchmark): fusing causes %sx slowdown)r   r   r-   r  r  r  r  r   r  r  r  rr  r  r  r,  r   r  r  r  rJ   rL   )rr   rq  subkernel_nodesr  r  
path1_listr  r  r  r  r{  r  r&  r  	ms2_clone_path2_listsmall_kernelr  s                    @rs   rZ  z!Scheduler.speedup_by_combo_kernelZ%  s(      #..0K?KK 	
J	
K ,,;rZ!#!/2 	$HAu)I ##I.  R55i@D13T
&u-::b>$$U ! " 2ICd#9	$<	*.*E*E!7+'CK Y,9c	""7==1SyL  E#)C2
   Ic	#0
 Y$44Q $ *c!f4$$]      	&#a&0  Y 	s=   AF9G G""G
GGH"H ?H  Hc                p    | j                   |   }|j                  J |j                  j                         S r~   )r  r   
get_layout)rr   r  r   s      rs   get_buffer_layoutzScheduler.get_buffer_layout%  s5    x(xx###xx""$$ru   c                   | j                   D ]  }|j                         s|j                  j                  D ]  }t        j
                  j                  j                  |j                        }|s9t        |      dk(  sHt        |j                  t        t        f      ri|j                         g k(  s}t        j
                  j                  j!                  |j                           y r  )rq  r[   r   r   ra   r   r  r  r   r@   r   r  rD   rC   rc  zero_dim_cpu_tensor_listr  )rr   r   r  r"  s       rs   ru  z$Scheduler.update_zero_dim_cpu_tensor%  s    JJ 	HD{{} ,,22 
HDWW3377		BF+F3u< *"MMJ8I+J! #OO-388<<TYYG
H	Hru   c                H    | j                   | j                   j                  S y)z:CUDA Stream index that current scheduler node assigned to.N)r  r  rq   s    rs   r
  zScheduler.current_stream_idx%  s%     ##/++666ru   c                6    | j                   x}t        |      S y)z9CUDA Stream name that current scheduler node assigned to.N)r
  r%   )rr   r  s     rs   current_stream_namezScheduler.current_stream_name%  s#     111J>":..ru   c                    t        |t              rJ | j                  |   }t        j                  j
                  j                  |      | _        y)z6Code-gen to enter the Stream context assigned to node.)r  N)r   r  rv  ra   r   r  codegen_cuda_stream_enterr  )rr   r   node_streams      rs   generate_stream_ctx_enterz#Scheduler.generate_stream_ctx_enter%  sI    d$:;;;))$/#$77#7#7#Q#Q" $R $
 ru   c                ~    | j                   J t        j                  j                  j	                          d| _         y)z1Code-gen to exit from the current Stream context.N)r  ra   r   r  codegen_cuda_stream_exitrq   s    rs   r  z"Scheduler.generate_stream_ctx_exit%  s2    ''333	557#' ru   c                &   || j                   v sJ t        |t              rdn| j                   |   }| j                  |k(  ry| j                  |y| j                  || j	                  |       y| j                          | j	                  |       y)am  Generate stream entering and exiting to properly run node in a multi-stream scenario.

        Stream context switching is only generated if ``node``'s assigned stream is different from
        the previous node's stream. NopKernelSchedulerNodes have stream=None and inherit the
        enclosing stream context (or do nothing if no context is active yet).
        N)rv  r   r  r
  r1  r  )rr   r   rQ  s      rs   r  z'Scheduler.generate_stream_ctx_switching%  s     t***** $ 67 $$T* 	
 ""f, $$0V^ $$,1C**40 ))+**40ru   )rq  zlist[ir.Operation]r   r  )r   z!dict[str, SchedulerDonatedBuffer]r  r  )r  r,  r   r   )r  r,  r   rf   r   r}   r  )r  r  r   r  )r  r,  r   r  )r   r  r   rf   )r  
str | Noner   r}   r  )r  rf   r   r  )r   r  rq  r  r   tuple[float, str]r~   rq  r  rz  r}   ry  r  r   r,  )r~  r   r  r  r   r7  )r  ir.MultiTemplateBufferr   r}   )
r  ir.OperationBufferr  r9  r  r   r   r   r   r  )r  r  r   r}   )rq  r  ry  r  r   z&tuple[LambdaFuture | None, ModuleType])r   rf   r   rf   r   rk   )r   rf   r   rf   )r   rf   r   rf   r\  OrderedSet[BaseSchedulerNode]r   rf   )r   rf   r   rf   r  r   r\  r;  )r  ,dict[BaseSchedulerNode, list[PendingFusion]]r\  r;  r   r  )
r.  1list[tuple[BaseSchedulerNode, BaseSchedulerNode]]r,  &dict[BaseSchedulerNode, PendingFusion]r/  r<  r\  r;  rj  r}   )r\  r;  r,  r>  )r8  r=  r9  r=  )rq  r  rj  r}   r   r  )rq  r  r   r   rE  r   r   z!Iterator[list[BaseSchedulerNode]])rX  r  r   r  )r   r   )r{  r  rd  r   r[  r}   r   z-tuple[ForeachKernelSchedulerNode | None, int])r  r  rP  r   rd  r   r[  r}   rS  zJCallable[[ForeachKernelSchedulerNode, list[BaseSchedulerNode], int], None]r   r  r  )rq  r  rj  r}   r   r=  r+  )r   rf   r   rf   r  r   r   r}   )r   rf   r   rf   r  z!tuple[str, ...] | OrderedSet[str]r   r,  r/  r*  )rF	  rf   r  rf   rS  r  r   r}   )r   rf   r   rf   r   z,tuple[int, SchedulerNode, sympy.Expr] | None)rE  rf   rJ  rf   r   r   )r   rf   r   rf   r   OrderedSet[str] | None)
r7  rf   r  rf   ri  zCSequence[tuple[BaseSchedulerNode, NestedReduction.PointwiseDomain]]r%  r}   r   r}   )FT)
r   rf   r   rf   r%  r}   r  r}   r   r}   )FTN)r   rf   r   rf   r%  r}   r  r}   ro	  r?  r   r}   )r   rf   r   rf   ro	  r?  r   r}   )r	  r;   r   rf   r   rf   r   r}   )r  r9   r  r9   r   r}   )r  r8   r  r9   r	  r}   r   r}   )r  r8   r  r:   r	  r*  r   r}   )r	  r8   r	  r8   r   r}   r  )r   r8   r  r}   r   r   )TN)
rE  rf   rJ  rf   r  r}   ro	  r?  r   r   )
r   rf   r   rf   r  r}   ro	  r?  r   r   )...)r   rf   r   rf   r  r}   r	  zLiteral[False]r  r}   r   r   )r   rf   r   rf   r  r}   r	  zLiteral[True]r  r}   r   ztuple[int, int, bool])TFT)r   rf   r   rf   r  r}   r	  r}   r  r}   r   zint | tuple[int, int, bool])r8  r=  r   r=  )rq  r   r   r   )r*
  rf   r   r  )r  r  r   BaseScheduling)r  r  r   r@  r>  )r   r,  r  r   r   r}   )r   rf   r   r5  )r   ro
  )r   6dict[str, ir.IRNode | ir.TorchBindObject | sympy.Expr])r[
  list[GraphPartitionSignature]r   r  )rq
  rg   rX
  rA  r   ro
  )r|
  list[PartitionType]r}
  z
list[bool]r   rB  )r_
  rA   r  r   r   rA   )r   z9tuple[list[PartitionType], list[GraphPartitionSignature]])r|
  rC  r[
  rB  r   r  )rq
  rg   r_
  rA   r   r  )r|
  rC  r[
  rB  r   z'contextlib.AbstractContextManager[None]r  r  r   z%tuple[float, float, list[str | None]])rq  r  r   r}   )r  r,  r   z	ir.Layoutr  )r   r5  )r   r   r   r   r:  rN  rw  r  r  r  r  propertyr  setterr  r  rr  r  rd   r  r|  rC  r{  r  rV  rt  r}  r~  r\  r  r  rr  rw  r}  r  r  r  r  r  r  r  r  r  r&  r2  r6  r>  rk  r0  rJ  r  rY  r  r]  r  rA  r  r  r  r  r  r	  r	  r	  r	  r  rJ	  r^	  rf	  rh	  r&  r  rs	  rp	  r	  r  r	  r	  r	  r	  r	  r2	  r	  r	  r	  r   r  r	  r	  r  r  r  r#
  r%
  r+
  r4
  r#  r9
  r  rA
  rO   rM
  rC
  rT
  rb
  rt
  r
  r
  r
  r  r  r  r
  r
  r  r
  r
  r
  r
  r
  r  rZ  r(  ru  r
  r-  r1  r  r  r  r  s   @rs   r  r    s   
h9T	#6
p(S
Q & & ( (7#,"HMP^KZ+#Z ,	 6	S*4#&/<$6!F	808	8, %)	*  "	
 
$
> 
>*6
>	
>'0'	'RIDV'8&'8 +'8 	'8
 '8 
'8R
 OS0AK	/ Z&Z/@Z	Zx>  ! 3	
 
0  ! '	
 3"D2$PD2 3D2 
	D2LO?PO? @O?  L	O?
 3O? O?bS2S @S0$K$ $U$(C&C C 
!	CJ &1  
+	 8e.N(
TC&,C& *C& 	C&
 
7C&J'$*'$ '$ *	'$ '$
'$ 
'$R?6 &6  6  
;	6 p,&,/@,	,\7&7/@7	7r.2&.2/@.2MP.2	.2`$&$/@$	$6< < !< <	<
 
<|M&M/@M	M^76&76/@76	76r]K ]K !]K 
	]K~^ ^ !^ 
	^@ ;(; ); 	;
 
;z`&`/@`	5`D
)
5F
	
J J !J 
 	J"!
'!
 !!

	!
 !
 
!
N "*.  ! 	
 $( 
4 "*.
 
 !
 	

 $(
 

, "*.=AuM uM !uM 	uM
 $(uM %;uM 
uMx	 >BG G !G
 %;G 
GR3)3)(93)BS3)	3)j 
 
 OT== )=GK=	=~==&/=	=* OI OIh '7J	 4 4"; !=A(#( $( 	(
 %;( 
(\ +/=A  ! $(	
 %; 
* 
  8;*-  ! 	
 (6 $( 
  
  7:*-$ $ !$ 	$
 (5$ $($ 
$ $ !.3*.^K ^K !^K 	^K
 (,^K $(^K 
%^K@d d !d 
	dLR
&R
/@R
	R
h6 Q6	:6@4@4	4	8*4
) 
&'*%5$

+:
	
=~ ! !F%	"	? '1' 
'RBH BH LBH 
"	BHHK -K @JK 	&K Z"
*"
 )"
 
!	"
H?&? 
!?B& 
!B%,%	 %@5&	B5&n";'"; 2"; 
	";H>01
 1
 +1
 
	1
f-;X	06-A--A;X-A	-A^@rhQ4Q	.QN5`%
H    
(1ru   c                  D    e Zd Zd f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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'dZ xZS )(r@  c                0    t         |           || _        y r~   )r  r:  r  )rr   r  r  s     rs   r:  zBaseScheduling.__init__%  s    "ru   c                R    | j                   r| j                   j                          y y r~   )r  r#
  rq   s    rs   free_buffers_in_schedulerz(BaseScheduling.free_buffers_in_scheduler%  s    >>NN'') ru   c                    t               S )z0Return a set of .codegen.common.BackendFeature()r   r  s     rs   get_backend_featuresz#BaseScheduling.get_backend_features&  s
    |ru   c                    t         )zO
        Check whether node1 and node2 can be vertically fused or not.
        r  r  s      rs   r	  z BaseScheduling.can_fuse_vertical&  
     "!ru   c                    t         )zQ
        Check whether node1 and node2 can be horizontally fused or not.
        r  r  s      rs   r	  z"BaseScheduling.can_fuse_horizontal&  rN  ru   c                   |j                         }t        |t        j                        sy|j	                         syt        |j
                  t        j                        rt        |j
                  j                        dk(  xrk t        |j
                  j                  d   t        j                        xr8 |j
                  j                  d   j                         |j                         k(  S y)av  
        A Multi-Output Template (referenced in #144012) is a template node
        with MultiOutputLayout, and its output buffers are instances of MultiOutput.
        In this context, we verify whether node1 represents the Multi-Output Template
        and node2 corresponds to one of its outputs. If so, we further check if
        backend supports this fusion.

        Fr   r   )r  r   r0   r!  r\   r   rB   r   r  rl  r  )rr   r   r   r  s       rs   r	  z.BaseScheduling.can_fuse_multi_outputs_template&  s     ..0,(9(9:557ejj"..1EJJ%%&!+ Ouzz003RYY?OJJ%%a(113|7L7L7NN ru   c                   |j                         s|j                         rt        j                  ||      S t        j	                  ||      r"t        j                  ||      rt        ||      S t        j                  ||      rt        ||      S t        |t              r|j                  |      S t        |t              r|j                  |      S t        |t              rLt        |t              r<t        |j                  t        j                         sJ t"        j%                  ||      S t&        j                  ||      S )z 
        Fuse two nodes
        )r  rB  ry   r2  r8  r  r  r   r  r  r   r  rR  r   r   r0   r  r*  r4  r   r  s      rs   ry   zBaseScheduling.fuse1&  s    !1!1!3-225%@@995
&&ue4(6677uE*5%8845??5))67??5))89j=?
 ejj"*D*DEEE7EEeUSS%**5%88ru   c                    t         )z[
        Process the iteration sizes in case a transformation needs to be applied.
        r  )rr   rC  s     rs   r$  zBaseScheduling.group_fnK&  rN  ru   c                    t         )z
        Given a template node, generate a kernel.

        This function is only available for triton now. If the third-party backend behaves as a sub-class
        of TritonScheduling, it can override it or reuse it.
        r  )rr   r  epilogue_nodesr	  s       rs   r  zBaseScheduling.codegen_templateS&  s
     "!ru   c                    t         rv  r  )rr   rq  rz  ry  s       rs   rw  z.BaseScheduling.generate_kernel_code_from_nodesa&  s
     "!ru   c                    t         rV  r  r  s     rs   r  zBaseScheduling.codegen_nodel&  
     "!ru   c                    t         r~   r  r  s     rs   r  z*BaseScheduling.codegen_mix_order_reductionr&  r  ru   c                    t         r~   r  r  s     rs   r  z'BaseScheduling.codegen_nested_reductionu&  r  ru   c                    t         )zt
        Generate synchronization code for the kernel. This method depends on the hardware characteristics.
        r  rq   s    rs   r  zBaseScheduling.codegen_syncx&  rX  ru   c                     y)z
        Check whether the backend is requesting the scheduler to flush the generated kernel.
        If not supported, please return False.
        Fr   rq   s    rs   r  zBaseScheduling.ready_to_flush~&  s    
 ru   c                    t         )z]
        Flush the generated kernel and python wrapper code to the source code file.
        r  rq   s    rs   r%
  zBaseScheduling.flush&  rX  ru   c                    t         )rq  r  rO  s     rs   rr  z$BaseScheduling.benchmark_fused_nodes&  
     "!ru   c                    t         )r|  r  )rr   r~  s     rs   r}  z)BaseScheduling.benchmark_codegened_module&  s
    
 "!ru   c                     y)z
        Return an unsigned integer which represents the priority of this fusion pair.
        The smaller is with higher priority.
        r   r   r  s      rs   r
  z'BaseScheduling.get_fusion_pair_priority&  s     ru   c                    t         )z
        Benchmark the list of nodes to combine and return the execution time
        and memory copy time in milliseconds on randomly generated inputs.
        r  )rr   r  r  s      rs   r  z%BaseScheduling.benchmark_combo_kernel&  r_  ru   c                |    |r:ddl m}  |||      }t        j                  j                  j                  ||       y y )Nr   )'set_kernel_post_grad_provenance_tracing)r  rd  ra   r   r  write_provenance_debug_handle)rr   r  r  rd  debug_handles        rs   codegen_commentzBaseScheduling.codegen_comment&  s>    
 UBL GG  >>\ ru   )r  zScheduler | Noner  )r  r  r   zOrderedSet[BackendFeature]r+  r  )rC  r  r   z"tuple[tuple[sympy.Expr, ...], ...])r  rf   rT  r  r	  r  r   r5  r~   r8  )r   z"FusedSchedulerNode | SchedulerNoder   r  )r   r  r   r  )r   r  r   r  r  r6  )r~  r   r   r7  r/  rD  )r  r  r  r5  r   r  )r   r   r   r:  rJ  rL  r	  r	  r	  ry   r$  r  rw  r  r  r  r  r  r%
  rr  r}  r
  r  rg  r  r  s   @rs   r@  r@  %  s   #*"&"/@"	""&"/@"	"&/@	49&9/@9	94"3"	+""(" 4" 4	"
 
"$ %)		"*	" 	" "		"
 
	"""""""0"	""&/@	"4"	." #'2   
	ru   r@  r+  )r   z$torch._inductor.codecache.LocalCache)r  rf   r   r,  )r  rf   r   zCallable[[Any], Any] | None)r  rf   r   r/  )r  r   r   r,  )r   rf   r  r  r  r-  r   r  )r  )FusedSchedulerNode | GroupedSchedulerNoder   r  )r  rh  r  r  r   r  r   r  )r   )r  zlist[list[int]]rC  r  r  r  r   r   )r  r9  r]  r:  r   r  r  )r  r   r  r   r  r   r  r   r  rH   r   ztuple[int, int])r  rq  r  rq  r  r   r  r   r  r   r  r   r  rH   r   r}   )r   z	ir.IRNoder   ro
  )r   rf   r   ro
  )r  rf   r   r}   )r   rf   r  r}   r   r}   )r  rf   r  rf   r   r}   )r   rf   r   rf   )
__future__r   rW  r
  r   rC  r  r1
  r  r  r  r  r  r  ry  r  r[  r
  r  r   r   concurrent.futuresr   r   r   r	   r
   r   r   r   r   r   typing_extensionsr   torch.utils._ordered_setr   r0   r   r   collections.abcr   r   r   typesr   torch._inductor.codegen.wrapperr   r  r   r  r   r  r`  torch._inductor.async_compiletorch.utils._pytreer  _pytreer  torch._dynamo.utilsr   r    torch._inductor.autotune_processr   torch._inductor.codecacher    r!   torch._inductor.irr"   torch._inductor.metricsr#   r$   torch._inductor.stream_utilsr%   %torch.fx.experimental.symbolic_shapesr&   torch.utils._sympy.functionsr'   torch.utils._sympy.symbolr(   r)   r*   torch.utils._tritonr+   rX  r,   r-   r.   r/   r1   analyze_preserves_zero_maskr2   codegen.commonr3   r4   r5   comm_analysisr6   r7   r8   r9   r:   r;   excr<   r=   fx_utilsr>   r?   r@   rA   rB   rC   rD   r$  rE   r  rF   rG   runtime.hintsrH   rI   runtime.runtime_utilsrJ   rK   rL   r   rM   rN   rO   rP   rQ   rR   rS   rT   rU   rV   rW   rX   rY   rZ   r[   r\   r]   r^   r_   r`   virtualizedra   	getLoggerr   rT  _logginggetArtifactLoggerr  ra  r  r
  r   rg   r   rh   ri   r  rk   r   r   r   r   r   r2  r  r   rf   ry  ro  rn  r  ru  r  r  r  r  rR  r  r   r  r  r   r  r  r*  rB  r  r  r  r  r  r  r1  rn  rm  r  r  r  r  r   r$  r&  r)  r+  r-  r/  r@  r  r@  r   ru   rs   <module>r     s   "           	  
     , 3	 	 	 ( / ) << J51   $ $ $ 6 E ? 7 M 8 > 1 O O * D D D M M ; : 2 $    J : F F &     *  g!^^--hA
NN44XO  >>;;$  11(LI 34y 4T]t_ D D D* ( ( (
* 
 T"N N #N*			 
	_ _D	m m` h8 h8 h8V 1_ 1 1T1 T1n 2 2, ,#L T"
 
 #
*  *K
*K4*K ,*K 
	*KZ"* 1 "*J5. 5A*% A*H:	$: $ 
	8k** k*\hG0 hGV\<. \<~<-+= <-~L:!3 L:^b, bP #%+#++  + 	+\0%01C0	08
1((( ( 	(
 #( (6		  	
   # 
@ 
 
 
> +9??, $
&" 6:
.2	4:$:5F:	:"P@ /2 /2 /2d = = =@{X1 {X1|qA Aru   