This is an automated email from the ASF dual-hosted git repository. tqchen pushed a commit to branch script/canonical-parser-df in repository https://gitbox.apache.org/repos/asf/tvm.git
commit c067f3967cf8267e7b54fbef43a1acc7e73a11b4 Author: Tianqi Chen <[email protected]> AuthorDate: Wed Sep 23 04:16:24 2026 +0000 Preserve lexical scopes and source locations in direct parser lowering --- python/tvm/script/parser/__init__.py | 5 + python/tvm/script/parser/entry.py | 7 +- python/tvm/script/parser/prescan.py | 82 +++++--- python/tvm/script/parser/transpile.py | 211 +++++++++++++++------ tests/parser/dummy_builder.py | 7 + tests/parser/test_parser.py | 154 +++++++++++++++ tests/python/relax/test_tvmscript_parser.py | 4 +- tests/python/tirx/test_parser_call_scopes.py | 2 +- tests/python/tirx/test_parser_constexpr.py | 2 +- tests/python/tirx/test_parser_emitted_spans.py | 2 +- tests/python/tirx/test_parser_scope_id_identity.py | 4 +- tests/python/tirx/test_parser_scope_id_names.py | 2 +- tests/python/unittest/test_parser_native_frames.py | 86 +++++++++ 13 files changed, 471 insertions(+), 97 deletions(-) diff --git a/python/tvm/script/parser/__init__.py b/python/tvm/script/parser/__init__.py index 121096e6a7..90e22e6b0a 100644 --- a/python/tvm/script/parser/__init__.py +++ b/python/tvm/script/parser/__init__.py @@ -74,6 +74,11 @@ def _initialize(): ir.ir_module = entry.ir_module ir.pyfunc = entry.pyfunc for namespace in (tir_namespace, relax_namespace, ir): + # Materialize declared public entry points after builder bootstrap. + # Source-text decorators need the same registered metadata as Python + # decorator lookup, including entries owned lazily by a dialect. + for name in vars(namespace).get("_ENTRY_EXPORTS", ()): + getattr(namespace, name) namespace.__all__ = sorted( { *vars(namespace).get("__all__", ()), diff --git a/python/tvm/script/parser/entry.py b/python/tvm/script/parser/entry.py index 4f7212e8b4..ad8afc4dd2 100644 --- a/python/tvm/script/parser/entry.py +++ b/python/tvm/script/parser/entry.py @@ -528,7 +528,7 @@ def _prepare_transpiler( # Give prescan its registered syntax policy before allocating injected names. if inspect.isfunction(source): tree.body[-1]._tvm_function_info = syntax_protocol.function_info(source) - prescan = PrescanCollector(metadata).collect(tree) + prescan = PrescanCollector(metadata, filename=filename).collect(tree) names = dict.fromkeys([*namespace, *prescan.reserved_names], 0) def fresh(prefix="_t"): @@ -590,7 +590,6 @@ def _prepare_transpiler( infrastructure_name, span, fresh, - name_map=names, track_span=track_span, definition_scopes_name=definition_scopes_name, prescan=prescan, @@ -666,9 +665,7 @@ def _build(tree, source, environment, definition_scope, filename, flags, *, trac transformer, namespace = _prepare_transpiler( tree, source, environment, definition_scope, filename, track_span=track_span ) - original_name = transformer.fresh("_original") - namespace[original_name] = source if inspect.isclass(source) else None - transformed, result = transformer.program(tree, original_name, namespace) + transformed, result = transformer.program(tree) runnable = recompose_builder( transformed, source_fn=source, diff --git a/python/tvm/script/parser/prescan.py b/python/tvm/script/parser/prescan.py index 1b7bd2987b..aa5004a50f 100644 --- a/python/tvm/script/parser/prescan.py +++ b/python/tvm/script/parser/prescan.py @@ -79,14 +79,17 @@ class PrescanContext: namespaces: frozenset # Explicit source dataflow outputs preserve native export identity on exit. with_outputs: object + # Only source functions referencing themselves need early standalone refs. + recursive_functions: frozenset class PrescanCollector(ast.NodeVisitor): """Collect binding syntax once, then discard all traversal accumulators.""" - def __init__(self, environment): + def __init__(self, environment, *, filename="<str>"): # Fixed lookup inputs last for this scan; never updated by assignments. self.environment = environment + self.filename = filename # Accumulators are frozen by collect(); no native values are stored. self.names = set(environment) self.bindings = {} @@ -94,6 +97,10 @@ class PrescanCollector(ast.NodeVisitor): self.outputs = {} self.namespaces = set() self.exports = {} + self.recursive = set() + # Active source declarations, used only to recognize self references. + self.functions = [] + self.module_name = None # Temporary lexical with-stack routes explicit output calls, then resets. self.regions = [] # Lexical scope and dialect restore on function/class exit. direct marks @@ -140,11 +147,15 @@ class PrescanCollector(ast.NodeVisitor): MappingProxyType(dict(self.outputs)), frozenset(self.namespaces), MappingProxyType({node: tuple(names) for node, names in self.exports.items()}), + frozenset(self.recursive), ) + def _error(self, node, message): + raise SyntaxError(message, (self.filename, node.lineno, node.col_offset + 1, None)) + def _binding(self, name, node, kind="ordinary", annotation=None, dtype=None): if name in self.namespaces: - raise SyntaxError(f"Script namespace {name!r} cannot be rebound or shadowed") + self._error(node, f"Script namespace {name!r} cannot be rebound or shadowed") self.names.add(name) item = Binding(name, node, kind, annotation, dtype, self.direct) self.bindings[self.scope].append(item) @@ -152,33 +163,39 @@ class PrescanCollector(ast.NodeVisitor): def visit_Name(self, node): if isinstance(node.ctx, ast.Store) and node.id in self.namespaces: - raise SyntaxError(f"Script namespace {node.id!r} cannot be rebound or shadowed") + self._error(node, f"Script namespace {node.id!r} cannot be rebound or shadowed") self.names.add(node.id) + if isinstance(node.ctx, ast.Load): + for function in reversed(self.functions): + if node.id == function.name: + self.recursive.add(function) + break def visit_arg(self, node): if node.arg in self.namespaces: - raise SyntaxError(f"Script namespace {node.arg!r} cannot be rebound or shadowed") + self._error(node, f"Script namespace {node.arg!r} cannot be rebound or shadowed") self.names.add(node.arg) self.generic_visit(node) def visit_alias(self, node): name = node.asname or node.name.split(".")[0] if isinstance(self.scope, ast.FunctionDef) and name in self.namespaces: - raise SyntaxError(f"Script namespace {name!r} cannot be rebound or shadowed") + self._error(node, f"Script namespace {name!r} cannot be rebound or shadowed") self.names.add(name) def visit_ClassDef(self, node): self.names.add(node.name) - old = self.scope - self.scope = node + old, old_module = self.scope, self.module_name + self.scope, self.module_name = node, node.name self.bindings[node] = [] self.generic_visit(node) - self.scope = old + self.scope, self.module_name = old, old_module def visit_FunctionDef(self, node): self._binding(node.name, node, "function") old_scope, old_builder, old_direct = self.scope, self.builder, self.direct self.scope, self.direct = node, True + self.functions.append(node) original_kind = getattr(node, "_tvm_function_info", None) if original_kind is not None: self.builder = original_kind.builder @@ -203,7 +220,8 @@ class PrescanCollector(ast.NodeVisitor): arg.arg, arg, "mutable_parameter" - if protocol.is_mutable_var_decl(constructor, syntax="parameter") + if getattr(self.builder, "supports_mutable_declarations", True) + and protocol.is_mutable_var_decl(constructor, syntax="parameter") else "parameter", arg.annotation, (declaration.dtype if declaration is not None else dtype) @@ -218,6 +236,7 @@ class PrescanCollector(ast.NodeVisitor): self.visit(decorator) if node.returns: self.visit(node.returns) + self.functions.pop() self.scope, self.builder, self.direct = old_scope, old_builder, old_direct visit_AsyncFunctionDef = visit_FunctionDef @@ -232,17 +251,24 @@ class PrescanCollector(ast.NodeVisitor): declaration = getattr(constructor, "__tvm_type_var_decl__", None) if declaration is not None and not value.args and not value.keywords: self._binding(target.id, target, "symbol", value, declaration.dtype) - elif protocol.is_mutable_var_decl(constructor, syntax="call") or ( - annotation is not None - and protocol.is_mutable_var_decl( - resolve_syntax( - annotation.value if isinstance(annotation, ast.Subscript) else annotation, - self.environment, - ), - syntax="annotation", + elif getattr(self.builder, "supports_mutable_declarations", True) and ( + protocol.is_mutable_var_decl(constructor, syntax="call") + or ( + annotation is not None + and protocol.is_mutable_var_decl( + resolve_syntax( + annotation.value + if isinstance(annotation, ast.Subscript) + else annotation, + self.environment, + ), + syntax="annotation", + ) ) ): self._binding(target.id, target, "mutable", annotation) + elif isinstance(value, ast.Name) and value.id == self.module_name: + self._binding(target.id, target, "module_alias") else: self._binding(target.id, target, annotation=annotation) elif isinstance(target, ast.Tuple | ast.List): @@ -270,6 +296,7 @@ class PrescanCollector(ast.NodeVisitor): def visit_If(self, node): old_direct, self.direct = self.direct, False + self.generic_visit(node) marker = ( isinstance(node.test, ast.Call) and isinstance(node.test.func, ast.Attribute) @@ -287,13 +314,24 @@ class PrescanCollector(ast.NodeVisitor): return last.targets[0].id if isinstance(last, ast.AnnAssign) and isinstance(last.target, ast.Name): return last.target.id - return None + return self.outputs.get(last) then, otherwise = ending(node.body), ending(node.orelse) - if then is None or then != otherwise: - raise SyntaxError("IR conditional branches must end with the same named output") - self.outputs[node] = then - self.generic_visit(node) + # Effect-only branches have no Python output. Native branch frames + # still reject a non-void expression used as an effect-only ending. + effects = bool( + node.body + and node.orelse + and isinstance(node.body[-1], ast.Expr) + and isinstance(node.orelse[-1], ast.Expr) + ) + if not effects: + if then is None or then != otherwise: + location = node.orelse[-1] if node.orelse else node + self._error( + location, "IR conditional branches must end with the same named output" + ) + self.outputs[node] = then self.direct = old_direct def visit_For(self, node): diff --git a/python/tvm/script/parser/transpile.py b/python/tvm/script/parser/transpile.py index 4974825923..ab18c87811 100644 --- a/python/tvm/script/parser/transpile.py +++ b/python/tvm/script/parser/transpile.py @@ -34,7 +34,7 @@ class IRBuilderTranspiler(ast.NodeTransformer): """A single statement/expression visitor over the entry-owned AST. ``environment``, ``prescan``, ``filename`` and ``span`` are fixed inputs for - one translation. ``name_map`` is the unit-wide collision-free name allocator; + one translation. ``fresh`` allocates names in the entry-owned name map; ``bindings`` receives only injected host namespaces and helper functions. ``dialect_prefix``, ``current_scope``, ``host_expression`` and annotation substitutions are saved/restored at their lexical visitor boundaries. They @@ -48,10 +48,9 @@ class IRBuilderTranspiler(ast.NodeTransformer): builder_name, infrastructure_name, span, - fresh=None, + fresh, *, prescan=None, - name_map=None, track_span=True, definition_scopes_name=None, current_scope=None, @@ -67,8 +66,7 @@ class IRBuilderTranspiler(ast.NodeTransformer): self.track_span = track_span self.definition_scopes_name = definition_scopes_name # Allocation/injection last for this source unit, including nested code. - self.name_map = name_map if name_map is not None else dict.fromkeys(environment, 0) - self.fresh = fresh or self.fresh_unique_name + self.fresh = fresh self.bindings = bindings if bindings is not None else {} # Lexical syntax context, restored on nested function/host/annotation exit. self.dialect_prefix = builder_name @@ -81,16 +79,6 @@ class IRBuilderTranspiler(ast.NodeTransformer): self.module_name = None self.module_functions = frozenset() - def fresh_unique_name(self, prefix="_t"): - """Reserve a generated identifier for the lifetime of this translation.""" - counter = self.name_map.get(prefix, 0) - while f"{prefix}{counter}" in self.name_map: - counter += 1 - name = f"{prefix}{counter}" - self.name_map[prefix] = counter + 1 - self.name_map[name] = 0 - return name - def _inject(self, value, prefix="_host"): name = self.fresh(prefix) self.bindings[name] = value @@ -177,6 +165,25 @@ class IRBuilderTranspiler(ast.NodeTransformer): self.host_expression = old def _resolve(self, node): + # Fixed namespace meanings coexist with Python lexical value bindings. + # A local ``range`` or callable hides the ambient binding for the whole + # source function, including reads before its assignment. + root = node + while isinstance(root, ast.Attribute): + root = root.value + scope = self.current_scope + if isinstance(root, ast.Name) and self.prescan is not None: + while isinstance(scope, ast.FunctionDef): + if any(item.name == root.id for item in self.prescan.bindings.get(scope, ())): + return None + scope = next( + ( + parent + for parent, items in self.prescan.bindings.items() + if any(item.node is scope and item.kind == "function" for item in items) + ), + None, + ) return resolve_syntax(node, self.environment) def _constexpr_operand(self, node): @@ -228,7 +235,11 @@ class IRBuilderTranspiler(ast.NodeTransformer): ], node, ) - return node + return ( + self._at(node, node) + if isinstance(node.ctx, ast.Load) and not self.host_expression + else node + ) def visit_Attribute(self, node): # Source: Module.f; Builder: f (the same reserved native GlobalVar). @@ -237,8 +248,13 @@ class IRBuilderTranspiler(ast.NodeTransformer): and node.value.id == self.module_name and node.attr in self.module_functions ): - return ast.copy_location(ast.Name(node.attr, node.ctx), node) - return self.generic_visit(node) + return self._at(ast.copy_location(ast.Name(node.attr, node.ctx), node), node) + result = self.generic_visit(node) + return ( + self._at(result, node) + if isinstance(node.ctx, ast.Load) and not self.host_expression + else result + ) def visit_Lambda(self, node): # Source: lambda n: n + outer; Builder: preserve the lambda's locals @@ -280,7 +296,51 @@ class IRBuilderTranspiler(ast.NodeTransformer): visit_DictComp = visit_ListComp visit_GeneratorExp = visit_ListComp - def visit_Call(self, node): + def visit_Constant(self, node): + # Source: literal; Builder: I.at_(literal_loc, literal). + return node if self.host_expression else self._at(node, node) + + def visit_List(self, node): + # Source: [a,b] / (a,b) / {a,b}; Builder: preserve the Python container + # and source-location every original child through this same visitor. + result = self.generic_visit(node) + return result if self.host_expression else self._at(result, node) + + visit_Tuple = visit_List + visit_Set = visit_List + visit_Dict = visit_List + + def visit_JoinedStr(self, node): + # Source: f"value={expr}"; Builder: keep literal fragments, visit expr. + # Python requires these fragments to remain Constant/FormattedValue. + for child in node.values: + if isinstance(child, ast.FormattedValue): + child.value = self.visit(child.value) + if child.format_spec is not None: + child.format_spec = self.visit_JoinedStr(child.format_spec) + return node + + def visit_Subscript(self, node): + # Source: buffer[index]; Builder: I.at_(load_loc, buffer[index]). + result = self.generic_visit(node) + if not self.host_expression and isinstance(node.ctx, ast.Load): + return self._at(result, node) + return result + + def _module_owner(self, node): + """Recognize fixed source module aliases from existing binding records.""" + if not isinstance(node, ast.Name): + return False + if node.id == self.module_name: + return True + records = [ + item + for item in self.prescan.bindings.get(self.current_scope, ()) + if item.name == node.id + ] + return bool(records) and all(item.kind == "module_alias" for item in records) + + def visit_Call(self, node, *, callee=None): # Source: X.Tensor(("n",), vdevice="cuda:0") # Builder: X.Tensor((X.resolve_type_var_("n"),), # vdevice=I.resolve_global_info("cuda:0")) @@ -297,11 +357,18 @@ class IRBuilderTranspiler(ast.NodeTransformer): isinstance(node.func, ast.Name) and node.func.id in self.module_functions ) or ( isinstance(node.func, ast.Attribute) - and isinstance(node.func.value, ast.Name) - and node.func.value.id == self.module_name + and self._module_owner(node.func.value) and node.func.attr in self.module_functions ) - node = self.generic_visit(node) + # Generated callees are already builder syntax. Only their original + # arguments need the main visitor; do not give injected names fake spans. + if callee is not None: + node.func = callee + elif not generated: + node.func = self.visit(node.func) + node.args = [self.visit(value) for value in node.args] + for keyword in node.keywords: + keyword.value = self.visit(keyword.value) # Source: declared_global(x, y) # Builder: X.call_global_var_(declared_global, [x, y]) if global_call and not self.host_expression: @@ -773,12 +840,17 @@ class IRBuilderTranspiler(ast.NodeTransformer): return self.generic_visit(node) if node.orelse: self._error(node, "A construction loop does not support an else clause") - iterable = node.iter - if isinstance(iterable, ast.Call) and self._resolve(iterable.func) is range: - iterable.func = ast.copy_location( - ast.Attribute(ast.Name(self.dialect_prefix, ast.Load()), "range_", ast.Load()), - iterable.func, + if isinstance(node.iter, ast.Call) and self._resolve(node.iter.func) is range: + # Normalize only the known builtin; a lexical range binding follows + # the ordinary source-call path. Arguments keep that one call scope. + iterable = self.visit_Call( + node.iter, + callee=ast.Attribute( + ast.Name(self.dialect_prefix, ast.Load()), "range_", ast.Load() + ), ) + else: + iterable = self.visit(node.iter) if isinstance(node.target, ast.Name): names = ast.Constant(node.target.id) elif isinstance(node.target, ast.Tuple | ast.List): @@ -791,7 +863,7 @@ class IRBuilderTranspiler(ast.NodeTransformer): ) else: self._error(node.target, "Loop targets must be names or a flat tuple of names") - context = self._operation("for_", [self.visit(iterable)], node, names=names) + context = self._operation("for_", [iterable], node, names=names) body = self.transform_statements(node.body) return ast.copy_location( ast.With([ast.withitem(context, node.target)], body or [ast.Pass()]), node @@ -913,7 +985,7 @@ class IRBuilderTranspiler(ast.NodeTransformer): return protocol.FunctionDecoratorInfo(None, python=True), ast.Dict([], []) self._error(node, f"Function {node.name!r} has no registered construction kind") - def function_program(self, node, *, local=False): + def function_program(self, node, *, local=False, declare=True): """Declare one native frame and emit one body function taking that frame. Annotation aliases retain actual definition-local Python values; source @@ -1045,19 +1117,22 @@ class IRBuilderTranspiler(ast.NodeTransformer): ) for item in facts: if item.direct and item.dtype is not None and item.kind in ("symbol", "parameter"): - declaration.append( - ast.copy_location( - ast.Expr( - self._operation( - "resolve_type_var_", - [ast.Constant(item.name)], - item.node, - dtype=ast.Constant(item.dtype), - ) - ), - item.node, - ) + symbol = self._operation( + "resolve_type_var_", + [ast.Constant(item.name)], + item.node, + dtype=ast.Constant(item.dtype), ) + if item.kind == "parameter": + alias = self.fresh("_symbol") + symbol_aliases[item.name] = alias + declaration.append( + ast.copy_location( + ast.Assign([ast.Name(alias, ast.Store())], symbol), item.node + ) + ) + else: + declaration.append(ast.copy_location(ast.Expr(symbol), item.node)) self.annotation_aliases = {**aliases, **symbol_aliases} constexpr_aliases = {} ordered = sorted( @@ -1142,7 +1217,9 @@ class IRBuilderTranspiler(ast.NodeTransformer): # Builder: with X.function(decl=True) as fn: signature with self._host(): options = self.visit(options) - keywords = [ast.keyword(None, options), ast.keyword("decl", ast.Constant(True))] + keywords = [ast.keyword(None, options)] + if declare: + keywords.append(ast.keyword("decl", ast.Constant(True))) if local: keywords.append(ast.keyword("local", ast.Constant(True))) if self.track_span: @@ -1153,21 +1230,21 @@ class IRBuilderTranspiler(ast.NodeTransformer): ), node, ) - statements.append( - ast.copy_location( - ast.With([ast.withitem(constructor, ast.Name(frame, ast.Store()))], declaration), - node, + if declare: + declaration_scope = ast.With( + [ast.withitem(constructor, ast.Name(frame, ast.Store()))], declaration ) - ) - statements.append( - ast.copy_location( - ast.Assign( - [ast.Name(node.name, ast.Store())], - ast.Attribute(ast.Name(frame, ast.Load()), "reference", ast.Load()), - ), - node, + statements.append(ast.copy_location(declaration_scope, node)) + reference = ast.Attribute(ast.Name(frame, ast.Load()), "reference", ast.Load()) + statements.append( + ast.copy_location(ast.Assign([ast.Name(node.name, ast.Store())], reference), node) + ) + else: + # Source: a standalone nonrecursive function. + # Builder: fn = X.function(); build(fn) enters once for signature+body. + statements.append( + ast.copy_location(ast.Assign([ast.Name(frame, ast.Store())], constructor), node) ) - ) # The body re-enters the exact native frame. Runtime parameter storage # stays native; the short iterator is consumed once by source parameters. iterator = self.fresh("_arguments") @@ -1237,7 +1314,11 @@ class IRBuilderTranspiler(ast.NodeTransformer): body.extend(self.transform_statements(node.body)) self.body_annotation_aliases = old_body_annotations resumed = ast.copy_location( - ast.With([ast.withitem(ast.Name(frame, ast.Load()))], body), node + ast.With( + [ast.withitem(ast.Name(frame, ast.Load()))], + body if declare else [*declaration, *body], + ), + node, ) definition = self._definition(body_name, [frame], [resumed], node) if not local: @@ -1251,11 +1332,10 @@ class IRBuilderTranspiler(ast.NodeTransformer): self.dialect_prefix, self.current_scope, self.annotation_aliases = old return statements, frame, body_name - def program(self, tree, original_name, bindings): + def program(self, tree): """Emit direct native module construction, declarations, then bodies.""" from .entry import make_opaque_function - self.bindings = bindings root, prefix = tree.body[-1], tree.body[:-1] if not isinstance(root, ast.ClassDef | ast.FunctionDef): self._error(root, "Source must contain one function or module class") @@ -1266,9 +1346,10 @@ class IRBuilderTranspiler(ast.NodeTransformer): self._error(root, "Duplicate function declaration") self.module_name = root.name if is_module else None self.module_functions = frozenset(item.name for item in functions) + declare = is_module or root in self.prescan.recursive_functions builder, result = self.fresh("_builder"), self.fresh("_result") body = [] - for function in functions: + for function in functions if declare else (): body.append( ast.copy_location( ast.Assign( @@ -1371,7 +1452,7 @@ class IRBuilderTranspiler(ast.NodeTransformer): ) python_functions.append((function.name, host)) continue - declaration, frame, callback = self.function_program(function) + declaration, frame, callback = self.function_program(function, declare=declare) body.extend(declaration) definitions.append( ast.copy_location( @@ -1390,7 +1471,13 @@ class IRBuilderTranspiler(ast.NodeTransformer): # result = builder.get() module = ast.copy_location( ast.With( - [ast.withitem(self._call(self.infrastructure_name, "ir_module", [], root))], body + [ + ast.withitem( + self._call(self.infrastructure_name, "ir_module", [], root), + ast.Name(root.name, ast.Store()) if is_module else None, + ) + ], + body, ), root, ) @@ -1402,7 +1489,7 @@ class IRBuilderTranspiler(ast.NodeTransformer): ast.Name(builder, ast.Store()), ) ], - [module], + [module] if declare else body, ), root, ) diff --git a/tests/parser/dummy_builder.py b/tests/parser/dummy_builder.py index 7e4f25cb58..87823e10ae 100644 --- a/tests/parser/dummy_builder.py +++ b/tests/parser/dummy_builder.py @@ -113,6 +113,11 @@ class Frame: def __getitem__(self, name): return self.language.functions[name] + def __getattr__(self, name): + if self.kind == "module" and name in self.language.references: + return self.language.references[name] + raise AttributeError(name) + class Language: """One recording builder namespace X and shared infrastructure namespace I.""" @@ -131,12 +136,14 @@ class Language: with_at_group_=self.with_at_group, resolve_global_info=self.resolve_global_info, reserve_function=self.reserve_function, + module_member_=lambda name, value: value, require_defined=self.require_defined, annotation_value_=lambda name, value: value, MISSING=self.missing, constexpr=protocol.constexpr, ) self.X = SimpleNamespace( + supports_mutable_declarations=True, function=lambda **kwargs: Frame(self, "function", **kwargs), func_name=self.func_name, arg=self.arg, diff --git a/tests/parser/test_parser.py b/tests/parser/test_parser.py index e66cebf1fc..b3d1a671dd 100644 --- a/tests/parser/test_parser.py +++ b/tests/parser/test_parser.py @@ -532,3 +532,157 @@ def main(): """, marker=protocol.constexpr, ) + + [email protected]("recursive", [False, True]) +def test_source_function_uses_declaration_only_when_reference_is_needed(language, recursive): + # Before: @X.script def main(x: ...): X.record(x) (or main(x)) + # Expected builder program: + # ordinary: with X.function(): x = X.arg(...); X.emit_(X.record(x)) + # recursive: with X.function(decl=True) as fn: X.arg(...) + # with fn: X.emit_(X.call_global_var_(fn.reference, [fn.params[0]])) + statement = "main(x)" if recursive else "X.record(x)" + result = language.parse(f"@X.script\ndef main(x: X.tensor((4,))):\n {statement}\n") + entries = [event for event in language.events if event[:2] == ("enter", "function")] + assert [event[2] for event in entries] == ([True, False] if recursive else [False]) + assert len([event for event in language.events if event[0] == "arg"]) == 1 + if recursive: + assert entries[0][3] is entries[1][3] + call = result.body[0][1] + assert call.op == "call" and call.args[0] is language.references["main"] + assert call.args[1] is result.params[0] + else: + assert result.body[0][1] is result.params[0] + + [email protected]("alias", [False, True]) +def test_module_alias_keeps_frame_identity_and_caller_dialect(language, alias): + # Before: cls = Module; cls.callee(x) + # Expected builder program: + # with I.ir_module() as Module: ... + # cls = X.bind_(Module, name="cls"); X.emit_(X.call_global_var_(cls.callee, [x])) + setup, owner = ("cls = Module", "cls") if alias else ("pass", "Module") + result = language.parse(f""" [email protected]_module +class Module: + @X.script + def caller(x: X.tensor((4,))): + {setup} + {owner}.callee(x) + @X.script + def callee(y: X.tensor((4,))): + X.record(y) +""") + call = result["caller"].body[0][1] + assert call.op == "call" + assert call.args == (language.references["callee"], result["caller"].params[0]) + modules = [event[3] for event in language.events if event[:2] == ("enter", "module")] + assert len(modules) == 1 + if alias: + bound = next(event[2] for event in language.events if event[:2] == ("bind", "cls")) + assert bound is modules[0] + assert bound.callee is language.references["callee"] + + [email protected]( + "body, line, column, message", + [ + (" X = 1\n", 3, 5, "namespace"), + (" for X in range(2):\n pass\n", 3, 9, "namespace"), + ( + " if condition:\n y = left\n else:\n z = right\n", + 6, + 9, + "same named output", + ), + ], +) +def test_syntax_diagnostics_point_to_the_offending_binding(language, body, line, column, message): + # Before: X = 1; or if condition: y = left; else: z = right + # Expected builder program: reject the offending namespace/output binding + # at its original filename, line and column, before running the builder. + language.X.__tvm_value_if__ = True + with pytest.raises(Exception, match=message) as error: + language.parse("@X.script\ndef main():\n" + body) + cause = error.value.__cause__ + assert isinstance(cause, SyntaxError) + assert (cause.filename, cause.lineno, cause.offset) == ("dummy.py", line, column) + assert f"dummy.py:{line}: SyntaxError:" in str(error.value) + assert f" {line} | {body.splitlines()[line - 3]}" in str(error.value) + assert not language.events + + +def test_void_branch_statements_need_no_synthetic_named_output(language): + # Before: if condition: X.record(None); else: X.record(None) + # Expected builder program: + # with X.If(condition): + # with X.Then(): X.emit_(X.record(None)) + # with X.Else(): X.emit_(X.record(None)) + language.X.__tvm_value_if__ = True + result = language.parse( + """ [email protected] +def main(): + if condition: + X.record(None) + else: + X.record(None) + X.record(9) +""", + condition=Value("condition"), + ) + assert [value for kind, value in result.body] == [None, None, 9] + assert not any(event[0] == "bind" for event in language.events) + + +def test_lexical_range_binding_calls_the_custom_iterator_once(language): + # Before: range = custom_range; for i in range(4): X.record(i) + # Expected builder program: range = X.bind_(custom_range, name="range") + # with X.for_(range(4), names="i") as i: X.emit_(X.record(i)) + calls = [] + + def custom_range(extent): + calls.append(extent) + return language.X.grid(2) + + result = language.parse( + "@X.script\ndef main():\n range = custom_range\n" + " for i in range(4):\n X.record(i)\n", + custom_range=custom_range, + ) + assert calls == [4] + variable = result.body[0][1] + assert variable.op == "loop" and variable.args == (2,) and variable.name == "i" + + [email protected]("expression", ["value", "holder.item", "items[0]", "7"]) +def test_non_call_expression_reads_keep_their_source_range(language, expression): + # Before: value; holder.item; items[0]; 7 + # Expected builder program: X.emit_(I.at_(expression_loc, expression)). + # Native expression identity survives; host constants stay ordinary values. + value, reads, located = Value("read"), [], [] + + class Holder: + @property + def item(self): + reads.append("attribute") + return value + + def __getitem__(self, index): + reads.append(index) + return value + + def at(location, result): + located.append(((str(location[0].name), *location[1:]), result)) + return language.at(location, result) + + language.I.at_ = at + holder = Holder() + result = language.parse( + f"@X.script\ndef main():\n {expression}\n", value=value, holder=holder, items=holder + ) + expected = 7 if expression == "7" else value + location = ("dummy.py", 3, 3, 5, 5 + len(expression)) + assert result.body[0][1] is expected + assert [(loc, result) for loc, result in located if loc == location] == [(location, expected)] + assert reads == {"holder.item": ["attribute"], "items[0]": [0]}.get(expression, []) diff --git a/tests/python/relax/test_tvmscript_parser.py b/tests/python/relax/test_tvmscript_parser.py index 9cda2610d3..875d2619e6 100644 --- a/tests/python/relax/test_tvmscript_parser.py +++ b/tests/python/relax/test_tvmscript_parser.py @@ -1239,8 +1239,8 @@ def test_if_branch_with_match_cast(): @R.function def func(A: R.Tensor([16, 16]), is_bfloat16: R.Prim("bool")): if is_bfloat16: - A = R.match_cast(A, R.Tensor([16, 16], "bfloat16")) - B = A.astype("float16") + matched = R.match_cast(A, R.Tensor([16, 16], "bfloat16")) + B = matched.astype("float16") else: B = R.match_cast(A, R.Tensor([16, 16], "float16")) return B diff --git a/tests/python/tirx/test_parser_call_scopes.py b/tests/python/tirx/test_parser_call_scopes.py index 3510c021b9..f111a3af6c 100644 --- a/tests/python/tirx/test_parser_call_scopes.py +++ b/tests/python/tirx/test_parser_call_scopes.py @@ -20,9 +20,9 @@ import pytest from tvm import ir, tirx from tvm.script import parser +from tvm.script import tirx as T from tvm.script.ir_builder import IRBuilder from tvm.script.ir_builder import ir as I -from tvm.script import tirx as T def loc(line, name="calls.py"): diff --git a/tests/python/tirx/test_parser_constexpr.py b/tests/python/tirx/test_parser_constexpr.py index 4e8822b408..2d0ffc958e 100644 --- a/tests/python/tirx/test_parser_constexpr.py +++ b/tests/python/tirx/test_parser_constexpr.py @@ -20,8 +20,8 @@ import pytest from tvm import ir, tirx from tvm.script import parser -from tvm.script.ir_builder import ir as I from tvm.script import tirx as T +from tvm.script.ir_builder import ir as I @pytest.mark.parametrize("marker", ["I.constexpr", "T.constexpr"]) diff --git a/tests/python/tirx/test_parser_emitted_spans.py b/tests/python/tirx/test_parser_emitted_spans.py index 1d07c9366b..9996a7256c 100644 --- a/tests/python/tirx/test_parser_emitted_spans.py +++ b/tests/python/tirx/test_parser_emitted_spans.py @@ -19,10 +19,10 @@ import pytest from tvm import ir, tirx +from tvm.script import tirx as T from tvm.script.ir_builder import IRBuilder from tvm.script.ir_builder import ir as I from tvm.script.ir_builder.base import BypassEmit -from tvm.script import tirx as T from tvm.tirx.script.builder import ir as native diff --git a/tests/python/tirx/test_parser_scope_id_identity.py b/tests/python/tirx/test_parser_scope_id_identity.py index bf02752e3b..d4815f18df 100644 --- a/tests/python/tirx/test_parser_scope_id_identity.py +++ b/tests/python/tirx/test_parser_scope_id_identity.py @@ -20,10 +20,10 @@ import pytest import tvm from tvm.script import parser +from tvm.script import tirx as T from tvm.script.ir_builder import IRBuilder -from tvm.script.ir_builder.base import BypassBind from tvm.script.ir_builder import ir as I -from tvm.script import tirx as T +from tvm.script.ir_builder.base import BypassBind @pytest.mark.parametrize( diff --git a/tests/python/tirx/test_parser_scope_id_names.py b/tests/python/tirx/test_parser_scope_id_names.py index 327ea41a8f..027f7c60cb 100644 --- a/tests/python/tirx/test_parser_scope_id_names.py +++ b/tests/python/tirx/test_parser_scope_id_names.py @@ -20,9 +20,9 @@ import pytest import tvm from tvm.script import parser +from tvm.script import tirx as T from tvm.script.ir_builder import IRBuilder from tvm.script.ir_builder.base import BypassBind -from tvm.script import tirx as T @pytest.mark.parametrize( diff --git a/tests/python/unittest/test_parser_native_frames.py b/tests/python/unittest/test_parser_native_frames.py index ae38929e00..389326e8e2 100644 --- a/tests/python/unittest/test_parser_native_frames.py +++ b/tests/python/unittest/test_parser_native_frames.py @@ -162,3 +162,89 @@ def test_body_attribute_overrides_default_after_declaration(): T.func_attr({"global_symbol": "custom"}) T.evaluate(0) assert builder.get()["main"].attrs["global_symbol"] == "custom" + + [email protected]("dialect", [T, R]) +def test_module_alias_is_the_native_frame_and_lookup_keeps_reference(dialect): + # Before: cls = Module; cls.callee + # Expected builder program: cls = X.bind_(Module, name="cls") + # The binding returns the same module frame and its plain native GlobalVar. + with IRBuilder() as builder, I.ir_module() as module: + with dialect.function(decl=True) as function: + dialect.func_name("callee") + annotation = T.int32() if dialect is T else R.Tensor((4,), "float32") + argument = dialect.arg("x", annotation) + depth = len(builder.frames) + alias = dialect.bind_(module, name="cls") + assert alias is module + assert len(builder.frames) == depth + assert alias.callee.same_as(function.reference) + assert isinstance(alias.callee, ir.GlobalVar) + with pytest.raises(AttributeError, match="missing"): + alias.missing + with function: + if dialect is T: + T.evaluate(argument) + else: + R.func_ret_value(argument) + assert builder.get().get_global_var("callee").same_as(function.reference) + + [email protected]("alias", [False, True]) [email protected]( + "dialect, decorator, annotation", + [(T, "T.prim_func", "T.int32"), (R, "R.function", 'R.Tensor((4,), "float32")')], +) +def test_module_member_calls_use_the_callers_native_dialect( + monkeypatch, dialect, decorator, annotation, alias +): + # Before: cls = Module; return cls.callee(x) + # Expected builder program: + # cls = X.bind_(Module, name="cls") + # X.return_(X.call_global_var_(cls.callee, [x])) + calls = [] + original = dialect.call_global_var_ + + def call(reference, arguments): + calls.append((reference, arguments)) + return original(reference, arguments) + + monkeypatch.setattr(dialect, "call_global_var_", call) + setup, owner = ("cls = Module", "cls") if alias else ("pass", "Module") + module = parser.parse(f""" [email protected]_module +class Module: + @{decorator} + def caller(x: {annotation}) -> {annotation}: + {setup} + return {owner}.callee(x) + @{decorator} + def callee(y: {annotation}) -> {annotation}: + return y +""") + assert len(calls) == 1 + reference, arguments = calls[0] + assert reference.same_as(module.get_global_var("callee")) + assert len(arguments) == 1 + assert arguments[0].same_as(module["caller"].params[0]) + + +def test_lexical_range_binding_preserves_custom_native_loop_extent(): + # Before: range = custom_range; for i in range(4): X.evaluate(i) + # Expected builder program: range = X.bind_(custom_range, name="range") + # with X.for_(range(4), names="i") as i: X.emit_(X.evaluate(i)) + calls = [] + + def custom_range(extent): + calls.append(extent) + return T.serial(2) + + function = parser.parse( + "@T.prim_func\ndef main():\n range = custom_range\n" + " for i in range(4):\n T.evaluate(i)\n", + extra_vars={"custom_range": custom_range}, + ) + assert calls == [4] + assert isinstance(function.body, tirx.For) + assert int(function.body.extent) == 2 + assert function.body.body.value.same_as(function.body.loop_var)
