Skip to content

Commit 1c6b008

Browse files
committed
Rename ArrayConstraint.stride_static to stride_constant
Signed-off-by: Greg Bonik <gbonik@nvidia.com>
1 parent 058b912 commit 1c6b008

6 files changed

Lines changed: 27 additions & 27 deletions

File tree

cext/tile_kernel.cpp

Lines changed: 4 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -699,8 +699,8 @@ static PyPtr parse_array_constraint(ConstantCursor& cursor) {
699699
PyPtr dtype = dtype_to_python(arrty.dtype);
700700
if (!dtype) return {};
701701

702-
PyPtr static_strides = steal(PyTuple_New(arrty.ndim));
703-
if (!static_strides) return {};
702+
PyPtr constant_strides = steal(PyTuple_New(arrty.ndim));
703+
if (!constant_strides) return {};
704704

705705
PyPtr stride_divisible_by = steal(PyTuple_New(arrty.ndim));
706706
if (!stride_divisible_by) return {};
@@ -726,7 +726,7 @@ static PyPtr parse_array_constraint(ConstantCursor& cursor) {
726726

727727
for (size_t i = 0; i < arrty.ndim; ++i) {
728728
PyObject* obj = special_bits.is_stride_one(i) ? one.get() : Py_None;
729-
PyTuple_SET_ITEM(static_strides.get(), i, Py_NewRef(obj));
729+
PyTuple_SET_ITEM(constant_strides.get(), i, Py_NewRef(obj));
730730

731731
obj = special_bits.is_stride_16byte_divisible(i) ? stride_divisor.get() : one.get();
732732
PyTuple_SET_ITEM(stride_divisible_by.get(), i, Py_NewRef(obj));
@@ -739,7 +739,7 @@ static PyPtr parse_array_constraint(ConstantCursor& cursor) {
739739
"{sO sI sO sO s() sO sO sO sO}",
740740
"dtype", dtype.get(),
741741
"ndim", static_cast<unsigned>(arrty.ndim),
742-
"stride_static", static_strides.get(),
742+
"stride_constant", constant_strides.get(),
743743
"stride_lower_bound_incl", zero.get(),
744744
"alias_groups",
745745
"may_alias_internally", special_bits.disjoint_elements ? Py_False : Py_True,

src/cuda/tile/_compile.py

Lines changed: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -151,7 +151,7 @@ def _create_kernel_parameters(parameter_constraints: Sequence[ParameterConstrain
151151

152152

153153
def _get_array_ty(param: ArrayConstraint):
154-
for static_stride, bound in zip(param.stride_static, param.stride_lower_bound_incl,
154+
for static_stride, bound in zip(param.stride_constant, param.stride_lower_bound_incl,
155155
strict=True):
156156
if static_stride is not None:
157157
continue
@@ -161,7 +161,7 @@ def _get_array_ty(param: ArrayConstraint):
161161

162162
return ArrayTy(param.dtype,
163163
shape=(None,) * param.ndim,
164-
strides=param.stride_static,
164+
strides=param.stride_constant,
165165
elements_disjoint=not param.may_alias_internally,
166166
base_ptr_div_by=param.base_addr_divisible_by,
167167
stride_div_by=param.stride_divisible_by,

src/cuda/tile/compilation/_name_mangling.py

Lines changed: 9 additions & 9 deletions
Original file line numberDiff line numberDiff line change
@@ -188,7 +188,7 @@ def _mangle_array_constraint(a: ArrayConstraint,
188188

189189
axis_predicates = OrderedDict()
190190
_collect_axis_predicate(a.shape_divisible_by, "i", 1, axis_predicates)
191-
_collect_axis_predicate(a.stride_static, "t", None, axis_predicates)
191+
_collect_axis_predicate(a.stride_constant, "t", None, axis_predicates)
192192
_collect_axis_predicate(a.stride_divisible_by, "v", 1, axis_predicates)
193193
_collect_axis_predicate(a.stride_lower_bound_incl, "l", None, axis_predicates)
194194

@@ -221,7 +221,7 @@ def _demangle_array_constraint(cursor: _Cursor,
221221

222222
# Read axis predicates
223223
shape_divisible_by = [1] * ndim
224-
stride_static = [None] * ndim
224+
stride_constant = [None] * ndim
225225
stride_divisible_by = [1] * ndim
226226
stride_lower_bound_incl = [None] * ndim
227227
while True:
@@ -241,9 +241,9 @@ def _demangle_array_constraint(cursor: _Cursor,
241241
if cursor.read("i") is not None:
242242
axis_shape_div_by = _demangle_divisibility(cursor)
243243

244-
axis_stride_static = None
244+
axis_stride_constant = None
245245
if cursor.read("t") is not None:
246-
axis_stride_static = _demangle_signed_int(cursor)
246+
axis_stride_constant = _demangle_signed_int(cursor)
247247

248248
axis_stride_div_by = 1
249249
if cursor.read("v") is not None:
@@ -262,11 +262,11 @@ def _demangle_array_constraint(cursor: _Cursor,
262262
f"Shape divisibility specified more than once for axis #{i}")
263263
shape_divisible_by[i] = axis_shape_div_by
264264

265-
if axis_stride_static is not None:
266-
if stride_static[i] is not None:
265+
if axis_stride_constant is not None:
266+
if stride_constant[i] is not None:
267267
raise mask_cursor.make_error(
268268
f"Static stride specified more than once for axis #{i}")
269-
stride_static[i] = axis_stride_static
269+
stride_constant[i] = axis_stride_constant
270270

271271
if axis_stride_div_by != 1:
272272
if stride_divisible_by[i] != 1:
@@ -283,7 +283,7 @@ def _demangle_array_constraint(cursor: _Cursor,
283283
axis_mask &= ~(1 << i)
284284

285285
for i in range(ndim):
286-
if stride_static[i] is not None:
286+
if stride_constant[i] is not None:
287287
if stride_divisible_by[i] != 1:
288288
raise orig_cursor.make_error(f"Stride divisibility specified together"
289289
f" with static stride for axis {i}")
@@ -309,7 +309,7 @@ def _demangle_array_constraint(cursor: _Cursor,
309309
stride_lower_bound_incl=stride_lower_bound_incl,
310310
alias_groups=alias_groups,
311311
may_alias_internally=may_alias_internally,
312-
stride_static=stride_static,
312+
stride_constant=stride_constant,
313313
stride_divisible_by=stride_divisible_by,
314314
shape_divisible_by=shape_divisible_by,
315315
base_addr_divisible_by=base_addr_div_by)

src/cuda/tile/compilation/_signature.py

Lines changed: 9 additions & 9 deletions
Original file line numberDiff line numberDiff line change
@@ -55,8 +55,8 @@ class ArrayConstraint:
5555
array has a zero stride. For most arrays produced by major tensor libraries,
5656
this can be assumed to be false. Setting this to True may disable certain
5757
optimizations of loads and stores to/from this array.
58-
stride_static (Sequence[int | None] | None):
59-
For each dimension of the array, an optional statically known value of its stride.
58+
stride_constant (Sequence[int | None] | None):
59+
For each dimension of the array, an optional constant value of its stride.
6060
For example, if the array is known to have a C-contiguous layout, the stride of
6161
the last dimension can be set to 1, which may enable certain optimizations of loads
6262
and stores from/to this array. Can be set to `None` if none of the dimensions
@@ -84,7 +84,7 @@ class ArrayConstraint:
8484
stride_lower_bound_incl: tuple[int | None, ...]
8585
alias_groups: tuple[str, ...]
8686
may_alias_internally: bool
87-
stride_static: tuple[int | None, ...]
87+
stride_constant: tuple[int | None, ...]
8888
stride_divisible_by: tuple[int, ...]
8989
shape_divisible_by: tuple[int, ...]
9090
base_addr_divisible_by: int
@@ -96,7 +96,7 @@ def __init__(self,
9696
stride_lower_bound_incl: Sequence[int | None] | int | None,
9797
alias_groups: Sequence[str],
9898
may_alias_internally: bool,
99-
stride_static: Sequence[int | None] | None = None,
99+
stride_constant: Sequence[int | None] | None = None,
100100
stride_divisible_by: Sequence[int] | int = 1,
101101
shape_divisible_by: Sequence[int] | int = 1,
102102
base_addr_divisible_by: int = 1):
@@ -110,21 +110,21 @@ def __init__(self,
110110
if ndim < 0:
111111
raise ValueError("`ndim` cannot be negative")
112112

113-
# stride_static
114-
stride_static = _parse_assumption_tuple(
115-
stride_static, ndim, "stride_static", None, _check_optional_int)
113+
# stride_constant
114+
stride_constant = _parse_assumption_tuple(
115+
stride_constant, ndim, "stride_constant", None, _check_optional_int)
116116

117117
# stride_lower_bound
118118
stride_lower_bound_incl = _parse_assumption_tuple(
119119
stride_lower_bound_incl, ndim, "stride_lower_bound_incl", None, _check_optional_int)
120120
stride_lower_bound_incl = _remove_redundant_lower_bounds(
121-
stride_static, stride_lower_bound_incl, "stride_lower_bound_incl")
121+
stride_constant, stride_lower_bound_incl, "stride_lower_bound_incl")
122122

123123
# stride_divisible_by
124124
stride_divisible_by = _parse_assumption_tuple(
125125
stride_divisible_by, ndim, "stride_divisible_by", 1, _check_divisibility)
126126
stride_divisible_by = _remove_redundant_divisibility_constraints(
127-
stride_static, stride_divisible_by, "stride_static")
127+
stride_constant, stride_divisible_by, "stride_constant")
128128

129129
# shape_divisible_by
130130
shape_divisible_by = _parse_assumption_tuple(

test/test_export_compat.py

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -105,7 +105,7 @@ def test_export_compat_cutile_python_v1():
105105
ct.compilation.ArrayConstraint(ct.int32, 2, stride_lower_bound_incl=0,
106106
alias_groups=(), may_alias_internally=False,
107107
stride_divisible_by=(4, 1),
108-
stride_static=(None, 1)),
108+
stride_constant=(None, 1)),
109109
ct.compilation.ArrayConstraint(ct.float32, 3, stride_lower_bound_incl=0,
110110
alias_groups=(), may_alias_internally=False),
111111
],

test/test_name_mangling.py

Lines changed: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -47,15 +47,15 @@
4747
id="array_simple",
4848
),
4949
50-
# 3D array with stride_static, stride_divisible_by, shape_divisible_by
50+
# 3D array with stride_constant, stride_divisible_by, shape_divisible_by
5151
# (dims 0 and 1 share shape_divisible_by=16), stride_lower_bound_incl,
5252
# and base_addr_divisible_by
5353
pytest.param(
5454
[ArrayConstraint(float32, 3,
5555
stride_lower_bound_incl=0,
5656
alias_groups=(),
5757
may_alias_internally=False,
58-
stride_static=[None, None, 1],
58+
stride_constant=[None, None, 1],
5959
stride_divisible_by=[8, 1, 1],
6060
shape_divisible_by=[16, 16, 1],
6161
base_addr_divisible_by=16)],

0 commit comments

Comments
 (0)