Skip to content

Commit 9dee58a

Browse files
gbonikxiaoqiqi177
authored andcommitted
Rename TupleConstraint.elements to items
Signed-off-by: Greg Bonik <gbonik@nvidia.com>
1 parent 37fbea2 commit 9dee58a

5 files changed

Lines changed: 39 additions & 45 deletions

File tree

cext/tile_kernel.cpp

Lines changed: 6 additions & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -1691,12 +1691,12 @@ static PyPtr parse_param_constraint(ConstantCursor& cursor,
16911691
Cursor<RefPtr<LeafAnnotationNode>>* annotation_cursor) {
16921692
ParameterKind pk = param_cursor->next();
16931693
if (pk == ParameterKind::TupleBegin) {
1694-
PyPtr elements_list = steal(PyList_New(0));
1695-
if (!elements_list) return {};
1694+
PyPtr items_list = steal(PyList_New(0));
1695+
if (!items_list) return {};
16961696
while (param_cursor->peek() != ParameterKind::TupleEnd) {
1697-
PyPtr elem = parse_param_constraint(cursor, param_cursor, annotation_cursor);
1698-
if (!elem) return {};
1699-
if (PyList_Append(elements_list.get(), elem.get()) < 0) return {};
1697+
PyPtr item = parse_param_constraint(cursor, param_cursor, annotation_cursor);
1698+
if (!item) return {};
1699+
if (PyList_Append(items_list.get(), item.get()) < 0) return {};
17001700
}
17011701
ParameterKind tuple_end = param_cursor->next();
17021702
CHECK(tuple_end == ParameterKind::TupleEnd);
@@ -1705,7 +1705,7 @@ static PyPtr parse_param_constraint(ConstantCursor& cursor,
17051705
if (!signature_module) return {};
17061706
PyPtr constraint_class = getattr(signature_module, "TupleConstraint");
17071707
if (!constraint_class) return {};
1708-
return steal(PyObject_CallOneArg(constraint_class.get(), elements_list.get()));
1708+
return steal(PyObject_CallOneArg(constraint_class.get(), items_list.get()));
17091709
}
17101710
LeafAnnotationNode* annotation = annotation_cursor->next().get();
17111711
return parse_element_constraint(cursor, pk, *annotation);

src/cuda/tile/_compile.py

Lines changed: 5 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -177,21 +177,21 @@ def _create_parameter(
177177

178178
if isinstance(constraint, TupleConstraint):
179179
if isinstance(annotation, LeafAnnotationNode):
180-
item_nodes = [annotation] * len(constraint.elements)
180+
item_nodes = [annotation] * len(constraint.items)
181181
elif isinstance(annotation, HomogeneousTupleNode):
182-
item_nodes = [annotation.each] * len(constraint.elements)
182+
item_nodes = [annotation.each] * len(constraint.items)
183183
elif isinstance(annotation, HeterogeneousTupleNode):
184-
if len(annotation.items) != len(constraint.elements):
184+
if len(annotation.items) != len(constraint.items):
185185
raise _make_constraint_error(
186-
f"Received a tuple of length {len(constraint.elements)}"
186+
f"Received a tuple of length {len(constraint.items)}"
187187
f" but the annotation implies length {len(annotation.items)}.",
188188
path)
189189
item_nodes = annotation.items
190190
else:
191191
assert False
192192

193193
item_vars = []
194-
for i, (item, node) in enumerate(zip(constraint.elements, item_nodes, strict=True)):
194+
for i, (item, node) in enumerate(zip(constraint.items, item_nodes, strict=True)):
195195
item_var = var.ctx.make_var(var.name + f"_{i}", var.loc)
196196
_create_parameter(item, node, path.tuple_item(i), item_var, nonconstant_flat_vars)
197197
item_vars.append(item_var)

src/cuda/tile/_passes/dataflow_analysis.py

Lines changed: 11 additions & 11 deletions
Original file line numberDiff line numberDiff line change
@@ -68,21 +68,21 @@ def _register_leaf_param(state, constraint: ArrayConstraint | ScalarConstraint,
6868

6969
def _register_tuple_params(state, constraint: TupleConstraint, flat_params, offset: int,
7070
alias_set_mapper) -> int:
71-
for elem in constraint.elements:
72-
if isinstance(elem, (ArrayConstraint, ScalarConstraint)):
73-
n = 1 + 2 * elem.ndim if isinstance(elem, ArrayConstraint) else 1
74-
_register_leaf_param(state, elem, flat_params[offset:offset + n], alias_set_mapper)
71+
for item in constraint.items:
72+
if isinstance(item, (ArrayConstraint, ScalarConstraint)):
73+
n = 1 + 2 * item.ndim if isinstance(item, ArrayConstraint) else 1
74+
_register_leaf_param(state, item, flat_params[offset:offset + n], alias_set_mapper)
7575
offset += n
76-
elif isinstance(elem, TupleConstraint):
77-
offset = _register_tuple_params(state, elem, flat_params, offset, alias_set_mapper)
78-
elif isinstance(elem, ListConstraint):
79-
assert isinstance(elem.element, ArrayConstraint)
76+
elif isinstance(item, TupleConstraint):
77+
offset = _register_tuple_params(state, item, flat_params, offset, alias_set_mapper)
78+
elif isinstance(item, ListConstraint):
79+
assert isinstance(item.element, ArrayConstraint)
8080
base_ptr, size_var = flat_params[offset], flat_params[offset + 1]
8181
state.tracker.update(base_ptr,
82-
DataPredicate(alias_set=alias_set_mapper(elem.alias_groups),
82+
DataPredicate(alias_set=alias_set_mapper(item.alias_groups),
8383
div_by=1,
84-
may_alias_internally=elem.elements_may_alias))
85-
elt_predicates = _get_array_predicates(elem.element, alias_set_mapper)
84+
may_alias_internally=item.elements_may_alias))
85+
elt_predicates = _get_array_predicates(item.element, alias_set_mapper)
8686
state.list_array_tracker.update(base_ptr,
8787
_AggregatePredicate(dict(enumerate(elt_predicates))))
8888
state.set_always_true(size_var)

src/cuda/tile/compilation/_name_mangling.py

Lines changed: 5 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -395,17 +395,17 @@ def _demangle_list_constraint(cursor: _Cursor,
395395

396396
def _mangle_tuple_constraint(constraint: TupleConstraint, alias_group_map: dict[str, int],
397397
cconv: CallingConvention) -> str:
398-
# Format: {count}{elem0_mangling}{elem1_mangling}...
399-
return f"{len(constraint.elements)}" + "".join(
400-
_mangle_constraint(e, alias_group_map, cconv) for e in constraint.elements)
398+
# Format: {count}{item0_mangling}{item1_mangling}...
399+
return f"{len(constraint.items)}" + "".join(
400+
_mangle_constraint(e, alias_group_map, cconv) for e in constraint.items)
401401

402402

403403
def _demangle_tuple_constraint(cursor: _Cursor,
404404
alias_group_demangler: _AliasGroupDemangler,
405405
cconv: CallingConvention) -> TupleConstraint:
406406
count = int(cursor.expect("[0-9]+", "Expected element count"))
407-
elements = [_demangle_constraint(cursor, alias_group_demangler, cconv) for _ in range(count)]
408-
return TupleConstraint(elements)
407+
items = [_demangle_constraint(cursor, alias_group_demangler, cconv) for _ in range(count)]
408+
return TupleConstraint(items)
409409

410410

411411
def _mangle_dtype(dtype: DType):

src/cuda/tile/compilation/_signature.py

Lines changed: 12 additions & 18 deletions
Original file line numberDiff line numberDiff line change
@@ -227,25 +227,19 @@ class TupleConstraint:
227227
Describes a tuple kernel parameter.
228228
229229
Args:
230-
elements: Per-element parameter constraints. Elements may be
231-
:class:`ScalarConstraint`, :class:`ArrayConstraint`,
232-
:class:`ConstantConstraint`, :class:`ListConstraint`, or nested
233-
:class:`TupleConstraint`.
230+
items: Per-item constraints.
234231
"""
235-
elements: "tuple[ParameterConstraint, ...]"
236-
237-
def __init__(
238-
self,
239-
elements: "Sequence[ParameterConstraint]",
240-
):
241-
for i, e in enumerate(elements):
242-
if not isinstance(e, (ScalarConstraint, ArrayConstraint, ConstantConstraint,
243-
TupleConstraint, ListConstraint)):
232+
items: "tuple[ParameterConstraint, ...]"
233+
234+
def __init__(self,
235+
items: "Sequence[ParameterConstraint]"):
236+
items = tuple(items)
237+
for i, item_constraint in enumerate(items):
238+
if not isinstance(item_constraint, ParameterConstraint):
244239
raise TypeError(
245-
f"TupleConstraint element {i} must be a ScalarConstraint,"
246-
f" ArrayConstraint, ConstantConstraint, TupleConstraint, or ListConstraint,"
247-
f" got {type(e).__name__}")
248-
object.__setattr__(self, "elements", tuple(elements))
240+
f"TupleConstraint item #{i} must be a ParameterConstraint,"
241+
f" got {type(item_constraint).__name__}")
242+
object.__setattr__(self, "items", items)
249243

250244

251245
@dataclass(frozen=False, eq=False)
@@ -445,7 +439,7 @@ def _collect_alias_groups(parameters: Sequence[ParameterConstraint]
445439
yield p, p.alias_groups
446440
yield from _collect_alias_groups([p.element])
447441
elif isinstance(p, TupleConstraint):
448-
yield from _collect_alias_groups(list(p.elements))
442+
yield from _collect_alias_groups(p.items)
449443

450444

451445
def _check_optional_int(i: int, val, param_name: str):

0 commit comments

Comments
 (0)