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