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

Reply via email to