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}))
