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 77ffe0934b4fb44ad903fabc41ce9320519da49e Author: Tianqi Chen <[email protected]> AuthorDate: Tue Sep 22 22:02:05 2026 +0000 [FR][TVMScript] Preserve public validation and diagnostic boundaries Retain decorator factories, function validation, Optional specialization boundaries, and source diagnostic coordinates through canonical construction. --- python/tvm/script/ir_builder/construction.py | 9 +-- python/tvm/script/parser/diagnostics.py | 1 + python/tvm/script/parser/frontend.py | 91 +++++++++++++++++++++++++--- python/tvm/script/parser/jit.py | 8 +-- 4 files changed, 90 insertions(+), 19 deletions(-) diff --git a/python/tvm/script/ir_builder/construction.py b/python/tvm/script/ir_builder/construction.py index 69a8b091c6..3c3096af44 100644 --- a/python/tvm/script/ir_builder/construction.py +++ b/python/tvm/script/ir_builder/construction.py @@ -41,7 +41,7 @@ _ABSENT_PARAMETERS = ContextVar("tvm_builder_absent_parameters", default=None) @contextmanager def specialization_context(name, bindings): """Pass validated JIT bindings to one root builder execution.""" - token = _SPECIALIZATION.set((name, dict(bindings))) + token = _SPECIALIZATION.set(None if bindings is None else (name, dict(bindings))) try: yield finally: @@ -148,9 +148,8 @@ class FunctionRecord: ): self.symbols.bind(captured_name, value) context = _SPECIALIZATION.get() - self.specialization = ( - context[1] if context is not None and context[0] == name and not local else {} - ) + self.is_specialization = context is not None and context[0] == name and not local + self.specialization = context[1] if self.is_specialization else {} self.captured_bindings = dict(captures or {}) if not local else {} absent = _ABSENT_PARAMETERS.get() self.absent_parameters = ( @@ -191,6 +190,8 @@ class FunctionRecord: # Optional is an annotation wrapper, owned by the JIT entry point. unwrap = getattr(annotation, "__tvm_optional_annotation__", None) if unwrap is not None: + if not self.is_specialization: + raise TypeError("T.Optional is only supported by @T.jit") annotation = unwrap() return self.parameter(name, annotation, location) diff --git a/python/tvm/script/parser/diagnostics.py b/python/tvm/script/parser/diagnostics.py index af25eebc35..35987d38d3 100644 --- a/python/tvm/script/parser/diagnostics.py +++ b/python/tvm/script/parser/diagnostics.py @@ -37,6 +37,7 @@ def diagnostic_error(error, compiler): location = getattr(error, "__tvm_script_location__", None) if location is not None: filename, start, end, column, end_column = location + column, end_column = column - 1, end_column - 1 elif isinstance(error, SyntaxError) and error.lineno: start = error.lineno end = error.end_lineno or start diff --git a/python/tvm/script/parser/frontend.py b/python/tvm/script/parser/frontend.py index 9a40aa60ad..347148d809 100644 --- a/python/tvm/script/parser/frontend.py +++ b/python/tvm/script/parser/frontend.py @@ -327,7 +327,12 @@ def make_decorator(builder, *, option_map=None, defaults=None): function.__tvm_function_options__ = options if deferred: return function - return parse(function) + result = parse( + function, + check_well_formed=options.get("check_well_formed", True), + ) + result.__name__ = function.__name__ + return result return apply(function) if function is not None else apply @@ -607,8 +612,8 @@ class Compiler: for value in ( node.lineno, node.end_lineno, - node.col_offset, - node.end_col_offset, + node.col_offset + 1, + node.end_col_offset + 1, ) ], ], @@ -921,17 +926,81 @@ def parse(source, extra_vars=None, *, filename=None, track_span: bool = True, ** try: root = compiler.tree.body[-1] root_name = root.name if isinstance(root, ast.FunctionDef) else None - with construction.specialization_context( - root_name, options.get("_specialization_bindings", {}) - ): + specialization = options.get("_specialization_bindings") + if specialization is None and options.get("absent_params") is not None: + specialization = {} + check_well_formed = options.get("check_well_formed") + if check_well_formed is None: + check_well_formed = True + for decorator in getattr(root, "decorator_list", ()): + if isinstance(decorator, ast.Call): + for keyword in decorator.keywords: + if keyword.arg == "check_well_formed": + check_well_formed = eval( + compile(ast.Expression(keyword.value), compiler.filename, "eval"), + compiler.env, + ) + with construction.specialization_context(root_name, specialization): with construction.absent_parameters(root_name, options.get("absent_params")): - return compiler.build() + result = compiler.build() + if check_well_formed: + _check_well_formed(result) + return result except DiagnosticError: raise except Exception as error: raise diagnostic_error(error, compiler) from error +def _check_well_formed(result): + """Apply the public entry point's default validation to constructed IR.""" + from tvm import ir, relax, s_tir, tirx + + message = ( + "Program is not well-formed. If this is deliberate, set " + "check_well_formed=False in the top-level decorator." + ) + if isinstance(result, ir.IRModule | relax.Function): + if not relax.analysis.check_well_formed(result): + raise ValueError(message) + if not isinstance(result, ir.IRModule | relax.Function | tirx.PrimFunc): + return + module = result if isinstance(result, ir.IRModule) else ir.IRModule.from_expr(result) + try: + s_tir.analysis.verify_well_formed(module) + for function in module.functions.values(): + if isinstance(function, tirx.PrimFunc) and not function.attrs.get("s_tir", False): + tirx.analysis.verify_tirx_well_formed(function) + except Exception as error: + raise ValueError(f"{message}\n{error}") from error + + +class _PyModuleFactory: + """Keep executable Python attachments on each fresh module instance.""" + + def __init__(self, module, original_class): + self.ir_module = module + self.original_class = original_class + self.pyfunc_methods = list(getattr(module, "pyfuncs", {})) + self.__name__ = original_class.__name__ + + def __call__(self, device=None, target=None): + from tvm import cpu, ir + from tvm.relax.base_py_module import BasePyModule + + source = self.ir_module + instance_module = ir.IRModule( + source.functions, attrs=source.attrs, global_infos=source.global_infos + ) + instance = BasePyModule(instance_module, device or cpu(0), target) + for name in self.pyfunc_methods: + instance.add_python_function(name, getattr(self.original_class, name)) + return instance + + def __getattr__(self, name): + return getattr(self.ir_module, name) + + def ir_module(module=None, **options): """Decorate a Python class with two-phase module construction. @@ -971,7 +1040,13 @@ def ir_module(module=None, **options): definition_scope = _definition_scope(frame) finally: del frame - return parse(module, _definition_scope=definition_scope, **options) + result = parse(module, _definition_scope=definition_scope, **options) + from tvm.relax.base_py_module import BasePyModule + + if issubclass(module, BasePyModule): + return _PyModuleFactory(result, module) + result.__name__ = module.__name__ + return result return apply(module) if module is not None else apply diff --git a/python/tvm/script/parser/jit.py b/python/tvm/script/parser/jit.py index 655726e74c..df522d5d6a 100644 --- a/python/tvm/script/parser/jit.py +++ b/python/tvm/script/parser/jit.py @@ -207,14 +207,8 @@ class TIRJit: self._closure_vars, _definition_scope=self._definition_scope, _specialization_bindings={**effective, **absent_params}, + check_well_formed=self.check_well_formed, ) - if self.check_well_formed: - from tvm.s_tir.analysis import verify_well_formed - from tvm.tirx.analysis import verify_tirx_well_formed - - verify_well_formed(prim_func) - if not prim_func.attrs.get("s_tir", False): - verify_tirx_well_formed(prim_func) setattr(prim_func, "__name__", self.func.__name__) self._cache[cache_key] = prim_func return prim_func
