
    ^j"                       d dl mZ d dlZd dlZd dlZd dlmZmZ d dlZd dl	m
Z
 ej                  j                  rdnddddd	d
ZdZdZdZ	 d'	 	 	 	 	 d(dZd)dZ G d de      Z G d de      Z e
       rjd dlZd dlZd dlZ eej2                  j4                  d      rd dlmZ 	 	 	 d*dZnL eej4                  j4                  d      rd dlmZ 	 	 	 d*dZn!	 	 	 d*dZn ej:                  dg dg d      Z G d de      Z G d de      Z G d  d!ej@                        Z!d+d"Z" G d# d$ej@                        Z# G d% d&ej@                        Z$y),    )annotationsN)autoEnum)has_triton_packagei    i      i   i   )XYZR0_R1_@   i      c                    | j                  dd      | j                  dd      z  | j                  dd      z  || j                  dd      z  S |z  S )NXBLOCK   YBLOCKZBLOCKR0_BLOCK)get)kwargsr0_blocks     h/var/www/ramen.bs-engineer-server.com/venv/lib/python3.12/site-packages/torch/_inductor/runtime/hints.pynative_matmul_block_numelr      si     	

8Q
**Xq
!	"
**Xq
!	" )1(86::j!$	H ?G	H    c                "    t        | t              S N)maxTRITON_DOT_MIN_BLOCK)r   s    r   native_matmul_persistent_rblockr   &   s    x-..r   c                      e Zd ZdZdZdZdZy)ReductionHintr   r         N)__name__
__module____qualname__INNEROUTER
OUTER_TINYDEFAULT r   r   r!   r!   *   s    EEJGr   r!   c                      e Zd ZdZdZy)TileHintr   r   N)r$   r%   r&   SQUAREr*   r+   r   r   r-   r-   1   s    FGr   r-   AttrsDescriptorr/   c                    | |d}t        j                  |t         j                  d      }|j                  d   dk(  sJ |j                  d   dk(  sJ |S )N)tt.divisibilitytt.equal_to)arg_propertiesclsr2   r   r3   r   )r/   	from_dictr$   property_values)divisible_by_16
equal_to_1pointer_range_32r   ress        r   AttrsDescriptorWrapperr<   @   sk     $3)F "++#)/2J2JKC &&'89R???&&}5:::Jr   c                "    | |d}t        di |S )N)r8   r9   r+   r0   )r8   r9   r:   r   s       r   r<   r<   W   s!     $3(F #,V,,r   c                    | xs dD ci c]	  }|fddgg }}|xs dD ](  }|f}||v r||   j                  ddg       !ddgg||<   * |S c c}w )Nr+   r2   r   ztt.pointer_range    )append)r8   r9   r:   xresultkeys         r   r<   r<   k   s     AP@USUW!qd/455WFW%+ =d&=3K&&(:B'?@$6#;"<F3K= M Xs   A)r8   r9   r:   )r+   r+   r+   )defaultsc                  n    e Zd Z e       Z e       Z e       Z e       Z e       Z e       Z	 e       Z
y)HeuristicTypeN)r$   r%   r&   r   PERSISTENT_REDUCTION	POINTWISE	REDUCTION
SPLIT_SCANTEMPLATEUSER_AUTOTUNEFIXEDr+   r   r   rF   rF      s4    6IIJvHFMFEr   rF   c                  (    e Zd ZdZej
                  Zy)AutotuneHintr   N)r$   r%   r&   ONE_ELEMENT_PER_THREADr   __str____repr__r+   r   r   rO   rO      s     ||Hr   rO   c                      e Zd ZU dZded<   ded<   ded<   ded<   dZd	ed
<   dZd	ed<   dZd	ed<   dZd	ed<   dZ	d	ed<   e
dd       Zeej                  dd              Zy)DevicePropertieszOCopy device properties into a data structure not requiring torch to be importedstrtypeintindexmulti_processor_countccN
int | Nonemajorregs_per_multiprocessormax_threads_per_multi_processormax_threads_per_block	warp_sizec                6    | j                   | j                   S dS )Nr?   )r`   selfs    r   warp_size_or_defaultz%DeviceProperties.warp_size_or_default   s    !%!;t~~CCr   c                   dd l }ddlm} |j                  }|j                  j
                  r|dk(  rd} ||      }|j                  |      }	 |j                  } | ||j                  ||j                  |      t        |dd       t        |d	d       t        |d
d       t        |dd      t        |d|dk7  rd      	      S d       	      S # t        $ r |dk(  r|j                  }n	|dk(  rd}n Y w xY w)Nr   )get_interface_for_devicecudahipxpumtiar   r\   r]   r^   r_   r   r`   cpur?   )	rV   rX   rY   rZ   r\   r]   r^   r_   r`   )torchtorch._dynamo.device_interfacerf   rV   versionrh   get_device_propertiesrY   AttributeErrorgpu_subslice_countrX   get_compute_capabilitygetattr)r5   devicerl   rf   device_typedevice_interfacepropsrY   s           r   createzDeviceProperties.create   s    	Kkk==!6K3F; 66v>	$)$?$?! ,,"766v>%$/$+E3Ld$S,38$- #*%1H$"Oe[u8L"W
 	
 SWW
 	
  	e#(-(@(@%&(*%	s   C $C21C2returnrW   )rz   rT   )r$   r%   r&   __doc____annotations__r\   r]   r^   r_   r`   propertyrd   classmethod	functoolscacherx   r+   r   r   rT   rT      s    Y
IJGE:*.Z.26#Z6(,:, Iz D D __ 
   
r   rT   c                @    t         j                  |       j                  S )a@  Return the wave/warp size in threads for the given device.

    Reads from torch.cuda.get_device_properties(device).warp_size via the
    cached DeviceProperties.create(). Correct on both AMD (64 for CDNA/gfx9,
    32 for RDNA/gfx10+) and NVIDIA (always 32). Falls back to 32 only when
    the field is unavailable.
    )rT   rx   rd   )rt   s    r   get_warp_sizer      s     ""6*???r   c                  z    e Zd ZU ded<   ded<   dZded<   dZ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)HalideInputSpecrU   ctypenameNzlist[str] | Noneshapestride
str | Noneoffsetalias_ofc                8    | j                   dv ry| j                   S )N)	at::Half*at::BFloat16*z	uint16_t*)r   rb   s    r   bindings_typezHalideInputSpec.bindings_type   s    ::77zzr   c                    | j                   dk(  ry| j                   dk(  ryd| j                   j                  dd       dS )	Nr   z$halide_type_t(halide_type_float, 16)r   z%halide_type_t(halide_type_bfloat, 16)zhalide_type_of<* z>())r   replacerb   s    r   halide_typezHalideInputSpec.halide_type   sA    ::$9::(: !3!3C!< =SAAr   c                    | j                   d u S r   r   rb   s    r   	is_scalarzHalideInputSpec.is_scalar   s    zzT!!r   c                    | j                   d uS r   r   rb   s    r   	is_bufferzHalideInputSpec.is_buffer   s    zz%%r   )rz   rU   rz   bool)r$   r%   r&   r|   r   r   r   r   r   r   r   r   r+   r   r   r   r      sK    J
I"E"#F#FJHj
B"&r   r   c                  \    e Zd ZU ded<   ded<   dZded<   dZded	<   dZd
ed<   ddZddZy)
HalideMetazlist[HalideInputSpec]argtypesrU   targetNr   	schedulerzdict[str, int | str] | Nonescheduler_flagsr[   cuda_devicec                    d| j                    g}| j                  r|j                  d| j                          | j                  rG| j                  sJ | j                  j	                         D ]  \  }}|j                  d| d|         |S )z-Command line args to pass to halide generatorztarget=zautoscheduler=zautoscheduler.=)r   r   r@   r   items)rc   argskvs       r   r   zHalideMeta.args   s    $++'(>>KK.(89:>>!>,,224 61nQCq456r   c                    | j                   d uS r   )r   rb   s    r   is_cudazHalideMeta.is_cuda  s    t++r   )rz   z	list[str]r   )	r$   r%   r&   r|   r   r   r   r   r   r+   r   r   r   r      s6    ##K Iz 37O07"K"	,r   r   r   )r   ztyping.Mapping[str, int]r   r[   rz   rW   )r   rW   rz   rW   )NNNry   )%
__future__r   collectionsr   typingenumr   r   rl   torch.utils._tritonr   rn   rh   TRITON_MAX_BLOCKTRITON_MAX_RSPLITTRITON_MAX_TENSOR_NUMELr   r   r   r!   r-   tritontriton.backends.compilertriton.compiler.compilerhasattrbackendscompilerr/   r<   
namedtuplerF   rO   
NamedTuplerT   r   r   r   r+   r   r   <module>r      s{   "      2 ""		   !   >B$0:/D t  ##v''):;< !!	& 
))+<	=< !!	-* !!	& 4[33=	D 4 3
v(( 3
l@&f'' &6,"" ,r   