This is an automated email from the ASF dual-hosted git repository.

masahi 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 a8f3a22cdd [Unity][TuningAPI] Temporary patch for large models  
(#14691)
a8f3a22cdd is described below

commit a8f3a22cddacf711e3022197137ade93f69f4c45
Author: Sunghyun Park <[email protected]>
AuthorDate: Thu Apr 20 23:35:44 2023 -0700

    [Unity][TuningAPI] Temporary patch for large models  (#14691)
    
    patch
---
 python/tvm/relax/transform/__init__.py             |  1 +
 src/relax/transform/meta_schedule.cc               | 21 ++++++++-----
 .../relax/test_transform_meta_schedule_tuning.py   | 35 +++++++++++++++++++---
 3 files changed, 46 insertions(+), 11 deletions(-)

diff --git a/python/tvm/relax/transform/__init__.py 
b/python/tvm/relax/transform/__init__.py
index 78f450b25c..977a03b96a 100644
--- a/python/tvm/relax/transform/__init__.py
+++ b/python/tvm/relax/transform/__init__.py
@@ -21,3 +21,4 @@ from .transform import *
 
 # Import to register the legalization functions.
 from . import legalize_ops
+from . import tuning_api
diff --git a/src/relax/transform/meta_schedule.cc 
b/src/relax/transform/meta_schedule.cc
index bb9c3579e7..03456d0ef8 100644
--- a/src/relax/transform/meta_schedule.cc
+++ b/src/relax/transform/meta_schedule.cc
@@ -41,23 +41,27 @@ class MetaScheduleTuner {
         work_dir_(work_dir),
         max_trials_global_(max_trials_global),
         params_(params) {
-    candgen_func_ = 
runtime::Registry::Get("relax.tuning_api.default_generate_candidate");
-    ICHECK(candgen_func_) << "Default candidate generation function is not 
found.";
+    // candgen_func_ = 
runtime::Registry::Get("relax.tuning_api.default_generate_candidate");
+    // ICHECK(candgen_func_) << "Default candidate generation function is not 
found.";
     normalize_mod_func_ = 
runtime::Registry::Get("tvm.meta_schedule.normalize_mod");
     ICHECK(normalize_mod_func_) << "Normalization function is not found.";
   }
 
   // TODO(@sunggg): Currently, only supports basic arguments.
   IRModule TuneIRMod(IRModule mod, transform::PassContext ctx) {
-    Trace trace = Downcast<Trace>(ctx->GetCurrentTrace());
-    ctx->PopTrace();
     Choice choice("tvm.meta_schedule.tune_relax", {params_, target_, 
work_dir_, max_trials_global_},
                   "relax.tuning_api.Choice.default_constr_func", {});
     Knob knob("meta_schedule.tune_irmod", {{"0", choice}});
+    knob->Apply(mod, "0");
+    /*
+    // TODO(@sunggg): revisit when we have a solution for large params
+    Trace trace = Downcast<Trace>(ctx->GetCurrentTrace());
+    ctx->PopTrace();
     Array<Trace> candidates = (*candgen_func_)(Array<Knob>({knob}), trace);
     ICHECK(candidates.size() == 1);
     Trace best_trace = candidates[0];
     ctx->PushTrace(best_trace);
+    */
     // since we separate tuning from application, return original IRModule
     return mod;
   }
@@ -66,13 +70,16 @@ class MetaScheduleTuner {
   tir::PrimFunc TuneTIR(tir::PrimFunc f, transform::PassContext ctx) {
     // TODO(@sunggg): Whenever we tune tir, assume we start a new trace w/o 
pushing to the trace
     // stack. Revisit later when we collect more usecases.
-    Trace trace = Trace((*normalize_mod_func_)(f), {}, {});
-
     Choice choice("tvm.meta_schedule.tune_tir", {target_, work_dir_, 
max_trials_global_},
                   "relax.tuning_api.Choice.default_constr_func", {});
     Knob knob("meta_schedule.tune_primfunc", {{"0", choice}});
+    knob->Apply((*normalize_mod_func_)(f), "0");
+    /*
+    // TODO(@sunggg): revisit when we have a solution for large params
+    Trace trace = Trace((*normalize_mod_func_)(f), {}, {});
     Array<Trace> candidates = (*candgen_func_)(Array<Knob>({knob}), trace);
     ICHECK(candidates.size() == 1);
+    */
     // since we separate tuning from application, return original IRModule
     return f;
   }
@@ -82,7 +89,7 @@ class MetaScheduleTuner {
   String work_dir_;
   Integer max_trials_global_;
   Map<String, runtime::NDArray> params_;
-  const runtime::PackedFunc* candgen_func_;
+  // const runtime::PackedFunc* candgen_func_;
   const runtime::PackedFunc* normalize_mod_func_;
 };
 
diff --git a/tests/python/relax/test_transform_meta_schedule_tuning.py 
b/tests/python/relax/test_transform_meta_schedule_tuning.py
index 39331548e4..267fd98e76 100644
--- a/tests/python/relax/test_transform_meta_schedule_tuning.py
+++ b/tests/python/relax/test_transform_meta_schedule_tuning.py
@@ -24,7 +24,9 @@ from tvm import relax
 from tvm.ir import transform
 from tvm.ir.module import IRModule
 from tvm.ir.transform import PassContext
-from tvm.relax.transform.tuning_api import Trace
+
+# TODO(@sunggg): re-enable Trace when we have a solution for large params
+# from tvm.relax.transform.tuning_api import Trace
 from tvm.script import relax as R
 from tvm.script import tir as T
 
@@ -78,7 +80,9 @@ def test_ms_tuning_irmodule():
     assert isinstance(mod, IRModule)
 
     with tempfile.TemporaryDirectory() as work_dir:
-        with target, transform.PassContext(trace=Trace(mod), opt_level=0):
+        """
+        # TODO(@sunggg): revisit when ready
+        with target, PassContext(trace=Trace(mod), opt_level=0):
             tuning_pass = relax.transform.MetaScheduleTuneIRMod(
                 params={}, work_dir=work_dir, max_trials_global=4
             )
@@ -86,6 +90,13 @@ def test_ms_tuning_irmodule():
             assert PassContext.current().get_trace_stack_size() == 1
             assert PassContext.current().get_current_trace().size == 1
             tvm.ir.assert_structural_equal(mod, out_mod)
+        """
+
+        with target, PassContext(opt_level=0):
+            tuning_pass = relax.transform.MetaScheduleTuneIRMod(
+                params={}, work_dir=work_dir, max_trials_global=4
+            )
+            out_mod = tuning_pass(mod)
 
             application_pass = 
relax.transform.MetaScheduleApplyDatabase(work_dir)
 
@@ -97,7 +108,9 @@ def test_ms_tuning_primfunc():
     mod = InputModule
     assert isinstance(mod, IRModule)
     with tempfile.TemporaryDirectory() as work_dir:
-        with target, transform.PassContext(trace=Trace(mod), opt_level=0):
+        """
+        # TODO(@sunggg): revisit when ready
+        with target, PassContext(trace=Trace(mod), opt_level=0):
             tuning_pass = relax.transform.MetaScheduleTuneTIR(
                 work_dir=work_dir, max_trials_global=4
             )
@@ -106,6 +119,12 @@ def test_ms_tuning_primfunc():
             # TODO (@sunggg): Need to determine how to track subgraph-level 
tuning traces.
             # Currently, we don't track this so the trace size. Revisit this 
later.
             tvm.ir.assert_structural_equal(mod, out_mod)
+        """
+        with target, PassContext(opt_level=0):
+            tuning_pass = relax.transform.MetaScheduleTuneIRMod(
+                params={}, work_dir=work_dir, max_trials_global=4
+            )
+            out_mod = tuning_pass(mod)
 
             application_pass = 
relax.transform.MetaScheduleApplyDatabase(work_dir)
             out_mod = application_pass(mod)
@@ -172,12 +191,20 @@ def test_ms_database_apply_fallback():
     target_cuda = tvm.target.Target("nvidia/geforce-rtx-3090-ti")
     assert isinstance(mod, IRModule)
     with tempfile.TemporaryDirectory() as work_dir:
-        with target_cuda, transform.PassContext(trace=Trace(mod), opt_level=0):
+        """
+        # TODO(@sunggg): Revisit when ready
+        with target_cuda, PassContext(trace=Trace(mod), opt_level=0):
             tuning_pass = relax.transform.MetaScheduleTuneTIR(
                 work_dir=work_dir, max_trials_global=0
             )
             out_mod = tuning_pass(mod)
             tvm.ir.assert_structural_equal(mod, out_mod)
+        """
+        with target_cuda, PassContext(opt_level=0):
+            tuning_pass = relax.transform.MetaScheduleTuneTIR(
+                work_dir=work_dir, max_trials_global=0
+            )
+            out_mod = tuning_pass(mod)
             default_pass = tvm.tir.transform.DefaultGPUSchedule()
             out_mod = default_pass(mod)
             tvm.ir.assert_structural_equal(out_mod, DefaultScheduledModule)

Reply via email to