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"

Reply via email to