This is an automated email from the ASF dual-hosted git repository. tqchen pushed a commit to branch script/canonical-parser-df in repository https://gitbox.apache.org/repos/asf/tvm.git
commit 370ba142d64cdc4a6d2f4dff38fca9dfe15159ca Author: Tianqi Chen <[email protected]> AuthorDate: Wed Sep 23 05:09:04 2026 +0000 Preserve dependent buffer metadata in printed signatures --- include/tvm/tirx/script/builder/ir.h | 5 +- python/tvm/tirx/script/builder/__init__.py | 9 +- python/tvm/tirx/script/builder/ir.py | 9 ++ src/tirx/script/builder/ir.cc | 4 +- src/tirx/script/printer/buffer.cc | 61 +++++----- src/tirx/script/printer/function.cc | 33 +++++ tests/python/tirx/test_parser_printer_type_vars.py | 135 +++++++++++++++++++++ 7 files changed, 221 insertions(+), 35 deletions(-) diff --git a/include/tvm/tirx/script/builder/ir.h b/include/tvm/tirx/script/builder/ir.h index 055f5478c5..52f09a16e8 100644 --- a/include/tvm/tirx/script/builder/ir.h +++ b/include/tvm/tirx/script/builder/ir.h @@ -117,13 +117,16 @@ Type FuncRet(Type ret_type); * \param storage_scope The optional storage scope of buffer data pointer. * \param align The alignment requirement of data pointer in bytes. * \param offset_factor The factor of elem_offset field. + * \param layout The buffer layout. + * \param allocated_addr Addresses assigned to the buffer allocation. * \return The matched buffer. */ BufferVar MatchBuffer(ffi::ObjectRef param, ffi::Array<PrimExpr> shape, PrimType dtype = PrimType::Float(32), ffi::Optional<Expr> data = std::nullopt, ffi::Array<PrimExpr> strides = {}, PrimExpr elem_offset = PrimExpr(), ffi::String storage_scope = "global", int align = -1, int offset_factor = 0, - ffi::Optional<Layout> layout = std::nullopt); + ffi::Optional<Layout> layout = std::nullopt, + ffi::Array<PrimExpr> allocated_addr = {}); /*! * \brief The block declaration statement. diff --git a/python/tvm/tirx/script/builder/__init__.py b/python/tvm/tirx/script/builder/__init__.py index ff13d25ad5..5f00f6f767 100644 --- a/python/tvm/tirx/script/builder/__init__.py +++ b/python/tvm/tirx/script/builder/__init__.py @@ -81,6 +81,7 @@ def type_var(name, *, dtype=None, span=None): "strides", "elem_offset", "byte_offset", + "allocated_addr", introduce=True, compound_declarations=True, as_type=True, @@ -510,6 +511,7 @@ def shared_scalar(dtype="float32"): "shape", "strides", "elem_offset", + "allocated_addr", introduce=True, compound_declarations=True, ) @@ -569,6 +571,9 @@ def match_buffer(*args, **kwargs): layout: Optional[Union[str, Layout]] The layout of the buffer. + allocated_addr : Expr or int or tuple of Expr or int, optional + Addresses assigned to the buffer allocation. + Returns ------- res : Buffer @@ -576,8 +581,8 @@ def match_buffer(*args, **kwargs): Notes ----- - Shape, stride and element-offset expression strings are resolved by the - construction protocol before the native buffer match is created. + Shape, stride, element-offset and allocation-address expression strings are + resolved by the construction protocol before the native buffer match is created. """ return _native.match_buffer(*args, **kwargs) diff --git a/python/tvm/tirx/script/builder/ir.py b/python/tvm/tirx/script/builder/ir.py index 4ea7c0f15d..e7091bb1c1 100644 --- a/python/tvm/tirx/script/builder/ir.py +++ b/python/tvm/tirx/script/builder/ir.py @@ -489,6 +489,7 @@ def match_buffer( align: int = -1, offset_factor: int = 0, layout: str | Layout | None = "default", + allocated_addr: Expr | int | tuple[Expr | int, ...] | None = None, ) -> Buffer: """The buffer match function. @@ -544,6 +545,9 @@ def match_buffer( layout: Optional[Union[str, Layout]] The layout of the buffer. + allocated_addr : Expr or int or tuple of Expr or int, optional + Addresses assigned to the buffer allocation. + Returns ------- res : Buffer @@ -562,6 +566,10 @@ def match_buffer( strides = [Var(s, "int64") if isinstance(s, str) else s for s in strides] else: strides = [] + if allocated_addr is None: + allocated_addr = [] + if not isinstance(allocated_addr, list | tuple): + allocated_addr = [allocated_addr] result = _ffi_api.MatchBuffer( # type: ignore[attr-defined] # pylint: disable=no-member param, shape, @@ -573,6 +581,7 @@ def match_buffer( align, offset_factor, _get_layout(layout, shape, scope), + allocated_addr, ) return result diff --git a/src/tirx/script/builder/ir.cc b/src/tirx/script/builder/ir.cc index 8106d5b844..4b0ae2c01a 100644 --- a/src/tirx/script/builder/ir.cc +++ b/src/tirx/script/builder/ir.cc @@ -149,9 +149,9 @@ tvm::Type FuncRet(tvm::Type ret_type) { BufferVar MatchBuffer(ffi::ObjectRef param, ffi::Array<PrimExpr> shape, PrimType dtype, ffi::Optional<Expr> data, ffi::Array<PrimExpr> strides, PrimExpr elem_offset, ffi::String storage_scope, int align, int offset_factor, - ffi::Optional<Layout> layout) { + ffi::Optional<Layout> layout, ffi::Array<PrimExpr> allocated_addr) { BufferVar buffer = BufferDecl(shape, dtype, "", data, strides, elem_offset, storage_scope, align, - offset_factor, layout, {}); + offset_factor, layout, allocated_addr); if (auto var = param.as<tvm::tirx::Var>()) { PrimFuncFrame frame = FindPrimFuncFrame("T.match_buffer"); Var v = var.value(); diff --git a/src/tirx/script/printer/buffer.cc b/src/tirx/script/printer/buffer.cc index d9d866e79a..7f4647db91 100644 --- a/src/tirx/script/printer/buffer.cc +++ b/src/tirx/script/printer/buffer.cc @@ -63,6 +63,23 @@ ffi::Map<ffi::String, ExprDoc> BufferAttrs( ffi::StructuralWalk<ffi::WalkOrder::kPostOrder>(data.value(), count_data_var); } auto is_new_var = [&](const Expr& e) { return e->IsInstance<VarNode>() && !d->IsVarDefined(e); }; + // All expression-string annotation fields use the same Python-binding rule. + // Bare TypeVars are real annotation bindings; compound expressions involving + // them, or any expression referring to a later parameter, must be quoted. + auto expression_doc = [&](const PrimExpr& e, const AccessPath& e_p, + bool was_undefined = false) -> ExprDoc { + bool needs_quote = stringify_undefined_shape && was_undefined; + auto walk_fn = [&](const Var& var) -> ffi::Expected<ffi::WalkResult> { + needs_quote = needs_quote || + (stringify_undefined_shape && + (!d->IsVarDefined(var) || stringify_shape_vars.count(var))) || + (stringify_compound_shape_vars.count(var) && !e.same_as(var)); + return ffi::WalkResult::Advance(); + }; + ffi::StructuralWalk<ffi::WalkOrder::kPostOrder>(e, walk_fn); + ExprDoc result = d->AsDoc<ExprDoc>(e, e_p); + return needs_quote ? ExprDoc(ExprStringDoc(result, e_p)) : result; + }; auto add_out_of_line_var_def = [&](const Var& var, const AccessPath& var_p) { TVM_FFI_ICHECK(!d->IsVarDefined(var)); ExprDoc lhs = DefineVar(var, frame, d); @@ -92,26 +109,11 @@ ffi::Map<ffi::String, ExprDoc> BufferAttrs( for (int i = 0; i < n; ++i) { PrimExpr e = shape[i]; AccessPath e_p = shape_p->ArrayItem(i); - bool contains_new_var = false; - bool contains_compound_shape_var = false; - auto walk_fn = [&](const Var& var) -> ffi::Expected<ffi::WalkResult> { - contains_new_var = - contains_new_var || !d->IsVarDefined(var) || stringify_shape_vars.count(var); - contains_compound_shape_var = - contains_compound_shape_var || stringify_compound_shape_vars.count(var); - return ffi::WalkResult::Advance(); - }; - ffi::StructuralWalk<ffi::WalkOrder::kPostOrder>(e, walk_fn); - if (is_new_var(e)) { + bool was_undefined = is_new_var(e); + if (was_undefined) { add_out_of_line_var_def(e.as_or_throw<Var>(), e_p); } - ExprDoc result = d->AsDoc<ExprDoc>(e, e_p); - bool is_bare_compound_shape_var = - e.as<VarNode>() && stringify_compound_shape_vars.count(e.as_or_throw<Var>()); - bool stringify_compound_expr = contains_compound_shape_var && !is_bare_compound_shape_var; - results.push_back((stringify_undefined_shape && contains_new_var) || stringify_compound_expr - ? ExprStringDoc(result, e_p) - : result); + results.push_back(expression_doc(e, e_p, was_undefined)); } kwargs.Set("shape", TupleDoc(results)); } @@ -164,7 +166,7 @@ ffi::Map<ffi::String, ExprDoc> BufferAttrs( continue; } } - results.push_back(d->AsDoc<ExprDoc>(e, e_p)); + results.push_back(expression_doc(e, e_p)); } kwargs.Set("strides", TupleDoc(results)); } @@ -173,18 +175,14 @@ ffi::Map<ffi::String, ExprDoc> BufferAttrs( if (const auto* int_imm = buffer->elem_offset.as<IntImmNode>()) { if (int_imm->value != 0 || int_imm->ty.as_or_throw<PrimType>()->dtype != buffer->DefaultIndexType()) { - kwargs.Set("elem_offset", - d->AsDoc<ExprDoc>(buffer->elem_offset, // - buffer_p->Attr("elem_offset"))); + kwargs.Set("elem_offset", expression_doc(buffer->elem_offset, buffer_p->Attr("elem_offset"))); } } else if (is_new_var(buffer->elem_offset)) { try_inline_def(buffer->elem_offset, buffer_p->Attr("elem_offset"), [=]() { return d->AsDoc<ExprDoc>(buffer, buffer_p)->Attr("elem_offset"); }); needs_print_factor = true; } else { - kwargs.Set("elem_offset", - d->AsDoc<ExprDoc>(buffer->elem_offset, // - buffer_p->Attr("elem_offset"))); + kwargs.Set("elem_offset", expression_doc(buffer->elem_offset, buffer_p->Attr("elem_offset"))); } // Step 6. Handle `buffer.scope` { @@ -234,12 +232,15 @@ ffi::Map<ffi::String, ExprDoc> BufferAttrs( // Unwrap single-element array: DeclBuffer expects Optional<PrimExpr>, not Array. // Use the normal expression printer so a bound scalar alias stays a scalar // load, while an ordinary buffer load retains its indices. - kwargs.Set("allocated_addr", - d->AsDoc<ExprDoc>(buffer->allocated_addr[0], - buffer_p->Attr("allocated_addr")->ArrayItem(0))); + kwargs.Set("allocated_addr", expression_doc(buffer->allocated_addr[0], + buffer_p->Attr("allocated_addr")->ArrayItem(0))); } else { - kwargs.Set("allocated_addr", - d->AsDoc<ExprDoc>(buffer->allocated_addr, buffer_p->Attr("allocated_addr"))); + ffi::Array<ExprDoc> addresses; + for (size_t i = 0; i < buffer->allocated_addr.size(); ++i) { + addresses.push_back(expression_doc(buffer->allocated_addr[i], + buffer_p->Attr("allocated_addr")->ArrayItem(i))); + } + kwargs.Set("allocated_addr", TupleDoc(addresses)); } } diff --git a/src/tirx/script/printer/function.cc b/src/tirx/script/printer/function.cc index 8cb05af4a3..404250e8a9 100644 --- a/src/tirx/script/printer/function.cc +++ b/src/tirx/script/printer/function.cc @@ -93,6 +93,32 @@ TVM_FFI_STATIC_INIT_BLOCK() { AccessPath var_p = p->Attr("params")->ArrayItem(i); if (var->ty.as<tirx::BufferTypeNode>()) { tirx::BufferVar buffer(var); + // Layout expressions have no expression-string syntax. + // Materialize a dependent buffer after its parameters + // are bound, using the existing handle/match_buffer form. + bool needs_body_declaration = false; + auto check_annotation_var = + [&](const tirx::Var& annotation_var) -> ffi::Expected<ffi::WalkResult> { + if (!bound_signature_vars.count(annotation_var)) { + needs_body_declaration = true; + } + return ffi::WalkResult::Advance(); + }; + if (buffer->layout.has_value() && + !ffi::StructuralEqual()(buffer->layout, + tirx::TileLayoutNode::DefaultLayout(buffer->shape))) { + ffi::StructuralWalk<ffi::WalkOrder::kPostOrder>(buffer->layout, check_annotation_var); + } + if (needs_body_declaration) { + tirx::Var handle(var->name + "_handle", PointerType::VoidPointerTy()); + ExprDoc handle_doc = DefineVar(handle, *f, d); + args.push_back(AssignDoc(handle_doc, std::nullopt, TIR(d, "handle"))); + IdDoc lhs = DefineBuffer(buffer, *f, d); + ExprDoc rhs = BufferDecl(buffer, "match_buffer", {handle_doc}, var_p->Attr("ty"), *f, + d, BufferVarDefinition::MatchBuffer); + (*f)->stmts.push_back(AssignDoc(lhs, rhs, std::nullopt)); + continue; + } std::unordered_set<tirx::Var> stringify_shape_vars; std::unordered_set<tirx::Var> stringify_compound_shape_vars; auto walk_fn = [&](const tirx::Var& shape_var) -> ffi::Expected<ffi::WalkResult> { @@ -108,6 +134,13 @@ TVM_FFI_STATIC_INIT_BLOCK() { for (const PrimExpr& shape : buffer->shape) { ffi::StructuralWalk<ffi::WalkOrder::kPostOrder>(shape, walk_fn); } + for (const PrimExpr& stride : buffer->strides) { + ffi::StructuralWalk<ffi::WalkOrder::kPostOrder>(stride, walk_fn); + } + ffi::StructuralWalk<ffi::WalkOrder::kPostOrder>(buffer->elem_offset, walk_fn); + for (const PrimExpr& address : buffer->allocated_addr) { + ffi::StructuralWalk<ffi::WalkOrder::kPostOrder>(address, walk_fn); + } IdDoc lhs = DefineBuffer(buffer, *f, d); ExprDoc annotation = BufferAttn(buffer, var_p->Attr("ty"), *f, d, std::move(stringify_shape_vars), diff --git a/tests/python/tirx/test_parser_printer_type_vars.py b/tests/python/tirx/test_parser_printer_type_vars.py new file mode 100644 index 0000000000..da70e8afdd --- /dev/null +++ b/tests/python/tirx/test_parser_printer_type_vars.py @@ -0,0 +1,135 @@ +# 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. + +"""Printed symbolic buffer metadata survives parsing and Python decoration.""" + +import sys +import types +from typing import TypeVar + +import pytest + +import tvm +from tvm.script import ir as I +from tvm.script import tirx as T + +_PRINT_MODES = [False, True] if sys.version_info >= (3, 12) else [False] + + +def _symbolic_buffer_function(field): + annotations = { + "layout": 'T.Buffer(("n", "n + 1", N), "float32")', + "layout_allocated_addr": 'T.Buffer(("n", "n + 1", N), "float32", allocated_addr=0)', + "layout_allocated_addr_tuple": ( + 'T.Buffer(("n", "n + 1", N), "float32", allocated_addr=(0, 16))' + ), + "strides": 'T.Buffer((N, 4), "float32", strides=("n + 1", 1), layout=None)', + "elem_offset": 'T.Buffer((N, 4), "float32", elem_offset="n + 1", layout=None)', + "allocated_addr": 'T.Buffer((N, 4), "float32", allocated_addr="n + 1", layout=None)', + } + return tvm.script.from_source( + 'N = TypeVar("N")\[email protected]_func\n' + f"def main(A: {annotations[field]}, n: T.int32):\n T.evaluate(n)\n" + ) + + +def _execute_printed(source, tmp_path, monkeypatch): + # The printer leaves TVM imports as comments. Supply those documented names, + # while compiling the unchanged source with its own future-import semantics. + path = tmp_path / "printed_buffer.py" + path.write_text(source) + module = types.ModuleType("_printed_buffer") + module.__file__ = str(path) + module.__dict__.update(T=T, I=I, TypeVar=TypeVar) + monkeypatch.setitem(sys.modules, module.__name__, module) + exec(compile(source, str(path), "exec", dont_inherit=True), module.__dict__) + return module + + [email protected]( + "field", + [ + "layout", + "layout_allocated_addr", + "layout_allocated_addr_tuple", + "strides", + "elem_offset", + "allocated_addr", + ], +) [email protected]( + "pep695", _PRINT_MODES, ids=lambda value: "pep695" if value else "portable" +) [email protected]("in_module", [False, True], ids=["function", "module"]) [email protected]("execute", [False, True], ids=["from_source", "eager_exec"]) +def test_symbolic_buffer_fields_roundtrip_with_later_scalar_parameter( + tmp_path, monkeypatch, field, pep695, in_module, execute +): + # Before: def main[N](A: X.Buffer(("n", "n + 1", N), ...), n: X.int32): ... + # Expected builder program: declare n in the native symbol map, then build + # X.arg("A", ...) and X.arg("n", ...) without an early Python n binding. + # Printed non-string layout expressions use A_handle: X.handle followed by + # A = X.match_buffer(A_handle, ...) in the body. Other symbolic fields stay + # quoted in X.Buffer annotations. Both forms preserve the exact native IR. + function = _symbolic_buffer_function(field) + has_layout = field.startswith("layout") + original = tvm.IRModule({"main": function}) if in_module else function + printed = original.script(extra_config={"script.use_pep695": pep695}) + assert ("from __future__ import annotations" in printed) is pep695 + assert ("def main[N](" in printed) is pep695 + assert ('N = TypeVar("N")' in printed) is not pep695 + if has_layout: + assert "A_handle: T.handle" in printed + assert "A = T.match_buffer(A_handle," in printed + assert "layout=T.TileLayout(" in printed + else: + assert "A: T.Buffer(" in printed + assert "T.match_buffer(" not in printed + assert f"{field}=" in printed + assert '"n + 1"' in printed + if execute: + namespace = _execute_printed(printed, tmp_path, monkeypatch) + actual = namespace.Module if in_module else namespace.main + else: + actual = tvm.script.from_source(printed) + tvm.ir.assert_structural_equal(original, actual) + actual_function = actual["main"] if in_module else actual + buffer, n = actual_function.params + assert str(n.ty.dtype) == "int32" + assert actual_function.body.value.same_as(n) + generic = buffer.ty.shape[2] if has_layout else buffer.ty.shape[0] + assert str(generic.ty.dtype) == "int64" + assert not generic.same_as(n) + if has_layout: + assert buffer.ty.shape[0].same_as(n) + expression = buffer.ty.shape[1] + assert buffer.ty.layout is not None + expected_addresses = { + "layout": [], + "layout_allocated_addr": [0], + "layout_allocated_addr_tuple": [0, 16], + }[field] + assert [int(address) for address in buffer.ty.allocated_addr] == expected_addresses + elif field == "strides": + expression = buffer.ty.strides[0] + elif field == "elem_offset": + expression = buffer.ty.elem_offset + else: + expression = buffer.ty.allocated_addr[0] + assert expression.a.same_as(n) + assert int(expression.b) == 1 + assert str(expression.ty.dtype) == "int32"
