9595 format_var ,
9696 LocalArrayContextManagerValue ,
9797)
98- from .._stub import TensorMapSwizzle
98+ from .._stub import TensorMapSwizzle , MbarrierScope
9999from cuda .tile ._ir import hir_stubs
100100from cuda .tile ._ir .typing_support import I32_TY , U64_TY , BOOL_TY
101101
@@ -1413,42 +1413,56 @@ def mbarrier_invalidate_impl(mbar: Var) -> Var:
14131413 )
14141414
14151415
1416- def _memory_space_to_mbar_suffix ( memory_space : MemorySpace ) -> str :
1417- match memory_space :
1416+ def _mbar_space_scope_suffix ( scope : MbarrierScope , space : MemorySpace ) -> str :
1417+ match space :
14181418 case MemorySpace .SHARED :
1419- return 'cta'
1419+ space_str = 'cta'
14201420 case MemorySpace .SHARED_CLUSTER :
1421- return 'cluster'
1422-
1423- raise TileCompilerError (f"Unexpected { memory_space = } " )
1424-
1425-
1426- def _mbar_space_scope_suffix (scope : MemorySpace , space : MemorySpace ):
1421+ space_str = 'cluster'
1422+ case _:
1423+ raise TileCompilerError (f"Unexpected { space = } " )
14271424 return (
14281425 ".scope."
1429- + _memory_space_to_mbar_suffix ( scope )
1426+ + scope . value
14301427 + ".space."
1431- + _memory_space_to_mbar_suffix ( space )
1428+ + space_str
14321429 )
14331430
14341431
1432+ def require_mbarrier_ordering (
1433+ ordering_var : Var ,
1434+ valid_orderings : tuple [MemoryOrder , ...],
1435+ ) -> MemoryOrder :
1436+ ordering = require_constant_enum (ordering_var , MemoryOrder )
1437+ if ordering not in valid_orderings :
1438+ formatted = ", " .join (str (o ) for o in valid_orderings )
1439+ raise TileTypeError (
1440+ f"Invalid mbarrier memory order { ordering } , expected one of { formatted } "
1441+ )
1442+ return ordering
1443+
1444+
1445+ ARRIVE_ORDERINGS = (MemoryOrder .RELEASE , MemoryOrder .RELAXED )
1446+ WAIT_ORDERINGS = (MemoryOrder .ACQUIRE , MemoryOrder .RELAXED )
1447+
1448+
14351449@impl (stub .mbarrier_arrive )
14361450def mbarrier_arrive_impl (
14371451 mbar : Var ,
14381452 count : Var ,
14391453 drop : Var ,
14401454 scope : Var ,
1441- relaxed : Var ,
1455+ ordering : Var ,
14421456) -> Var | None :
14431457 count = astype (count , datatype .int32 )
14441458 drop = require_constant_bool (drop )
1445- scope = require_constant_enum (scope , MemorySpace )
1446- relaxed = require_constant_bool ( relaxed )
1459+ scope = require_constant_enum (scope , MbarrierScope )
1460+ ordering = require_mbarrier_ordering ( ordering , ARRIVE_ORDERINGS )
14471461 space = require_mbarrier_ptr (mbar ).memory_space
14481462 intrinsic = "llvm.nvvm.mbarrier.arrive"
14491463 if drop :
14501464 intrinsic += '.drop'
1451- if relaxed :
1465+ if ordering is MemoryOrder . RELAXED :
14521466 intrinsic += '.relaxed'
14531467 intrinsic += _mbar_space_scope_suffix (scope , space )
14541468
@@ -1468,18 +1482,18 @@ def mbarrier_arrive_expect_tx_impl(
14681482 bytes : Var ,
14691483 drop : Var ,
14701484 scope : Var ,
1471- relaxed : Var ,
1485+ ordering : Var ,
14721486) -> Var | None :
14731487 bytes = astype (bytes , datatype .int32 )
14741488 drop = require_constant_bool (drop )
1475- scope = require_constant_enum (scope , MemorySpace )
1476- relaxed = require_constant_bool ( relaxed )
1489+ scope = require_constant_enum (scope , MbarrierScope )
1490+ ordering = require_mbarrier_ordering ( ordering , ARRIVE_ORDERINGS )
14771491 space = require_mbarrier_ptr (mbar ).memory_space
14781492 intrinsic = "llvm.nvvm.mbarrier.arrive"
14791493 if drop :
14801494 intrinsic += '.drop'
14811495 intrinsic += '.expect.tx'
1482- if relaxed :
1496+ if ordering is MemoryOrder . RELAXED :
14831497 intrinsic += '.relaxed'
14841498 intrinsic += _mbar_space_scope_suffix (scope , space )
14851499
@@ -1497,7 +1511,7 @@ def mbarrier_arrive_expect_tx_impl(
14971511def mbarrier_expect_tx_impl (mbar : Var , bytes : Var , scope : Var ) -> Var :
14981512 space = require_mbarrier_ptr (mbar ).memory_space
14991513 bytes = astype (bytes , datatype .int32 )
1500- scope = require_constant_enum (scope , MemorySpace )
1514+ scope = require_constant_enum (scope , MbarrierScope )
15011515 intrinsic = "llvm.nvvm.mbarrier.expect.tx"
15021516 intrinsic += _mbar_space_scope_suffix (scope , space )
15031517 add_operation (
@@ -1512,7 +1526,7 @@ def mbarrier_expect_tx_impl(mbar: Var, bytes: Var, scope: Var) -> Var:
15121526def mbarrier_complete_tx_impl (mbar : Var , bytes : Var , scope : Var ) -> Var :
15131527 space = require_mbarrier_ptr (mbar ).memory_space
15141528 bytes = astype (bytes , datatype .int32 )
1515- scope = require_constant_enum (scope , MemorySpace )
1529+ scope = require_constant_enum (scope , MbarrierScope )
15161530 intrinsic = "llvm.nvvm.mbarrier.complete.tx"
15171531 intrinsic += _mbar_space_scope_suffix (scope , space )
15181532 add_operation (
@@ -1525,14 +1539,14 @@ def mbarrier_complete_tx_impl(mbar: Var, bytes: Var, scope: Var) -> Var:
15251539
15261540@impl (stub .mbarrier_test_wait )
15271541def mbarrier_test_wait_impl (
1528- mbar : Var , state : Var , scope : Var , relaxed : Var
1542+ mbar : Var , state : Var , scope : Var , ordering : Var
15291543) -> Var :
1530- scope = require_constant_enum (scope , MemorySpace )
1544+ scope = require_constant_enum (scope , MbarrierScope )
15311545 state = astype (state , datatype .int64 )
15321546 require_mbarrier_ptr (mbar , (MemorySpace .SHARED ,))
1533- relaxed = require_constant_bool ( relaxed )
1547+ ordering = require_mbarrier_ordering ( ordering , WAIT_ORDERINGS )
15341548 intrinsic = "llvm.nvvm.mbarrier.test.wait"
1535- if relaxed :
1549+ if ordering is MemoryOrder . RELAXED :
15361550 intrinsic += ".relaxed"
15371551 intrinsic += _mbar_space_scope_suffix (scope , MemorySpace .SHARED )
15381552 results = add_operation (
@@ -1546,14 +1560,14 @@ def mbarrier_test_wait_impl(
15461560
15471561@impl (stub .mbarrier_test_wait_parity )
15481562def mbarrier_test_wait_parity_impl (
1549- mbar : Var , parity : Var , scope : Var , relaxed : Var
1563+ mbar : Var , parity : Var , scope : Var , ordering : Var
15501564) -> Var :
15511565 require_mbarrier_ptr (mbar , (MemorySpace .SHARED ,))
15521566 parity = astype (parity , datatype .int32 )
1553- scope = require_constant_enum (scope , MemorySpace )
1554- relaxed = require_constant_bool ( relaxed )
1567+ scope = require_constant_enum (scope , MbarrierScope )
1568+ ordering = require_mbarrier_ordering ( ordering , WAIT_ORDERINGS )
15551569 intrinsic = "llvm.nvvm.mbarrier.test.wait.parity"
1556- if relaxed :
1570+ if ordering is MemoryOrder . RELAXED :
15571571 intrinsic += ".relaxed"
15581572 intrinsic += _mbar_space_scope_suffix (scope , MemorySpace .SHARED )
15591573 results = add_operation (
@@ -1575,19 +1589,19 @@ def mbarrier_try_wait_impl(
15751589 state : Var ,
15761590 time_hint : Var ,
15771591 scope : Var ,
1578- relaxed : Var ,
1592+ ordering : Var ,
15791593) -> Var :
15801594 require_mbarrier_ptr (mbar , (MemorySpace .SHARED ,))
15811595 state = astype (state , datatype .int64 )
1582- scope = require_constant_enum (scope , MemorySpace )
1583- relaxed = require_constant_bool ( relaxed )
1596+ scope = require_constant_enum (scope , MbarrierScope )
1597+ ordering = require_mbarrier_ordering ( ordering , WAIT_ORDERINGS )
15841598 intrinsic = "llvm.nvvm.mbarrier.try.wait"
15851599 args = (mbar , state )
15861600 if not _is_none (time_hint ):
15871601 intrinsic += ".tl"
15881602 time_hint = astype (time_hint , datatype .int32 )
15891603 args = (* args , time_hint )
1590- if relaxed :
1604+ if ordering is MemoryOrder . RELAXED :
15911605 intrinsic += ".relaxed"
15921606 intrinsic += _mbar_space_scope_suffix (scope , MemorySpace .SHARED )
15931607 results = add_operation (
@@ -1605,19 +1619,19 @@ def mbarrier_try_wait_parity_impl(
16051619 parity : Var ,
16061620 time_hint : Var ,
16071621 scope : Var ,
1608- relaxed : Var ,
1622+ ordering : Var ,
16091623) -> Var :
16101624 require_mbarrier_ptr (mbar , (MemorySpace .SHARED ,))
16111625 parity = astype (parity , datatype .int32 )
1612- scope = require_constant_enum (scope , MemorySpace )
1613- relaxed = require_constant_bool ( relaxed )
1626+ scope = require_constant_enum (scope , MbarrierScope )
1627+ ordering = require_mbarrier_ordering ( ordering , WAIT_ORDERINGS )
16141628 intrinsic = "llvm.nvvm.mbarrier.try.wait.parity"
16151629 args = (mbar , parity )
16161630 if not _is_none (time_hint ):
16171631 time_hint = astype (time_hint , datatype .int32 )
16181632 args = (* args , time_hint )
16191633 intrinsic += ".tl"
1620- if relaxed :
1634+ if ordering is MemoryOrder . RELAXED :
16211635 intrinsic += ".relaxed"
16221636 intrinsic += _mbar_space_scope_suffix (scope , MemorySpace .SHARED )
16231637 results = add_operation (
0 commit comments