Ë
    g^(hÈ4  ã                  óö   — d dl mZ d dlZd dlmZmZ d dlZddlmZ ddlm	Z	 ddl
mZmZ ddlmZmZ dd	lmZmZmZ dd
lmZ erd dlZd dlmZ ddlmZ ddlmZ  G d„ dej:                  «      Z G d„ d«      Zy)é    )ÚannotationsN)ÚAnyÚTYPE_CHECKINGé   )Úconfig)Ú
write_text)Úget_metric_tableÚis_metric_table_enabled)ÚDevicePropertiesÚReductionHint)ÚBaseSchedulerNodeÚ	SchedulerÚ	WhyNoFuse)ÚV)Ú
OrderedSet)ÚSIMDKernelFeatures)ÚTritonKernelc                  ó   — e Zd ZdZdd„Zy)ÚSortablez>Anything that can be used as a list.sort() key (int/tuple/etc)c                 ó   — y ©N© )ÚselfÚothers     úU/var/www/skyplay_api_hub/venv/lib/python3.12/site-packages/torch/_inductor/choices.pyÚ__lt__zSortable.__lt__   s   � ó    N)r   ztyping.SelfÚreturnÚbool)Ú__name__Ú
__module__Ú__qualname__Ú__doc__r   r   r   r   r   r      s   „ ÙHä5r   r   c                  ó(  — e Zd Z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
e	 	 	 	 	 	 	 	 	 	 dd„«       Ze	 	 	 	 	 	 	 	 	 	 dd	„«       Ze	 	 	 	 	 	 	 	 dd
„«       Zy)ÚInductorChoicesax  
    This class contains a collection of default heuristics that effect performance of our generated
    code.  We try to not put correctness requirements in this file.

    You can override the choices made here by doing:

            class MyHeuristics(InductorChoices):
                ...

            torch._inductor.virtualized.V.set_choices_handler(MyHeuristics())
    c                ó   — |S )zTHook to change the kwargs passed to TritonKernel, used to apply fixed configurationsr   )r   Ú
kernel_clsÚfeaturesÚgroupsÚkernel_kwargss        r   Útriton_kernel_kwargsz$InductorChoices.triton_kernel_kwargs+   s
   € ð Ðr   c                ó¾  — t         j                  j                  ryt         j                  j                  r+t        j
                  j                  «       j                  dk(  ryt        j
                  j                  j                  | j                  d¬«      }|dk  rd|z  }n	|dk  rd	}nyt        j
                  j                  j                  | j                  |«      S )
z>Heuristic to decide if a cooperative reduction should be used.TÚcpuFé   )Úfallbacké   i €  é   i    )r   ÚtritonÚforce_cooperative_reductionsÚcooperative_reductionsr   ÚgraphÚget_current_device_or_throwÚtypeÚsizevarsÚ	size_hintÚnumelÚstatically_known_geqÚreduction_numel)r(   ÚxhintÚ	thresholds      r   Ú should_use_cooperative_reductionz0InductorChoices.should_use_cooperative_reduction5   s©   € ô �=‰=×5Ò5Øä—‘×4Ò4Ü�w‰w×2Ñ2Ó4×9Ñ9¸UÒBàä—‘× Ñ ×*Ñ*¨8¯>©>ÀAÐ*ÓFˆØ�AŠ:Ø ™‰IØ�bŠ[Ø‰Iàä�w‰w×Ñ×4Ñ4Ø×$Ñ$ ió
ð 	
r   c                óè  — t         j                  j                  syt        j                  dij                  | j                  «       d«      }|rD	 |dt        t        j                  j                  j                  | j                  «      d«      z  z  }t         j                  j                  r|dz  }t        j                  j                  j                  | j                   |«      S # t        $ r Y Œ^w xY w)zO
        Heuristic to decide if a persistent reduction should be used.
        Fi   é@   é    r1   )r   r2   Úpersistent_reductionsr   ÚINNERÚgetÚget_reduction_hintÚminr   r5   r8   r9   r:   Ú
ValueErrorÚmulti_kernelÚstatically_known_leqr<   )r(   Úcooperative_reductionr>   s      r   Úshould_use_persistent_reductionz/InductorChoices.should_use_persistent_reductionL   sÌ   € ô �}‰}×2Ò2Øä×Ñ ð
ç
‰#ˆh×)Ñ)Ó+¨RÓ
0ð 	ñ !ðØ˜R¤3¤q§w¡w×'7Ñ'7×'AÑ'AÀ(Ç.Á.Ó'QÐSUÓ#VÑVÑV�	ô �=‰=×%Ò%Ø˜‰OˆIÜ�w‰w×Ñ×4Ñ4Ø×$Ñ$ ió
ð 	
øô ò Ùðús   ÁAC% Ã%	C1Ã0C1c                ó°   — | j                  «       t        j                  k(  xr4 t        j                  j
                  j                  | j                  d«      S )a  
        Heuristic to decide if we should drop the X dimension from a persistent reduction kernel.
        So the [XBLOCK, RBLOCK] block becomes a [RBLOCK] block and XBLOCK is forced to be always 1.
        Strangely this is faster than a [1, RBLOCK] block in some cases.
        é   )rF   r   rD   r   r5   r8   r;   r<   )r(   s    r   Úwant_no_x_dimzInductorChoices.want_no_x_dimj   sG   € ð ×'Ñ'Ó)¬]×-@Ñ-@Ñ@ò UÜ—‘× Ñ ×5Ñ5°h×6NÑ6NÐPSÓTð	
r   c                óì  ‡‡— t        j                  | «      }|j                  }d}dŠd}||z  |z  }‰|z  |z  }	d}
d|
z  }|rÛ|d|z  k\  ry|dk  ry||z  |k  r|}n°||z  |	k  rm||z  d|z  z  }||z   dz
  |z  }|||z  z   dz
  ||z  z  Št        j                  |«      }t        |ˆfd„¬	«      }t        |‰z
  «      d
k  rt        ||«      }n>‰}n;t        j                  |«      }t        |ˆfd„¬	«      }t        |‰z
  «      dk  r|}n‰}|||z  z   dz
  ||z  z  S d}d}||z   dz
  |z  }||z  |k  r|}n­||z  |	k  rj||z  |z  }||z   dz
  |z  }|||z  z   dz
  ||z  z  Št        j                  |«      }t        |ˆfd„¬	«      }t        ‰|z
  «      dk  rt        ||«      }n>‰}n;t        j                  |«      }t        |ˆfd„¬	«      }t        |‰z
  «      dk  r|}n‰}|||z  z   dz
  ||z  z  S )zÄHeuristic to decide the RSPLIT used for split reductions.
        When a reduction has a small number of outputs there is not enough parallelism,
        so we will do the reduction in two phases.rB   i   i   r0   r.   r   i    c                ó    •— t        | ‰z
  «      S r   ©Úabs©ÚxÚtmp_split_sizes    €r   ú<lambda>z8InductorChoices.reduction_split_factor.<locals>.<lambda>š   ó   ø€ ´c¸!¸nÑ:LÓ6M€ r   )Úkeyé   c                ó    •— t        | ‰z
  «      S r   rR   ©rU   Úmax_elements_per_threads    €r   rW   z8InductorChoices.reduction_split_factor.<locals>.<lambda>¢   ó   ø€ ´c¸!Ð>UÑ:UÓ6V€ r   é2   é   é€   c                ó    •— t        | ‰z
  «      S r   rR   rT   s    €r   rW   z8InductorChoices.reduction_split_factor.<locals>.<lambda>º   rX   r   é   c                ó    •— t        | ‰z
  «      S r   rR   r\   s    €r   rW   z8InductorChoices.reduction_split_factor.<locals>.<lambda>Á   r^   r   )r   ÚcreateÚmulti_processor_countÚsympyÚdivisorsrG   rS   Úmax)ÚdeviceÚreduction_numel_hintÚ
numel_hintÚinner_reductionÚpropsÚnum_smÚmin_elements_per_threadÚthreads_per_smÚmin_elements_per_deviceÚmax_elements_per_deviceÚ	num_warpsÚnum_threadsÚ
split_sizeÚtarget_blocksÚblocks_per_outputrh   ÚclosestÚrvals_per_threadÚxvals_per_blockÚxblocksr]   rV   s                       @@r   Úreduction_split_factorz&InductorChoices.reduction_split_factorv   s°  ù€ ô !×'Ñ'¨Ó/ˆØ×,Ñ,ˆØ"$ÐØ"%ÐØˆØ"9¸FÑ"BÀ^Ñ"SÐØ"9¸FÑ"BÀ^Ñ"SÐØˆ	Ø˜9‘nˆáð ˜Q ™ZÒ'ØØ# tÒ+ØØ# jÑ0Ð4KÒKØ4‘
Ø%¨
Ñ2Ð5LÒLØ &¨Ñ 7¸AÀ¹OÑ L�Ø%2°ZÑ%?À!Ñ%CÈ
Ñ$RÐ!à(¨;Ð9JÑ+JÑJÈQÑNØ!Ð$5Ñ5ñ"7�ô !Ÿ>™>Ð*>Ó?�Ü˜hÓ,MÔN�Ü�w Ñ/Ó0°2Ò5ä!$ WÐ.EÓ!F‘Jà!/‘Jä Ÿ>™>Ð*>Ó?�Ü˜hÓ,VÔW�Ü�wÐ!8Ñ8Ó9¸BÒ>à!(‘Jà!8�JØ(¨:¸Ñ+CÑCÀaÑGØ˜[Ñ(ñð ð  !ÐØ!ˆOØ! OÑ3°aÑ7¸OÑKˆGØ# jÑ0Ð3JÒJØ4‘
Ø%¨
Ñ2Ð5LÒLØ &¨Ñ 7¸KÑ H�Ø!.°Ñ!8¸1Ñ!<ÀÑ H�à(Ð+;¸mÑ+KÑKÈaÑOØ&¨Ñ6ñ"8�ô !Ÿ>™>Ð*>Ó?�Ü˜hÓ,MÔN�Ü�~¨Ñ/Ó0°2Ò5Ü!$ WÐ.EÓ!F‘Jà!/‘Jä Ÿ>™>Ð*>Ó?�Ü˜hÓ,VÔW�Ü�wÐ!8Ñ8Ó9¸BÒ>à!(‘Jà!8�Jà(Ð+;¸jÑ+HÑHÈ1ÑLØ  :Ñ-ñð r   c                óÖ  ‡ ‡‡‡— |dk(  rÄt         j                  r ‰j                  «       s‰j                  «       r”t        d«      rv‰j                  j                  «       ‰j                  j                  «       z  Št        ‰«      dkD  r3t        d«      j                  ˆˆˆˆ fd„«        t        ‰‰«      d«       y t        ‰‰«      d«       y‰j                  «       sg‰j                  «       sWt        ‰j                  «       «      t        ‰j                  «       «      z   t         j                  kD  r t        ‰‰«      d«       y‰ j                  ‰‰«      r t        ‰‰«      d«       yy	)
aï  
        Heuristics to prevent fusion applied to both horizontal and vertical fusions.  Heuristics here should not
        be needed for correctness and tweaking them may yield additional performance.

        See also some related heuristics that can be changed via config:
            - config.triton.tiling_prevents_pointwise_fusion
            - config.triton.tiling_prevents_reduction_fusion
            - config.aggressive_fusion (will cause this function to be called more times)
        r   Ú'fusion_failure_due_to_indexing_mismatchc                 óB  •— t         j                  j                  t         j                  j                  ‰j	                  «       ‰j	                  «       t        ‰j                  «       «      t        ‰j                  «       «      t        ‰ «      ‰j                  ‰‰‰ «      dœS )N)Úpre_grad_graph_idÚpost_grad_graph_idÚ
node1_nameÚ
node2_nameÚnode1_debug_strÚnode2_debug_strÚcommon_buffer_namesÚfailure_reason)	r   r5   Úgraph_idr‚   Úget_namer   Ú	debug_strÚlistÚdecide_fusion_fail_reason)Úcommon_buf_namesÚnode1Únode2Ú	schedulers   €€€€r   rW   z*InductorChoices.can_fuse.<locals>.<lambda>å   su   ø€ Ü12·±×1AÑ1AÜ23·'±'×2LÑ2LØ*/¯.©.Ó*:Ø*/¯.©.Ó*:Ü/9¸%¿/¹/Ó:KÓ/LÜ/9¸%¿/¹/Ó:KÓ/LÜ37Ð8HÓ3IØ.7×.QÑ.QØ % uÐ.>ó/ñ!€ r   z'no shared data due to indexing mismatchFzno shared datazexceeds max fusionz Fusion will increase peak memoryT)r   Úaggressive_fusionÚis_reductionr
   Úread_writesÚbuffer_namesÚlenr	   Úadd_rowr   Ú
is_foreachÚ	get_nodesÚmax_fusion_sizeÚcan_fusion_increase_peak_memory)r‘   r�   r�   Úshared_data_scorerŽ   s   ``` @r   Úcan_fusezInductorChoices.can_fuseÌ   s7  û€ ð   Ò!Ü×(Ò(¨E×,>Ñ,>Ô,@ÀE×DVÑDVÔDXä&Ð'PÔQà×%Ñ%×2Ñ2Ó4°u×7HÑ7H×7UÑ7UÓ7WÑWð !ô Ð'Ó(¨1Ò,Ü$Ð%NÓO×WÑWöôð ,”I˜e UÓ+Ð,UÔVØ Ø#ŒI�e˜UÓ#Ð$4Ô5Øð × Ñ Ô"Ø×$Ñ$Ô&Ü�E—O‘OÓ%Ó&¬¨U¯_©_Ó->Ó)?Ñ?Ä&×BXÑBXÒXà#ŒI�e˜UÓ#Ð$8Ô9Øà×4Ñ4°U¸EÔBØ#ŒI�e˜UÓ#Ð$FÔGØàr   c                 ó   — y)zCHook for heuristics to prevent vertical (producer/consumer) fusionsTr   ©r‘   r�   r�   rœ   s       r   Úcan_fuse_verticalz!InductorChoices.can_fuse_vertical  s   € ð r   c                óš   — |t         j                  k  r t        ||«      d«       y| j                  ||«      r t        ||«      d«       yy)zEHook for heuristics to prevent horizontal (consumer/consumer) fusionsÚscore_fusion_memory_thresholdFz=Nodes are too far away. Fusing them may increase peak memory.T)r   r¢   r   Úare_long_distant_nodesrŸ   s       r   Úcan_fuse_horizontalz#InductorChoices.can_fuse_horizontal  sS   € ð œv×CÑCÒCØ#ŒI�e˜UÓ#Ð$CÔDØØ×+Ñ+¨E°5Ô9Ø#ŒI�e˜UÓ#ØOôð Ør   c                ó”  — | j                  ||«      }t        t        |j                  |j                  z
  «      t        |j                  |j                  z
  «      «       }|j                  «       rd}n+d|j                  «       t        j                  k(  xr |dkD  z   }||j                  «       |j                  «       k(  xr |dkD  ||fS )a—  
        Assign a score (higher comes first) to the fusion of node1 and node2.
        When different fusions conflict with each other, this is the way we
        decide what order to run them in.

        Our current score is based on:
        - The type of fusion (template/reduction/etc)
        - Estimate of the saved memory operations
        - Fusions closer together in original graph order
        r   r   )	Úscore_fusion_memoryri   rS   Ú	min_orderÚ	max_orderÚis_templater   Úepilogue_fusion_firstr“   )r‘   r�   r�   Úmemory_scoreÚproximity_scoreÚtemplate_scores         r   Úscore_fusionzInductorChoices.score_fusion"  sÉ   € ð  !×4Ñ4°U¸EÓBˆÜÜ�—‘ %§/¡/Ñ1Ó2Ü�—‘ %§/¡/Ñ1Ó2ó
ð 
ˆð ×ÑÔØ‰NàØ×"Ñ"Ó$¬×(DÑ(DÑDò %Ø  1Ñ$ñˆNð Ø×ÑÓ  E×$6Ñ$6Ó$8Ñ8ÒM¸\ÈAÑ=MØØð	
ð 	
r   N)
r'   ztype[TritonKernel]r(   r   r)   zlist[sympy.Expr]r*   údict[str, Any]r   r¯   )r(   r   r   r   )r(   r   rK   r   r   r   )
rj   ztorch.devicerk   Úintrl   r°   rm   r   r   r°   )
r‘   r   r�   r   r�   r   rœ   r°   r   r   )r‘   r   r�   r   r�   r   r   r   )r    r!   r"   r#   r+   Ústaticmethodr?   rL   rO   r}   r�   r    r¤   r®   r   r   r   r%   r%      sÙ  „ ñ
ðà&ðð %ðð !ð	ð
 &ðð 
óð ò
ó ð
ð, ð
Ø$ð
Ø=Að
à	ò
ó ð
ð: ò	
ó ð	
ð ðSØðSà!ðSð ðSð ð	Sð
 
òSó ðSðj ð7Øð7à ð7ð !ð7ð ð	7ð
 
ò7ó ð7ðr ðØðà ðð !ðð ð	ð
 
òó ðð ðØðà ðð !ðð ð	ð
 
òó ðð" ð#
Øð#
à ð#
ð !ð#
ð 
ò	#
ó ñ#
r   r%   ) Ú
__future__r   Útypingr   r   rg   Ú r   Ú	codecacher   Úmetricsr	   r
   Úruntime.hintsr   r   r‘   r   r   r   Úvirtualizedr   ÚtorchÚtorch.utils._ordered_setr   Úcodegen.simd_kernel_featuresr   Úcodegen.tritonr   ÚProtocolr   r%   r   r   r   ú<module>r¾      sS   ðÝ "ã ß %ã å Ý !ß >ß :ß >Ñ >Ý ñ ÛÝ3å@Ý,ô6ˆv�‰ô 6÷h
ò h
r   