
    ^jH                         d dl Z d dlmc mc mZ d dlmZmZ d dl	m
Z
 d dlmZmZ g dZ G d de
      Zeddej                   d	ej                   fd
       Zedd       Zedd       Zedd       Zedddd       Zy)    N)constexpr_functionjit)SwizzledSharedLayout)builtin_unwrap_if_constexpr)allocate_mbarrierarriveinit
invalidateMBarrierLayoutwaitc                   J     e Zd ZdZd fd	Zeeddedefd              Z	 xZ
S )r   z
    Layout for mbarrier synchronization in Ampere and later architectures.

    Args:
        cga_layout (List[List[int]]): CGA 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     /var/www/ramen.bs-engineer-server.com/venv/lib/python3.12/site-packages/triton/experimental/gluon/language/nvidia/ampere/mbarrier.pyr   zMBarrierLayout.__init__   s$    Q!qPZP`^`a    num_ctastwo_ctac                    t        j                  |       } t        j                  |      }|r| dz  dk(  sJ d       | dkD  sJ d       | | dz
  z  dk(  sJ d       g }|r|j                  dg       | dz  } t        t	        t        j                  |                   D ]  }|j                  d|z  g        t        |      S )z
        Create a multi-CTA mbarrier layout.

        Args:
            num_ctas (int): Number of CTAs.
            two_cta (bool): Whether the barrier should synchronize every other CTA
           r   z&num_ctas must be even for two-CTA modeznum_ctas must be positiver   znum_ctas must be a power of two)ttglr   appendrangeintmathlog2r   )r   r   basesis       r   multictazMBarrierLayout.multicta   s     ,,X6++G4a<1$N&NN$!|888|HqL)a/R1RR/LL!NHs499X./0 	!ALL!Q$ 	!e$$r   N)F)__name__
__module____qualname____doc__r   staticmethodr   r#   boolr(   __classcell__)r   s   @r   r   r      s8    b %3 % %  %r   r   batchtwo_ctasc                 ,   t        j                         }|s|n|dz  }t        j                  | du xs t        | j                  t
                     | |gn| |g}t        j                  t         j                  |t        j                  ||            }|S )z
    Helper function to allocate an mbarrier

    Args:
        two_ctas (bool): Whether the barrier should synchronize every other CTA
    r   N)r   r   )
r    r   static_assert
isinstancevaluer#   allocate_shared_memoryint64r   r(   )r1   r2   r   	num_elemsshapebars         r   r   r   1   s      $}}H08h!mIu}D
5;;(DE+0=YKui>PE

%
%

8DC
 Jr   c                 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
   E   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@   rB   s     r   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depsrB   xs         r   r   r   ]   sg     &Et$D"#AHH#D#**8??ELL$++W[\ $s   A9)rI   rB   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)rF   r=   create_mbarrier_arriver?   )r@   rI   rB   rA   s       r   r	   r	   n   s9     Et$D,,X__eT[[Qr   )NFr)   )T N)r$   "triton.experimental.gluon.languageexperimentalgluonlanguager    "triton.experimental.gluon._runtimer   r   +triton.experimental.gluon.language._layoutsr   (triton.experimental.gluon.language._corer   r   __all__r   	constexprr   r
   r   r   r	   rN   r   r   <module>rX      s     1 1 F L R
Y#%) #%L T^^ dnn  & 		C 		C 	= 	= 	] 	]  	!T 
R 	
Rr   