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", [
