Skip to content

Commit bfff83e

Browse files
committed
Complete ordering to memory_order rename in pointer load/store
Signed-off-by: Qiqi Xiao <qiqix@nvidia.com>
1 parent 6c1ecd3 commit bfff83e

8 files changed

Lines changed: 73 additions & 73 deletions

File tree

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

Lines changed: 4 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -71,9 +71,9 @@ class StorePointer(Operation, opcode="store_pointer", memory_effect=MemoryEffect
7171
value: Var = operand()
7272
alignment: Optional[int] = attribute()
7373
volatile: bool = attribute(default=False)
74-
ordering: Optional[MemoryOrder] = attribute(default=None)
74+
memory_order: Optional[MemoryOrder] = attribute(default=None)
7575

76-
valid_orderings = (
76+
valid_memory_orders = (
7777
None,
7878
MemoryOrder.WEAK,
7979
MemoryOrder.RELAXED,
@@ -86,9 +86,9 @@ class LoadPointer(Operation, opcode="load_pointer", memory_effect=MemoryEffect.L
8686
pointer: Var = operand()
8787
alignment: Optional[int] = attribute()
8888
volatile: bool = attribute(default=False)
89-
ordering: Optional[MemoryOrder] = attribute(default=None)
89+
memory_order: Optional[MemoryOrder] = attribute(default=None)
9090

91-
valid_orderings = (
91+
valid_memory_orders = (
9292
None,
9393
MemoryOrder.WEAK,
9494
MemoryOrder.RELAXED,

experimental/cuda-lang/src/cuda/lang/_ir/op_impl/mbarrier_impl.py

Lines changed: 16 additions & 16 deletions
Original file line numberDiff line numberDiff line change
@@ -61,21 +61,21 @@ def _mbar_space_scope_suffix(scope: MbarrierScope, space: MemorySpace) -> str:
6161
return ".scope." + scope.value + ".space." + space_str
6262

6363

64-
def require_mbarrier_ordering(
65-
ordering_var: Var,
66-
valid_orderings: tuple[MemoryOrder, ...],
64+
def require_mbarrier_memory_order(
65+
memory_order_var: Var,
66+
valid_memory_orders: tuple[MemoryOrder, ...],
6767
) -> MemoryOrder:
68-
ordering = require_constant_enum(ordering_var, MemoryOrder)
69-
if ordering not in valid_orderings:
70-
formatted = ", ".join(str(o) for o in valid_orderings)
68+
memory_order = require_constant_enum(memory_order_var, MemoryOrder)
69+
if memory_order not in valid_memory_orders:
70+
formatted = ", ".join(str(o) for o in valid_memory_orders)
7171
raise TypeCheckingError(
72-
f"Invalid mbarrier memory order {ordering}, expected one of {formatted}"
72+
f"Invalid mbarrier memory order {memory_order}, expected one of {formatted}"
7373
)
74-
return ordering
74+
return memory_order
7575

7676

77-
ARRIVE_ORDERINGS = (MemoryOrder.RELEASE, MemoryOrder.RELAXED)
78-
WAIT_ORDERINGS = (MemoryOrder.ACQUIRE, MemoryOrder.RELAXED)
77+
ARRIVE_MEMORY_ORDERS = (MemoryOrder.RELEASE, MemoryOrder.RELAXED)
78+
WAIT_MEMORY_ORDERS = (MemoryOrder.ACQUIRE, MemoryOrder.RELAXED)
7979

8080

8181
@impl(mbarrier.mbarrier_arrive)
@@ -89,7 +89,7 @@ def mbarrier_arrive_impl(
8989
count = astype(count, datatype.int32)
9090
drop = require_constant_bool(drop)
9191
scope = require_constant_enum(scope, MbarrierScope)
92-
memory_order = require_mbarrier_ordering(memory_order, ARRIVE_ORDERINGS)
92+
memory_order = require_mbarrier_memory_order(memory_order, ARRIVE_MEMORY_ORDERS)
9393
space = require_mbarrier_ptr(mbar).memory_space
9494
intrinsic = "llvm.nvvm.mbarrier.arrive"
9595
if drop:
@@ -119,7 +119,7 @@ def mbarrier_arrive_expect_transaction_impl(
119119
bytes = astype(bytes, datatype.int32)
120120
drop = require_constant_bool(drop)
121121
scope = require_constant_enum(scope, MbarrierScope)
122-
memory_order = require_mbarrier_ordering(memory_order, ARRIVE_ORDERINGS)
122+
memory_order = require_mbarrier_memory_order(memory_order, ARRIVE_MEMORY_ORDERS)
123123
space = require_mbarrier_ptr(mbar).memory_space
124124
intrinsic = "llvm.nvvm.mbarrier.arrive"
125125
if drop:
@@ -176,7 +176,7 @@ def mbarrier_test_wait_impl(
176176
scope = require_constant_enum(scope, MbarrierScope)
177177
state = astype(state, datatype.int64)
178178
require_mbarrier_ptr(mbar, (MemorySpace.SHARED,))
179-
memory_order = require_mbarrier_ordering(memory_order, WAIT_ORDERINGS)
179+
memory_order = require_mbarrier_memory_order(memory_order, WAIT_MEMORY_ORDERS)
180180
intrinsic = "llvm.nvvm.mbarrier.test.wait"
181181
if memory_order is MemoryOrder.RELAXED:
182182
intrinsic += ".relaxed"
@@ -196,7 +196,7 @@ def mbarrier_test_wait_parity_impl(
196196
require_mbarrier_ptr(mbar, (MemorySpace.SHARED,))
197197
parity = astype(parity, datatype.int32)
198198
scope = require_constant_enum(scope, MbarrierScope)
199-
memory_order = require_mbarrier_ordering(memory_order, WAIT_ORDERINGS)
199+
memory_order = require_mbarrier_memory_order(memory_order, WAIT_MEMORY_ORDERS)
200200
intrinsic = "llvm.nvvm.mbarrier.test.wait.parity"
201201
if memory_order is MemoryOrder.RELAXED:
202202
intrinsic += ".relaxed"
@@ -220,7 +220,7 @@ def mbarrier_try_wait_impl(
220220
require_mbarrier_ptr(mbar, (MemorySpace.SHARED,))
221221
state = astype(state, datatype.int64)
222222
scope = require_constant_enum(scope, MbarrierScope)
223-
memory_order = require_mbarrier_ordering(memory_order, WAIT_ORDERINGS)
223+
memory_order = require_mbarrier_memory_order(memory_order, WAIT_MEMORY_ORDERS)
224224
intrinsic = "llvm.nvvm.mbarrier.try.wait"
225225
args = (mbar, state)
226226
if not is_none(time_hint):
@@ -249,7 +249,7 @@ def mbarrier_try_wait_parity_impl(
249249
require_mbarrier_ptr(mbar, (MemorySpace.SHARED,))
250250
parity = astype(parity, datatype.int32)
251251
scope = require_constant_enum(scope, MbarrierScope)
252-
memory_order = require_mbarrier_ordering(memory_order, WAIT_ORDERINGS)
252+
memory_order = require_mbarrier_memory_order(memory_order, WAIT_MEMORY_ORDERS)
253253
intrinsic = "llvm.nvvm.mbarrier.try.wait.parity"
254254
args = (mbar, parity)
255255
if not is_none(time_hint):

experimental/cuda-lang/src/cuda/lang/_ir/op_impl/pointer_impl.py

Lines changed: 11 additions & 11 deletions
Original file line numberDiff line numberDiff line change
@@ -210,7 +210,7 @@ def pointer_getitem(object: Var[PointerTy], key: Var[Type]):
210210
pointer=pointer,
211211
volatile=False,
212212
alignment=None,
213-
ordering=None,
213+
memory_order=None,
214214
)
215215

216216

@@ -226,7 +226,7 @@ def pointer_setitem(object: Var[PointerTy], key: Var[Type], value: Var[Type]):
226226
value=value,
227227
alignment=None,
228228
volatile=False,
229-
ordering=None,
229+
memory_order=None,
230230
)
231231

232232

@@ -258,7 +258,7 @@ def array_setitem(object: Var, key: Var, value: Var):
258258
value=value,
259259
alignment=None,
260260
volatile=False,
261-
ordering=None,
261+
memory_order=None,
262262
)
263263

264264

@@ -277,14 +277,14 @@ def pointer_load(
277277
count: Var,
278278
alignment: Var,
279279
volatile: Var,
280-
ordering: Var,
280+
memory_order: Var,
281281
) -> Var:
282282
pointee_dtype = require_pointer_type(pointer).pointee_dtype
283283
count = require_optional_constant_int(count)
284284
volatile = require_constant_bool(volatile)
285285
alignment = require_optional_alignment(alignment)
286-
ordering = require_pointer_memory_order(LoadPointer, ordering)
287-
is_atomic = ordering not in (None, MemoryOrder.WEAK)
286+
memory_order = require_pointer_memory_order(LoadPointer, memory_order)
287+
is_atomic = memory_order not in (None, MemoryOrder.WEAK)
288288
if count is None or count == 1:
289289
if is_atomic:
290290
_require_atomic_scalar_dtype(pointee_dtype, "load")
@@ -303,7 +303,7 @@ def pointer_load(
303303
pointer=pointer,
304304
volatile=volatile,
305305
alignment=alignment,
306-
ordering=ordering,
306+
memory_order=memory_order,
307307
)
308308

309309

@@ -312,18 +312,18 @@ def pointer_store(
312312
value: Var,
313313
alignment: Var,
314314
volatile: Var,
315-
ordering: Var,
315+
memory_order: Var,
316316
) -> None:
317317
pointer_ty = require_pointer_type(pointer)
318318
volatile = require_constant_bool(volatile)
319319
alignment = require_optional_alignment(alignment)
320-
ordering = require_pointer_memory_order(StorePointer, ordering)
320+
memory_order = require_pointer_memory_order(StorePointer, memory_order)
321321

322322
pointee_dtype = pointer_ty.pointee_dtype
323323
value = implicit_cast(value, pointee_dtype,
324324
"Stored value type is incompatible with pointer type")
325325

326-
is_atomic = ordering not in (None, MemoryOrder.WEAK)
326+
is_atomic = memory_order not in (None, MemoryOrder.WEAK)
327327
if is_atomic:
328328
if isinstance(value.get_type(), VectorTy):
329329
raise TypeCheckingError(
@@ -340,7 +340,7 @@ def pointer_store(
340340
value=value,
341341
volatile=volatile,
342342
alignment=alignment,
343-
ordering=ordering,
343+
memory_order=memory_order,
344344
)
345345

346346

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

Lines changed: 6 additions & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -235,19 +235,19 @@ def require_optional_alignment(alignment: Var) -> int | None:
235235

236236
def require_pointer_memory_order(
237237
operation: type[LoadPointer] | type[StorePointer],
238-
ordering_var: Var,
238+
memory_order_var: Var,
239239
):
240-
ordering = require_optional_constant_enum(ordering_var, MemoryOrder)
241-
if ordering in operation.valid_orderings:
242-
return ordering
240+
memory_order = require_optional_constant_enum(memory_order_var, MemoryOrder)
241+
if memory_order in operation.valid_memory_orders:
242+
return memory_order
243243

244244
formatted_expected = ", ".join(
245-
"None" if order is None else str(order) for order in operation.valid_orderings
245+
"None" if order is None else str(order) for order in operation.valid_memory_orders
246246
)
247247
operation_name = "load" if operation is LoadPointer else "store"
248248
raise make_type_checking_error(
249249
f"Invalid memory order for Pointer.{operation_name}. "
250-
f"Got {ordering}, expected one of {formatted_expected}"
250+
f"Got {memory_order}, expected one of {formatted_expected}"
251251
)
252252

253253

experimental/cuda-lang/src/cuda/lang/_passes/ir2mlir/pass_definition.py

Lines changed: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -834,7 +834,7 @@ def lower_vector_getitem(self, operation: ops.VectorGetItem) -> Sequence[mlir.Va
834834
@lower_operation.register
835835
def lower_load_pointer(self, operation: ops.LoadPointer) -> Sequence[mlir.Value]:
836836
ptr_dtype = operation.pointer.get_type().pointer_dtype
837-
ordering = _get_llvm_memory_ordering(operation.ordering)
837+
ordering = _get_llvm_memory_ordering(operation.memory_order)
838838
info = PointerInfo(ptr_dtype)
839839
assert not info.opaque, f"Expected a typed pointer, got {ptr_dtype}"
840840
result_type = ir_type_to_mlir_type(operation.result_var.get_type())
@@ -851,7 +851,7 @@ def lower_load_pointer(self, operation: ops.LoadPointer) -> Sequence[mlir.Value]
851851
@lower_operation.register
852852
def lower_store_pointer(self, operation: ops.StorePointer) -> Sequence[mlir.Value]:
853853
pointer = self.get_var(operation.pointer)
854-
ordering = _get_llvm_memory_ordering(operation.ordering)
854+
ordering = _get_llvm_memory_ordering(operation.memory_order)
855855
value = self.get_var(operation.value)
856856
mlir.llvm.add_StoreOp(
857857
value=value,

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

Lines changed: 8 additions & 8 deletions
Original file line numberDiff line numberDiff line change
@@ -10,8 +10,8 @@
1010
from cuda.tile._memory_model import MemoryOrder
1111

1212

13-
ArriveOrdering = Literal[MemoryOrder.RELAXED, MemoryOrder.RELEASE]
14-
WaitOrdering = Literal[MemoryOrder.RELAXED, MemoryOrder.ACQUIRE]
13+
ArriveMemoryOrder = Literal[MemoryOrder.RELAXED, MemoryOrder.RELEASE]
14+
WaitMemoryOrder = Literal[MemoryOrder.RELAXED, MemoryOrder.ACQUIRE]
1515

1616

1717
@stub
@@ -42,7 +42,7 @@ def mbarrier_arrive(
4242
*,
4343
drop: bool = False,
4444
scope: MbarrierScope = MbarrierScope.BLOCK,
45-
memory_order: ArriveOrdering = MemoryOrder.RELEASE,
45+
memory_order: ArriveMemoryOrder = MemoryOrder.RELEASE,
4646
) -> "uint64 | None":
4747
"""Arrive at ``mbar``. When the mbarrier resides in ``MemorySpace.SHARED``,
4848
an opaque 64-bit value capturing the phase of the mbarrier object _prior_
@@ -69,7 +69,7 @@ def mbarrier_arrive_expect_transaction(
6969
*,
7070
drop: bool = False,
7171
scope: MbarrierScope = MbarrierScope.BLOCK,
72-
memory_order: ArriveOrdering = MemoryOrder.RELEASE,
72+
memory_order: ArriveMemoryOrder = MemoryOrder.RELEASE,
7373
) -> "uint64 | None":
7474
"""Arrive at ``mbar`` and add expected transaction bytes.
7575
@@ -127,7 +127,7 @@ def mbarrier_test_wait(
127127
state,
128128
*,
129129
scope: MbarrierScope = MbarrierScope.BLOCK,
130-
memory_order: WaitOrdering = MemoryOrder.ACQUIRE,
130+
memory_order: WaitMemoryOrder = MemoryOrder.ACQUIRE,
131131
) -> "bool_":
132132
"""Non-blocking test whether ``mbar`` has completed.
133133
@@ -148,7 +148,7 @@ def mbarrier_test_wait_parity(
148148
parity: int,
149149
*,
150150
scope: MbarrierScope = MbarrierScope.BLOCK,
151-
memory_order: WaitOrdering = MemoryOrder.ACQUIRE,
151+
memory_order: WaitMemoryOrder = MemoryOrder.ACQUIRE,
152152
) -> "bool_":
153153
"""Phase-parity variant of ``mbarrier_test_wait``.
154154
``parity`` is the 0/1 integer parity of the phase to test for.
@@ -171,7 +171,7 @@ def mbarrier_try_wait(
171171
*,
172172
time_hint: int | None = None,
173173
scope: MbarrierScope = MbarrierScope.BLOCK,
174-
memory_order: WaitOrdering = MemoryOrder.ACQUIRE,
174+
memory_order: WaitMemoryOrder = MemoryOrder.ACQUIRE,
175175
) -> "bool_":
176176
"""Bounded-wait test whether ``mbar`` has completed.
177177
@@ -202,7 +202,7 @@ def mbarrier_try_wait_parity(
202202
*,
203203
time_hint: int | None = None,
204204
scope: MbarrierScope = MbarrierScope.BLOCK,
205-
memory_order: WaitOrdering = MemoryOrder.ACQUIRE,
205+
memory_order: WaitMemoryOrder = MemoryOrder.ACQUIRE,
206206
) -> "bool_":
207207
"""Phase-parity variant of ``mbarrier_try_wait``.
208208
``parity`` is the 0/1 integer parity of the phase to test for.

0 commit comments

Comments
 (0)