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 d99d3f42abce07ec9ad3be98e4a9ab43bb8c6ed0
Author: Tianqi Chen <[email protected]>
AuthorDate: Sun Sep 20 22:35:44 2026 +0000

    [TVMScript] Add concrete primitive builder conventions and declarations
---
 include/tvm/tirx/script/builder/frame.h       |  10 +-
 include/tvm/tirx/script/builder/ir.h          |   3 +
 python/tvm/tirx/script/builder_v2/__init__.py | 413 ++++++++++++++++++++++++++
 src/tirx/script/builder/frame.cc              |  28 +-
 src/tirx/script/builder/ir.cc                 |   7 +
 5 files changed, 453 insertions(+), 8 deletions(-)

diff --git a/include/tvm/tirx/script/builder/frame.h 
b/include/tvm/tirx/script/builder/frame.h
index 677781956b..1954b5b36d 100644
--- a/include/tvm/tirx/script/builder/frame.h
+++ b/include/tvm/tirx/script/builder/frame.h
@@ -93,6 +93,11 @@ class PrimFuncFrameNode : public TIRFrameNode {
   bool s_tir;
   /*! \brief Whether it is a persistent kernel. */
   bool persistent;
+  /*! \brief Whether this frame declares a bodyless signature. */
+  bool is_declaration{false};
+  /*! \brief Finalized function and its module identity. */
+  ffi::Optional<tvm::tirx::PrimFunc> function;
+  ffi::Optional<GlobalVar> global_var;
 
   static void RegisterReflection() {
     namespace refl = tvm::ffi::reflection;
@@ -106,7 +111,10 @@ class PrimFuncFrameNode : public TIRFrameNode {
         .def_ro("env_threads", &PrimFuncFrameNode::env_threads)
         .def_ro("root_alloc_buffers", &PrimFuncFrameNode::root_alloc_buffers)
         .def_ro("s_tir", &PrimFuncFrameNode::s_tir)
-        .def_ro("persistent", &PrimFuncFrameNode::persistent);
+        .def_ro("persistent", &PrimFuncFrameNode::persistent)
+        .def_ro("is_declaration", &PrimFuncFrameNode::is_declaration)
+        .def_ro("function", &PrimFuncFrameNode::function)
+        .def_ro("global_var", &PrimFuncFrameNode::global_var);
   }
   TVM_FFI_DECLARE_OBJECT_INFO_FINAL("script.ir_builder.tirx.PrimFuncFrame", 
PrimFuncFrameNode,
                                     TIRFrameNode);
diff --git a/include/tvm/tirx/script/builder/ir.h 
b/include/tvm/tirx/script/builder/ir.h
index 48250fbfa4..dcaff15b52 100644
--- a/include/tvm/tirx/script/builder/ir.h
+++ b/include/tvm/tirx/script/builder/ir.h
@@ -68,6 +68,9 @@ BufferVar BufferDecl(ffi::Array<PrimExpr> shape, PrimType 
dtype, ffi::String buf
  */
 PrimFuncFrame PrimFunc(bool is_private, bool s_tir = false, bool persistent = 
false);
 
+/*! \brief Construct a bodyless function signature using the ordinary 
signature operations. */
+PrimFuncFrame DeclFunction(bool is_private = false, bool s_tir = false, bool 
persistent = false);
+
 /*!
  * \brief The PrimFunc variable arguments adding function.
  * \param name The name of the variable.
diff --git a/python/tvm/tirx/script/builder_v2/__init__.py 
b/python/tvm/tirx/script/builder_v2/__init__.py
new file mode 100644
index 0000000000..0d79bd9f0b
--- /dev/null
+++ b/python/tvm/tirx/script/builder_v2/__init__.py
@@ -0,0 +1,413 @@
+# 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 TIRx construction operations over the shared native IRBuilder 
stack."""
+
+import builtins as _python
+from functools import partial as _partial
+from functools import wraps as _wraps
+
+from tvm import ir as _ir
+from tvm import tirx as _tir
+from tvm.script.ir_builder import IRBuilder as _IRBuilder
+from tvm.script.ir_builder import ir as _I
+from tvm.script.ir_builder.base import IRBuilderFrame as _NativeFrame
+from tvm.script.ir_builder.protocol import MISSING as _MISSING
+from tvm.script.ir_builder.protocol import at as _at
+from tvm.script.ir_builder.protocol import expression_args as _expression_args
+from tvm.script.ir_builder.protocol import span_context as _span_context
+from tvm.tirx.script import builder as _T
+from tvm.tirx.script.builder import *  # pylint: 
disable=wildcard-import,unused-wildcard-import
+from tvm.tirx.script.builder import _ffi_api
+from tvm.tirx.script.builder import frame as _frame
+
+
+def type_var(name, *, dtype=None, span=None):
+    """Construct a signature symbol; shape symbols default to int64."""
+    return _ir.Var(name, "int64" if dtype is None else dtype, span)
+
+
+@_expression_args("shape", "strides", "elem_offset", "byte_offset", 
introduce=True)
+def Buffer(
+    shape,
+    dtype="float32",
+    data=None,
+    strides=None,
+    elem_offset=None,
+    byte_offset=None,
+    scope="global",
+    align=0,
+    offset_factor=0,
+    layout="default",
+    allocated_addr=None,
+    buffer_name="",
+    *,
+    span=None,
+):
+    """Construct a concrete buffer from resolved shape expressions."""
+    with _span_context(span):
+        return _at(
+            span,
+            _T.Buffer(
+                shape,
+                dtype,
+                data,
+                strides,
+                elem_offset,
+                byte_offset,
+                scope,
+                align,
+                offset_factor,
+                layout,
+                allocated_addr,
+                buffer_name,
+            ),
+        )
+
+
+buffer = Buffer
+
+
+def Ptr(dtype, storage_scope="global", *, span=None):
+    """Construct a concrete pointer variable usable as a function 
annotation."""
+    if callable(dtype) and not isinstance(dtype, _ir.Expr):
+        dtype = dtype()
+    if isinstance(dtype, _ir.Expr):
+        dtype = dtype.ty
+    if isinstance(dtype, _ir.PrimType):
+        dtype = dtype.dtype
+    with _span_context(span):
+        return _at(span, _T.ptr(dtype, storage_scope))
+
+
+class _Frame:
+    """Preserve a frame's source range through native finalization."""
+
+    def __init__(self, native, span=None):
+        self.native = native
+        self.span = span
+        self.result = {}
+
+    def __enter__(self):
+        with _span_context(self.span):
+            value = self.native.__enter__()
+        return self if value is self.native else value
+
+    def __exit__(self, *exc):
+        with _span_context(self.span):
+            return self.native.__exit__(*exc)
+
+    @property
+    def reference(self):
+        """Return the stable module reference after signature finalization."""
+        return self.native.global_var
+
+    def __getattr__(self, name):
+        return getattr(self.native, name)
+
+
+def function(*, private=False, s_tir=False, persistent=False, span=None):
+    """Enter a native primitive-function definition frame."""
+    with _span_context(span):
+        return _Frame(_T.prim_func(private=private, s_tir=s_tir, 
persistent=persistent), span)
+
+
+def decl_function(*, private=False, s_tir=False, persistent=False, span=None):
+    """Declare a bodyless signature using the native function frame."""
+    with _span_context(span):
+        return _Frame(_ffi_api.DeclFunction(private, s_tir, persistent), span)
+
+
+def arg(name, annotation, *, span=None):
+    """Use the same concrete parameter object in declaration and definition."""
+    if callable(annotation) and not isinstance(annotation, _ir.Expr):
+        annotation = annotation()
+    if isinstance(annotation, _ir.Type):
+        annotation = _ir.Var(name, annotation)
+    with _span_context(span):
+        if _tir.is_buffer_var(annotation) and annotation.ty.layout is not None:
+            frames = _IRBuilder.current().frames
+            if _python.any(
+                isinstance(frame, _frame.PrimFuncFrame) and frame.s_tir for 
frame in frames
+            ):
+                ty = annotation.ty
+                annotation = _T.Buffer(
+                    ty.shape,
+                    ty.dtype,
+                    strides=ty.strides,
+                    elem_offset=ty.elem_offset,
+                    scope=ty.storage_scope,
+                    align=ty.data_alignment,
+                    offset_factor=ty.offset_factor,
+                    layout=None,
+                    allocated_addr=list(ty.allocated_addr),
+                    buffer_name=name,
+                )
+        return _T.arg(name, _at(span, annotation))
+
+
+def func_ret_type(annotation, *, span=None):
+    """Set the signature's concrete return type."""
+    if callable(annotation) and not isinstance(annotation, _ir.Expr):
+        annotation = annotation()
+    if isinstance(annotation, _ir.Expr):
+        annotation = annotation.ty
+    with _span_context(span):
+        return _T.func_ret(annotation)
+
+
+def _name(value, name, span):
+    if name is not None:
+        _IRBuilder.name(name, value)
+    return _at(span, value)
+
+
+def _enter_concise(frame):
+    native = frame.native if isinstance(frame, _Frame) else frame
+    native.add_callback(_partial(frame.__exit__, None, None, None))
+    return frame.__enter__()
+
+
+def _as_expr(value):
+    if isinstance(value, _ir.Expr):
+        return value
+    if isinstance(value, str):
+        return _ir.StringImm(value)
+    if isinstance(value, list | tuple):
+        return _ir.Tuple([_as_expr(item) for item in value])
+    return _tir.const(value)
+
+
+def _check_unterminated():
+    frames = _IRBuilder.current().frames
+    if not frames or not isinstance(frames[-1], _frame.TIRFrame):
+        return
+    statements = frames[-1].stmts
+    while statements:
+        last = statements[-1]
+        if isinstance(last, _tir.Return | _tir.Break | _tir.Continue):
+            raise ValueError("An operation cannot follow an unconditional 
terminator")
+        if not isinstance(last, _tir.SeqStmt):
+            break
+        statements = last.seq
+
+
+def bind_(value=_MISSING, *, ty=None, name=None, span=None, name_span=None, 
previous=_MISSING):
+    """Bind concrete values, preserving existing mutable scalar storage."""
+    name_span = span if name_span is None else name_span
+    _check_unterminated()
+    with _span_context(span):
+        if previous is not _MISSING and isinstance(previous, _ir.TensorLoad):
+            if value is _MISSING:
+                raise ValueError("A reassignment requires an initializer")
+            _T.buffer_store(previous.source, value, previous.indices)
+            return previous
+        if isinstance(value, _I.meta_var):
+            return value.value
+        if isinstance(ty, _T.LocalVectorAnnotation):
+            if value is not _MISSING:
+                raise ValueError("Vector annotation does not support an 
initializer")
+            return _name(_T.alloc_local(ty.shape, ty.dtype), name, name_span)
+        if isinstance(ty, _T.LetAnnotation):
+            if value is _MISSING:
+                raise ValueError("An immutable binding requires an 
initializer")
+            value = _as_expr(value)
+            variable = _name(ty.as_var(rhs_dtype=value.ty), name, name_span)
+            _T.Bind(value, var=variable)
+            return variable
+        if value is _MISSING:
+            raise ValueError("Uninitialized scalar bindings are not supported")
+        if ty is not None:
+            annotation = ty() if callable(ty) and not isinstance(ty, _ir.Expr) 
else ty
+            annotation = annotation.ty if isinstance(annotation, _ir.Expr) 
else annotation
+            if not isinstance(annotation, _ir.PrimType) or str(annotation) == 
"handle":
+                raise TypeError("Mutable scalar annotations require a 
primitive scalar type")
+            result = _T.local_scalar(str(annotation)).scalar
+            _name(result.source, name, name_span)
+            _T.buffer_store(result.source, value, [0])
+            return result
+        if (
+            isinstance(value, _ir.TensorLoad)
+            and _tir.is_buffer_var(value.source)
+            and not value.source.name
+            and len(value.source.ty.shape) == 1
+            and isinstance(value.source.ty.shape[0], _tir.IntImm)
+            and value.source.ty.shape[0].value == 1
+        ):
+            _name(value.source, name, name_span)
+            return value
+        if isinstance(value, _T.scalar_wrapper):
+            _name(value.scalar.source, name, name_span)
+            return value.scalar
+        if isinstance(value, _NativeFrame | _Frame):
+            return _name(_enter_concise(value), name, name_span)
+        if isinstance(value, list | tuple):
+            for index, item in enumerate(value):
+                bind_(item, name=None if name is None else f"{name}_{index}", 
span=span)
+            return value
+        if getattr(type(value), "_is_meta_class", False):
+            if name is not None:
+                _T.name_meta_class_value(name, value)
+            return value
+        if _tir.is_buffer_var(value) or isinstance(value, _tir.IterVar | 
_tir.Layout):
+            return _name(value, name, name_span)
+        if isinstance(value, _ir.Var) and not value.name:
+            return _name(value, name, name_span)
+        if isinstance(value, _ir.TensorRegion):
+            return value
+        if not isinstance(value, _ir.Expr | _python.int | _python.float | 
_python.bool | str):
+            return value
+        value = _as_expr(value)
+        if _ir.is_prim_expr(value):
+            result = _T.local_scalar(str(value.ty.dtype)).scalar
+            _name(result.source, name, name_span)
+            _T.buffer_store(result.source, value, [0])
+            return result
+        return _name(_T.Bind(value), name, name_span)
+
+
+def emit_(value, *, span=None):
+    """Consume one expression statement, including effect-only calls."""
+    if value is None or isinstance(value, str | _ir.Var):
+        return
+    _check_unterminated()
+    with _span_context(span):
+        if isinstance(value, _NativeFrame | _Frame):
+            _enter_concise(value)
+        elif hasattr(value, "frames"):
+            for frame in value.frames:
+                _enter_concise(frame)
+        elif isinstance(value, _tir.BufferStore):
+            _T.buffer_store(value.buffer, value.value, value.indices)
+        else:
+            _T.evaluate(value)
+
+
+def setitem(target, key, value, *, span=None):
+    """Construct an indexed store after the caller has evaluated its 
operands."""
+    _check_unterminated()
+    with _span_context(span):
+        _T.buffer_store(target, value, key)
+
+
+def return_(value, *, span=None):
+    """Construct an IR return without exiting the Python construction 
helper."""
+    _check_unterminated()
+    if value is None:
+        raise TypeError("A primitive function return requires an expression")
+    with _span_context(span):
+        _T.Return(_as_expr(value))
+
+
+def _require_loop():
+    for frame in reversed(_IRBuilder.current().frames):
+        if isinstance(frame, _frame.ForFrame | _frame.WhileFrame):
+            return
+        if isinstance(frame, _frame.PrimFuncFrame):
+            break
+    raise ValueError("Loop control requires an enclosing primitive loop")
+
+
+def break_(*, span=None):
+    """Construct a break targeting the nearest primitive loop."""
+    _require_loop()
+    _check_unterminated()
+    with _span_context(span):
+        _T.Break()
+
+
+def continue_(*, span=None):
+    """Construct a continue targeting the nearest primitive loop."""
+    _require_loop()
+    _check_unterminated()
+    with _span_context(span):
+        _T.Continue()
+
+
+def assert_(condition, message="", *, span=None):
+    """Emit the native flat assertion with its own source range."""
+    _check_unterminated()
+    kind = "RuntimeError"
+    if isinstance(message, tuple):
+        if len(message) != 2 or not isinstance(message[0], str):
+            raise TypeError("Assertion metadata must be (error_kind, 
message_parts)")
+        kind, message = message
+    if isinstance(message, list | tuple):
+        message = [str(part) for part in message]
+    if not isinstance(message, list | tuple):
+        message = [message]
+    with _span_context(span):
+        with _T.Assert(condition, message, error_kind=kind):
+            pass
+
+
+def If(condition, *, span=None):
+    with _span_context(span):
+        return _Frame(_T.If(condition), span)
+
+
+def Then(*, span=None):
+    with _span_context(span):
+        return _Frame(_T.Then(), span)
+
+
+def Else(*, span=None):
+    with _span_context(span):
+        return _Frame(_T.Else(), span)
+
+
+def For(iterable, *, span=None):
+    """Adapt a native loop frame, or a Python range, to construction scope."""
+    if isinstance(iterable, _python.range):
+        iterable = _T.serial(iterable.start, iterable.stop, step=iterable.step)
+    if not isinstance(iterable, _frame.ForFrame):
+        raise TypeError("A primitive for loop requires a native loop frame or 
range")
+    return _Frame(iterable, span)
+
+
+def While(condition, *, span=None):
+    with _span_context(span):
+        return _Frame(_T.While(condition), span)
+
+
+def unpack(value):
+    """Project a concrete IR tuple of known arity, preserving Python 
iteration."""
+    if isinstance(value, _ir.Tuple):
+        return _python.tuple(value.fields)
+    if isinstance(value, _ir.Expr) and isinstance(value.ty, _ir.TupleType):
+        return _python.tuple(_ir.TupleGetItem(value, i) for i in 
range(len(value.ty.fields)))
+    return value
+
+
+def alloc_scalar(dtype="float32", scope="global"):
+    """Allocate scalar storage and return its concrete load expression."""
+    value = _T.alloc_scalar(dtype, scope)
+    return value.scalar if isinstance(value, _T.scalar_wrapper) else value
+
+
+def local_scalar(dtype="float32"):
+    return alloc_scalar(dtype, "local")
+
+
+def shared_scalar(dtype="float32"):
+    return alloc_scalar(dtype, "shared")
+
+
+@_expression_args("shape", "strides", "elem_offset", introduce=True)
+@_wraps(_T.match_buffer)
+def match_buffer(*args, **kwargs):
+    """Construct a native buffer match with resolved symbolic shape fields."""
+    return _T.match_buffer(*args, **kwargs)
diff --git a/src/tirx/script/builder/frame.cc b/src/tirx/script/builder/frame.cc
index 9ae7e19154..5c008ec1a5 100644
--- a/src/tirx/script/builder/frame.cc
+++ b/src/tirx/script/builder/frame.cc
@@ -119,7 +119,9 @@ void PrimFuncFrameNode::ExitWithScope() {
   // s_tir-mode normalization: drop stale default layouts (see comment on
   // STirBufferLayoutNormalizer above) and rewrite body references coherently.
   ffi::Array<tvm::tirx::BufferVar> effective_root_alloc_buffers = 
root_alloc_buffers;
-  tvm::tirx::Stmt body = AsStmt(stmts);
+  TVM_FFI_CHECK(!is_declaration || (stmts.empty() && 
root_alloc_buffers.empty()), ValueError)
+      << "A function declaration cannot contain body statements";
+  tvm::tirx::Stmt body = is_declaration ? tvm::tirx::Stmt() : AsStmt(stmts);
   auto normalizer = ffi::make_object<STirBufferLayoutNormalizer>();
   ffi::Array<tvm::tirx::Var> effective_args;
   ffi::Map<tvm::tirx::Var, tvm::Expr> param_replacements;
@@ -151,14 +153,16 @@ void PrimFuncFrameNode::ExitWithScope() {
     }
   }
   if (!normalizer->Empty()) {
-    body = normalizer->Mutate(body, 
InplaceMode::kAllow).ValueOrUnchanged(body);
+    if (!is_declaration) {
+      body = normalizer->Mutate(body, 
InplaceMode::kAllow).ValueOrUnchanged(body);
+    }
     ffi::Array<tvm::tirx::BufferVar> new_root_alloc_buffers;
     for (const tvm::tirx::BufferVar& buffer : root_alloc_buffers) {
       new_root_alloc_buffers.push_back(normalizer->Lookup(buffer));
     }
     effective_root_alloc_buffers = std::move(new_root_alloc_buffers);
   }
-  if (!param_replacements.empty()) {
+  if (!is_declaration && !param_replacements.empty()) {
     auto f_substitute =
         [&param_replacements](
             const tvm::tirx::Var& var) -> 
ffi::Expected<ffi::UnchangedOr<ffi::Any>> {
@@ -175,8 +179,11 @@ void PrimFuncFrameNode::ExitWithScope() {
       /*body=*/body,
       /*ret_type=*/ret_type.value_or(TupleType::Empty()),
       /*attrs=*/attrs.defined() ? DictAttrs(attrs) : DictAttrs(),
-      /*span=*/tvm::Span());
-  func = tvm::tirx::ScriptComplete(func, effective_root_alloc_buffers, s_tir);
+      /*span=*/IRBuilder::Current()->GetCurrentSourceSpan());
+  if (!is_declaration) {
+    func = tvm::tirx::ScriptComplete(func, effective_root_alloc_buffers, 
s_tir);
+  }
+  function = func;
   IRBuilder builder = IRBuilder::Current();
   if (builder->frames.empty()) {
     TVM_FFI_CHECK(!builder->result.has_value(), ValueError)
@@ -190,11 +197,18 @@ void PrimFuncFrameNode::ExitWithScope() {
     const ffi::String& func_name = name.value_or("");
     if (!frame->global_var_map.count(func_name)) {
       // Case. First time visiting the function.
-      ir::DeclFunction(func_name, func);
+      global_var = ir::DeclFunction(func_name, func);
     }
     // Define the function.
     // Note we do checks to disallow redefinition of functions inside the 
`DefFunction`.
-    ir::DefFunction(func_name, func);
+    if (!global_var.has_value()) {
+      TVM_FFI_CHECK(!is_declaration, ValueError)
+          << "function " << func_name << " already exists";
+      global_var = frame->global_var_map.at(func_name);
+    }
+    if (!is_declaration) {
+      ir::DefFunction(func_name, func);
+    }
   } else {
     TVM_FFI_THROW(ValueError) << "Cannot find where to insert PrimFunc";
   }
diff --git a/src/tirx/script/builder/ir.cc b/src/tirx/script/builder/ir.cc
index f87d6d7cf5..59d3ee728d 100644
--- a/src/tirx/script/builder/ir.cc
+++ b/src/tirx/script/builder/ir.cc
@@ -84,6 +84,12 @@ PrimFuncFrame PrimFunc(bool is_private, bool s_tir, bool 
persistent) {
   return PrimFuncFrame(n);
 }
 
+PrimFuncFrame DeclFunction(bool is_private, bool s_tir, bool persistent) {
+  PrimFuncFrame frame = PrimFunc(is_private, s_tir, persistent);
+  frame->is_declaration = true;
+  return frame;
+}
+
 Var Arg(ffi::String name, Var var) {
   PrimFuncFrame frame = FindPrimFuncFrame("T.Arg");
   details::Namer::Name(var, name);
@@ -959,6 +965,7 @@ TVM_FFI_STATIC_INIT_BLOCK() {
                                      ffi::Optional<PrimExpr>, ffi::String, 
int, int,
                                      ffi::Optional<Layout>, 
ffi::Array<PrimExpr>)>(BufferDecl))
       .def("script.ir_builder.tirx.PrimFunc", PrimFunc)
+      .def("script.ir_builder.tirx.DeclFunction", DeclFunction)
       .def("script.ir_builder.tirx.Arg",
            [](ffi::String name, ffi::ObjectRef obj) -> ffi::ObjectRef {
              using namespace tvm::tirx;

Reply via email to