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

tqchen 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 42bffc31ff [Target] Refine equality check on TargetKind instances 
(#17321)
42bffc31ff is described below

commit 42bffc31ff2aa14b18275f70a3d658156dbed2a2
Author: wrongtest <[email protected]>
AuthorDate: Tue Sep 3 22:51:42 2024 +0800

    [Target] Refine equality check on TargetKind instances (#17321)
    
    refine target kind identity
    
    Co-authored-by: wrongtest <[email protected]>
---
 src/target/target_kind.cc                 | 15 ++++++++++++++-
 tests/python/target/test_target_target.py | 16 ++++++++++++++++
 2 files changed, 30 insertions(+), 1 deletion(-)

diff --git a/src/target/target_kind.cc b/src/target/target_kind.cc
index fced74c3a5..979b755af8 100644
--- a/src/target/target_kind.cc
+++ b/src/target/target_kind.cc
@@ -35,7 +35,20 @@
 
 namespace tvm {
 
-TVM_REGISTER_NODE_TYPE(TargetKindNode);
+// helper to get internal dev function in objectref.
+struct TargetKind2ObjectPtr : public ObjectRef {
+  static ObjectPtr<Object> Get(const TargetKind& kind) { return 
GetDataPtr<Object>(kind); }
+};
+
+TVM_REGISTER_NODE_TYPE(TargetKindNode)
+    .set_creator([](const std::string& name) {
+      auto kind = TargetKind::Get(name);
+      ICHECK(kind.defined()) << "Cannot find target kind \'" << name << '\'';
+      return TargetKind2ObjectPtr::Get(kind.value());
+    })
+    .set_repr_bytes([](const Object* n) -> std::string {
+      return static_cast<const TargetKindNode*>(n)->name;
+    });
 
 TVM_STATIC_IR_FUNCTOR(ReprPrinter, vtable)
     .set_dispatch<TargetKindNode>([](const ObjectRef& obj, ReprPrinter* p) {
diff --git a/tests/python/target/test_target_target.py 
b/tests/python/target/test_target_target.py
index e977ef10aa..1a52a46da1 100644
--- a/tests/python/target/test_target_target.py
+++ b/tests/python/target/test_target_target.py
@@ -559,5 +559,21 @@ def test_target_from_device_opencl(input_device):
     assert target.thread_warp_size == dev.warp_size
 
 
+def test_module_dict_from_deserialized_targets():
+    target = Target("llvm")
+
+    from tvm.script import tir as T
+
+    @T.prim_func
+    def func():
+        T.evaluate(0)
+
+    func = func.with_attr("Target", target)
+    target2 = tvm.ir.load_json(tvm.ir.save_json(target))
+    mod = tvm.IRModule({"main": func})
+    lib = tvm.build({target2: mod}, target_host=target)
+    lib["func"]()
+
+
 if __name__ == "__main__":
     tvm.testing.main()

Reply via email to