
    ^j                       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
mZ ddlmZmZmZmZmZ ddlmZmZmZmZ dd	lmZmZmZmZmZmZ dd
lm Z  ddl!m"Z" erd dl#m$Z$ ddlm%Z%  ejL                  e'      Z(ejR                  jT                  Z*ed        Z+ed        Z,ed        Z- ede- e"d      d      Z.dZ/	  ede+de/z   dz   e/z   dz         Z0dZ1 ede,de1z   dz   e1z   d z         Z2 eejf                  d!d"e*jf                  jh                  #      Z5d$ Z6 ee6d      Z7 G d% d&e      Z8	 	 	 	 	 	 	 	 	 	 	 	 	 	 	 	 	 	 	 	 d9d'Z9d( Z:d) Z; ee*jf                        	 	 	 	 	 	 	 	 	 	 	 	 	 	 	 	 	 	 d:d*       Z3 ee*jx                        d+        Z<d, Z= ee*jf                  e=       	 	 	 	 	 	 	 	 	 	 	 	 	 	 	 	 	 	 	 	 d;d-Z>	 	 	 	 	 	 	 	 	 	 	 	 	 	 	 	 	 	 	 	 d;d.Z?d/ Z@ ee@d      ZAd0 ZB eeBd      ZCed1        ZD ed2eD e"d3            ZEed4        ZF ed5eF e"d6            ZG ee*j                  jh                        ZI ee*j                  jh                        	 	 	 	 	 	 	 	 	 	 	 	 	 	 	 	 	 	 	 	 	 	 d<d7       ZJd8 ZK ee*j                  eK       y)=    )annotationsN)TYPE_CHECKING	TypedDict)CKGroupedConvFwdTemplate   )configir)add_layout_constraintconstrain_to_fx_stridesfallback_handler	loweringsregister_lowering)autotune_select_algorithmExternKernelChoiceSymbolicGridFnTritonTemplate)is_onesis_zerospad_listlikesympy_productuse_ck_conv_templateuse_triton_template)V   )load_kernel_template)Sequence)	TensorBoxc               F     || |z  |z  |d          |||d         |d   fS NBLOCK_MBLOCK_NGROUPS )nchwmetacdivs         f/var/www/ramen.bs-engineer-server.com/venv/lib/python3.12/site-packages/torch/_inductor/kernel/conv.pyconv2d_gridr+   /   s9     	QUQYY(QY X     c               L     || |z  |z  |z  |d          |||d         |d   fS r   r#   )r$   r%   dr&   r'   r(   r)   s          r*   conv3d_gridr/   8   s=     	QUQY]DO,QY X r,   c               H     || |d          |||d          |||d         fS )Nr!   BLOCK_LBLOCK_Cr#   )r$   r%   lr(   r)   s        r*   depthwise_conv1d_gridr4   H   s8     	QY QY QY  r,   depthwise_conv1dtriton_depthwise_convT)namegridsource"cache_codegen_enabled_for_templatea  
        idx_x_h = i - PADDING_H + idx_y_h * STRIDE_H
        idx_x_w = j - PADDING_W + idx_y_w * STRIDE_W
        idx_x_c = tl.arange(0, BLOCK_K) + k

        x_ptrs = x_base + (
            (idx_x_h * stride_xh)[:, None]
            + (idx_x_w * stride_xw)[:, None]
            + (idx_x_c * stride_xc)[None, :]
        )
        mask_x = (
            (idx_n < BATCH)[:, None]
            & (idx_x_h >= 0)[:, None]
            & (idx_x_h < IN_H)[:, None]
            & (idx_x_w >= 0)[:, None]
            & (idx_x_w < IN_W)[:, None]
            & (idx_x_c < GROUP_IN_C)[None, :]
        )
        matrix_x = tl.load(x_ptrs, mask=mask_x, other=0.0)

        w_ptrs = w_base + (
            (idx_x_c * stride_wc_in)[:, None] + (i * stride_wh) + (j * stride_ww)
        )
        mask_w = (idx_x_c[:, None] < GROUP_IN_C) & (idx_y_c[None, :] < GROUP_OUT_C)
        matrix_w = tl.load(w_ptrs, mask=mask_w, other=0.0)
        acc += tl.dot(matrix_x, matrix_w, allow_tf32=ALLOW_TF32)
convolution2da  
{{def_kernel("X", "W")}}
    # Tensor dimensions
    BATCH = {{size("X", 0)}}
    IN_C = {{size("X", 1)}}
    IN_H = {{size("X", 2)}}
    IN_W = {{size("X", 3)}}
    OUT_C = {{size(None, 1)}}
    OUT_H = {{size(None, 2)}}
    OUT_W = {{size(None, 3)}}

    # Strides:
    stride_xn = {{stride("X", 0)}}
    stride_xc = {{stride("X", 1)}}
    stride_xh = {{stride("X", 2)}}
    stride_xw = {{stride("X", 3)}}
    stride_wc_out = {{stride("W", 0)}}
    stride_wc_in = {{stride("W", 1)}}
    stride_wh = {{stride("W", 2)}}
    stride_ww = {{stride("W", 3)}}

    nhw = tl.program_id(0).to(INDEX_DTYPE) * BLOCK_M + tl.arange(0, BLOCK_M)
    idx_y_w = nhw % OUT_W
    nh = nhw // OUT_W
    idx_y_h = nh % OUT_H
    idx_n = nh // OUT_H
    idx_y_c = tl.program_id(1).to(INDEX_DTYPE) * BLOCK_N + tl.arange(0, BLOCK_N)

{% if GROUPS == 1 %}
    group = 0
    GROUP_IN_C = IN_C
    GROUP_OUT_C = OUT_C
{% else %}
    group = tl.program_id(2).to(INDEX_DTYPE)
    GROUP_IN_C = IN_C // GROUPS
    GROUP_OUT_C = OUT_C // GROUPS
{% endif %}

    x_base = X + (group * stride_xc * GROUP_IN_C + idx_n * stride_xn)[:, None]
    w_base = (
        W + (group * stride_wc_out * GROUP_OUT_C + idx_y_c * stride_wc_out)[None, :]
    )

    acc = tl.zeros((BLOCK_M, BLOCK_N), dtype=tl.float32)

{% if UNROLL %}
{% for i in range(KERNEL_H) %}
{% for j in range(KERNEL_W) %}
    i = {{i}}
    j = {{j}}
    for k in range(0, GROUP_IN_C, BLOCK_K):
        a  
{% endfor %}
{% endfor %}
{% else %}
    # Could be simplified, but slightly slower:
    # for i in range(KERNEL_H):
    #     for j in range(KERNEL_W):
    #         for k in range(0, GROUP_IN_C, BLOCK_K):
    BLOCK_K_COUNT = (GROUP_IN_C + BLOCK_K - 1) // BLOCK_K
    for ijk in range(KERNEL_H * KERNEL_W * BLOCK_K_COUNT):
        k = (ijk % BLOCK_K_COUNT) * BLOCK_K
        ij = ijk // BLOCK_K_COUNT
        i = ij // KERNEL_W
        j = ij % KERNEL_W
        a  
{% endif %}

    mask = (
        (idx_n < BATCH)[:, None]
        & (idx_y_h < OUT_H)[:, None]
        & (idx_y_w < OUT_W)[:, None]
        & (idx_y_c < GROUP_OUT_C)[None, :]
    )
    idx_n = idx_n[:, None]
    idx_c = idx_y_c[None, :] + group * GROUP_OUT_C
    idx_h = idx_y_h[:, None]
    idx_w = idx_y_w[:, None]

    # inductor generates a suffix
    {{store_output(("idx_n", "idx_c", "idx_h", "idx_w"), "acc", "mask", val_shape=("BLOCK_M", "BLOCK_N"))}}
)r7   r8   r9   a  
        idx_x_d = d - PADDING_D + idx_y_d * STRIDE_D
        idx_x_h = i - PADDING_H + idx_y_h * STRIDE_H
        idx_x_w = j - PADDING_W + idx_y_w * STRIDE_W
        idx_x_c = tl.arange(0, BLOCK_K) + k

        x_ptrs = x_base + (
            (idx_x_d * stride_xd)[:, None]
            + (idx_x_h * stride_xh)[:, None]
            + (idx_x_w * stride_xw)[:, None]
            + (idx_x_c * stride_xc)[None, :]
        )
        mask_x = (
            (idx_n < BATCH)[:, None]
            & (idx_x_d >= 0)[:, None]
            & (idx_x_d < IN_D)[:, None]
            & (idx_x_h >= 0)[:, None]
            & (idx_x_h < IN_H)[:, None]
            & (idx_x_w >= 0)[:, None]
            & (idx_x_w < IN_W)[:, None]
            & (idx_x_c < GROUP_IN_C)[None, :]
        )
        matrix_x = tl.load(x_ptrs, mask=mask_x, other=0.0)

        w_ptrs = w_base + (
            (idx_x_c * stride_wc_in)[:, None] +
            (d * stride_wd) + (i * stride_wh) + (j * stride_ww)
        )
        mask_w = (idx_x_c[:, None] < GROUP_IN_C) & (idx_y_c[None, :] < GROUP_OUT_C)
        matrix_w = tl.load(w_ptrs, mask=mask_w, other=0.0)
        acc += tl.dot(matrix_x, matrix_w, allow_tf32=ALLOW_TF32)
convolution3dax  
{{def_kernel("X", "W")}}
    # Tensor dimensions
    BATCH = {{size("X", 0)}}
    IN_C = {{size("X", 1)}}
    IN_D = {{size("X", 2)}}
    IN_H = {{size("X", 3)}}
    IN_W = {{size("X", 4)}}
    OUT_C = {{size(None, 1)}}
    OUT_D = {{size(None, 2)}}
    OUT_H = {{size(None, 3)}}
    OUT_W = {{size(None, 4)}}

    # Strides:
    stride_xn = {{stride("X", 0)}}
    stride_xc = {{stride("X", 1)}}
    stride_xd = {{stride("X", 2)}}
    stride_xh = {{stride("X", 3)}}
    stride_xw = {{stride("X", 4)}}
    stride_wc_out = {{stride("W", 0)}}
    stride_wc_in = {{stride("W", 1)}}
    stride_wd = {{stride("W", 2)}}
    stride_wh = {{stride("W", 3)}}
    stride_ww = {{stride("W", 4)}}

    ndhw = tl.program_id(0).to(INDEX_DTYPE) * BLOCK_M + tl.arange(0, BLOCK_M)
    idx_y_w = ndhw % OUT_W
    ndh = ndhw // OUT_W
    idx_y_h = ndh % OUT_H
    nd = ndh // OUT_H
    idx_y_d = nd % OUT_D
    idx_n = nd // OUT_D
    idx_y_c = tl.program_id(1).to(INDEX_DTYPE) * BLOCK_N + tl.arange(0, BLOCK_N)

{% if GROUPS == 1 %}
    group = 0
    GROUP_IN_C = IN_C
    GROUP_OUT_C = OUT_C
{% else %}
    group = tl.program_id(2).to(INDEX_DTYPE)
    GROUP_IN_C = IN_C // GROUPS
    GROUP_OUT_C = OUT_C // GROUPS
{% endif %}

    x_base = X + (group * stride_xc * GROUP_IN_C + idx_n * stride_xn)[:, None]
    w_base = (
        W + (group * stride_wc_out * GROUP_OUT_C + idx_y_c * stride_wc_out)[None, :]
    )

    acc = tl.zeros((BLOCK_M, BLOCK_N), dtype=tl.float32)

{% if UNROLL %}
{% for d in range(KERNEL_D) %}
{% for i in range(KERNEL_H) %}
{% for j in range(KERNEL_W) %}
    d = {{d}}
    i = {{i}}
    j = {{j}}
    for k in range(0, GROUP_IN_C, BLOCK_K):
        aF  
{% endfor %}
{% endfor %}
{% endfor %}
{% else %}
    # Could be simplified, but slightly slower:
    # for d in range(KERNEL_D):
    #   for i in range(KERNEL_H):
    #     for j in range(KERNEL_W):
    #         for k in range(0, GROUP_IN_C, BLOCK_K):
    BLOCK_K_COUNT = (GROUP_IN_C + BLOCK_K - 1) // BLOCK_K
    for dijk in range(KERNEL_D * KERNEL_H * KERNEL_W * BLOCK_K_COUNT):
        k = (dijk % BLOCK_K_COUNT) * BLOCK_K
        dij = dijk // BLOCK_K_COUNT
        j = dij % KERNEL_W
        di = dij // KERNEL_W
        i = di % KERNEL_H
        d = di // KERNEL_H
        a  
{% endif %}

    mask = (
        (idx_n < BATCH)[:, None]
        & (idx_y_d < OUT_D)[:, None]
        & (idx_y_h < OUT_H)[:, None]
        & (idx_y_w < OUT_W)[:, None]
        & (idx_y_c < GROUP_OUT_C)[None, :]
    )
    idx_n = idx_n[:, None]
    idx_c = idx_y_c[None, :] + group * GROUP_OUT_C
    idx_d = idx_y_d[:, None]
    idx_h = idx_y_h[:, None]
    idx_w = idx_y_w[:, None]

    # inductor generates a suffix
    {{store_output(("idx_n", "idx_c", "idx_d", "idx_h", "idx_w"), "acc", "mask", val_shape=("BLOCK_M", "BLOCK_N"))}}
zat::convolutionF)has_out_variantop_overloadc          
         t        j                  t        j                  |d      d      }t        j                  | j                  dddd      |j                  dd      |j                  dddd            S )Nr   r      r   )out)torchsqueezematmulpermute)xr'   rB   s      r*   conv1x1_via_mmrH   f  s]    emmAr*B/A<<			!Q1qyyACKK1a4K r,   c                  J    e Zd ZU ded<   ded<   ded<   ded<   ded<   ded	<   y
)ConvLayoutParamstuple[int, ...]stridepaddingdilationbool
transposedoutput_paddingintgroupsN)__name__
__module____qualname____annotations__r#   r,   r*   rJ   rJ   p  s%    ##Kr,   rJ   c	                   t         j                  j                  j                  }	t         j                  j                  5  t
        j                  j                  j                  t        j                  |       t        j                  |      t        j                  |       |	|       |	|       |	|      | |	|      |	      }
t        j                  |
j                               }t        j                  |
j                               }ddd       t        j                  | j                         | j!                         |      S # 1 sw Y   =xY w)z)Determine output layout for a convolutionN)r   graphsizevarsguard_int_seq	fake_moderC   opsatenconvolutionr	   ir_node_to_tensorconvert_shape_to_inductorsizerL   FixedLayoutget_device_or_error	get_dtype)rG   weightbiasrL   rM   rN   rP   rQ   rS   guardoutputsizess               r*   conv_layoutrk   y  s    & GG**E	
		 ?++  #  (  &&M'N(O.!

 ,,V[[];--fmmo>? >>			 ? ?s   C	EEc                    t        t        t        |                   }|j                  d|j	                  d             |S )Nr   r@   )listreversedrangeinsertpop)rankorders     r*   channels_last_orderrt     s0    %+&'E	LLEIIbM"Lr,   c                   t        |j                               }t        |dz
        D ]   }t        t        j
                     |d      }" t        t        j                     |ddg      }t        j                  j                  | t        |            } t        t        |            }|j                  |j                  d             t        t        j                     | |      } | j                         ^ }}t        t        j                     | t        |      |g      } |t        t        j                      | |      }nt        t        j"                     || |      }t        t        j                     |g |d      }t        t        |            }	|	j%                  d|	j                  d             t        t        j                     ||	      S )Nr   r@   dimr   r   )lenget_sizero   Lr^   rD   rF   r	   ExternKernelrequire_stride_orderrt   rm   appendrq   reshaper   mmaddmmrp   )
rG   rf   rg   rr   _	x_permuterj   in_chanresultresult_permutes
             r*   convert_1x1_conv_to_mmr     sa   v !D4!8_ 14<<R01t||_VaV,F
,,Q0CD0IJAU4[!IY]]1%&	$,,9%AjjlOUG	$,,M%0':;A|477Av&4::tQ/t||_V\u\b\2F%+&N!^//34T\\?6>22r,   c	                b    t        |      }t        |      }t        |      }t        |      }t        |t              s)t        j                  j
                  j                  |      }t        |t              sJ t        t        j                  j
                  j                  |            }t        t        j                  j
                  j                  |            }||||||dt        j                         }	t         j                               t        j                               dz
  k(  rVt        t        j                     t        t        t        j                       dg j                               |fi d      S t        j                  j
                  j                  j                               ^}
}}t         j                               dk(  rt        |      dk(  r|	dk(  rj#                  d|z   d|z   d|z   d|z   d	       t        t        j$                      d
       t        t        j$                     d
      t        t        j                     t         |fi d
      S t        |      t'        |      }t'        |      }t'        |      }t'        |      } fd}t(        j*                  xs t(        j,                  }t(        j.                  s	|r |       rt1        |      rt1        |      rvt3        |      rkt1        |      r`|s^t3        |      rS|dk(  rNt        j                  j
                  j5                  t7         j                               d      rt9         |      S ||	dk7  rt         dfi }t        j                  j
                  j;                  |j                         d   d      r|S t        t        j<                     |t        t        j>                     ||j                         d   gdgz  z               S  jA                          jA                          t        j                  jB                  rud
k(  rpt        j                  xjD                  dz  c_"        t        jF                  jI                          t        jF                  jI                        tK         dfi }ntK         dfi }t        jL                  t        j                  j
                  jO                  |jP                              }t        jF                  jS                   |       t        jF                  jS                  |      g d}| g}dd<   |jU                  dd       not        jF                  jW                  |      }|J  |g}|jY                          t        j                  j
                  j                  |j                                g }tZ        j\                  j^                  ja                  d      rtc        jd                  |||fi g}tZ        j\                  j^                  ja                  d      rtg        |      rt1        |      r|st3        |      rt        j                  j
                  j;                  ||z   j                         d         rt1        |      r@t1        |      r5t3        |      r*|dk(  r%|ji                  tj        je                  ||             |dkD  xr |dk(  xr |
|k(  }|rrdk(  rmt        jl                  jo                  |	      }|D ]I  }tq        jr                  |f f||d   |d   |d   |jt                  |jv                  d|jx                   K t        jl                  j{                  |	      } j}                         j~                  } |t7         j                         d   g j                         d
d       |
||      D ]K  }t1        |      }|r|jv                  nt        |jv                  d      }d
k(  rrt        jr                  |f f||d   |d   |d   |d   |d   |d   ||tZ        j                  j                  j                  dk(  |jt                  |d|jx                   dk(  st        jr                  |fi d fd|d|d   d|d   d|d
   d|d   d|d   d|d
   d|d   d|d   d |d
   d!|d"|d#tZ        j                  j                  j                  dk(  d$|jt                  d%||jx                   N t        |      r/t        j                  || f||fn	t               z   ||||&       t        d'|||      \  }}|S )(zGLower aten.convolution using Inductor convolution kernels or fallbacks.rL   rM   rN   rP   rQ   rS   r   r   rv   rA   xpu)r   )r   )rL   rM   rN   rQ   r   c                    t         j                  j                  rdk(  ryt        d fi } t	        j
                  t         j                  j                  j                  | j                              }|t        j                  k(  S )Nr   T)
r   rY   
layout_optrk   r	   get_stride_orderrZ   guarding_hints_or_throwrL   NHWC_STRIDE_ORDER)layoutreq_stride_orderkwargsndimrf   rG   s     r*   channels_last_convz'convolution.<locals>.channels_last_conv  sn    77$!)Q77 ..GG44V]]C
  2#7#777r,   Ncpurg   ATENTRITON)input_nodesr   KERNEL_SIZECONV_STRIDEPADDING
num_stages	num_warps
dtype_size   tf32)r   r   KERNEL_HKERNEL_WSTRIDE_HSTRIDE_W	PADDING_H	PADDING_Wr"   UNROLL
ALLOW_TF32r   r   r   r   KERNEL_Dr   r   STRIDE_Dr   r   	PADDING_Dr   r   r"   r   r   r   r   )r   rL   rM   rN   rS   n_spatial_dimensionsr_   )Jtuple
isinstancerR   r   rY   rZ   	guard_intr[   r	   get_device_typerx   ry   rz   r^   rD   r_   expandupdate	unsqueezer   r   max_autotunemax_autotune_gemmconv_1x1_as_mmr   r   statically_known_gtr   r   statically_known_equalsaddviewrealizer   num_channels_last_convr{   require_channels_lastrk   r   r   rL   r|   rp   realize_inputfreeze_layoutrC   	_inductorutils_use_conv_autotune_backendaten_convolutionbindr   r}   aten_conv1x1_via_mmchoicesget_depthwise_conv_configsdepthwise_conv1d_templatemaybe_append_choicer   r   r   get_conv_configsre   itemsizeminconv2d_templatebackendscudnnfp32_precisionconv3d_templater   r   add_ck_conv_choicesr   ) rG   rf   rg   rL   rM   rN   rP   rQ   rS   device_typeout_chanr   kernel_shaper   autotuning_gemmr   r   r   ordered_kwargs_for_cpp_kernelargsr   is_depthwisedepthwise_configscfgconv_configsr   unrollr   noder   r   r   s    ``                            @@r*   r_   r_     s    6]FGnGXH>*Nfc"!!++F3fc""" 177##11&9:FAGG$$227;<G  ( F $$Q'K
1::<C 12Q66$++q1*<qzz|*<=vtVvV
 	

 ()ww'7'7'E'EfooFW'X$Hg
 1::<A#l"3q"8[E=Q-'> 8O"&"7		
 dnnaQ'4>>"6q164262
 	

 |D&$'F7D)GHd+H!.$7N
8 ))EV-E-EO 
		?7I7KL!FOWH^$aKGG00qzz|1LaP%a66K50Q777733FOO4Ea4H!LM{AdiiL(9!(<'=s
'JK
 	
 IIK
NN
 	wwdai	&&!+&OO11!4 66v>Q77Q77 ..GG44V]]C
 OO004DE55f>NO%! |6{v%,,Q7,,T264 	&&t}}7G77?!!- 	
 	88B'H^$GG44Wv5Eqzz|TUW L!!!NN.33D&ABzIglIx67IDAI !		 D D[ Q( )==
!"F! ,Q &q	#AJ"~~!mm
 jj
 yy11+>[[]++
1::<?>QZZ\!"-=>?!	
 5	C \*F *0S5JIqy33!"F!)!_)!_#AY#AY%aj%aj! "$~~33BBfL"~~'!" jj#& 33!"F " *!_	
 *!_ *!_ $AY $AY $AY &aj &aj &aj "  "!"  %~~33BBfL#$  #~~%& (jj)A5	l F# 44F$2BwP!%		
 (wfMGD!Kr,   c                (    t        | ||||||||	      S N)r_   )rG   rf   rg   rL   rM   rN   rP   rQ   rS   	benchmarkdeterministiccudnn_enabled
allow_tf32s                r*   _convolutionr     s%      	64(JPV r,   c                    | j                   t        j                  j                  j                  j
                  u sJ t        j                  j                  r||fS t        | g|i |S r   )
targetrC   r]   r^   r_   defaultr   rY   r   r   fx_noder   r   s      r*   constrain_conv_to_fx_stridesr     sT    >>UYY^^77?????wwV|&w@@@@r,   c	                   t         j                  j                  j                  }	t         j                  j                  5  t        j                  |       }
t        j                  |      }t        j                  |      }t        j                  j                  j                  |
||d  |	|       |	|       |	|      | |	|      |d      \  }}}t        j                  |j                               }t        j                  |j                               }d d d        t        j                  |j                         |j!                               S # 1 sw Y   =xY w)NTFFr   rY   rZ   r[   r\   r	   r`   rC   r]   r^   convolution_backwardra   rb   rL   rc   rd   re   )grad_outinputrf   rL   rM   rN   rP   rQ   rS   rh   gorG   r'   dxr   rj   stride_s                    r*   conv_bwd_input_layoutr     s    GG**E	
		 <!!(+  '  (99>>66&M'N(O.! 
Aq ,,RWWY7..ryy{;'<* >>!!#	 +< <   CEEc	                   t         j                  j                  j                  }	t         j                  j                  5  t        j                  |       }
t        j                  |      }t        j                  |      }t        j                  j                  j                  |
||d  |	|       |	|       |	|      | |	|      |d      \  }}}t        j                  |j                               }t        j                  |j                               }d d d        t        j                  |j                         |j!                               S # 1 sw Y   =xY w)NFTFr   )r   r   rf   rL   rM   rN   rP   rQ   rS   rh   r   rG   r'   r   dwrj   r   s                    r*   conv_bwd_weight_layoutr     s    GG**E	
		 <!!(+  '  (99>>66&M'N(O.! 
2q ,,RWWY7..ryy{;'<* >>""$	 +< <r   c               d   | j                  t        j                        rt        j                  }
nt        j                  }
t        j                  ||	j
                  | j                  |
      }t        j                  j                  j                  j                  d |	d || |d ||||||d       |	S )Nmemory_formatdtypedevicer   r   out1out2out3grad_outputr   rf   
bias_sizesrL   rM   rN   rP   rQ   rS   output_maskis_contiguousrC   channels_lastcontiguous_formatemptyr   r   r]   r^   r   rB   )x_tgo_tw_shaperL   rM   rN   rP   rQ   rS   rB   
memory_fmtdummy_weights               r*   call_aten_dwr  D  s     u':':;((
,,
;;syy:L 
IINN''++%( ,   Jr,   c               d   | j                  t        j                        rt        j                  }
nt        j                  }
t        j                  ||	j
                  | j                  |
      }t        j                  j                  j                  j                  |	d d | ||d ||||||d       |	S )Nr   r   r   r  r  )r  w_tx_shaperL   rM   rN   rP   rQ   rS   rB   r  dummy_inputs               r*   call_aten_dxr  p  s     (;(;<((
,,
++syyJK 
IINN''++%( ,   Jr,   c               P    |d   } |||z  |z  |d          || |z  |d         |fS Nr"   r    r!   r#   )coutcinkhkwr(   r)   gs          r*   conv2d_bwd_weight_gridr    s@    XAS2X]DO,TQYY(	 r,   convolution2d_bwd_weighttriton_conv2d_bwd_weightc               P    |d   } || |z  |z  |d          |||z  |d         |fS r  r#   )r$   r  r&   r'   r(   r)   r  s          r*   conv2d_bwd_input_gridr#    s@    XAQUQYY(SAXtI'	 r,   convolution2d_bwd_inputtriton_conv2d_bwd_inputc                   t        |      }t        |      }t        |      }t        |      }t        |	t              s)t        j                  j
                  j                  |	      }	t        j                  j
                  j                  |j                               ^}}}t        t        j                  j
                  j                  |            }t        t        j                  j
                  j                  |            }t        t        j                  j
                  j                  |            }|j                          |j                          | j                          ||||||	d}t        |      }t        ||      }t        ||      }t        ||      }t        ||      }t        j                  |      }t        j                  j                  |      }|j!                         j"                  }d}d}g }g }t%        | ||fi |}|
d   r#t        j                  j&                  ru|dk(  rpt        j                  xj(                  dz  c_        t        j*                  j-                  |      }t        j*                  j-                  |       } t%        | ||fi |}nt        j                  j
                  j                  }t        j.                   ||j0                              }t        j*                  j3                  ||      }t        j*                  j3                  | |      } || g}t4        j6                  j8                  j;                  d      rt=        |      r|st?        |      r |tA        |j                         d   g|j                         dd       |||      D ]  }|dk(  s	d	}tC        jD                  |f|| f||d   |d   |d   |d   |d   |d   |d   |d   |	t4        jF                  jH                  jJ                  |jL                  |jN                  d
|jP                    d}d}g }g }tS        | ||fi |}|
d   r#t        j                  j&                  ru|dk(  rpt        j                  xj(                  dz  c_        t        j*                  j-                  |       } t        j*                  j-                  |      }tS        | ||fi |}nt        j                  j
                  j                  }t        j.                   ||j0                              }t        j*                  j3                  | |      } t        j*                  j3                  ||      }| |g}t4        j6                  j8                  jU                  d      rt=        |      r|st?        |      r |tA        |j                         d   g|j                         dd       |||      D ]  }|dk(  s	d	}tW        jD                  |f| |f||d   |d   |d   |d   |d   |d   |d   |d   |	t4        jF                  jH                  jJ                  |jL                  |jN                  d
|jP                    |s|stY        | |||||||||	|
      S |
d   rzt4        j6                  j8                  j;                  d      s|s>|j[                  t\        j_                  ||g d|j                         ||||||	
             ta        d|||      \  }} |
d   rzt4        j6                  j8                  jU                  d      s|s>|j[                  tb        j_                  ||g d|j                         ||||||	
             ta        d|||      \  }} d}!|
d   r:|8te        tf        jh                     | dgtk        tm        d|dz               z         }!|||!fS )a  
    Lowering function for backward convolution operator.

    TRITON kernels are only registered for supported configurations (currently 2D convolutions).
    For unsupported dimensions or configurations, the choices list remains empty,
    triggering an automatic fallback to the ATen reference implementation.
    This ensures correctness for all cases while enabling TRITON optimizations only where implemented.
    r   FNr   r   r   r   r   T)r   r   r   r   r   r   r   r   
DILATION_H
DILATION_Wr"   r   r   r   r   )r  rL   rM   rN   rP   rQ   rS   )
r   r   r   r  rL   rM   rN   rP   rQ   rS   convolution_bwd_weight)r  rL   rM   rN   rP   rQ   rS   )
r   r   r   r  rL   rM   rN   rP   rQ   rS   convolution_bwd_input)axis)7r   r   rR   r   rY   rZ   r   r[   ry   r   rx   r   r	   r   r   r   re   r   r   r   r   r{   r   r   rL   r|   rC   r   r   %_use_conv_bwd_weight_autotune_backendr   r   r   conv2d_bwd_weight_templater   r   r   r   r   r   r   r   $_use_conv_bwd_input_autotune_backendconv2d_bwd_input_template"aten_convolution_backward_fallbackr}   ext_kn_aten_dwr   r   ext_kn_aten_dxrz   r^   sumrm   ro   )"r   r   rf   r  rL   rM   rN   rP   rQ   rS   r  r   r   r   r   r   r   r   r   has_triton_dw_choicesr   
choices_dwargs_w	layout_dwrh   stride_orderr   has_triton_dx_choicesr   
choices_dxargs_x	layout_dxr   dbs"                                     r*   convolution_backward_loweringr>    s   , 6]FGnGXH>*Nfc"!!++F3'(ww'7'7'E'EfooFW'X$Hg 177##11&9:FAGG$$227;<GQWW%%33H=>H	MMO
NN  ( F |D&$'F7D)GHd+H!.$7N$$U+K99--k:L"++J!	BJF&xI&II1~77$!)GG**a/*OO99%@E<<XFH.xQ&QIGG$$22E..uY5E5E/FGLOO88ME;;HlSH" OO!!GGQ#I.(#u~~/2JU^^5Eab5IJK%	  19,0).BB"%*H$5(!-a!-a")!*")!*!'!'#+A;#+A;%#(>>#7#7#B#B#&>>"%--  **!: "	BJF%hvHHI1~77$!)GG**a/*<<XFH__::6BF-hvPPIGG$$22E..uY5E5E/FGL;;HlSH__99&,OFF# OO!!FFxP#I.( $u~~/2JU^^5Eab5IJK%	  19,0)-AA"%-v$6(!-a!-a")!*")!*!'!'#+A;#+A;%#(>>#7#7#B#B#&>>"%--  **!< !)>1
 	
 1~OO!!GGO(## &$3 #OO-!#%)#1!% $ 2 *$j&)
A 1~OO!!FFvN(## &$3 "NN,!#%)#1!% $ 2 *#Z
A 
B1~*0txx[d5D1H3E.F(FGB<r,   c                    | j                   t        j                  j                  j                  j
                  k(  sJ t        j                  j                  r||fS t        | g|i |S r   )
r   rC   r]   r^   r   r   r   rY   r   r   r   s      r*    constrain_conv_bwd_to_fx_stridesr@    sT    >>UYY^^@@HHHHHwwV|&w@@@@r,   )rG   r   rf   r   rg   TensorBox | NonerL   Sequence[int]rM   rK   rN   rK   rP   rO   rQ   rK   rS   rR   return	ir.Layout)rG   r   rf   r   rg   rA  rL   rB  rM   rB  rN   rB  rP   rO   rQ   rB  rS   rR   )r   r   r   r   rf   r   rL   rB  rM   rK   rN   rK   rP   rO   rQ   rK   rS   rR   rC  rD  )r   r   r   r   rf   r   r  zSequence[int] | NonerL   rB  rM   rB  rN   rB  rP   rO   rQ   rB  rS   rR   r  zSequence[bool])L
__future__r   loggingtypingr   r   rC   -torch._inductor.codegen.rocm.ck_conv_templater    r   r	   loweringr
   r   r   r   rz   r   select_algorithmr   r   r   r   r   r   r   r   r   r   r   virtualizedr   	mm_commonr   collections.abcr   r   	getLoggerrT   logr]   r^   r+   r/   r4   r   LOOP_BODY_2Dr   LOOP_BODY_3Dr   r_   r   r   rH   r   rJ   rk   rt   r   r   r   r   r   r  r1  r  r2  r  r-  r#  r/  r   r0  r>  r@  r#   r,   r*   <module>rS     s   "  +  R      + (g! yy~~       +		 78'+	 8
 !		3h i4jkCH IDJKUYvB !		;x y<z{O` aPbccgR &	  ((	  )> y ((( ( 	(
 ( ( ( $( ( (V3. 4##$MMM M 	M
 M M M "M M %M` 4$$% &(A d&&(D E&&& & 	&
 & & & $& & &R&&& & 	&
 & & & $& & &R&R $L$7&R $L$7   ,	#	 :;    +	"	 9:  &6d6O6O6W6W%X " 4,,445CCC C %	C
 C C C C "C C  C 6CLA d//1Q Rr,   