
    Wi                    x   d dl mZ d dlmZmZmZ d dlmZ d dlm	c m
c mc mZ d dlmZmZ d dlmZmZ erd dlmZ d dlmZ g d	Z ed
       G d dej.                               Ze G d dej2                               Ze	 d	 	 	 	 	 	 	 dd       Ze	 d	 	 	 	 	 dd       Ze	 d	 dd       Zeddd       Zy)    )annotations)ListTupleTYPE_CHECKING)	dataclassN)PaddedSharedLayoutSwizzledSharedLayout)builtin_unwrap_if_constexpr)ir)shared_memory_descriptor)
async_load
async_waitmake_tensor_descriptortensor_descriptortensor_descriptor_typeT)eqc                  b    e Zd ZU dZded<   ded<   ded<   ded<   dd	Zdd
ZddZddZddZ	y)r   z!The type for a tensor descriptor.zttgl.block_type
block_typezttgl.tuple_type
shape_typestrides_type)PaddedSharedLayout | SwizzledSharedLayoutlayoutc                <    d| j                    d| j                   dS )Nztensor_descriptor<z, >)r   r   selfs    ~/home/sietch6/trending-topics-pipeline/venv/lib/python3.12/site-packages/triton/experimental/gluon/language/amd/gfx1250/tdm.py__str__ztensor_descriptor_type.__str__   s     #DOO#4Bt{{m1EE    c                    ||   }|dz  }| j                   j                  ||      \  }}| j                  j                  ||      \  }}t        ||||       }||fS )N   )r   _unflatten_irr   r   )r   handlescursorhandleshapestridesvalues          r   r#   z$tensor_descriptor_type._unflatten_ir   sd    !55gvFv++99'6J!&%$?f}r    c                    | j                   j                  j                         }|j                  | j                   j	                  |      || j
                  j                  |            S N)r   
element_tyis_int_signed!get_tensor_descriptor_layout_typeto_irr   _to_ir)r   builder	is_signeds      r   r0   ztensor_descriptor_type._to_ir$   sT    OO..<<>	88OO!!'*KKw'
 	
r    c                    |j                  | j                  |             | j                  j                  ||       | j                  j                  ||       y r+   )appendr0   r   _flatten_ir_typesr   )r   r1   outs      r   r5   z(tensor_descriptor_type._flatten_ir_types,   sA    

4;;w'())'37++GS9r    c           	         d| j                   j                          d| j                  j                          d| j                  j                          d| j                  j                          d	S )NTD_)r   mangler   r   r   r   s    r   r:   ztensor_descriptor_type.mangle1   sb    DOO**,-Qt/E/E/G.H$J[J[JbJbJdIeefgkgrgrgygyg{f||~r    N)returnstr)r$   List[ir.value]r%   intr;   zTuple[tensor_descriptor, int])r1   
ir.builderr;   zir.type)r1   r?   r6   zList[ir.type]r;   None)
__name__
__module____qualname____doc____annotations__r   r#   r0   r5   r:    r    r   r   r      s8    +!!55F
:
@r    r   c                      e Zd ZU dZded<   ded<   ded<   ded<   dd	Zed
        Zed        Zed        Z	ed        Z
y)r   z4A descriptor representing a tensor in global memory.zir.valuer&   z
ttgl.tupler'   r(   r   typec                    |j                  | j                         | j                  j                  |       | j                  j                  |       y r+   )r4   r&   r'   _flatten_irr(   )r   r$   s     r   rJ   ztensor_descriptor._flatten_ir>   s6    t{{#

w'  )r    c                .    | j                   j                  S r+   )rH   r   r   s    r   r   ztensor_descriptor.block_typeC   s    yy###r    c                B    | j                   j                  j                  S r+   )rH   r   r'   r   s    r   block_shapeztensor_descriptor.block_shapeG   s    yy##)))r    c                B    | j                   j                  j                  S r+   )rH   r   r,   r   s    r   dtypeztensor_descriptor.dtypeK   s    yy##...r    c                .    | j                   j                  S r+   )rH   r   r   s    r   r   ztensor_descriptor.layoutO   s    yyr    N)r$   r=   r;   r@   )rA   rB   rC   rD   rE   rJ   propertyr   rM   rO   r   rF   r    r   r   r   5   sr    >
  *
 $ $ * * / /    r    r   c                   t        |      }d|cxk  rdk  sn J d| d       t        |      |k(  sJ d| dt        |              t        |      |k(  sJ d| dt        |              t        | j                  t        j                        sJ d	       t        |      }t        |t        t        f      sJ d
       t        |t              r|j                  dk(  sJ d       | j                  }|j                  |d      }|j                  |d      }	t        j                  |      }t        j                  |      }t        j                  | j                  j                  |      }
t        |
|j                  |j                  |      }|j!                  d      }|j"                  j%                  |j'                  |j"                        |||	|      }t)        ||||      S )a  Make a tensor descriptor object.

    Args:
        base (tensor): base pointer of the tensor in global memory.
        shape (List[int]): shape of the tensor.
        strides (List[int]): strides of the tensor.
        block_shape (List[int]): block shape of the tensor.
        layout (PaddedSharedLayout | SwizzledSharedLayout): the layout of the tensor in shared memory.

    Returns:
        tensor_descriptor: the created tensor descriptor object
    r"      z Expected 1 <= ndim <= 5 but got z dimensionsz	Expected z strides but got zExpected block_shape to have z dimensions but got zExpected base to be a pointerzBExpected layout to be a PaddedSharedLayout or SwizzledSharedLayoutz3Expected max_phase to be 1 for SwizzledSharedLayoutFrequire_i64Tzero)len
isinstancerO   ttglpointer_typer   r   r	   	max_phaser&   _convert_to_ir_valuestupler   rH   r,   r   _str_to_padding_optionr1   create_make_tensor_descriptorr0   r   )baser'   r(   rM   r   	_semanticndimbase_handleshape_handlesstride_handlesr   rH   paddingr&   s                 r   r   r   T   s     u:D>>O=dV;OO>w<4R9TF2CCL>!RR{t#m'DTFJ^_bcj_k^l%mm#djj$"3"34U6UU4!&)Ff13GHI MLMI&./1$[&[[$++K33Eu3MM44W$4ONJJuEjj!G!5!5{CJ!*ejj',,OD..v6G<<T[[IZIZ=[]hjw=KWVF VUGT::r    c                8   |j                  |d      }|j                  |      }|j                  }t        |      }||j                  nt        j
                  j                         }|j                  j                  | j                  ||j                  ||       y)a-  Load a block of tensor specified in tensor descriptor from global memory to shared memory asynchronously.

    Args:
        src (tensor_descriptor): the source tensor descriptor.
        offsets (List[int]): the offsets from the base pointer in the tensor descriptor.
        dest (shared_memory_descriptor): the shared memory destination to store the loaded data.
        pred (bool, optional): Predicate to enable or disable the load. Defaults to True.
        mbarrier (shared_memory_descriptor, optional): The barrier object to signal "arrive" on.
    FrT   N)	r\   	to_tensorr&   r   rY   r   r)   r1   %create_async_tdm_copy_global_to_local)	srcoffsetsdestpredmbarrierra   offset_handlespred_handlembarrier_handles	            r   r   r      s     44W%4PNt$D++K#H-H)1)=hoo477==?O;;CJJX\XcXcep<KMr    c                    |j                  |d      }|j                  j                  | j                  ||j                         y)ak  Store a block of tensor specified in tensor descriptor from shared memory to global memory asynchronously.

    Args:
        dest (tensor_descriptor): the destination tensor descriptor.
        offsets (List[int]): the offsets from the base pointer in the tensor descriptor.
        src (shared_memory_descriptor): the shared memory source to load the data.
    FrT   N)r\   r1   %create_async_tdm_copy_local_to_globalr&   )rl   rk   rj   ra   ro   s        r   async_storert      s<     44W%4PN;;DKKY\YcYcdr    c                P    t        |       } |j                  j                  |        y)zWait for the outstanding asynchronous tensor operations to complete.

    Args:
        num_outstanding (int): number of outstanding async tensor operations to wait for.
    N)r   r1   create_async_tdm_wait)num_outstandingra   s     r   r   r      s"     +?;O++O<r    r+   )r`   zttgl.tensorr'   "List[ttgl.constexpr | ttgl.tensor]r(   rx   rM   zList[ttgl.constexpr]r   r   r;   r   )TNN)rj   r   rk   rx   rl   r   rm   boolrn   r   r;   r@   )rl   r   rk   rx   rj   r   r;   r@   )r   N)r;   r@   )
__future__r   typingr   r   r   dataclassesr   (triton.experimental.gluon.language._coreexperimentalgluonlanguage_corerY   +triton.experimental.gluon.language._layoutsr   r	   r
   r   	triton._Cr   r   __all__	base_typer   
base_valuer   r   r   rt   r   rF   r    r   <module>r      s   " - - ! 7 7 ` RQ
o d!@T^^ !@ !@H      < 	 Y](;$F(;Ui(;#L(;ar(; 	(;V 	W[MM,DM`dM 	M( 	
e#'
e 	
e 	= 	=r    