
    Wi	                         d dl mZ d dlmZmZ g dZ G d de      Zedd       Zedd       Zedd
       Z	ed	ddd       Z
y)    )SwizzledSharedLayout)builtin_unwrap_if_constexpr)arriveinit
invalidateMBarrierLayoutwaitc                   $     e Zd ZdZd fd	Z xZS )r	   z
    Layout for mbarrier synchronization in Ampere and later architectures.

    Args:
        cga_layout (List[List[int]]): CTA layout bases. Defaults to [].
    c                 8    t         |   ddddg|xs g        y )N   r   )vec	per_phase	max_phaseorder
cga_layout)super__init__)selfr   	__class__s     /home/sietch6/trending-topics-pipeline/venv/lib/python3.12/site-packages/triton/experimental/gluon/language/nvidia/ampere/mbarrier.pyr   zMBarrierLayout.__init__   s$    Q!qPZP`^`a    N)__name__
__module____qualname____doc__r   __classcell__)r   s   @r   r	   r	      s    b br   r	   Nc                 f    t        |      }|j                  j                  | j                  |       y)z
    Initialize an mbarrier with a specified count.

    Args:
        mbarrier (shared_memory_descriptor): The barrier object to initialize.
        count (int): The initial count for the barrier.
    N)r   buildercreate_mbarrier_inithandle)mbarriercount	_semantics      r   r   r      s(     !'E**8??EBr   c                 N    |j                   j                  | j                         y)z
    Invalidate an mbarrier, resetting its state.

    Args:
        mbarrier (shared_memory_descriptor): The barrier object to invalidate.
    N)r    create_mbarrier_invalr"   )r#   r%   s     r   r   r       s     ++HOO<r   Tc                     |j                  |      }|j                  |      }|D cg c]  }|j                   }}|j                  j                  | j                  |j                  |j                  |       yc c}w )a  
    Wait until the mbarrier object completes its current phase.

    Args:
        mbarrier (shared_memory_descriptor): The barrier object to wait on.
        phase (int): The phase index to wait for.
        pred (bool): Predicate. Operation is skipped if predicate is False. Defaults to True.
        deps (Sequence[shared_memory_descriptor]): Dependent allocations barrier is waiting on. Used to track liveness of dependent allocations. Defaults to ().
    N)	to_tensorr"   r    create_mbarrier_wait)r#   phasepreddepsr%   xs         r   r
   r
   +   sg     &Et$D"#AHH#D#**8??ELL$++W[\ $s   A9)r,   r%   c                    d}|j                  |      }|j                  j                  | j                  ||j                         y)a  
    Arrive on an mbarrier, signaling that a thread has reached the barrier.

    Args:
        mbarrier (shared_memory_descriptor): The barrier object to arrive on.
        pred (bool): Predicate. Operation is skipped if predicate is False. Defaults to True.
    r   N)r)   r    create_mbarrier_arriver"   )r#   r,   r%   r$   s       r   r   r   <   s9     Et$D,,X__eT[[Qr   r   )T N)+triton.experimental.gluon.language._layoutsr   (triton.experimental.gluon.language._corer   r   __all__r	   r   r   r
   r   r1   r   r   <module>r5      sz    L R
D	b) 	b 		C 		C 	= 	= 	] 	]  	!T 
R 	
Rr   