Skip to content
Merged
Show file tree
Hide file tree
Changes from 1 commit
Commits
Show all changes
27 commits
Select commit Hold shift + click to select a range
ea31e4b
[CMAKE] Upgrade TVM build baseline to C++20
Ubospica Jun 11, 2026
f69ef25
[CMAKE] Keep downstream helper standards unchanged
Ubospica Jun 11, 2026
a93a250
[DOCS] Fix from source dependency list formatting
Ubospica Jun 14, 2026
f3840d7
[ARITH] Add optional Z3-backed proving to Analyzer
Ubospica Jun 3, 2026
254010f
update
Ubospica Jun 10, 2026
a6b45d6
Drop accidentally committed build artifacts
Ubospica Jun 10, 2026
551df5d
[CMAKE] Link Z3 statically from the z3-staticlib package by default
Ubospica Jun 10, 2026
5c9a546
Stop ignoring python/tvm/_version.py
Ubospica Jun 10, 2026
6c114da
[CMAKE] Enable static z3 package by default
Ubospica Jun 14, 2026
84a2b4a
[CMAKE] Clean up LLVM include propagation
Ubospica Jun 14, 2026
048b5d0
[CMAKE] Auto-enable Z3 when available
Ubospica Jun 14, 2026
80fc0ef
[CMAKE] Use MSVC include dirs for LLVM headers
Ubospica Jun 14, 2026
f46a5bd
[CMAKE] Document MSVC LLVM include handling
Ubospica Jun 14, 2026
f211a32
[CMAKE] Clarify MSVC LLVM include handling
Ubospica Jun 14, 2026
010b079
[CI] Retry GPU infrastructure check
Ubospica Jun 14, 2026
1d995cc
[TEST] Match reflected structural equal access paths
Ubospica Jun 14, 2026
6495bd5
[TEST] Lift rlimit for Z3 monotonicity proof
Ubospica Jun 14, 2026
6d25638
[TEST] Lift rlimit for Z3 index bound proof
Ubospica Jun 14, 2026
eccf6d1
[ARITH] Pin Z3 context lifetime in prover
Ubospica Jun 15, 2026
9c9f152
[TEST] Tolerate reflected access path formatting
Ubospica Jun 15, 2026
c94829f
[CMAKE] Use z3-static config module for Z3 lookup
Ubospica Jun 16, 2026
a4df6ba
[CMAKE] Require z3-static config package
Ubospica Jun 16, 2026
5ddcbfc
[LINT] Apply Z3 formatting fixes
Ubospica Jun 16, 2026
7b89ce7
[FFI] Restore tvm-ffi submodule pointer
Ubospica Jun 17, 2026
f028cee
[TEST] Restore Relax structural path checks
Ubospica Jun 17, 2026
c857a26
[CI] Retry infrastructure checks
Ubospica Jun 17, 2026
fde41c9
[CI] Retry stuck Windows setup
Ubospica Jun 17, 2026
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
Prev Previous commit
Next Next commit
[ARITH] Add optional Z3-backed proving to Analyzer
Add an optional Z3 SMT solver backend to tvm::arith::Analyzer for
stronger integer arithmetic proving. The integration is guarded by a
new USE_Z3 CMake option (default OFF). When enabled, Analyzer::CanProve
runs the existing analysis path first and only falls back to Z3 when the
existing analyzers cannot prove the predicate. When disabled, a stub
implementation keeps the C++ and Python APIs available without Z3.
  • Loading branch information
Ubospica committed Jun 16, 2026
commit f3840d7a45d3d241305efbb867a923148388c5c2
2 changes: 2 additions & 0 deletions CMakeLists.txt
Original file line number Diff line number Diff line change
Expand Up @@ -89,6 +89,7 @@ tvm_option(COMPILER_RT_PATH "Path to COMPILER-RT" "3rdparty/compiler-rt")
# Contrib library options
tvm_option(USE_BLAS "The blas library to be linked" none)
tvm_option(USE_AMX "Enable Intel AMX" OFF)
tvm_option(USE_Z3 "Build with Z3 SMT solver support" OFF)
tvm_option(USE_MKL "MKL root path when use MKL blas" OFF)
tvm_option(USE_DNNL "Enable DNNL codegen" OFF)
tvm_option(USE_CUDNN "Build with cuDNN" OFF)
Expand Down Expand Up @@ -459,6 +460,7 @@ include(cmake/modules/contrib/AMX.cmake)
include(cmake/modules/contrib/CUTLASS.cmake)
include(cmake/modules/contrib/Random.cmake)
include(cmake/modules/contrib/Sort.cmake)
include(cmake/modules/contrib/Z3.cmake)
include(cmake/modules/contrib/CoreML.cmake)
include(cmake/modules/contrib/TensorRT.cmake)
include(cmake/modules/contrib/NNAPI.cmake)
Expand Down
76 changes: 76 additions & 0 deletions cmake/modules/contrib/Z3.cmake
Original file line number Diff line number Diff line change
@@ -0,0 +1,76 @@
# 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.

if(NOT USE_Z3)
list(APPEND COMPILER_SRCS src/target/z3/z3_prover_off.cc)
return()
endif()

find_package(Z3 QUIET)
set(Z3_PYTHON_RESULT 1)

if(NOT Z3_FOUND)
find_package(Python3 COMPONENTS Interpreter QUIET)
if(Python3_EXECUTABLE)
execute_process(
COMMAND "${Python3_EXECUTABLE}" -c "import z3; print(z3.__path__[0])"
OUTPUT_VARIABLE Z3_PYTHON_PACKAGE_DIR
OUTPUT_STRIP_TRAILING_WHITESPACE
RESULT_VARIABLE Z3_PYTHON_RESULT
)
endif()

if(Z3_PYTHON_RESULT EQUAL 0 AND NOT Z3_PYTHON_PACKAGE_DIR STREQUAL "")
find_path(Z3_INCLUDE_DIR NO_DEFAULT_PATH NAMES z3++.h PATHS "${Z3_PYTHON_PACKAGE_DIR}/include")
find_library(
Z3_LIBRARY
NO_DEFAULT_PATH
NAMES z3 libz3
PATHS "${Z3_PYTHON_PACKAGE_DIR}/bin" "${Z3_PYTHON_PACKAGE_DIR}/lib"
"${Z3_PYTHON_PACKAGE_DIR}/lib64"
)
endif()
endif()

if(TARGET z3::libz3 OR TARGET Z3::libz3)
if(TARGET z3::libz3)
set(Z3_TARGET z3::libz3)
else()
set(Z3_TARGET Z3::libz3)
endif()
get_target_property(Z3_TARGET_INCLUDE_DIRS ${Z3_TARGET} INTERFACE_INCLUDE_DIRECTORIES)
if(Z3_TARGET_INCLUDE_DIRS)
include_directories(SYSTEM ${Z3_TARGET_INCLUDE_DIRS})
endif()
list(APPEND TVM_LINKER_LIBS ${Z3_TARGET})
elseif(Z3_FOUND OR (Z3_INCLUDE_DIR AND Z3_LIBRARY))
if(NOT Z3_INCLUDE_DIR AND Z3_CXX_INCLUDE_DIRS)
set(Z3_INCLUDE_DIR ${Z3_CXX_INCLUDE_DIRS})
endif()
if(NOT Z3_LIBRARY AND Z3_LIBRARIES)
set(Z3_LIBRARY ${Z3_LIBRARIES})
endif()
if(NOT Z3_INCLUDE_DIR OR NOT Z3_LIBRARY)
message(FATAL_ERROR "USE_Z3 is ON, but Z3 include directory or library was not found.")
endif()
include_directories(SYSTEM ${Z3_INCLUDE_DIR})
list(APPEND TVM_LINKER_LIBS ${Z3_LIBRARY})
else()
message(FATAL_ERROR "USE_Z3 is ON, but Z3 was not found. Install Z3 or PyPI z3-solver.")
endif()

list(APPEND COMPILER_SRCS src/target/z3/z3_prover_on.cc)
125 changes: 122 additions & 3 deletions include/tvm/arith/analyzer.h
Original file line number Diff line number Diff line change
Expand Up @@ -27,6 +27,7 @@
#include <tvm/arith/int_set.h>
#include <tvm/ffi/cast.h>
#include <tvm/ffi/reflection/registry.h>
#include <tvm/ffi/string.h>
#include <tvm/ir/expr.h>
#include <tvm/ir/with_context.h>

Expand Down Expand Up @@ -299,7 +300,7 @@ class RewriteSimplifier {
*
* \return an exit function that must be called to cleanup the constraint can be nullptr.
*/
TVM_DLL std::function<void()> EnterConstraint(const PrimExpr& constraint);
TVM_DLL std::function<void()> EnterConstraint(const PrimExpr& constraint, bool is_assume = false);

/*! \brief Flags to enable more computationally-intensive simplifications
*
Expand Down Expand Up @@ -588,6 +589,103 @@ class IntSetAnalyzer {
Impl* impl_;
};

class Z3Prover {
public:
/*!
* \brief Update binding of var to a new expression.
*
* \param var The variable of interest.
* \param new_range The range of allowed values for this var.
* \param allow_override whether we allow override of existing information.
*/
TVM_DLL void Bind(const Var& var, const Range& new_range, bool allow_override = false);

/*!
* \brief Update binding of var to a new expression.
*
* \param var The variable of interest.
* \param expr The bound expression.
* \param allow_override whether we allow override of existing information.
*/
TVM_DLL void Bind(const Var& var, const PrimExpr& expr, bool allow_override = false);

/*!
* \brief Whether can we prove expr is always true.
*
* \param expr The expression.
* \return Whether we can prove it.
*/
TVM_DLL bool CanProve(const PrimExpr& expr);

/*!
* \brief Update the internal state to enter constraint.
*
* \param constraint A constraint expression.
* \param is_assume Whether the constraint comes from an assumption.
* \return an exit function that must be called to cleanup the constraint can be nullptr.
*/
std::function<void()> EnterConstraint(const PrimExpr& constraint, bool is_assume = false);

/*!
* \brief Get the SMTLIB2 representation of the current context.
*
* \param expr The optional expression to check.
* \return The SMTLIB2 string.
*/
ffi::String GetSMTLIB2(const ffi::Optional<PrimExpr> expr);

/*!
* \brief Get statistics about Z3 prover.
*
* \return The statistics string.
*/
ffi::String GetStats();

/*!
* \brief Set timeout in milliseconds for Z3 prover.
*
* \param timeout_ms The timeout in milliseconds.
*/
void SetTimeoutMs(unsigned timeout_ms);

/*!
* \brief Set resource limitation for Z3 prover.
*
* \param rlimit the resource limitation.
*/
void SetRLimit(unsigned rlimit);

/*!
* \brief Get the Z3 model for the given expression if satisfiable.
*
* \param expr The expression to get the model for.
* \return The model as a string.
*/
ffi::String GetModel(const PrimExpr& expr);

/*!
* \brief Count the number of integer values that satisfy the current constraints.
*
* This method uses Z3's model enumeration to count how many distinct values of
* the given variable satisfy all current constraints.
*
* \param var The variable to count satisfying values for.
* \param max_count Maximum number of solutions to enumerate.
* \param min_consecutive Minimum consecutive count requirement.
* \return The number of distinct values that satisfy the constraints, or a negative error code.
*/
TVM_DLL int64_t CountSatisfyingValues(const Var& var, int64_t max_count = 2048,
int64_t min_consecutive = 1);

private:
friend class Analyzer;
explicit Z3Prover(AnalyzerObj* parent);
TVM_DLL ~Z3Prover();
void CopyFrom(const Z3Prover& other);
class Impl;
Impl* impl_;
};

/*!
* \brief Analyzer that contains bunch of sub-analyzers.
*
Expand All @@ -612,6 +710,8 @@ class TVM_DLL AnalyzerObj : public ffi::Object {
IntSetAnalyzer int_set;
/*! \brief sub-analyzer transitive comparisons */
TransitiveComparisonAnalyzer transitive_comparisons;
/*! \brief sub-analyzer using Z3 */
Z3Prover z3_prover;
/*! \brief constructor */
AnalyzerObj();
/*!
Expand Down Expand Up @@ -810,7 +910,16 @@ class ConstraintContext {
* \param constraint The constraint to be applied.
*/
ConstraintContext(const Analyzer& analyzer, PrimExpr constraint)
: analyzer_(analyzer), constraint_(constraint) {}
: ConstraintContext(analyzer, std::move(constraint), false) {}
/*!
* \brief Construct a constraint context.
* \param analyzer The analyzer whose context is updated. The context
* keeps a reference to the analyzer while the scope is active.
* \param constraint The constraint to be applied.
* \param is_assume Whether the constraint comes from an assumption.
*/
ConstraintContext(const Analyzer& analyzer, PrimExpr constraint, bool is_assume)
: analyzer_(analyzer), constraint_(std::move(constraint)), is_assume_(is_assume) {}
/*!
* \brief Construct a constraint context from a borrowed analyzer object.
* \param analyzer The borrowed analyzer object.
Expand All @@ -819,7 +928,15 @@ class ConstraintContext {
* This overload is for internal callers that already operate on AnalyzerObj*.
*/
ConstraintContext(AnalyzerObj* analyzer, PrimExpr constraint)
: ConstraintContext(ffi::GetRef<Analyzer>(analyzer), std::move(constraint)) {}
: ConstraintContext(ffi::GetRef<Analyzer>(analyzer), std::move(constraint), false) {}
/*!
* \brief Construct a constraint context from a borrowed analyzer object.
* \param analyzer The borrowed analyzer object.
* \param constraint The constraint to be applied.
* \param is_assume Whether the constraint comes from an assumption.
*/
ConstraintContext(AnalyzerObj* analyzer, PrimExpr constraint, bool is_assume)
: ConstraintContext(ffi::GetRef<Analyzer>(analyzer), std::move(constraint), is_assume) {}
// enter the scope.
void EnterWithScope();
// exit the scope.
Expand All @@ -830,6 +947,8 @@ class ConstraintContext {
PrimExpr constraint_;
/*! \brief functions to be called in recovery */
std::vector<std::function<void()>> recovery_functions_;
/*! \brief Whether the constraint comes from an assumption. */
bool is_assume_;
};

} // namespace arith
Expand Down
40 changes: 40 additions & 0 deletions python/tvm/arith/analyzer.py
Original file line number Diff line number Diff line change
Expand Up @@ -128,6 +128,46 @@ class Analyzer(Object):
def __init__(self):
self.__init_handle_by_constructor__(_ffi_api.Analyzer)

def get_smtlib2(self, expr: tirx.PrimExpr = None) -> str:
"""Get the current Z3 problem in SMT-LIB2 format.

Parameters
----------
expr : Optional[PrimExpr]
The expression to prove. If provided, its negation is added to the problem.
"""
return _ffi_api.AnalyzerGetSMTLIB2(self, expr)

def set_z3_timeout_ms(self, timeout_ms: int) -> None:
"""Set Z3 timeout in milliseconds.

Parameters
----------
timeout_ms : int
The timeout in milliseconds.
"""
_ffi_api.AnalyzerSetZ3TimeoutMs(self, timeout_ms)

def set_z3_rlimit(self, rlimit: int) -> None:
"""Set Z3 resource limit.

Parameters
----------
rlimit : int
The resource limit.
"""
_ffi_api.AnalyzerSetZ3RLimit(self, rlimit)

def get_z3_stats(self) -> str:
"""Get Z3 solver statistics.

Returns
-------
stats : str
The Z3 statistics.
"""
return _ffi_api.AnalyzerGetZ3Stats(self)

def const_int_bound(self, expr: tirx.PrimExpr) -> ConstIntBound:
"""Find constant integer bound for expr.

Expand Down
25 changes: 23 additions & 2 deletions src/arith/analyzer.cc
Original file line number Diff line number Diff line change
Expand Up @@ -39,7 +39,8 @@ AnalyzerObj::AnalyzerObj()
modular_set(this),
rewrite_simplify(this),
canonical_simplify(this),
int_set(this) {}
int_set(this),
z3_prover(this) {}

void AnalyzerObj::Bind(const Var& var, const PrimExpr& expr, bool allow_override) {
PrimExpr new_expr = expr;
Expand All @@ -52,6 +53,7 @@ void AnalyzerObj::Bind(const Var& var, const PrimExpr& expr, bool allow_override
this->canonical_simplify.Update(var, new_expr, allow_override);
this->int_set.Update(var, this->int_set(new_expr), allow_override);
this->transitive_comparisons.Bind(var, expr, allow_override);
this->z3_prover.Bind(var, expr, allow_override);
}

void AnalyzerObj::Bind(const Var& var, const Range& range, bool allow_override) {
Expand All @@ -62,6 +64,7 @@ void AnalyzerObj::Bind(const Var& var, const Range& range, bool allow_override)
this->const_int_bound.Bind(var, range, allow_override);
this->int_set.Bind(var, range, allow_override);
this->transitive_comparisons.Bind(var, range, allow_override);
this->z3_prover.Bind(var, range, allow_override);
}
// skip modular_set
// skip rewrite simplify
Expand Down Expand Up @@ -128,9 +131,11 @@ void ConstraintContext::EnterWithScope() {
// entering the scope.
recovery_functions_.push_back(analyzer_->const_int_bound.EnterConstraint(constraint_));
recovery_functions_.push_back(analyzer_->modular_set.EnterConstraint(constraint_));
recovery_functions_.push_back(analyzer_->rewrite_simplify.EnterConstraint(constraint_));
recovery_functions_.push_back(
analyzer_->rewrite_simplify.EnterConstraint(constraint_, is_assume_));
recovery_functions_.push_back(analyzer_->int_set.EnterConstraint(constraint_));
recovery_functions_.push_back(analyzer_->transitive_comparisons.EnterConstraint(constraint_));
recovery_functions_.push_back(analyzer_->z3_prover.EnterConstraint(constraint_, is_assume_));
}

void ConstraintContext::ExitWithScope() {
Expand Down Expand Up @@ -231,6 +236,10 @@ bool AnalyzerObj::CanProve(const PrimExpr& expr, ProofStrength strength) {
}
}

if (z3_prover.CanProve(simplified)) {
return true;
}
}
return false;
}

Expand Down Expand Up @@ -334,6 +343,18 @@ TVM_FFI_STATIC_INIT_BLOCK() {
return static_cast<int64_t>(
analyzer->transitive_comparisons.TryCompare(lhs, rhs, propagate_inequalities));
})
.def("arith.AnalyzerGetSMTLIB2",
[](Analyzer analyzer, ffi::Optional<PrimExpr> expr) {
return analyzer->z3_prover.GetSMTLIB2(expr);
})
.def("arith.AnalyzerSetZ3TimeoutMs", [](Analyzer analyzer, int64_t timeout_ms) {
analyzer->z3_prover.SetTimeoutMs(static_cast<unsigned>(timeout_ms));
})
.def("arith.AnalyzerSetZ3RLimit", [](Analyzer analyzer, int64_t rlimit) {
analyzer->z3_prover.SetRLimit(static_cast<unsigned>(rlimit));
})
.def("arith.AnalyzerGetZ3Stats",
[](Analyzer analyzer) { return analyzer->z3_prover.GetStats(); })
.def("arith.AnalyzerGetEnabledExtensions",
[](Analyzer analyzer) {
return static_cast<std::int64_t>(analyzer->rewrite_simplify.GetEnabledExtensions());
Expand Down
Loading