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

Reply via email to