From e8dd97bca0eb69af55d6b48985733b72e220f90e Mon Sep 17 00:00:00 2001 From: StrongbodyStrongmind <143160405+StrongbodyStrongmind@users.noreply.github.com> Date: Thu, 3 Sep 2026 23:46:33 +0800 Subject: [PATCH] [TIR] Ignore None-valued pragma annotations --- src/s_tir/transform/lower_opaque_block.cc | 8 +++++--- src/tirx/transform/lower_tirx_opaque.cc | 14 ++++++++----- ...test_s_tir_transform_lower_opaque_block.py | 20 +++++++++++++++++++ .../test_tir_transform_flatten_buffer.py | 20 +++++++++++++++++++ 4 files changed, 54 insertions(+), 8 deletions(-) diff --git a/src/s_tir/transform/lower_opaque_block.cc b/src/s_tir/transform/lower_opaque_block.cc index 586eb727b7ff..d3dcf2c3a932 100644 --- a/src/s_tir/transform/lower_opaque_block.cc +++ b/src/s_tir/transform/lower_opaque_block.cc @@ -158,9 +158,7 @@ class OpaqueBlockLower : public StmtExprMutator { /*! \brief Convert attr value from annotation map into PrimExpr. */ PrimExpr ConvertAttrValue(const ffi::String& key, const Any& obj) { - if (obj == nullptr) { - return PrimExpr(); - } else if (auto expr = obj.try_cast()) { + if (auto expr = obj.try_cast()) { return expr.value(); } else if (auto str = obj.try_cast()) { return std::move(StringImm(str.value())); @@ -187,6 +185,10 @@ class OpaqueBlockLower : public StmtExprMutator { for (const auto& kv : annotations) { const ffi::String& key = kv.first; if (tirx::attr::IsPragmaKey(key)) { + if (kv.second == nullptr) { + continue; + } + pragma_attrs->emplace_back(key, ConvertAttrValue(key, kv.second)); } else if (!is_block) { // the loop annotation is preserved diff --git a/src/tirx/transform/lower_tirx_opaque.cc b/src/tirx/transform/lower_tirx_opaque.cc index 7c6ba2b2c940..99fcd7793e56 100644 --- a/src/tirx/transform/lower_tirx_opaque.cc +++ b/src/tirx/transform/lower_tirx_opaque.cc @@ -120,9 +120,7 @@ class TIRxOpaqueLower : public StmtExprMutator { /*! \brief Convert attr value from annotation map into PrimExpr. */ PrimExpr ConvertAttrValue(const ffi::String& key, const Any& obj) { - if (obj == nullptr) { - return PrimExpr(); - } else if (auto expr = obj.try_cast()) { + if (auto expr = obj.try_cast()) { return expr.value(); } else if (auto str = obj.try_cast()) { return std::move(prim::StringImm(str.value())); @@ -149,11 +147,17 @@ class TIRxOpaqueLower : public StmtExprMutator { for (const auto& kv : annotations) { const ffi::String& key = kv.first; if (key == "pragma_unroll") { - preserved_annotations.Set(key, kv.second); + if (kv.second != nullptr) { + preserved_annotations.Set(key, kv.second); + } } else if (tirx::attr::IsPragmaKey(key)) { + if (kv.second == nullptr) { + continue; + } + pragma_attrs->emplace_back(key, ConvertAttrValue(key, kv.second)); } else { - // loop annotations are always preserved (no SBlock annotation dropping here) + // Loop annotations are always preserved preserved_annotations.Set(key, kv.second); } } diff --git a/tests/python/s_tir/transform/test_s_tir_transform_lower_opaque_block.py b/tests/python/s_tir/transform/test_s_tir_transform_lower_opaque_block.py index 441074128e6d..273fa0bdc5ea 100644 --- a/tests/python/s_tir/transform/test_s_tir_transform_lower_opaque_block.py +++ b/tests/python/s_tir/transform/test_s_tir_transform_lower_opaque_block.py @@ -421,6 +421,26 @@ def after(A: T.Buffer(8, "float32"), B: T.Buffer(8, "float32")): tvm.ir.assert_structural_equal(mod["main"], after.with_attr("global_symbol", "main")) +def test_none_pragma_annotation(): + @T.prim_func(s_tir=True) + def before(A: T.Buffer(8, "float32"), B: T.Buffer(8, "float32")): + for i in T.serial(8, annotations={"pragma_unroll_explicit": None}): + with T.sblock("block"): + B[i] = A[i] + 1.0 + + @T.prim_func(s_tir=True) + def after(A: T.Buffer(8, "float32"), B: T.Buffer(8, "float32")): + for i in T.serial(8): + B[i] = A[i] + 1.0 + + mod = tvm.IRModule.from_expr(before.with_attr("global_symbol", "main")) + mod = tvm.s_tir.transform.LowerOpaqueBlock()(mod) + tvm.ir.assert_structural_equal( + mod["main"], + after.with_attr("global_symbol", "main"), + ) + + def test_boolean_handling(): _check(boolean_handling_before, boolean_handling_after) diff --git a/tests/python/tirx-transform/test_tir_transform_flatten_buffer.py b/tests/python/tirx-transform/test_tir_transform_flatten_buffer.py index d01400d3a5b2..7c2943879cb8 100644 --- a/tests/python/tirx-transform/test_tir_transform_flatten_buffer.py +++ b/tests/python/tirx-transform/test_tir_transform_flatten_buffer.py @@ -368,5 +368,25 @@ def main(): tvm.ir.assert_structural_equal(After, Expected) +def test_build_with_optional_pragma_unroll_explicit(): + def check(value): + @I.ir_module(s_tir=True) + class Module: + @T.prim_func(s_tir=True) + def main(A: T.Buffer((4, 5, 6), "int16"), B: T.Buffer((4, 5, 6), "int16")): + for ax0 in T.serial(4, annotations={"pragma_unroll_explicit": value}): + for ax1, ax2 in T.grid(5, 6): + with T.sblock("copy"): + v0, v1, v2 = T.axis.remap("SSS", [ax0, ax1, ax2]) + T.reads(A[v0, v1, v2]) + T.writes(B[v0, v1, v2]) + B[v0, v1, v2] = A[v0, v1, v2] + + tvm.build(Module, target="llvm") + + for value in [None, True, False]: + check(value) + + if __name__ == "__main__": tvm.testing.main()