This is an automated email from the ASF dual-hosted git repository.
lunderberg pushed a commit to branch main
in repository https://gitbox.apache.org/repos/asf/tvm.git
The following commit(s) were added to refs/heads/main by this push:
new 7131411f0a [Bugfix][TIR][VTA] Update host-side target, even without
device func (#14982)
7131411f0a is described below
commit 7131411f0a212cbcf8c16547e4b633001cba469e
Author: Eric Lunderberg <[email protected]>
AuthorDate: Tue May 30 15:43:21 2023 -0500
[Bugfix][TIR][VTA] Update host-side target, even without device func
(#14982)
This resolves an issue introduced by the combination of
https://github.com/apache/tvm/pull/14918 and
https://github.com/apache/tvm/pull/14945. The bug occurred for
targets that do not require device-side codegen, but do require a
`device_type` other than `kDLCPU`. It wasn't caught by CI, as the
issue only occurred with the combination of both PRs.
1. #14918 updated `SplitHostDevice` to only modify the `"target"`
attribute when a device-side function has been extracted.
2. For VTA, there is no device-side function, as everything is done
through host-side API calls.
3. From (1) and (2), the VTA examples kept the target
`T.target("ext_dev", host="llvm")` after the `SplitHostDevice`
pass, instead of being updated to `T.target("llvm")`.
4. #14945 restricted CombineContextCall to only apply to host-side
passes.
5. From (4) and (5), the `CombineContextCall` pass was no longer
applied to the VTA context calls.
This PR fixes `SplitHostDevice`, updating the target from
`T.target("ext_dev", host="llvm")` to `T.target("llvm")`, even if no
device sections have been extracted from the function.
---
src/tir/transforms/split_host_device.cc | 10 +++++-----
.../unittest/test_tir_transform_split_host_device.py | 16 ++++++++++++++++
2 files changed, 21 insertions(+), 5 deletions(-)
diff --git a/src/tir/transforms/split_host_device.cc
b/src/tir/transforms/split_host_device.cc
index 9270b356ba..2de831e8ad 100644
--- a/src/tir/transforms/split_host_device.cc
+++ b/src/tir/transforms/split_host_device.cc
@@ -108,12 +108,12 @@ PrimFunc SplitHostDevice(PrimFunc func, IRModule*
device_mod, const GlobalVar& g
HostDeviceSplitter splitter(device_mod, name_prefix);
- auto body = splitter(func->body);
-
- if (!body.same_as(func->body)) {
+ if (auto body = splitter(func->body); !body.same_as(func->body)) {
func.CopyOnWrite()->body = body;
- auto target_host = target->GetHost().value_or(Target("llvm"));
- func = WithAttr(std::move(func), tvm::attr::kTarget, target_host);
+ }
+
+ if (auto target_host = target->GetHost()) {
+ func = WithAttr(std::move(func), tvm::attr::kTarget, target_host.value());
}
return func;
diff --git a/tests/python/unittest/test_tir_transform_split_host_device.py
b/tests/python/unittest/test_tir_transform_split_host_device.py
index cf866ae005..1599b9a031 100644
--- a/tests/python/unittest/test_tir_transform_split_host_device.py
+++ b/tests/python/unittest/test_tir_transform_split_host_device.py
@@ -168,5 +168,21 @@ class
TestSplitHostDeviceWithoutFuncHostAttribute(BaseCompare):
return mod
+class TestSplitHostDevice(BaseCompare):
+ """Like TestSplitHostDevice, but no device regions to extract
+
+ Even if there are no device regions, the host-side function should
+ still have its "target" attribute updated.
+ """
+
+ def before():
+ T.func_attr({"target": T.target("ext_dev", host="llvm")})
+ T.evaluate(0)
+
+ def expected():
+ T.func_attr({"target": T.target("llvm")})
+ T.evaluate(0)
+
+
if __name__ == "__main__":
tvm.testing.main()