This is an automated email from the ASF dual-hosted git repository.
lunderberg pushed a commit to branch unity
in repository https://gitbox.apache.org/repos/asf/tvm.git
The following commit(s) were added to refs/heads/unity by this push:
new d88cc4267d [Unity][Transform] Implement UpdateParamStructInfo (#16305)
d88cc4267d is described below
commit d88cc4267dc5a9ef9a9a5de315c45d5de526ce95
Author: Eric Lunderberg <[email protected]>
AuthorDate: Thu Jan 4 11:02:06 2024 -0600
[Unity][Transform] Implement UpdateParamStructInfo (#16305)
* [Unity][Transform] Implement UpdateParamStructInfo
Provide a convenience method to update parameter struct info,
propagating any changes to internal bindings and return type.
* lint fix
* Update implementation to update params in relax::Function mutator
---
python/tvm/relax/transform/__init__.py | 1 +
python/tvm/relax/transform/transform.py | 27 ++++-
src/relax/transform/update_param_struct_info.cc | 111 +++++++++++++++++++++
.../test_transform_update_param_struct_info.py | 71 +++++++++++++
4 files changed, 209 insertions(+), 1 deletion(-)
diff --git a/python/tvm/relax/transform/__init__.py
b/python/tvm/relax/transform/__init__.py
index 19316c76b8..3a0460f99f 100644
--- a/python/tvm/relax/transform/__init__.py
+++ b/python/tvm/relax/transform/__init__.py
@@ -67,6 +67,7 @@ from .transform import (
StaticPlanBlockMemory,
ToMixedPrecision,
ToNonDataflow,
+ UpdateParamStructInfo,
UpdateVDevice,
VMBuiltinLower,
VMShapeLower,
diff --git a/python/tvm/relax/transform/transform.py
b/python/tvm/relax/transform/transform.py
index 9589f661d7..e4ba4b7d22 100644
--- a/python/tvm/relax/transform/transform.py
+++ b/python/tvm/relax/transform/transform.py
@@ -24,7 +24,7 @@ from typing import Callable, Dict, List, Mapping, Optional,
Sequence, Tuple, Uni
import numpy as np # type: ignore
import tvm.ir
-from tvm.relax import Expr, Var
+from tvm.relax import Expr, Var, StructInfo
from tvm.relax.dpl import DFPattern
from tvm.runtime import NDArray, Object
from tvm.tir import IndexMap, PrimFunc
@@ -1224,6 +1224,31 @@ def SplitCallTIRByPattern(patterns: List[PrimFunc],
fcodegen: Callable) -> tvm.i
return _ffi_api.SplitCallTIRByPattern(patterns, fcodegen) # type: ignore
+def UpdateParamStructInfo(sinfo_func: Callable[[Var], Optional[StructInfo]]):
+ """Update struct info of parameters
+
+ Update struct info of parameters. Internal bindings and function
+ return type will be updated using relax's struct inference rules.
+ Errors resulting from struct inference will be propagated to the
+ user.
+
+ Parameters
+ ----------
+ sinfo_func: Callable[[Var], Optional[StructInfo]]
+
+ A function that is called once for each function parameter,
+ and returns the updated struct info to be used for it. If the
+ function returns `None`, the parameter is not modified.
+
+ Returns
+ -------
+ ret : tvm.transform.Pass
+ The corresponding pass.
+
+ """
+ return _ffi_api.UpdateParamStructInfo(sinfo_func) # type: ignore
+
+
def CombineParallelMatmul(check=None):
"""Combine multiple matmul operators sharing the same LHS matrix into one,
followed by slicing. When all matmul branches in a tree have the same set
of fused ops,
diff --git a/src/relax/transform/update_param_struct_info.cc
b/src/relax/transform/update_param_struct_info.cc
new file mode 100644
index 0000000000..327185fd0b
--- /dev/null
+++ b/src/relax/transform/update_param_struct_info.cc
@@ -0,0 +1,111 @@
+/*
+ * 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.
+ */
+
+/*!
+ * \file tvm/relax/transform/update_param_struct_info.cc
+ * \brief Mutate IRModule to accept new parameters
+ */
+
+#include <tvm/relax/expr.h>
+#include <tvm/relax/expr_functor.h>
+#include <tvm/relax/transform.h>
+
+#include <optional>
+#include <regex>
+#include <unordered_map>
+#include <vector>
+
+#include "utils.h"
+
+namespace tvm {
+namespace relax {
+
+namespace {
+class ParamStructInfoMutator : public ExprMutator {
+ public:
+ explicit ParamStructInfoMutator(TypedPackedFunc<Optional<StructInfo>(Var)>
sinfo_func)
+ : sinfo_func_(sinfo_func) {}
+
+ using ExprMutator::VisitExpr_;
+ using ExprMutator::VisitVarDef_;
+
+ Expr VisitExpr_(const FunctionNode* op) override {
+ auto func = GetRef<Function>(op);
+
+ auto params = op->params.Map([this](Var param) {
+ if (auto new_sinfo = sinfo_func_(param)) {
+ auto new_param = WithStructInfo(param, new_sinfo.value());
+ var_remap_[param->vid] = new_param;
+ return new_param;
+ } else {
+ return param;
+ }
+ });
+
+ if (!params.same_as(func->params)) {
+ func.CopyOnWrite()->params = params;
+ }
+ return ExprMutator::VisitExpr_(func.get());
+ }
+
+ TypedPackedFunc<Optional<StructInfo>(Var)> sinfo_func_;
+};
+} // namespace
+
+namespace transform {
+Pass UpdateParamStructInfo(TypedPackedFunc<Optional<StructInfo>(Var)>
sinfo_func) {
+ auto pass_func = [=](IRModule mod, PassContext pc) {
+ ParamStructInfoMutator mutator(sinfo_func);
+
+ std::unordered_set<GlobalVar, ObjectPtrHash, ObjectPtrEqual> to_remove;
+ std::unordered_map<GlobalVar, Function, ObjectPtrHash, ObjectPtrEqual>
to_add;
+
+ for (const auto& [gvar, base_func] : mod->functions) {
+ if (auto func = base_func.as<Function>()) {
+ auto updated = Downcast<Function>(mutator(func.value()));
+ if (!updated.same_as(base_func)) {
+ GlobalVar new_gvar(gvar->name_hint);
+ UpdateStructInfo(new_gvar, GetStructInfo(updated));
+ to_add.insert({new_gvar, updated});
+ to_remove.insert(gvar);
+ }
+ }
+ }
+
+ if (to_remove.size() || to_add.size()) {
+ auto write_ptr = mod.CopyOnWrite();
+
+ for (const auto& gvar : to_remove) {
+ write_ptr->Remove(gvar);
+ }
+ for (const auto& [gvar, func] : to_add) {
+ write_ptr->Add(gvar, func);
+ }
+ }
+
+ return mod;
+ };
+ return tvm::transform::CreateModulePass(pass_func, 1,
"UpdateParamStructInfo", {});
+}
+
+TVM_REGISTER_GLOBAL("relax.transform.UpdateParamStructInfo").set_body_typed(UpdateParamStructInfo);
+
+} // namespace transform
+} // namespace relax
+} // namespace tvm
diff --git a/tests/python/relax/test_transform_update_param_struct_info.py
b/tests/python/relax/test_transform_update_param_struct_info.py
new file mode 100644
index 0000000000..6680580fe0
--- /dev/null
+++ b/tests/python/relax/test_transform_update_param_struct_info.py
@@ -0,0 +1,71 @@
+# 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.
+
+import inspect
+from typing import Optional
+
+import pytest
+
+import tvm.testing
+from tvm import relax
+from tvm.script import ir as I, relax as R
+
+
+class Base:
+ def test_compare(self):
+ transform = relax.transform.UpdateParamStructInfo(self.update_sinfo)
+
+ if inspect.isclass(self.Expected) and issubclass(self.Expected,
Exception):
+ with pytest.raises(self.Expected):
+ transform(self.Before)
+ else:
+ after = transform(self.Before)
+ tvm.ir.assert_structural_equal(self.Expected, after)
+
+ def update_sinfo(self, var: relax.Var) -> Optional[relax.StructInfo]:
+ """The struct info update function provided to the transform"""
+ raise NotImplementedError("Should be implemented in derived class")
+
+
+class TestSimple(Base):
+ def update_sinfo(self, var: relax.Var) -> Optional[relax.StructInfo]:
+ if var.name_hint == "weight":
+ return relax.TensorStructInfo([64, 16], "float32")
+
+ @I.ir_module
+ class Before:
+ @R.function
+ def main(
+ x: R.Tensor([16], "float32"),
+ weight: R.Tensor([32, 16], "float32"),
+ ) -> R.Tensor([32], "float32"):
+ out: R.Tensor([32], "float32") = R.matmul(weight, x)
+ return out
+
+ @I.ir_module
+ class Expected:
+ @R.function
+ def main(
+ x: R.Tensor([16], "float32"),
+ weight: R.Tensor([64, 16], "float32"),
+ ) -> R.Tensor([64], "float32"):
+ out: R.Tensor([64], "float32") = R.matmul(weight, x)
+ return out
+
+
+if __name__ == "__main__":
+ tvm.testing.main()