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
7 changes: 6 additions & 1 deletion src/relax/transform/bundle_model_params.cc
Original file line number Diff line number Diff line change
Expand Up @@ -74,7 +74,7 @@ class ModelParamBundler : public ExprMutator {
Expr VisitExpr_(const VarNode* op) override {
auto var = GetRef<Var>(op);
if (auto it = var_to_expr_.find(var); it != var_to_expr_.end()) {
return (*it).second;
return builder_->Emit((*it).second, op->name_hint());
} else {
return ExprMutator::VisitExpr_(op);
}
Expand All @@ -84,6 +84,11 @@ class ModelParamBundler : public ExprMutator {
Map<Var, Expr> var_to_expr_;
};

Function BundleModelParams(const Function& func) {
ModelParamBundler mutator;
return Downcast<Function>(mutator(func));
}

namespace transform {
Pass BundleModelParams() {
runtime::TypedPackedFunc<IRModule(IRModule, PassContext)> pass_func = [=](IRModule mod,
Expand Down
12 changes: 12 additions & 0 deletions src/relax/transform/utils.h
Original file line number Diff line number Diff line change
Expand Up @@ -421,6 +421,18 @@ Expr EliminateCommonSubexpr(const Expr& expr, bool call_only = false);
*/
Expr CanonicalizeBindings(const Expr& expr);

/* \brief Remove use of trivial bindings
*
* Utility for converting from individual model parameters to a single
* parameter with a tuple of parameters. If the `kNumInput` attribute
* is absent, no model parameters are present, so no updates are made.
*
* \param func The function to be updated.
*
* \ret The updated function.
*/
Function BundleModelParams(const Function& func);

} // namespace relax
} // namespace tvm

Expand Down
91 changes: 91 additions & 0 deletions tests/python/relax/test_transform_bundle_model_params.py
Original file line number Diff line number Diff line change
Expand Up @@ -102,5 +102,96 @@ def main(
tvm.ir.assert_structural_equal(after, Expected)


def test_dataflow():
"""Parameters can be substituted into a dataflow block"""

@tvm.script.ir_module
class Before:
@R.function
def main(
a: R.Tensor([16], "float32"),
b: R.Tensor([16], "float32"),
c: R.Tensor([16], "float32"),
) -> R.Tensor([16], "float32"):
R.func_attr({"num_input": 1})
with R.dataflow():
expr = a
expr = R.add(expr, b)
expr = R.add(expr, c)
R.output(expr)
return expr

@tvm.script.ir_module
class Expected:
@R.function
def main(
a: R.Tensor([16], "float32"),
params: R.Tuple(R.Tensor([16], "float32"), R.Tensor([16], "float32")),
) -> R.Tensor([16], "float32"):
R.func_attr({"num_input": 1})
with R.dataflow():
expr = a
b = params[0]
expr = R.add(expr, b)
c = params[1]
expr = R.add(expr, c)
R.output(expr)
return expr

mod = Before
after = relax.transform.BundleModelParams()(mod)
tvm.ir.assert_structural_equal(after, Expected)


def test_variable_names():
"""Parameters retain their names within the updated function

For readability, the parameter names should be used to generate
the new variable names.

Like `test_basic`, but explicitly checks the names of bound
variables.
"""

@tvm.script.ir_module
class Before:
@R.function
def main(
a: R.Tensor([16], "float32"),
b: R.Tensor([16], "float32"),
c: R.Tensor([16], "float32"),
) -> R.Tensor([16], "float32"):
R.func_attr({"num_input": 1})
expr = a
expr = R.add(expr, b)
expr = R.add(expr, c)
return expr

@tvm.script.ir_module
class Expected:
@R.function
def main(
a: R.Tensor([16], "float32"),
params: R.Tuple(R.Tensor([16], "float32"), R.Tensor([16], "float32")),
) -> R.Tensor([16], "float32"):
R.func_attr({"num_input": 1})
expr = a
b = params[0]
expr = R.add(expr, b)
c = params[1]
expr = R.add(expr, c)
return expr

mod = Before
after = relax.transform.BundleModelParams()(mod)
tvm.ir.assert_structural_equal(after, Expected)

for binding, expected_binding in zip(
after["main"].body.blocks[0].bindings,
Expected["main"].body.blocks[0].bindings,
):
assert binding.var.name_hint == expected_binding.var.name_hint


if __name__ == "__main__":
tvm.testing.main()