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 df523a35a2dde05ae7c4ac6a296a472992292347 Author: Tianqi Chen <[email protected]> AuthorDate: Wed Sep 23 08:49:41 2026 +0000 Keep wide bitwise mask bounds in the modular analyzer integer domain --- src/sym/modular_set.cc | 4 ++-- tests/python/sym/test_sym_modular_set.py | 22 ++++++++++++++++++++++ 2 files changed, 24 insertions(+), 2 deletions(-) diff --git a/src/sym/modular_set.cc b/src/sym/modular_set.cc index 16b129ad82..c2aafe1820 100644 --- a/src/sym/modular_set.cc +++ b/src/sym/modular_set.cc @@ -291,9 +291,9 @@ class ModularSetAnalyzer::Impl : public tvm::ExprFunctor<ModularSetAnalyzer::Ent Entry Dispatch_(const prim::BitwiseAndNode* op) final { Entry b = Dispatch(op->b); - if (b.is_const()) { + if (b.is_const() && b.base >= 0 && b.base < std::numeric_limits<int64_t>::max()) { int shift; - if (is_const_power_of_two_integer(IntImm::Int32(b.base + 1), &shift)) { + if (is_const_power_of_two_integer(IntImm::Int64(b.base + 1), &shift)) { return ModByConst(op->a, static_cast<int64_t>(1) << shift, true); } } diff --git a/tests/python/sym/test_sym_modular_set.py b/tests/python/sym/test_sym_modular_set.py index a8b52512ce..cf8db39e19 100644 --- a/tests/python/sym/test_sym_modular_set.py +++ b/tests/python/sym/test_sym_modular_set.py @@ -15,6 +15,8 @@ # specific language governing permissions and limitations # under the License. # ruff: noqa: F841 +import pytest + import tvm import tvm.testing from tvm import te @@ -230,5 +232,25 @@ def test_bitwise_and(): assert m.base == 0 [email protected]("bits", [31, 32, 40, 62]) +def test_bitwise_and_wide_mask(bits): + analyzer = tvm.sym.Analyzer() + x = tvm.tirx.Var("x", "int64") + mask = tvm.tirx.const((1 << bits) - 1, "int64") + result = analyzer.modular_set((x * 16 + 3) & mask) + assert result.coeff == 16 + assert result.base == 3 + + +def test_bitwise_and_max_int64_mask(): + analyzer = tvm.sym.Analyzer() + x = tvm.tirx.Var("x", "int64") + mask = tvm.tirx.const((1 << 63) - 1, "int64") + # The modulus is outside signed int64; retain a conservative result. + result = analyzer.modular_set(x & mask) + assert result.coeff == 1 + assert result.base == 0 + + if __name__ == "__main__": tvm.testing.main()
