Skip to content

Commit 454a5ce

Browse files
[lang] Use more narrow enums for mbarrier api
Signed-off-by: Asher Mancinelli <amancinelli@nvidia.com>
1 parent 88231c8 commit 454a5ce

6 files changed

Lines changed: 114 additions & 84 deletions

File tree

experimental/cuda-lang/src/cuda/lang/__init__.py

Lines changed: 2 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -70,6 +70,7 @@
7070
TensorMap,
7171
tensor_map_tiled,
7272
nanosleep,
73+
MbarrierScope,
7374
mbarrier_init,
7475
mbarrier_invalidate,
7576
mbarrier_arrive,
@@ -202,6 +203,7 @@
202203
"tensor_map_tiled",
203204
"nanosleep",
204205
"mbarrier",
206+
"MbarrierScope",
205207
"mbarrier_init",
206208
"mbarrier_invalidate",
207209
"mbarrier_arrive",

experimental/cuda-lang/src/cuda/lang/_enums.py

Lines changed: 7 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -14,3 +14,10 @@ class TensorMapSwizzle(enum.Enum):
1414
SWIZZLE_128B_ATOM_32B = _cext.CU_TENSOR_MAP_SWIZZLE_128B_ATOM_32B
1515
SWIZZLE_128B_ATOM_32B_FLIP_8B = _cext.CU_TENSOR_MAP_SWIZZLE_128B_ATOM_32B_FLIP_8B
1616
SWIZZLE_128B_ATOM_64B = _cext.CU_TENSOR_MAP_SWIZZLE_128B_ATOM_64B
17+
18+
19+
class MbarrierScope(enum.Enum):
20+
"""Scope of the threads that observe an mbarrier operation."""
21+
22+
BLOCK = "cta"
23+
CLUSTER = "cluster"

experimental/cuda-lang/src/cuda/lang/_ir/ops.py

Lines changed: 52 additions & 38 deletions
Original file line numberDiff line numberDiff line change
@@ -95,7 +95,7 @@
9595
format_var,
9696
LocalArrayContextManagerValue,
9797
)
98-
from .._stub import TensorMapSwizzle
98+
from .._stub import TensorMapSwizzle, MbarrierScope
9999
from cuda.tile._ir import hir_stubs
100100
from 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)
14361450
def 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(
14971511
def 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:
15121526
def 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)
15271541
def 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)
15481562
def 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(

experimental/cuda-lang/src/cuda/lang/_stub/__init__.py

Lines changed: 3 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -66,6 +66,8 @@
6666
tensor_map_tiled,
6767
)
6868

69+
from cuda.lang._enums import MbarrierScope
70+
6971
from .mbarrier import (
7072
mbarrier_init,
7173
mbarrier_invalidate,
@@ -138,6 +140,7 @@
138140
"TensorMap",
139141
"tensor_map_tiled",
140142
"nanosleep",
143+
"MbarrierScope",
141144
"mbarrier_init",
142145
"mbarrier_invalidate",
143146
"mbarrier_arrive",

experimental/cuda-lang/src/cuda/lang/_stub/mbarrier.py

Lines changed: 25 additions & 21 deletions
Original file line numberDiff line numberDiff line change
@@ -2,9 +2,16 @@
22
#
33
# SPDX-License-Identifier: Apache-2.0
44

5-
from cuda.lang._datatype import MemorySpace
5+
from typing import Literal
6+
7+
from cuda.lang._enums import MbarrierScope
68
from cuda.lang._execution import stub
79
from cuda.lang._datatype import uint64, bool_
10+
from cuda.tile._memory_model import MemoryOrder
11+
12+
13+
ArriveOrdering = Literal[MemoryOrder.RELAXED, MemoryOrder.RELEASE]
14+
WaitOrdering = Literal[MemoryOrder.RELAXED, MemoryOrder.ACQUIRE]
815

916

1017
@stub
@@ -23,8 +30,8 @@ def mbarrier_arrive(
2330
count: int = 1,
2431
*,
2532
drop: bool = False,
26-
scope: MemorySpace = MemorySpace.SHARED,
27-
relaxed: bool = False,
33+
scope: MbarrierScope = MbarrierScope.BLOCK,
34+
ordering: ArriveOrdering = MemoryOrder.RELEASE,
2835
) -> "uint64 | None":
2936
"""Arrive at ``mbar``. When the mbarrier resides in ``MemorySpace.SHARED``,
3037
an opaque 64-bit value capturing the phase of the mbarrier object _prior_
@@ -39,19 +46,18 @@ def mbarrier_arrive_expect_tx(
3946
bytes: int,
4047
*,
4148
drop: bool = False,
42-
scope: MemorySpace = MemorySpace.SHARED,
43-
relaxed: bool = False,
49+
scope: MbarrierScope = MbarrierScope.BLOCK,
50+
ordering: ArriveOrdering = MemoryOrder.RELEASE,
4451
) -> "uint64 | None":
45-
"""Arrive at ``mbar`` and set the expected transaction count to ``bytes``.
46-
"""
52+
...
4753

4854

4955
@stub
5056
def mbarrier_expect_tx(
5157
mbar,
5258
bytes: int,
5359
*,
54-
scope: MemorySpace = MemorySpace.SHARED,
60+
scope: MbarrierScope = MbarrierScope.BLOCK,
5561
) -> None:
5662
...
5763

@@ -61,7 +67,7 @@ def mbarrier_complete_tx(
6167
mbar,
6268
bytes: int,
6369
*,
64-
scope: MemorySpace = MemorySpace.SHARED,
70+
scope: MbarrierScope = MbarrierScope.BLOCK,
6571
) -> None:
6672
...
6773

@@ -71,20 +77,19 @@ def mbarrier_test_wait(
7177
mbar,
7278
state,
7379
*,
74-
scope: MemorySpace = MemorySpace.SHARED,
75-
relaxed: bool = False,
80+
scope: MbarrierScope = MbarrierScope.BLOCK,
81+
ordering: WaitOrdering = MemoryOrder.ACQUIRE,
7682
) -> "bool_":
77-
"""Non-blocking test whether ``mbar`` has completed.
78-
"""
83+
"""Non-blocking test whether ``mbar`` has completed."""
7984

8085

8186
@stub
8287
def mbarrier_test_wait_parity(
8388
mbar,
8489
parity: int,
8590
*,
86-
scope: MemorySpace = MemorySpace.SHARED,
87-
relaxed: bool = False,
91+
scope: MbarrierScope = MbarrierScope.BLOCK,
92+
ordering: WaitOrdering = MemoryOrder.ACQUIRE,
8893
) -> "bool_":
8994
"""Phase-parity variant of ``mbarrier_test_wait``.
9095
``parity`` is the 0/1 integer parity of the phase to test for.
@@ -97,11 +102,10 @@ def mbarrier_try_wait(
97102
state,
98103
*,
99104
time_hint: int | None = None,
100-
scope: MemorySpace = MemorySpace.SHARED,
101-
relaxed: bool = False,
105+
scope: MbarrierScope = MbarrierScope.BLOCK,
106+
ordering: WaitOrdering = MemoryOrder.ACQUIRE,
102107
) -> "bool_":
103-
"""Bounded-wait test whether ``mbar`` has completed.
104-
"""
108+
"""Bounded-wait test whether ``mbar`` has completed."""
105109

106110

107111
@stub
@@ -110,8 +114,8 @@ def mbarrier_try_wait_parity(
110114
parity: int,
111115
*,
112116
time_hint: int | None = None,
113-
scope: MemorySpace = MemorySpace.SHARED,
114-
relaxed: bool = False,
117+
scope: MbarrierScope = MbarrierScope.BLOCK,
118+
ordering: WaitOrdering = MemoryOrder.ACQUIRE,
115119
) -> "bool_":
116120
"""Phase-parity variant of ``mbarrier_try_wait``.
117121
``parity`` is the 0/1 integer parity of the phase to test for.

0 commit comments

Comments
 (0)