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()

Reply via email to