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)