This is an automated email from the ASF dual-hosted git repository. tqchen pushed a commit to branch tvmscript-generic-parser-builder in repository https://gitbox.apache.org/repos/asf/tvm.git
commit c5a0bd4bebcb9addc53a8b6b639bc4800a200d51 Author: Tianqi Chen <[email protected]> AuthorDate: Mon Sep 21 03:36:30 2026 +0000 Route attribute stores through registered construction operations --- python/tvm/script/parser_v2/transform.py | 36 +++++++++++++++++++++++++------- tests/python/tvmscript/test_parser_v2.py | 20 ++++++++++++++++++ 2 files changed, 48 insertions(+), 8 deletions(-) diff --git a/python/tvm/script/parser_v2/transform.py b/python/tvm/script/parser_v2/transform.py index 4746d8a788..b3270c45fc 100644 --- a/python/tvm/script/parser_v2/transform.py +++ b/python/tvm/script/parser_v2/transform.py @@ -331,6 +331,17 @@ class Transformer(ast.NodeTransformer): target.id, self._operation("bind_", [value], statement, **keywords), target ) ] + if isinstance(target, ast.Attribute): + return [ + self._statement( + self._operation( + "setattr", + [self._expression(target.value), ast.Constant(target.attr), value], + statement, + ), + statement, + ) + ] if isinstance(target, ast.Subscript): return [ self._statement( @@ -440,15 +451,24 @@ class Transformer(ast.NodeTransformer): value = self._located(ast.BinOp(previous, node.op, self._expression(node.value)), node) value = self._call(self.infrastructure_name, "_at", [self.span(node), value], node) return result + self._bind(node.target, value, node) - if not isinstance(node.target, ast.Subscript): - self._error(node.target, "An augmented assignment requires a name or index") + if not isinstance(node.target, ast.Subscript | ast.Attribute): + self._error(node.target, "An augmented assignment requires a name, attribute, or index") base_stmt, base = self._cache( self._expression(node.target.value), node.target.value, "base" ) - key_stmt, key = self._cache(self._index(node.target.slice), node.target.slice, "key") - load = self._located( - ast.Subscript(copy.deepcopy(base), copy.deepcopy(key), ast.Load()), node.target - ) + result.append(base_stmt) + if isinstance(node.target, ast.Attribute): + key, operation = ast.Constant(node.target.attr), "setattr" + load = self._located( + ast.Attribute(copy.deepcopy(base), node.target.attr, ast.Load()), node.target + ) + else: + key_stmt, key = self._cache(self._index(node.target.slice), node.target.slice, "key") + result.append(key_stmt) + operation = "setitem" + load = self._located( + ast.Subscript(copy.deepcopy(base), copy.deepcopy(key), ast.Load()), node.target + ) old_stmt, old = self._cache( self._call( self.infrastructure_name, "_at", [self.span(node.target), load], node.target @@ -456,10 +476,10 @@ class Transformer(ast.NodeTransformer): node.target, "old", ) - result.extend([base_stmt, key_stmt, old_stmt]) + result.append(old_stmt) value = self._located(ast.BinOp(old, node.op, self._expression(node.value)), node) value = self._call(self.infrastructure_name, "_at", [self.span(node), value], node) - result.append(self._statement(self._operation("setitem", [base, key, value], node), node)) + result.append(self._statement(self._operation(operation, [base, key, value], node), node)) return result def visit_Expr(self, node): diff --git a/tests/python/tvmscript/test_parser_v2.py b/tests/python/tvmscript/test_parser_v2.py index 1ab21bd0f5..1976f1098c 100644 --- a/tests/python/tvmscript/test_parser_v2.py +++ b/tests/python/tvmscript/test_parser_v2.py @@ -54,6 +54,9 @@ class _Recorder: def setitem(self, target, key, value, **metadata): target[key] = value + def setattr(self, target, name, value, **metadata): + setattr(target, name, value) + def unpack(self, value): return value @@ -124,6 +127,14 @@ def test_assignment_order_and_hierarchical_unpack(): events.append(("store", value)) super().__setitem__(key, value) + @property + def field(self): + return self[0] + + @field.setter + def field(self, value): + self[0] = value + target = Target() def base(): @@ -144,6 +155,8 @@ def test_assignment_order_and_hierarchical_unpack(): def f(): base()[index()] = value() base()[index()] += value() + base().field = value() + base().field += value() base()[index()], (a, b) = (1, (2,)) """, {"base": base, "index": index, "value": value}, @@ -160,6 +173,13 @@ def test_assignment_order_and_hierarchical_unpack(): "read", "value", ("store", 14), + "value", + "base", + ("store", 7), + "base", + "read", + "value", + ("store", 14), "base", "index", ("store", 1),
