This is an automated email from the ASF dual-hosted git repository.

tqchen pushed a commit to branch tvmscript-generic-parser-builder
in repository https://gitbox.apache.org/repos/asf/tvm.git

commit c96690797d4d75bf7abb7697bff2301620cd0883
Author: Tianqi Chen <[email protected]>
AuthorDate: Sun Sep 20 22:35:52 2026 +0000

    [TVMScript] Add concrete Relax builder conventions and signature frames
---
 include/tvm/relax/script/builder/frame.h       |  18 +-
 include/tvm/relax/script/builder/ir.h          |  18 ++
 python/tvm/relax/script/builder_v2/__init__.py | 328 +++++++++++++++++++++++++
 src/relax/script/builder/frame.cc              |  66 ++++-
 src/relax/script/builder/ir.cc                 |  57 ++++-
 src/relax/script/builder/utils.h               |  12 +-
 6 files changed, 481 insertions(+), 18 deletions(-)

diff --git a/include/tvm/relax/script/builder/frame.h 
b/include/tvm/relax/script/builder/frame.h
index a2f663fb50..357c8a31a3 100644
--- a/include/tvm/relax/script/builder/frame.h
+++ b/include/tvm/relax/script/builder/frame.h
@@ -36,6 +36,10 @@ namespace relax {
 /*! \brief The base ir_builder frame for the relax dialect. */
 class RelaxFrameNode : public IRBuilderFrameNode {
  public:
+  /*! \brief Source range captured when this frame is entered. */
+  Span source_span;
+
+  void EnterWithScope() override;
   static void RegisterReflection() {
     namespace refl = tvm::ffi::reflection;
     refl::ObjectDef<RelaxFrameNode>();
@@ -117,6 +121,13 @@ class FunctionFrameNode : public SeqExprFrameNode {
   ffi::Map<ffi::String, Any> attrs;
   /*! \brief The block builder to create Relax function. */
   tvm::relax::BlockBuilder block_builder;
+  /*! \brief Whether this frame constructs only a function signature. */
+  bool declaration = false;
+  bool local = false;
+  ffi::Optional<tvm::Var> local_var;
+  /*! \brief Finalized function and its stable module reference. */
+  ffi::Optional<tvm::relax::Function> function;
+  ffi::Optional<tvm::GlobalVar> global_var;
 
   static void RegisterReflection() {
     namespace refl = tvm::ffi::reflection;
@@ -125,7 +136,10 @@ class FunctionFrameNode : public SeqExprFrameNode {
         .def_ro("params", &FunctionFrameNode::params)
         .def_ro("ret_ty", &FunctionFrameNode::ret_ty)
         .def_ro("is_pure", &FunctionFrameNode::is_pure)
-        .def_ro("attrs", &FunctionFrameNode::attrs);
+        .def_ro("attrs", &FunctionFrameNode::attrs)
+        .def_ro("function", &FunctionFrameNode::function)
+        .def_ro("global_var", &FunctionFrameNode::global_var)
+        .def_ro("local_var", &FunctionFrameNode::local_var);
     // `binding_blocks` and `output` are inherited from SeqExprFrameNode.
     // `block_builder` is not registered as it's not visited.
   }
@@ -163,6 +177,8 @@ class BindingBlockFrameNode : public RelaxFrameNode {
    * \note Only used for a dataflow block.
    */
   ffi::Array<tvm::Var> output_vars;
+  /*! \brief Statement ranges for explicitly emitted bindings. */
+  ffi::Map<tvm::Var, Span> binding_spans;
 
   static void RegisterReflection() {
     namespace refl = tvm::ffi::reflection;
diff --git a/include/tvm/relax/script/builder/ir.h 
b/include/tvm/relax/script/builder/ir.h
index 2516dc0134..b52dcd0ae0 100644
--- a/include/tvm/relax/script/builder/ir.h
+++ b/include/tvm/relax/script/builder/ir.h
@@ -39,6 +39,15 @@ namespace relax {
  */
 TVM_DLL FunctionFrame Function(bool is_pure, bool is_private);
 
+/*! \brief Start a bodyless declaration using the normal signature operations. 
*/
+TVM_DLL FunctionFrame DeclFunction(bool is_pure, bool is_private, bool local);
+
+/*! \brief Define a local function under its cached declaration identity. */
+TVM_DLL FunctionFrame LocalFunction(bool is_pure, const tvm::Var& reference);
+
+/*! \brief Add a cached parameter without changing its identity. */
+TVM_DLL tvm::Var ArgVar(const ffi::String& name, const tvm::Var& var);
+
 /*!
  * \brief Add a parameter to the last function frame.
  * \param name The name of the parameter.
@@ -117,6 +126,15 @@ TVM_DLL tvm::Var EmitMatchCast(const tvm::relax::Expr& 
value, const tvm::Type& t
  */
 TVM_DLL tvm::Var EmitVarBinding(const tvm::relax::VarBinding& binding);
 
+/*! \brief Emit a binding with separate statement and variable-name ranges. */
+TVM_DLL tvm::Var EmitV2(const tvm::relax::Expr& value,
+                        const ffi::Optional<tvm::Type>& annotate_ty,
+                        const ffi::Optional<Span>& name_span);
+
+/*! \brief Emit a match cast with separate statement and variable-name ranges. 
*/
+TVM_DLL tvm::Var EmitMatchCastV2(const tvm::relax::Expr& value, const 
tvm::Type& ty,
+                                 const ffi::Optional<Span>& name_span);
+
 ///////////////////////////// If Then Else /////////////////////////////
 
 /*!
diff --git a/python/tvm/relax/script/builder_v2/__init__.py 
b/python/tvm/relax/script/builder_v2/__init__.py
new file mode 100644
index 0000000000..fafeb2fd47
--- /dev/null
+++ b/python/tvm/relax/script/builder_v2/__init__.py
@@ -0,0 +1,328 @@
+# 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.
+"""Concrete Relax construction operations over the shared native builder 
stack."""
+
+# pylint: disable=wildcard-import,redefined-builtin,invalid-name
+import builtins as _python
+import numbers as _numbers
+
+import tvm_ffi as _ffi
+
+from tvm import ir as _ir
+from tvm import relax as _relax
+from tvm.script.ir_builder import IRBuilder as _IRBuilder
+from tvm.script.ir_builder import ir as _I
+from tvm.script.ir_builder import protocol as _protocol
+
+from .. import builder as _legacy
+from ..builder import *
+from ..builder import _ffi_api
+from ..builder import frame as _frame
+
+
+@_protocol.expression_args("shape", introduce=True, dtype="int64", 
scalar_strings=False)
+def Tensor(shape=None, dtype=None, vdevice=None, ndim=-1, *, span=None):
+    """Construct a concrete tensor type from already resolved dimensions."""
+    if isinstance(shape, _python.str) and dtype is None:
+        dtype, shape = shape, None
+    if isinstance(vdevice, _python.str):
+        target, _, index = vdevice.partition(":")
+        vdevice = _I.lookup_vdevice(target, int(index) if index else 0)
+    return _relax.TensorType(shape, dtype, vdevice, ndim, span)
+
+
+@_protocol.expression_args("values", introduce=True, dtype="int64")
+def Shape(values=None, ndim=-1, *, span=None):
+    """Construct a concrete shape type."""
+    return _relax.ShapeType(values, ndim, span)
+
+
+def _type(value):
+    if value is None:
+        return _ir.TupleType([])
+    if callable(value):
+        value = value()
+    if _ir.is_prim_expr(value):
+        value = value.ty
+    if not isinstance(value, _ir.Type):
+        raise TypeError(f"Expected a concrete type, got 
{type(value).__name__}")
+    return value
+
+
+def Callable(params=None, ret=None, purity=None, derive_func=None, *, 
span=None):
+    """Construct a concrete function type."""
+    if purity is None:
+        purity = params is not None
+    if params is None:
+        return _relax.FuncType.opaque_func(
+            ret=None if ret is None else _type(ret),
+            derive_func=derive_func,
+            purity=purity,
+            span=span,
+        )
+    if derive_func is not None:
+        raise ValueError("A derivation function requires an opaque callable")
+    if not isinstance(params, list | _python.tuple):
+        params = [params]
+    return _relax.FuncType([_type(param) for param in params], _type(ret), 
purity, span)
+
+
+def Tuple(*fields, span=None):
+    """Construct a concrete tuple type."""
+    if len(fields) == 1 and isinstance(fields[0], list | _python.tuple):
+        fields = fields[0]
+    return _ir.TupleType([_type(field) for field in fields], span)
+
+
+def Prim(dtype, *, span=None):
+    """Construct a primitive type."""
+    return _ir.PrimType(dtype)
+
+
+def Object(*, span=None):
+    """Construct the unconstrained Relax value type."""
+    return _relax.AnyType(span)
+
+
+Any = Object
+
+
+def type_var(name, *, dtype=None, span=None):
+    """Construct a signature symbol under Relax's default shape dtype 
policy."""
+    return _ir.Var(name, "int64" if dtype is None else dtype, span)
+
+
+class _Frame:
+    """Retain source metadata and exports around an existing native frame."""
+
+    def __init__(self, native, span=None):
+        self.native = native
+        self.span = span
+        self.result = {}
+
+    def __getattr__(self, name):
+        return getattr(self.native, name)
+
+    @property
+    def reference(self):
+        """Return the stable module or local function reference after 
declaration."""
+        if isinstance(self.native, _frame.FunctionFrame):
+            local_var = self.native.local_var
+            return local_var if local_var is not None else 
self.native.global_var
+        raise AttributeError("This frame does not declare a function")
+
+    def __enter__(self):
+        with _protocol.span_context(self.span):
+            self.native.__enter__()
+        return self
+
+    def __exit__(self, exc_type, exc_value, traceback):
+        with _protocol.span_context(self.span):
+            self.native.__exit__(exc_type, exc_value, traceback)
+        if exc_type is None:
+            if isinstance(self.native, _frame.BindingBlockFrame):
+                self.result = {var.name: var for var in 
self.native.output_vars}
+            elif (
+                isinstance(self.native, _frame.FunctionFrame) and 
self.native.local_var is not None
+            ):
+                self.result = {self.native.name: self.native.local_var}
+            elif isinstance(self.native, _frame.IfFrame):
+                self.result = {self.native.var_name: self.native.var}
+        return False
+
+
+def function(is_pure=True, is_private=False, *, local=False, reference=None, 
span=None):
+    """Enter a definition using the native Relax function frame."""
+    if local:
+        if reference is None:
+            raise ValueError("A local function requires its declared 
reference")
+        return _Frame(_ffi_api.LocalFunction(is_pure, reference), span)
+    return _Frame(_legacy.function(is_pure, is_private), span)
+
+
+def decl_function(is_pure=True, is_private=False, *, local=False, span=None):
+    """Declare a bodyless function with the same signature operations as a 
definition."""
+    return _Frame(_ffi_api.DeclFunction(is_pure, is_private, local), span)
+
+
+def arg(name, ty, *, span=None):
+    """Add a parameter, retaining a cached parameter's identity when 
supplied."""
+    with _protocol.span_context(span):
+        if isinstance(ty, _ir.Var):
+            return _ffi_api.ArgVar(name, ty)
+        return _protocol.at(span, _legacy.arg(name, _type(ty)))
+
+
+def func_ret_type(ret_ty):
+    """Set the concrete return type of the active declaration or definition."""
+    return _legacy.func_ret_type(_type(ret_ty))
+
+
+func_ret_ty = func_ret_type
+
+
+def dataflow(*, span=None):
+    """Create a dataflow region whose result maps exported names to finalized 
vars."""
+    return _Frame(_legacy.dataflow(), span)
+
+
+def If(condition, *, span=None):
+    """Create a conditional region with finalized named exports."""
+    return _Frame(_legacy.If(condition), span)
+
+
+def Then(*, span=None):
+    """Create the true branch of a conditional."""
+    return _Frame(_legacy.Then(), span)
+
+
+def Else(*, span=None):
+    """Create the false branch of a conditional."""
+    return _Frame(_legacy.Else(), span)
+
+
+def _check_unterminated():
+    for frame in reversed(_IRBuilder.current().frames):
+        if isinstance(frame, _frame.FunctionFrame):
+            if frame.output is not None:
+                raise ValueError("A Relax operation cannot follow an 
unconditional return")
+            break
+
+
+def _value(value, ty=None):
+    if isinstance(value, _python.tuple):
+        return _relax.utils.convert_to_expr(value)
+    if isinstance(value, _numbers.Number):
+        if isinstance(ty, _ir.PrimType):
+            return _relax.prim_value(value, dtype=ty.dtype)
+        return _relax.const(value)
+    return value
+
+
+def bind_(
+    value=_protocol.MISSING,
+    *,
+    ty=None,
+    name=None,
+    span=None,
+    name_span=None,
+    previous=_protocol.MISSING,
+):
+    """Emit an immutable Relax binding and return the newly bound value."""
+    if value is _protocol.MISSING:
+        raise ValueError("Relax bindings require an initializer")
+    if isinstance(value, _I.meta_var):
+        return value.value
+    _check_unterminated()
+    ty = None if ty is None else _type(ty)
+    value = _value(value, ty)
+    with _protocol.span_context(span):
+        if isinstance(value, _relax.MatchCast):
+            if ty is not None and not _ffi.structural_equal(ty, value.ty):
+                raise TypeError("The binding annotation differs from the 
match-cast type")
+            result = _ffi_api.EmitMatchCastV2(value.value, value.ty, name_span)
+        elif isinstance(value, _relax.Expr):
+            result = _ffi_api.EmitV2(value, ty, name_span)
+        else:
+            return value
+    if name is not None:
+        _IRBuilder.name(name, result)
+    return _protocol.at(name_span if name_span is not None else span, result)
+
+
+def emit_(value, *, span=None):
+    """Consume an expression statement; effect-only operations return None."""
+    if value is not None:
+        bind_(value, span=span)
+
+
+def return_(value=None, *, span=None):
+    """Record the function result without exiting Python construction."""
+    _check_unterminated()
+    with _protocol.span_context(span):
+        if value is None:
+            value = _relax.Tuple([])
+        _legacy.func_ret_value(_value(value))
+
+
+def match_cast(value, ty, *, span=None):
+    """Construct a concrete match-cast binding for bind_ to consume."""
+    if value is None:
+        raise ValueError("The match-cast value cannot be None")
+    ty = _type(ty)
+    return _relax.MatchCast(_ir.Var("", ty), _value(value), ty, span)
+
+
+def unpack(value):
+    """Project an IR tuple with known arity; leave Python iteration 
unchanged."""
+    if isinstance(value, _relax.Tuple):
+        return _python.tuple(value.fields)
+    if isinstance(value, _relax.Expr) and isinstance(value.ty, _ir.TupleType):
+        return _python.tuple(_relax.TupleGetItem(value, i) for i in 
range(len(value.ty.fields)))
+    return value
+
+
+def assert_(condition, message="", *, span=None):
+    """Construct a runtime assertion with construction-time diagnostic text."""
+    if not isinstance(message, _python.str):
+        raise TypeError("An assertion message must be construction-time text")
+    with _protocol.span_context(span):
+        emit_(_protocol.at(span, _legacy.assert_op(condition, 
format=message)), span=span)
+
+
+def For(*args, span=None, **kwargs):
+    """Reject imperative loops in the Relax expression dialect."""
+    raise TypeError("Relax does not support imperative for loops")
+
+
+def break_(*, span=None):
+    """Reject imperative loop control in Relax."""
+    raise TypeError("Relax does not support break")
+
+
+def continue_(*, span=None):
+    """Reject imperative loop control in Relax."""
+    raise TypeError("Relax does not support continue")
+
+
+def setitem(target, index, value, *, span=None):
+    """Reject mutable stores, which are not a Relax binding operation."""
+    raise TypeError("Relax does not support indexed assignment")
+
+
+__all__ = [
+    *_legacy.ir.__all__,
+    "Any",
+    "Callable",
+    "For",
+    "Object",
+    "Prim",
+    "Shape",
+    "Tensor",
+    "Tuple",
+    "assert_",
+    "break_",
+    "continue_",
+    "bind_",
+    "decl_function",
+    "emit_",
+    "match_cast",
+    "return_",
+    "setitem",
+    "type_var",
+    "unpack",
+]
diff --git a/src/relax/script/builder/frame.cc 
b/src/relax/script/builder/frame.cc
index 798ce96f7d..acacdb7eef 100644
--- a/src/relax/script/builder/frame.cc
+++ b/src/relax/script/builder/frame.cc
@@ -41,6 +41,11 @@ TVM_FFI_STATIC_INIT_BLOCK() {
   ElseFrameNode::RegisterReflection();
 }
 
+void RelaxFrameNode::EnterWithScope() {
+  source_span = IRBuilder::Current()->GetCurrentSourceSpan();
+  IRBuilderFrameNode::EnterWithScope();
+}
+
 void SeqExprFrameNode::ExitWithScope() {
   // At this moment, there should be at most one BindingBlockFrame which 
hasn't ended. In this case,
   // call its `ExitBindingBlockFrame` and check if there is any more unended 
BindingBlockFrame.
@@ -60,20 +65,40 @@ void SeqExprFrameNode::EnterWithScope() {
 
 void FunctionFrameNode::EnterWithScope() {
   this->block_builder->BeginScope(params);
-  SeqExprFrameNode::EnterWithScope();
+  if (declaration) {
+    RelaxFrameNode::EnterWithScope();
+  } else {
+    SeqExprFrameNode::EnterWithScope();
+  }
 }
 
 void FunctionFrameNode::ExitWithScope() {
   using ir::IRModuleFrame;
   using tvm::relax::Expr;
   IRBuilder builder = IRBuilder::Current();
+  if (declaration) {
+    TVM_FFI_CHECK(name.has_value(), ValueError) << "A function declaration 
requires a name";
+    TVM_FFI_CHECK(local || builder->FindFrame<IRModuleFrame>().has_value(), 
ValueError)
+        << "A function declaration requires an IRModule frame";
+    RelaxFrameNode::ExitWithScope();
+    block_builder->EndScope();
+    function = tvm::relax::Function::CreateEmpty(
+        params, ret_ty.value_or(tvm::relax::AnyType()), is_pure.value_or(true),
+        DictAttrs(attrs), source_span);
+    if (local) {
+      local_var = tvm::Var(name.value(), 
tvm::relax::GetType(function.value()), source_span);
+    } else {
+      global_var = ir::DeclFunction(name.value(), function.value());
+    }
+    return;
+  }
   SeqExprFrameNode::ExitWithScope();
   // Step 1: Create the function.
   TVM_FFI_CHECK(output.has_value(), ValueError)
       << "A Relax function must have a return value. Please use "
          "`return` to return an Expr";
 
-  Expr body = 
this->block_builder->Normalize(tvm::relax::SeqExpr(binding_blocks, 
output.value()));
+  Expr body = 
this->block_builder->Normalize(tvm::relax::SeqExpr(binding_blocks, 
output.value(), source_span));
   // if the function is not private, add a global symbol to its attributes
   if (!is_private.value_or(false) && name.has_value() && 
!attrs.count(tvm::attr::kGlobalSymbol)) {
     attrs.Set(tvm::attr::kGlobalSymbol, name.value());
@@ -83,9 +108,15 @@ void FunctionFrameNode::ExitWithScope() {
                             /*body=*/body,
                             /*ret_ty=*/ret_ty,
                             /*is_pure=*/is_pure.value_or(true),
-                            /*attrs=*/DictAttrs(attrs));
+                            /*attrs=*/DictAttrs(attrs),
+                            /*span=*/source_span);
+  function = func;
   // Step 2: Update IRModule.
-  if (builder->frames.empty()) {
+  if (local) {
+    TVM_FFI_CHECK(local_var.has_value(), ValueError)
+        << "A local function definition requires its declared reference";
+    EmitVarBinding(tvm::relax::VarBinding(local_var.value(), func, 
source_span));
+  } else if (builder->frames.empty()) {
     // Case 0. No outer frame, return function directly
     TVM_FFI_CHECK(!builder->result.has_value(), ValueError)
         << "Builder.result has already been set";
@@ -104,6 +135,7 @@ void FunctionFrameNode::ExitWithScope() {
     // Define the function.
     // Note we do checks to disallow redefinition of functions inside the 
`DefFunction`.
     ir::DefFunction(func_name, func);
+    global_var = frame->global_var_map[func_name];
   } else {
     TVM_FFI_THROW(ValueError) << "Cannot find where to insert Relax.Function";
   }
@@ -169,8 +201,11 @@ 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) {
-      tvm::Var new_output_var(output_var->name, 
tvm::relax::GetType(output_var));
+      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)) {
+        binding_spans.Set(new_output_var, span.value());
+      }
       var_remap[output_var] = new_output_var;
     }
     VarReplacer mutator(std::move(var_remap));
@@ -188,6 +223,14 @@ void BindingBlockFrameNode::ExitWithScope() {
     }
   }
 
+  // Variable rewriting may rebuild bindings, so attach their own source 
ranges last.
+  block->span = source_span;
+  for (const auto& binding : block->bindings) {
+    if (auto span = binding_spans.Get(binding->var)) {
+      binding->span = span.value();
+    }
+  }
+
   // Step 3. Get the last frame from the IRBuilder frame stack.
   ffi::Optional<RelaxFrame> opt_last_frame = 
IRBuilder::Current()->GetLastFrame<RelaxFrame>();
   TVM_FFI_ICHECK(opt_last_frame.has_value());
@@ -216,8 +259,11 @@ void BindingBlockFrameNode::ExitWithScope() {
 
 void IfFrameNode::EnterWithScope() {
   const ffi::Array<IRBuilderFrame>& frames = IRBuilder::Current()->frames;
-  for (const IRBuilderFrame& frame : frames) {
-    const auto* block_frame = frame.as<BindingBlockFrameNode>();
+  for (auto it = frames.rbegin(); it != frames.rend(); ++it) {
+    if ((*it)->IsInstance<FunctionFrameNode>()) {
+      break;
+    }
+    const auto* block_frame = (*it).as<BindingBlockFrameNode>();
     if (block_frame && block_frame->is_dataflow) {
       TVM_FFI_THROW(ValueError) << "Cannot create an IfFrame inside a dataflow 
block.";
     }
@@ -229,10 +275,10 @@ void IfFrameNode::ExitWithScope() {
   RelaxFrameNode::ExitWithScope();
   TVM_FFI_CHECK(then_expr.has_value(), ValueError)
       << "The body of then part is expected to be defined before exiting.";
-  TVM_FFI_CHECK(then_expr.has_value(), ValueError)
+  TVM_FFI_CHECK(else_expr.has_value(), ValueError)
       << "The body of else part is expected to be defined before exiting.";
-  auto body = tvm::relax::If(condition, then_expr.value(), else_expr.value());
-  var = Emit(body);
+  auto body = tvm::relax::If(condition, then_expr.value(), else_expr.value(), 
source_span);
+  var = EmitV2(body, std::nullopt, std::nullopt);
   IRBuilder::Name(var_name, var);
 }
 
diff --git a/src/relax/script/builder/ir.cc b/src/relax/script/builder/ir.cc
index bebf2fbbea..cca4aef551 100644
--- a/src/relax/script/builder/ir.cc
+++ b/src/relax/script/builder/ir.cc
@@ -61,6 +61,32 @@ FunctionFrame Function(bool is_pure, bool is_private) {
   return FunctionFrame(n);
 }
 
+FunctionFrame DeclFunction(bool is_pure, bool is_private, bool local) {
+  FunctionFrame frame = Function(is_pure, is_private);
+  frame->declaration = true;
+  frame->local = local;
+  return frame;
+}
+
+FunctionFrame LocalFunction(bool is_pure, const tvm::Var& reference) {
+  FunctionFrame frame = Function(is_pure, true);
+  frame->local = true;
+  frame->local_var = reference;
+  return frame;
+}
+
+tvm::Var ArgVar(const ffi::String& name, const tvm::Var& var) {
+  FunctionFrame frame = FindFunctionFrame("R.arg");
+  TVM_FFI_CHECK(var->name == name, ValueError)
+      << "A cached parameter must retain its declaration name";
+  for (const auto& param : frame->params) {
+    TVM_FFI_CHECK(param->name != name, ValueError) << "Duplicate function 
parameter: " << name;
+  }
+  frame->params.push_back(var);
+  frame->block_builder->AddDefinitionToScope(var);
+  return var;
+}
+
 tvm::Var Arg(const ffi::String& name, const tvm::Type& ty) {
   FunctionFrame frame = FindFunctionFrame("R.Arg");
   tvm::Var var(name, ty);
@@ -142,6 +168,9 @@ TVM_FFI_STATIC_INIT_BLOCK() {
   namespace refl = tvm::ffi::reflection;
   refl::GlobalDef()
       .def("script.ir_builder.relax.Function", Function)
+      .def("script.ir_builder.relax.DeclFunction", DeclFunction)
+      .def("script.ir_builder.relax.LocalFunction", LocalFunction)
+      .def("script.ir_builder.relax.ArgVar", ArgVar)
       .def("script.ir_builder.relax.Arg", Arg)
       .def("script.ir_builder.relax.FuncName", FuncName)
       .def("script.ir_builder.relax.FuncAttrs", FuncAttrs)
@@ -239,12 +268,38 @@ tvm::Var EmitVarBinding(const tvm::relax::VarBinding& 
binding) {
   return binding->var;
 }
 
+namespace {
+
+tvm::Var RecordBindingSpan(tvm::Var var, const ffi::Optional<Span>& name_span) 
{
+  Span span = IRBuilder::Current()->GetCurrentSourceSpan();
+  if (span.defined()) {
+    CheckBindingBlockFrameExistAndUnended()->binding_spans.Set(var, span);
+  }
+  var->span = name_span.value_or(span);
+  return var;
+}
+
+}  // namespace
+
+tvm::Var EmitV2(const tvm::relax::Expr& value,
+                const ffi::Optional<tvm::Type>& annotate_ty,
+                const ffi::Optional<Span>& name_span) {
+  return RecordBindingSpan(Emit(value, annotate_ty), name_span);
+}
+
+tvm::Var EmitMatchCastV2(const tvm::relax::Expr& value, const tvm::Type& ty,
+                         const ffi::Optional<Span>& name_span) {
+  return RecordBindingSpan(EmitMatchCast(value, ty), name_span);
+}
+
 TVM_FFI_STATIC_INIT_BLOCK() {
   namespace refl = tvm::ffi::reflection;
   refl::GlobalDef()
       .def("script.ir_builder.relax.Emit", Emit)
       .def("script.ir_builder.relax.EmitMatchCast", EmitMatchCast)
-      .def("script.ir_builder.relax.EmitVarBinding", EmitVarBinding);
+      .def("script.ir_builder.relax.EmitVarBinding", EmitVarBinding)
+      .def("script.ir_builder.relax.EmitV2", EmitV2)
+      .def("script.ir_builder.relax.EmitMatchCastV2", EmitMatchCastV2);
 }
 
 /////////////////////////////// SeqExpr ///////////////////////////////
diff --git a/src/relax/script/builder/utils.h b/src/relax/script/builder/utils.h
index e4d63c62f6..0ffd80c748 100644
--- a/src/relax/script/builder/utils.h
+++ b/src/relax/script/builder/utils.h
@@ -108,7 +108,7 @@ inline tvm::relax::SeqExpr GetSeqExprForBranch(const 
SeqExprFrame& frame, ffi::S
                                                       
last_block->bindings.end() - 1);
 
   tvm::Var new_var(last_binding->var->name + output_var_suffix,
-                   tvm::relax::GetType(last_binding->var));
+                   tvm::relax::GetType(last_binding->var), 
last_binding->var->span);
   tvm::relax::Expr body;
 
   const auto* var_binding = last_binding.as<tvm::relax::VarBindingNode>();
@@ -116,21 +116,21 @@ inline tvm::relax::SeqExpr GetSeqExprForBranch(const 
SeqExprFrame& frame, ffi::S
   if (var_binding && tvm::relax::IsLeafOrTuple(var_binding->value)) {
     body = var_binding->value;
   } else if (var_binding) {
-    last_block_bindings.push_back(tvm::relax::VarBinding(new_var, 
var_binding->value));
+    last_block_bindings.push_back(tvm::relax::VarBinding(new_var, 
var_binding->value, var_binding->span));
     body = new_var;
   } else if (const auto* match_cast = 
last_binding.as<tvm::relax::MatchCastNode>()) {
     last_block_bindings.push_back(
-        tvm::relax::MatchCast(new_var, match_cast->value, match_cast->ty));
+        tvm::relax::MatchCast(new_var, match_cast->value, match_cast->ty, 
match_cast->span));
     body = new_var;
   } else {
     TVM_FFI_CHECK(false, TypeError) << "Unsupported binding type: " << 
last_binding->GetTypeKey();
   }
 
   new_blocks.push_back(last_block->IsInstance<tvm::relax::DataflowBlockNode>()
-                           ? tvm::relax::DataflowBlock(last_block_bindings)
-                           : tvm::relax::BindingBlock(last_block_bindings));
+                           ? tvm::relax::DataflowBlock(last_block_bindings, 
last_block->span)
+                           : tvm::relax::BindingBlock(last_block_bindings, 
last_block->span));
 
-  return tvm::relax::SeqExpr(new_blocks, body);
+  return tvm::relax::SeqExpr(new_blocks, body, frame->source_span);
 }
 
 }  // namespace relax

Reply via email to