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 = [¶m_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;
