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 af6dbafed2faa43eac7ce43d115dcceb271adb5a
Author: Tianqi Chen <[email protected]>
AuthorDate: Tue Sep 22 22:02:05 2026 +0000

    [FR][Relax] Preserve local function regions and recursive signatures
    
    Keep nested function declarations in their owning binding region and 
initialize recursive local function signatures before lowering their bodies.
---
 src/relax/script/builder/frame.cc           |  39 +++++++++-
 tests/script/test_parser_local_functions.py | 117 ++++++++++++++++++++++++++++
 2 files changed, 154 insertions(+), 2 deletions(-)

diff --git a/src/relax/script/builder/frame.cc 
b/src/relax/script/builder/frame.cc
index a443f2de86..c8e5678cd5 100644
--- a/src/relax/script/builder/frame.cc
+++ b/src/relax/script/builder/frame.cc
@@ -85,7 +85,10 @@ void FunctionFrameNode::ExitWithScope() {
     function = tvm::relax::Function::CreateEmpty(params, 
ret_ty.value_or(tvm::relax::AnyType()),
                                                  is_pure.value_or(true), 
DictAttrs(attrs), span);
     if (local) {
-      local_var = tvm::Var(name.value(), 
tvm::relax::GetType(function.value()), span);
+      auto ty = tvm::relax::GetType(function.value());
+      local_var = CheckBindingBlockFrameExistAndUnended()->is_dataflow
+                      ? tvm::relax::DataflowVar(name.value(), ty, span)
+                      : tvm::Var(name.value(), ty, span);
     } else {
       global_var = ir::DeclFunction(name.value(), function.value());
     }
@@ -115,9 +118,38 @@ void FunctionFrameNode::ExitWithScope() {
   if (local) {
     TVM_FFI_CHECK(local_var.has_value(), ValueError)
         << "A local function definition requires its declared reference";
+    bool recursive = false;
+    for (const tvm::Var& var : tvm::relax::FreeVars(func)) {
+      recursive = recursive || var.same_as(local_var.value());
+    }
     // Retain the declared reference identity while publishing its inferred 
type,
     // just as DefFunction refines a module's global reference after 
definition.
-    local_var.value()->ty = tvm::relax::GetType(func);
+    Type reference_type = tvm::relax::GetType(func);
+    if (recursive) {
+      // A recursive reference has its own provisional signature.  Its formal
+      // primitive parameters must not bind the definition's parameter objects.
+      // Keep lexical captures intact while renewing these signature binders.
+      ffi::Map<tvm::Var, tvm::Expr> signature_params;
+      for (const tvm::Var& param : params) {
+        if (param.as<PrimVar>()) {
+          signature_params.Set(param, param.CopyWithName(param->name));
+        }
+      }
+      if (!signature_params.empty()) {
+        auto signature = reference_type.as_or_throw<tvm::relax::FuncType>();
+        auto bind_param = [&](const Type& ty) { return tvm::relax::Bind(ty, 
signature_params); };
+        reference_type =
+            tvm::relax::FuncType(signature->params.value().Map(bind_param),
+                                 bind_param(signature->ret), 
signature->purity, signature->span);
+      }
+      if (local_var.value()->IsInstance<tvm::relax::DataflowVarNode>()) {
+        // The self-reference is captured by the nested function, so it must
+        // survive the dataflow region.  Region finalization rewrites every use
+        // together, including recursive calls, to the same ordinary Var.
+        
CheckBindingBlockFrameExistAndUnended()->output_vars.push_back(local_var.value());
+      }
+    }
+    local_var.value()->ty = reference_type;
     EmitVarBinding(tvm::relax::VarBinding(local_var.value(), func, span));
   } else if (!builder->HasConstructionFrames()) {
     // Case 0. No outer frame, return function directly
@@ -204,6 +236,9 @@ void BindingBlockFrameNode::ExitWithScope() {
     ffi::Array<tvm::Var> new_output_vars;
     std::unordered_map<tvm::Var, tvm::Var, ffi::ObjectPtrHash, 
ffi::ObjectPtrEqual> var_remap;
     for (const auto& output_var : output_vars) {
+      if (var_remap.count(output_var)) {
+        continue;
+      }
       tvm::Var new_output_var(output_var->name, 
tvm::relax::GetType(output_var), output_var->span);
       new_output_vars.push_back(new_output_var);
       if (auto span = binding_spans.Get(output_var)) {
diff --git a/tests/script/test_parser_local_functions.py 
b/tests/script/test_parser_local_functions.py
new file mode 100644
index 0000000000..ad1a6aacb4
--- /dev/null
+++ b/tests/script/test_parser_local_functions.py
@@ -0,0 +1,117 @@
+# Licensed to the Apache Software Foundation (ASF) under one
+# or more contributor license agreements.  See the NOTICE file
+# distributed with this work for additional information
+# regarding copyright ownership.  The ASF licenses this file
+# to you under the Apache License, Version 2.0 (the
+# "License"); you may not use this file except in compliance
+# with the License.  You may obtain a copy of the License at
+#
+#   http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing,
+# software distributed under the License is distributed on an
+# "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY
+# KIND, either express or implied.  See the License for the
+# specific language governing permissions and limitations
+# under the License.
+"""Local Relax functions retain their enclosing binding region and 
identities."""
+
+import textwrap
+
+import pytest
+
+import tvm
+from tvm import relax
+from tvm.script.parser import parse
+
+
+def _local_binding(function):
+    return next(
+        binding
+        for block in function.body.blocks
+        for binding in block.bindings
+        if isinstance(binding.value, relax.Function)
+    )
+
+
[email protected]("dataflow", [False, True])
[email protected]("recursive", [False, True])
+def test_local_function_binding_region_and_reference_identity(dataflow, 
recursive):
+    inner_return = "inner(y)" if recursive else "R.add(x, y)"
+    region = f"""\
[email protected]
+def inner(y: R.Tensor((2,), "float32")) -> R.Tensor((2,), "float32"):
+    return {inner_return}
+result = inner(x)
+"""
+    if dataflow:
+        region = "with R.dataflow():\n" + textwrap.indent(region + 
"R.output(result)\n", "    ")
+    source = (
+        '@R.function\ndef main(x: R.Tensor((2,), "float32")):\n'
+        + textwrap.indent(region, "    ")
+        + "    return result\n"
+    )
+    function = parse(source)
+    binding = _local_binding(function)
+    assert isinstance(binding.var, relax.DataflowVar) == (dataflow and not 
recursive)
+    call = function.body.blocks[0].bindings[-1].value
+    assert call.op.same_as(binding.var)
+    inner_call = binding.value.body.blocks[0].bindings[0].value
+    if recursive:
+        assert inner_call.op.same_as(binding.var)
+    else:
+        assert inner_call.args[0].same_as(function.params[0])
+        assert inner_call.args[1].same_as(binding.value.params[0])
+    relax.analysis.well_formed(tvm.IRModule({"main": function}))
+
+
+def test_explicitly_output_local_function_keeps_recursive_reference():
+    function = parse("""
[email protected]
+def main():
+    with R.dataflow():
+        @R.function
+        def inner(y: R.Tensor((2,), "float32")) -> R.Tensor((2,), "float32"):
+            return inner(y)
+        R.output(inner)
+    return inner
+""")
+    binding = _local_binding(function)
+    assert not isinstance(binding.var, relax.DataflowVar)
+    assert function.body.body.same_as(binding.var)
+    inner_call = binding.value.body.blocks[0].bindings[0].value
+    assert inner_call.op.same_as(binding.var)
+    relax.analysis.well_formed(tvm.IRModule({"main": function}))
+
+
[email protected]("recursive", [False, True])
+def 
test_local_dependent_signature_preserves_parameter_and_capture_scopes(recursive):
+    inner_return = "inner(current, value)" if recursive else "value"
+    function = parse(f"""
[email protected]
+def main(n: R.Prim("int64"), m: R.Prim("int64"), x: R.Tensor((n, m), 
"float32")):
+    @R.function
+    def inner(current: R.Prim("int64"), value: R.Tensor((current, m), 
"float32")) -> R.Tensor(
+        (current, m), "float32"
+    ):
+        return {inner_return}
+    return inner(n, x)
+""")
+    binding = _local_binding(function)
+    current, value = binding.value.params
+    assert value.ty.shape[0].same_as(current)
+    assert binding.value.ret_ty.shape[0].same_as(current)
+    signature_current = binding.var.ty.params[1].shape[0]
+    assert signature_current.same_as(current) != recursive
+    assert binding.var.ty.ret.shape[0].same_as(signature_current)
+    assert binding.var.ty.params[1].shape[1].same_as(function.params[1])
+    assert binding.var.ty.ret.shape[1].same_as(function.params[1])
+    tvm.ir.assert_structural_equal(binding.var.ty, binding.value.ty)
+    call = function.body.blocks[0].bindings[-1].value
+    assert call.op.same_as(binding.var)
+    assert call.ty.shape[0].same_as(function.params[0])
+    if recursive:
+        recursive_call = binding.value.body.blocks[0].bindings[0].value
+        assert recursive_call.op.same_as(binding.var)
+        assert recursive_call.ty.shape[0].same_as(current)
+    relax.analysis.well_formed(tvm.IRModule({"main": function}))

Reply via email to