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 02d727c054d6e33856a33e15109a079369a93a3b
Author: Tianqi Chen <[email protected]>
AuthorDate: Tue Sep 22 23:13:50 2026 +0000

    [FR] Preserve explicit stride dtypes in script printing
---
 src/tirx/script/printer/buffer.cc                  | 12 +++++++----
 .../script/test_parser_construction_regressions.py | 23 ++++++++++++++++++++++
 2 files changed, 31 insertions(+), 4 deletions(-)

diff --git a/src/tirx/script/printer/buffer.cc 
b/src/tirx/script/printer/buffer.cc
index e1b43645bc..0b37da6e3d 100644
--- a/src/tirx/script/printer/buffer.cc
+++ b/src/tirx/script/printer/buffer.cc
@@ -159,10 +159,14 @@ ffi::Map<ffi::String, ExprDoc> BufferAttrs(
       PrimExpr e = strides[i];
       AccessPath e_p = strides_p->ArrayItem(i);
       if (is_new_var(e)) {
-        if (try_inline_def(e, e_p, [=]() {
-              return d->AsDoc<ExprDoc>(buffer, buffer_p)
-                  ->Attr("strides")[{LiteralDoc::Int(i, std::nullopt)}];
-            })) {
+        // String stride declarations have int64 dtype.
+        PrimType stride_ty = e.ty();
+        if (!stride_ty.IsScalar() || 
!stride_ty.MatchesElementType(DLDataTypeCode::kDLInt, 64)) {
+          add_out_of_line_var_def(e.as_or_throw<Var>(), e_p);
+        } else if (try_inline_def(e, e_p, [=]() {
+                     return d->AsDoc<ExprDoc>(buffer, buffer_p)
+                         ->Attr("strides")[{LiteralDoc::Int(i, std::nullopt)}];
+                   })) {
           results.push_back(LiteralDoc::Str(e.as_or_throw<Var>()->name, e_p));
           continue;
         }
diff --git a/tests/script/test_parser_construction_regressions.py 
b/tests/script/test_parser_construction_regressions.py
index 0deffb7b62..9dd490d279 100644
--- a/tests/script/test_parser_construction_regressions.py
+++ b/tests/script/test_parser_construction_regressions.py
@@ -119,6 +119,29 @@ def main(a: T.handle):
     assert str(buffer.strides[0].ty.dtype) == dtype
 
 
[email protected]("dtype", ["int32", "int64"])
+def test_decl_buffer_stride_roundtrip_preserves_dtype(dtype):
+    # Standalone buffer declarations can contain free stride symbols.
+    function = parser.parse(
+        f"""
[email protected]_func(s_tir=True)
+def main(data: T.handle("float32")):
+    s0 = T.{dtype}()
+    A = T.decl_buffer((1,), "float32", data=data, strides=(s0,))
+    T.evaluate(A[0])
+""",
+        check_well_formed=False,
+    )
+    printed = function.script()
+    if dtype == "int64":
+        assert 'strides=("s0",)' in printed
+    else:
+        assert "s0 = T.int32()" in printed
+    tvm.ir.assert_structural_equal(
+        function, parser.parse(printed, check_well_formed=False), 
map_free_vars=True
+    )
+
+
 @pytest.mark.parametrize(
     "axes",
     [

Reply via email to