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
