Skip to content
Open
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
161 changes: 146 additions & 15 deletions src/relax/transform/fuse_ops.cc
Original file line number Diff line number Diff line change
Expand Up @@ -28,6 +28,7 @@
*/

#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/dataflow_matcher.h>
Expand Down Expand Up @@ -106,6 +107,110 @@ constexpr uint32_t kMaxFusedOps = 256;

TVM_REGISTER_PASS_CONFIG_OPTION("relax.FuseOps.max_depth", int64_t);

// Shape symbols are not ordinary dataflow variables. Record the match_cast that
// first defines each symbol, so both partitioning and scheduling see its uses.
class SymbolicDependencyCollector {
public:
using DependencyMap = std::unordered_map<const VarNode*, std::vector<Var>>;
using VisitResult = ffi::Expected<ffi::Optional<ffi::VisitInterrupt>>;

static DependencyMap Collect(const Function& func) {
SymbolicDependencyCollector collector;
ffi::StructuralVisit(
func,
[&](const Function& function, ffi::StructuralVisitorObj* visitor) -> VisitResult {
auto saved_producers = collector.producers_;
auto saved_binding = collector.current_binding_;
collector.current_binding_ = nullptr;
for (const Var& param : function->params) {
// Scalar parameters can themselves be used as shape symbols.
collector.producers_.emplace(param.get(), nullptr);
collector.DefineSymbols(GetType(param));
}
auto result = visitor->VisitExpected(function->body);
collector.current_binding_ = saved_binding;
collector.producers_ = std::move(saved_producers);
return result;
},
[&](const If& if_expr, ffi::StructuralVisitorObj* visitor) -> VisitResult {
visitor->Visit(if_expr->cond);
auto saved_producers = collector.producers_;
auto saved_binding = collector.current_binding_;
// Branch-local definitions are invisible to the other branch and to
// the enclosing binding, including uses in the branch result type.
collector.current_binding_ = nullptr;
visitor->Visit(if_expr->true_branch);
collector.producers_ = saved_producers;
visitor->Visit(if_expr->false_branch);
collector.producers_ = std::move(saved_producers);
collector.current_binding_ = saved_binding;
return visitor->VisitExpected(GetType(if_expr));
},
[&](const Binding& binding, ffi::StructuralVisitorObj* visitor) -> VisitResult {
auto saved_binding = collector.current_binding_;
collector.current_binding_ = binding->var.get();
visitor->Visit(GetBoundValue(binding));
visitor->Visit(GetType(binding->var));
if (const auto* match_cast = binding.as<MatchCastNode>()) {
// Existing symbols in the value or target type are uses, not redefinitions.
visitor->Visit(match_cast->ty);
collector.DefineSymbols(match_cast->ty);
}
collector.current_binding_ = saved_binding;
return std::nullopt;
},
[&](const Var& var, ffi::StructuralVisitorObj* visitor) -> VisitResult {
collector.UseVar(var);
return visitor->VisitExpected(GetType(var));
},
[](const FuncType&, ffi::StructuralVisitorObj*) -> VisitResult {
// Function-type symbols are locally bound, not caller dependencies.
return std::nullopt;
});
return collector.dependencies_;
}

private:
void UseVar(const Var& var) {
auto it = producers_.find(var.get());
if (current_binding_ && it != producers_.end() && it->second &&
it->second != current_binding_) {
auto& deps = dependencies_[current_binding_];
Var producer = ffi::GetRef<Var>(it->second);
if (std::none_of(deps.begin(), deps.end(),
[&](const Var& dep) { return dep.same_as(producer); })) {
deps.push_back(producer);
}
}
}

void DefineSymbols(const Type& ty) {
auto define_shape = [&](const ffi::Array<PrimExpr>& values) {
for (const PrimExpr& value : values) {
// Only a bare symbol is definable; compound expressions are constraints.
if (const auto* var = value.as<tirx::VarNode>()) {
producers_.emplace(var, current_binding_);
}
}
};
ffi::StructuralWalk<ffi::WalkOrder::kPreOrder>(
ty,
[](const FuncType&) -> ffi::Expected<ffi::WalkResult> { return ffi::WalkResult::Skip(); },
[&](const ShapeType& shape) -> ffi::Expected<ffi::WalkResult> {
if (shape->values.has_value()) define_shape(shape->values.value());
return ffi::WalkResult::Skip();
},
[&](const ShapeExpr& shape) -> ffi::Expected<ffi::WalkResult> {
define_shape(shape->values);
return ffi::WalkResult::Skip();
});
}

const VarNode* current_binding_{nullptr};
std::unordered_map<const VarNode*, const VarNode*> producers_;
DependencyMap dependencies_;
};

class GraphCreator : public ExprVisitor {
public:
/*!
Expand All @@ -130,7 +235,9 @@ class GraphCreator : public ExprVisitor {
func->GetAttr<ffi::String>(attr::kCodegen).has_value()) {
continue;
}
creator(ffi::GetRef<Function>(func));
auto function = ffi::GetRef<Function>(func);
creator.symbolic_deps_ = SymbolicDependencyCollector::Collect(function);
creator(function);
}

// The algorithm of the graph creator ensures that each created node will be added to the
Expand Down Expand Up @@ -167,7 +274,8 @@ class GraphCreator : public ExprVisitor {

void VisitBinding_(const MatchCastNode* binding) final {
IndexedForwardGraph::Node* node = CreateNode(binding->var.get());
SetNodePattern(node, OpPatternKind::kOpaque);
VisitUnsupportedNode(binding->value, node);
AddSymbolicDependencies(binding->var, node);
AddToPostDFSOrder(node, binding->var.get());
}

Expand All @@ -190,11 +298,18 @@ class GraphCreator : public ExprVisitor {
// Case 3. The type of the expression is not fusion-supported.
// In this case, we skip adding edges, adding an empty node into graph.
}
AddSymbolicDependencies(binding->var, node);
AddToPostDFSOrder(node, binding->var.get());
}

/********** Non-Leaf Expression Nodes **********/

void AddSymbolicDependencies(const Var& var, IndexedForwardGraph::Node* node) {
for (const Var& producer : symbolic_deps_[var.get()]) {
AddEdge(graph_.node_map.at(producer.get()), node, OpPatternKind::kOpaque);
}
}

void VisitCall(const CallNode* call, IndexedForwardGraph::Node* binding_var_node) {
TVM_FFI_ICHECK_NOTNULL(binding_var_node);

Expand Down Expand Up @@ -381,6 +496,8 @@ class GraphCreator : public ExprVisitor {
std::unordered_set<IndexedForwardGraph::Node*> initialized_nodes_;
/*! \brief The model params in the function input */
std::unordered_set<const VarNode*> input_params_;
/*! \brief Dependencies on bindings that define shape symbols. */
SymbolicDependencyCollector::DependencyMap symbolic_deps_;
};

/*!
Expand Down Expand Up @@ -839,6 +956,7 @@ class OperatorFusor : public ExprMutator {
if (func->IsInstance<relax::FunctionNode>() && !func->HasNonzeroAttr(attr::kPrimitive) &&
!func->GetAttr<ffi::String>(attr::kCodegen).has_value()) {
outer_bindings_ = AnalyzeVar2Value(func);
symbolic_deps_ = SymbolicDependencyCollector::Collect(func.as_or_throw<Function>());
auto updated_func = VisitExpr(func).as_or_throw<Function>();
builder_->UpdateFunction(gv, updated_func);
outer_bindings_ = {};
Expand All @@ -862,6 +980,7 @@ class OperatorFusor : public ExprMutator {

BindingBlock VisitBindingBlock_(const DataflowBlockNode* block) final {
group2func_.clear();
group_deps_.clear();

// Step 1. Collect the bindings for each grouped function.
CollectFuncBindings(block->bindings);
Expand Down Expand Up @@ -1011,20 +1130,11 @@ class OperatorFusor : public ExprMutator {
// - If the var's group is same as the binding's, the var is defined in the same group
// - If the var's group is different with the binding's, the var must be the output from
// another group. Mark it to be the group output.
auto update_boundary = [this, binding, &cur_group](const Expr& e) {
auto update_boundary = [this, &cur_group](const Expr& e) {
if (e->IsInstance<VarNode>() && obj2group_.count(e.get())) {
const Var& used_var = e.as_or_throw<Var>();
Group* producer_group = GetGroupFromVar(used_var);
// Only check those group defined before.
// Skip the vars from input or groups with single binding.
if (producer_group != cur_group) {
for (Group* depgroup : group_deps_[producer_group]) {
TVM_FFI_ICHECK(depgroup != cur_group)
<< "A cyclic dependency detected between the groups " << binding->var->name
<< " and " << used_var->name << " are in.";
}
group_deps_[cur_group].push_back(producer_group);
}
AddGroupDependency(cur_group, producer_group);

if (auto producer = group2func_.find(producer_group);
producer_group != cur_group && producer != group2func_.end()) {
Expand All @@ -1040,6 +1150,21 @@ class OperatorFusor : public ExprMutator {
TVM_FFI_ICHECK_NOTNULL(match_cast);
PostOrderVisit(match_cast->value, update_boundary);
}

// Shape dependencies constrain scheduling without adding tensor outputs or
// parameters to the grouped functions.
for (const Var& producer : symbolic_deps_[binding->var.get()]) {
if (obj2group_.count(producer.get())) {
AddGroupDependency(cur_group, GetGroupFromVar(producer));
}
}
}
}

void AddGroupDependency(Group* consumer, Group* producer) {
auto& deps = group_deps_[consumer];
if (consumer != producer && std::find(deps.begin(), deps.end(), producer) == deps.end()) {
deps.push_back(producer);
}
}

Expand Down Expand Up @@ -1094,14 +1219,18 @@ class OperatorFusor : public ExprMutator {
}

std::unordered_set<Group*> visited;
std::unordered_set<Group*> visiting;

std::function<void(Group*, std::function<void(Group*)>)> dfs_visit;
dfs_visit = [this, &visited, &dfs_visit](Group* g, auto leaf_fun) {
dfs_visit = [this, &visited, &visiting, &dfs_visit](Group* g, auto leaf_fun) {
if (!visited.count(g)) {
visited.insert(g);
TVM_FFI_ICHECK(visiting.insert(g).second)
<< "A cyclic dependency detected between fusion groups.";
for (auto dep : group_deps_[g]) {
dfs_visit(dep, leaf_fun);
}
visiting.erase(g);
visited.insert(g);
leaf_fun(g);
}
};
Expand Down Expand Up @@ -1129,6 +1258,8 @@ class OperatorFusor : public ExprMutator {
std::unordered_map<Group*, FunctionCreator> group2func_;
/*! \brief Bindings visible while rewriting the current Relax function. */
ffi::Map<Var, Expr> outer_bindings_;
/*! \brief Dependencies on bindings that define shape symbols. */
SymbolicDependencyCollector::DependencyMap symbolic_deps_;
/*!
* \brief A map from a group to its dependent groups, used to detect cyclic dependencies.
* \note Use vector so we can be deterministic, there won't be a lot of dep groups so
Expand Down
74 changes: 74 additions & 0 deletions tests/python/relax/test_frontend_onnx.py
Original file line number Diff line number Diff line change
Expand Up @@ -7636,6 +7636,80 @@ def main(
# )


@pytest.mark.parametrize("input_shape", [(1, 4, 5, 5), (2, 3, 4, 4)])
def test_size_slice_reshape_with_fusion(input_shape):
"""Regression for #20177: preserve the runtime Slice shape before allocation."""
x = helper.make_tensor_value_info("x", TensorProto.FLOAT, input_shape)
y = helper.make_tensor_value_info("y", TensorProto.FLOAT, input_shape)
initializers = [
helper.make_tensor("flat_shape", TensorProto.INT64, [1], [-1]),
helper.make_tensor("size_shape", TensorProto.INT64, [1], [1]),
helper.make_tensor("slice_starts", TensorProto.INT64, [1], [0]),
helper.make_tensor("slice_axes", TensorProto.INT64, [1], [0]),
helper.make_tensor("slice_steps", TensorProto.INT64, [1], [1]),
]
nodes = [
helper.make_node("Relu", ["x"], ["h0"]),
helper.make_node("Relu", ["h0"], ["positive"]),
helper.make_node("Reshape", ["x", "flat_shape"], ["flat"]),
helper.make_node("Size", ["flat"], ["size"]),
helper.make_node("Reshape", ["size", "size_shape"], ["size_1d"]),
helper.make_node(
"Slice",
["flat", "slice_starts", "size_1d", "slice_axes", "slice_steps"],
["projected"],
),
helper.make_node("Shape", ["x"], ["x_shape"]),
helper.make_node("Reshape", ["projected", "x_shape"], ["splice"]),
helper.make_node("Add", ["positive", "splice"], ["y"]),
]
graph = helper.make_graph(
nodes,
"vm_shape_lower_size_slice",
[x],
[y],
initializers,
)
# Keep the model compatible with the ONNX Runtime versions used in CI.
model = helper.make_model(graph, opset_imports=[helper.make_opsetid("", 18)], ir_version=8)

onnx.checker.check_model(model)
session = onnxruntime.InferenceSession(
model.SerializeToString(), providers=["CPUExecutionProvider"]
)
mod = from_onnx(model, opset=18, keep_params_in_input=True)
pipeline = tvm.transform.Sequential(
[
relax.backend.DispatchSampling(),
relax.backend.DispatchSortScan(),
relax.transform.LegalizeOps(),
relax.transform.AnnotateTIROpPattern(),
relax.transform.FoldConstant(),
relax.transform.FuseOps(fuse_opt_level=2),
relax.transform.FuseTIR(),
relax.transform.RewriteDataflowReshape(),
relax.transform.ToNonDataflow(),
relax.transform.RemovePurityChecking(),
relax.transform.CallTIRRewrite(),
relax.transform.StaticPlanBlockMemory(),
relax.transform.LowerAllocTensor(),
relax.transform.KillAfterLastUse(),
relax.transform.LowerRuntimeBuiltin(),
relax.transform.ComputePrimValue(),
relax.transform.VMShapeLower(emit_err_ctx=True),
relax.transform.AttachGlobalSymbol(),
]
)
mod, params = relax.frontend.detach_params(mod)
executable = tvm.compile(mod, target="llvm", relax_pipeline=pipeline)
vm = relax.VirtualMachine(executable, tvm.cpu())
data = (np.arange(np.prod(input_shape), dtype="float32") % 7 - 3).reshape(input_shape)
actual = vm["main"](tvm.runtime.tensor(data), *params.get("main", [])).numpy()
expected = session.run(None, {"x": data})[0]
tvm.testing.assert_allclose(actual, expected)
tvm.testing.assert_allclose(actual, np.maximum(data, 0) + data)


def test_slice_dynamic_inputs_ir():
slice_node = helper.make_node("Slice", ["x", "starts", "ends", "axes", "steps"], ["y"])

Expand Down
Loading
Loading