Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
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
16 changes: 0 additions & 16 deletions include/tvm/tirx/analysis.h
Original file line number Diff line number Diff line change
Expand Up @@ -106,22 +106,6 @@ TVM_DLL ffi::Array<Var> UndefinedVars(const PrimExpr& expr, const ffi::Array<Var
*/
TVM_DLL CallEffectKind SideEffect(const PrimExpr& expr);

/*!
* \brief Whether the given Stmt uses any var in the given variable set.
* \param stmt The Stmt to be checked.
* \param vset_contains The check function to see if a var is in the variable set.
* \return Whether `stmt` uses any var in the given variable set.
*/
TVM_DLL bool UsesVar(const Stmt& stmt, std::function<bool(const VarNode*)> vset_contains);

/*!
* \brief Whether the given PrimExpr uses any var in the given variable set.
* \param expr The PrimExpr to be checked.
* \param vset_contains The check function to see if var is in the variable set.
* \return Whether `expr` uses any var in the given variable set.
*/
TVM_DLL bool UsesVar(const PrimExpr& expr, std::function<bool(const VarNode*)> vset_contains);

/*!
* \brief Verifies whether the IR stmt or Expr is in SSA form.
* That is: each Var is defined and assigned once(in Let/For)
Expand Down
12 changes: 10 additions & 2 deletions src/arith/detect_linear_equation.cc
Original file line number Diff line number Diff line change
Expand Up @@ -110,7 +110,11 @@ class LinearEqDetector : public ExprFunctor<LinearEqEntry(const Expr&, const Pri
}
LinearEqEntry VisitExprDefault_(const ffi::Object* op, const PrimExpr& e) final {
if (fail_) return LinearEqEntry();
if (UsesVar(e, [this](const VarNode* var) { return var == var_.get(); })) {
auto walkfn = [this](const Var& var) -> ffi::Expected<ffi::WalkResult> {
return var.get() == var_.get() ? ffi::WalkResult::Interrupt(ffi::VisitInterrupt(var))
: ffi::WalkResult::Advance();
};
if (ffi::StructuralWalk<ffi::WalkOrder::kPreOrder>(e, walkfn).has_value()) {
fail_ = true;
return LinearEqEntry();
} else {
Expand Down Expand Up @@ -157,11 +161,15 @@ ffi::Array<PrimExpr> DetectLinearEquation(const PrimExpr& e, const ffi::Array<Pr

std::unordered_set<const VarNode*> vset;
auto vset_contains = [&](const VarNode* node) { return vset.count(node) != 0; };
auto walkfn = [&](const Var& var) -> ffi::Expected<ffi::WalkResult> {
return vset_contains(var.get()) ? ffi::WalkResult::Interrupt(ffi::VisitInterrupt(var))
: ffi::WalkResult::Advance();
};

for (size_t i = vars.size(); i > 1; --i) {
vset.insert(vars[i - 1].get());
// The previous coeff contains the variable
if (UsesVar(coeff[i - 2], vset_contains)) {
if (ffi::StructuralWalk<ffi::WalkOrder::kPreOrder>(coeff[i - 2], walkfn).has_value()) {
return ffi::Array<PrimExpr>();
}
}
Expand Down
10 changes: 7 additions & 3 deletions src/arith/int_set.cc
Original file line number Diff line number Diff line change
Expand Up @@ -24,6 +24,7 @@
#include <tvm/arith/int_set.h>
#include <tvm/arith/iter_affine_map.h>
#include <tvm/ffi/cast.h>
#include <tvm/ffi/extra/structural_visit.h>
#include <tvm/ffi/function.h>
#include <tvm/ffi/reflection/registry.h>
#include <tvm/ir/prim/expr.h>
Expand Down Expand Up @@ -599,10 +600,13 @@ class IntervalSetEvaluator : public ExprFunctor<IntervalSet(const Expr&)> {
}
// If the indices do not contain any variables to be relaxed, return the TensorLoad itself.
// Otherwise return `IntervalSet::everything()` since we have no knowledge on the buffer data.
auto walkfn = [dom_map = &this->dom_map_](const Var& var) -> ffi::Expected<ffi::WalkResult> {
return dom_map->find(var) != dom_map->end()
? ffi::WalkResult::Interrupt(ffi::VisitInterrupt(var))
: ffi::WalkResult::Advance();
};
for (const PrimExpr& index : op->indices) {
if (UsesVar(index, [dom_map = &this->dom_map_](const VarNode* var) {
return dom_map->find(ffi::GetRef<Var>(var)) != dom_map->end();
})) {
if (ffi::StructuralWalk<ffi::WalkOrder::kPreOrder>(index, walkfn).has_value()) {
return IntervalSet::Everything();
}
}
Expand Down
7 changes: 6 additions & 1 deletion src/arith/ir_mutator_with_analyzer.h
Original file line number Diff line number Diff line change
Expand Up @@ -26,6 +26,7 @@

#include <tvm/arith/analyzer.h>
#include <tvm/ffi/cast.h>
#include <tvm/ffi/extra/structural_visit.h>
#include <tvm/ir/scope_stack.h>
#include <tvm/ir/with_context.h>
#include <tvm/tirx/analysis.h>
Expand Down Expand Up @@ -108,8 +109,12 @@ class IRMutatorWithAnalyzer : public tirx::StmtExprMutator {
auto f_use_itervar = [&iter_var_nodes](const tirx::VarNode* v) {
return iter_var_nodes.count(v);
};
auto walkfn = [&](const tirx::Var& var) -> ffi::Expected<ffi::WalkResult> {
return f_use_itervar(var.get()) ? ffi::WalkResult::Interrupt(ffi::VisitInterrupt(var))
: ffi::WalkResult::Advance();
};
// simple heuristics for detecting predicate
if (tirx::UsesVar(condition, f_use_itervar)) {
if (ffi::StructuralWalk<ffi::WalkOrder::kPreOrder>(condition, walkfn).has_value()) {
iter_predicates_.push_back(condition);
callback();
iter_predicates_.pop_back();
Expand Down
28 changes: 22 additions & 6 deletions src/arith/iter_affine_map.cc
Original file line number Diff line number Diff line change
Expand Up @@ -23,6 +23,7 @@
#include <tvm/arith/analyzer.h>
#include <tvm/arith/iter_affine_map.h>
#include <tvm/ffi/cast.h>
#include <tvm/ffi/extra/structural_visit.h>
#include <tvm/ffi/reflection/registry.h>
#include <tvm/ir/prim/expr.h>
#include <tvm/tirx/analysis.h>
Expand Down Expand Up @@ -1348,27 +1349,35 @@ bool MatchBoundConstraints(PrimExpr pred, ffi::Map<PrimVar, Range>* input_iters,
auto f_use_itervar = [&input_iter_nodes](const VarNode* v) {
return input_iter_nodes.count(v);
};
auto walkfn = [&](const Var& var) -> ffi::Expected<ffi::WalkResult> {
return f_use_itervar(var.get()) ? ffi::WalkResult::Interrupt(ffi::VisitInterrupt(var))
: ffi::WalkResult::Advance();
};
bool lhs_uses_itervar =
ffi::StructuralWalk<ffi::WalkOrder::kPreOrder>(lhs_expr, walkfn).has_value();
bool rhs_uses_itervar =
ffi::StructuralWalk<ffi::WalkOrder::kPreOrder>(rhs_expr, walkfn).has_value();
bool bound_at_left;
if (UsesVar(lhs_expr, f_use_itervar) || UsesVar(rhs_expr, f_use_itervar)) {
if (lhs_uses_itervar || rhs_uses_itervar) {
// At least it uses one input iter
if (is_const_int(lhs_expr) || !UsesVar(lhs_expr, f_use_itervar)) {
if (is_const_int(lhs_expr) || !lhs_uses_itervar) {
bound_at_left = true;
} else if (is_const_int(rhs_expr) || !UsesVar(rhs_expr, f_use_itervar)) {
} else if (is_const_int(rhs_expr) || !rhs_uses_itervar) {
bound_at_left = false;
} else {
bound_at_left = false; // accumulate bound to rhs
PrimExpr sum_parts = lhs_expr - rhs_expr;
lhs_expr = 0;
rhs_expr = 0;
std::function<void(const PrimExpr&, bool)> f_extract =
[&lhs_expr, &rhs_expr, f_use_itervar, &f_extract](const PrimExpr& part, bool sign) {
[&lhs_expr, &rhs_expr, &walkfn, &f_extract](const PrimExpr& part, bool sign) {
if (const prim::AddNode* add = part.as<prim::AddNode>()) {
f_extract(add->a, sign);
f_extract(add->b, sign);
} else if (const prim::SubNode* sub = part.as<prim::SubNode>()) {
f_extract(sub->a, sign);
f_extract(sub->b, !sign);
} else if (UsesVar(part, f_use_itervar)) {
} else if (ffi::StructuralWalk<ffi::WalkOrder::kPreOrder>(part, walkfn).has_value()) {
lhs_expr = sign ? lhs_expr + part : lhs_expr - part;
} else {
rhs_expr = sign ? rhs_expr - part : rhs_expr + part;
Expand Down Expand Up @@ -1429,8 +1438,15 @@ bool IterRangeSanityCheck(const ffi::Map<PrimVar, Range>& iter_ranges) {
std::unordered_set<Var> iters;
for (const auto& it : iter_ranges) iters.insert(it.first);
auto f = [&](const VarNode* var) { return iters.count(ffi::GetRef<Var>(var)); };
auto walkfn = [&](const Var& var) -> ffi::Expected<ffi::WalkResult> {
return f(var.get()) ? ffi::WalkResult::Interrupt(ffi::VisitInterrupt(var))
: ffi::WalkResult::Advance();
};
for (const auto& it : iter_ranges) {
if (UsesVar(it.second->min, f) || UsesVar(it.second->extent, f)) return false;
if (ffi::StructuralWalk<ffi::WalkOrder::kPreOrder>(it.second->min, walkfn).has_value() ||
ffi::StructuralWalk<ffi::WalkOrder::kPreOrder>(it.second->extent, walkfn).has_value()) {
return false;
}
}
return true;
}
Expand Down
31 changes: 19 additions & 12 deletions src/relax/analysis/tir_op_pattern_kind.cc
Original file line number Diff line number Diff line change
Expand Up @@ -19,6 +19,7 @@

#include <tvm/arith/iter_affine_map.h>
#include <tvm/ffi/cast.h>
#include <tvm/ffi/extra/structural_visit.h>
#include <tvm/ffi/reflection/registry.h>
#include <tvm/relax/analysis.h>
#include <tvm/relax/op_attr_types.h>
Expand Down Expand Up @@ -235,10 +236,13 @@ class PatternKindAnalyzer : public StmtExprVisitor {
return false;
}
}
auto walkfn = [&vars](const tirx::Var& var) -> ffi::Expected<ffi::WalkResult> {
return !vars.count(var.get()) ? ffi::WalkResult::Interrupt(ffi::VisitInterrupt(var))
: ffi::WalkResult::Advance();
};
for (const PrimExpr& load_index : load->indices) {
// return false if there are vars used in load indices but not in store indices.
if (tirx::UsesVar(load_index,
[&vars](const tirx::VarNode* var) { return !vars.count(var); })) {
if (ffi::StructuralWalk<ffi::WalkOrder::kPreOrder>(load_index, walkfn).has_value()) {
return false;
}
}
Expand Down Expand Up @@ -318,17 +322,20 @@ class PatternKindAnalyzer : public StmtExprVisitor {
*/
static bool IsPureReducePattern(ffi::Array<tirx::Var> reduce_loops,
ffi::Array<PrimExpr> indices) {
auto walkfn = [&](const tirx::Var& var) -> ffi::Expected<ffi::WalkResult> {
return std::any_of(reduce_loops.begin(), reduce_loops.end(),
[&](const tirx::Var& loop) { return loop.same_as(var); })
? ffi::WalkResult::Interrupt(ffi::VisitInterrupt(var))
: ffi::WalkResult::Advance();
};
for (const PrimExpr& e : indices) {
int id = -1;
if (UsesVar(e, [&](const tirx::VarNode* var) {
for (size_t i = 0; i < reduce_loops.size(); ++i) {
if (reduce_loops[i].get() == var) {
id = i;
return true;
}
}
return false;
})) {
auto result = ffi::StructuralWalk<ffi::WalkOrder::kPreOrder>(e, walkfn);
if (result.has_value()) {
tirx::Var var = result.value()->value.cast<tirx::Var>();
int id =
std::distance(reduce_loops.begin(),
std::find_if(reduce_loops.begin(), reduce_loops.end(),
[&](const tirx::Var& loop) { return loop.same_as(var); }));
if (!reduce_loops[id].same_as(e)) {
return false;
}
Expand Down
8 changes: 6 additions & 2 deletions src/relax/transform/rewrite_dataflow_reshape.cc
Original file line number Diff line number Diff line change
Expand Up @@ -22,6 +22,7 @@
*/
#include <tvm/arith/analyzer.h>
#include <tvm/ffi/cast.h>
#include <tvm/ffi/extra/structural_visit.h>
#include <tvm/ffi/reflection/registry.h>
#include <tvm/relax/analysis.h>
#include <tvm/relax/expr_functor.h>
Expand All @@ -41,8 +42,11 @@ std::vector<size_t> GetUsedTensorArgIndices(const tirx::PrimFunc& fn, size_t num
for (size_t i = 0; i < num_args; ++i) {
if (auto buffer = fn->params[i].as<tirx::BufferVar>()) {
auto buffer_var = buffer.value().var();
if (tirx::UsesVar(fn->body,
[=](const tirx::VarNode* var) { return var == buffer_var.get(); })) {
auto walkfn = [=](const tirx::Var& var) -> ffi::Expected<ffi::WalkResult> {
return var.get() == buffer_var.get() ? ffi::WalkResult::Interrupt(ffi::VisitInterrupt(var))
: ffi::WalkResult::Advance();
};
if (ffi::StructuralWalk<ffi::WalkOrder::kPreOrder>(fn->body, walkfn).has_value()) {
indices.push_back(i);
}
}
Expand Down
7 changes: 6 additions & 1 deletion src/s_tir/meta_schedule/postproc/rewrite_reduction_block.cc
Original file line number Diff line number Diff line change
Expand Up @@ -16,6 +16,7 @@
* specific language governing permissions and limitations
* under the License.
*/
#include <tvm/ffi/extra/structural_visit.h>
#include <tvm/ffi/reflection/registry.h>
#include <tvm/s_tir/stmt.h>

Expand Down Expand Up @@ -67,14 +68,18 @@ struct ReductionBlockFinder : private StmtVisitor {
return true;
}
auto f_find = [this](const VarNode* var) -> bool { return thread_bound_loop_vars_.count(var); };
auto walkfn = [&](const Var& var) -> ffi::Expected<ffi::WalkResult> {
return f_find(var.get()) ? ffi::WalkResult::Interrupt(ffi::VisitInterrupt(var))
: ffi::WalkResult::Advance();
};
const SBlockNode* block = realize->block.get();
TVM_FFI_ICHECK_EQ(block->iter_vars.size(), realize->iter_values.size());
int n = block->iter_vars.size();
for (int i = 0; i < n; ++i) {
IterVar iter_var = block->iter_vars[i];
PrimExpr binding = realize->iter_values[i];
if (iter_var->iter_type == tirx::kCommReduce) {
if (UsesVar(binding, f_find)) {
if (ffi::StructuralWalk<ffi::WalkOrder::kPreOrder>(binding, walkfn).has_value()) {
return false;
}
}
Expand Down
15 changes: 11 additions & 4 deletions src/s_tir/schedule/analysis/analysis.cc
Original file line number Diff line number Diff line change
Expand Up @@ -17,6 +17,7 @@
* under the License.
*/
#include <tvm/ffi/cast.h>
#include <tvm/ffi/extra/structural_visit.h>
#include <tvm/ffi/reflection/registry.h>
#include <tvm/ir/op.h>
#include <tvm/s_tir/stmt.h>
Expand Down Expand Up @@ -1826,6 +1827,14 @@ ffi::Optional<TensorizeInfo> GetTensorizeLoopMapping(const s_tir::ScheduleState&
// C[i, j] += A[i, k] * B[k, j]

int next_block_ind = block_loops.size() - 1;
auto desc_walkfn = [&desc_loop_vars](const Var& var) -> ffi::Expected<ffi::WalkResult> {
return desc_loop_vars.count(var.get()) ? ffi::WalkResult::Interrupt(ffi::VisitInterrupt(var))
: ffi::WalkResult::Advance();
};
auto block_walkfn = [&block_loop_vars](const Var& var) -> ffi::Expected<ffi::WalkResult> {
return block_loop_vars.count(var.get()) ? ffi::WalkResult::Interrupt(ffi::VisitInterrupt(var))
: ffi::WalkResult::Advance();
};
for (int i_desc = n_desc_vars - 1; i_desc >= 0; --i_desc) {
// Step 3.1. Find the corresponding loop of the i_desc-th block var of desc
const PrimExpr& desc_bind = desc_block->iter_values[i_desc];
Expand All @@ -1834,8 +1843,7 @@ ffi::Optional<TensorizeInfo> GetTensorizeLoopMapping(const s_tir::ScheduleState&
for (int i = 0, n = desc_loops.size(); i < n; ++i) {
// Check if desc_bind = loops[i]->loop_var + stuff-irrelevant-of-loop-vars
PrimExpr residual = analyzer->Simplify(desc_bind - desc_loops[i]->loop_var);
if (!UsesVar(residual,
[&desc_loop_vars](const VarNode* var) { return desc_loop_vars.count(var); })) {
if (!ffi::StructuralWalk<ffi::WalkOrder::kPreOrder>(residual, desc_walkfn).has_value()) {
desc_loop = desc_loops[i];
iter_type_desc = iter_types_desc[i];
break;
Expand Down Expand Up @@ -1869,8 +1877,7 @@ ffi::Optional<TensorizeInfo> GetTensorizeLoopMapping(const s_tir::ScheduleState&
if (ret->loop_map.find(block_loop_sref) != ret->loop_map.end()) continue;

PrimExpr residual = analyzer->Simplify(block_bind - block_loops[i]->loop_var);
if (UsesVar(residual,
[&block_loop_vars](const VarNode* var) { return block_loop_vars.count(var); })) {
if (ffi::StructuralWalk<ffi::WalkOrder::kPreOrder>(residual, block_walkfn).has_value()) {
continue;
}
// padding is allowed only when the block has trivial bindings
Expand Down
10 changes: 7 additions & 3 deletions src/s_tir/schedule/analysis/reducer.cc
Original file line number Diff line number Diff line change
Expand Up @@ -17,6 +17,7 @@
* under the License.
*/
#include <tvm/ffi/cast.h>
#include <tvm/ffi/extra/structural_visit.h>
#include <tvm/te/operation.h>

#include "../utils.h"
Expand Down Expand Up @@ -553,10 +554,13 @@ bool ReductionIterNotIndexOutputBuffer(const SBlock& block) {
buffer_allocated.insert(buffer.get());
}

auto walkfn = [&](const Var& var) -> ffi::Expected<ffi::WalkResult> {
return reduction_block_iters.count(var.get())
? ffi::WalkResult::Interrupt(ffi::VisitInterrupt(var))
: ffi::WalkResult::Advance();
};
auto f_uses_reduction_block_var = [&](const PrimExpr& expr) -> bool {
return UsesVar(expr, [&](const VarNode* var) { //
return reduction_block_iters.count(var);
});
return ffi::StructuralWalk<ffi::WalkOrder::kPreOrder>(expr, walkfn).has_value();
};

std::unordered_map<const VarNode*, const VarNode*> match_buffer_sources;
Expand Down
Loading
Loading