tvm
Loading...
Searching...
No Matches
database.h
Go to the documentation of this file.
1/*
2 * Licensed to the Apache Software Foundation (ASF) under one
3 * or more contributor license agreements. See the NOTICE file
4 * distributed with this work for additional information
5 * regarding copyright ownership. The ASF licenses this file
6 * to you under the Apache License, Version 2.0 (the
7 * "License"); you may not use this file except in compliance
8 * with the License. You may obtain a copy of the License at
9 *
10 * http://www.apache.org/licenses/LICENSE-2.0
11 *
12 * Unless required by applicable law or agreed to in writing,
13 * software distributed under the License is distributed on an
14 * "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY
15 * KIND, either express or implied. See the License for the
16 * specific language governing permissions and limitations
17 * under the License.
18 */
19#ifndef TVM_S_TIR_META_SCHEDULE_DATABASE_H_
20#define TVM_S_TIR_META_SCHEDULE_DATABASE_H_
21
22#include <tvm/ffi/container/array.h>
23#include <tvm/ffi/function.h>
24#include <tvm/ffi/reflection/registry.h>
25#include <tvm/ffi/string.h>
26#include <tvm/ir/expr.h>
27#include <tvm/ir/module.h>
28#include <tvm/ir/prim/expr.h>
32#include <tvm/target/target.h>
33
34#include <filesystem>
35#include <memory>
36
37namespace tvm {
38namespace s_tir {
39using namespace tvm::prim;
40namespace meta_schedule {
41
42class ModuleEquality;
43
45class WorkloadNode : public ffi::Object {
46 public:
53
54 static void RegisterReflection() {
55 namespace refl = tvm::ffi::reflection;
56 refl::ObjectDef<WorkloadNode>().def_ro("mod", &WorkloadNode::mod);
57 }
58 TVM_FFI_DECLARE_OBJECT_INFO_FINAL("s_tir.meta_schedule.Workload", WorkloadNode, ffi::Object);
59
64 ffi::ObjectRef AsJSON() const;
65};
66
71class Workload : public ffi::ObjectRef {
72 public:
74 explicit Workload(ffi::ObjectPtr<WorkloadNode> data) : ffi::ObjectRef(data) {}
79 TVM_DLL explicit Workload(IRModule mod);
85 TVM_DLL explicit Workload(IRModule mod, THashCode shash);
91 TVM_DLL static Workload FromJSON(const ffi::ObjectRef& json_obj);
92
94};
95
98 size_t operator()(const Workload& a) const { return a->shash; }
99};
100
103 explicit WorkloadEqual(const ModuleEquality& mod_eq) : mod_eq_(mod_eq) {}
104
105 bool operator()(const Workload& a, const Workload& b) const;
106
107 private:
109 const ModuleEquality& mod_eq_;
110};
111
113class MeasureCandidate;
114
116class TuningRecordNode : public ffi::Object {
117 public:
121 Workload workload{ffi::UnsafeInit()};
123 ffi::Optional<ffi::Array<FloatImm>> run_secs;
125 ffi::Optional<Target> target;
127 ffi::Optional<ffi::Array<ArgInfo>> args_info;
128
129 static void RegisterReflection() {
130 namespace refl = tvm::ffi::reflection;
131 refl::ObjectDef<TuningRecordNode>()
132 .def_ro("trace", &TuningRecordNode::trace)
133 .def_ro("workload", &TuningRecordNode::workload)
134 .def_ro("run_secs", &TuningRecordNode::run_secs)
135 .def_ro("target", &TuningRecordNode::target)
136 .def_ro("args_info", &TuningRecordNode::args_info);
137 }
138 TVM_FFI_DECLARE_OBJECT_INFO_FINAL("s_tir.meta_schedule.TuningRecord", TuningRecordNode,
139 ffi::Object);
140
149 ffi::ObjectRef AsJSON() const;
154 bool IsValid() const;
155};
156
161class TuningRecord : public ffi::ObjectRef {
162 public:
171 TVM_DLL explicit TuningRecord(s_tir::Trace trace, Workload workload,
172 ffi::Optional<ffi::Array<FloatImm>> run_secs,
173 ffi::Optional<Target> target,
174 ffi::Optional<ffi::Array<ArgInfo>> args_info);
181 TVM_DLL static TuningRecord FromJSON(const ffi::ObjectRef& json_obj, const Workload& workload);
183};
184
185class Database;
186
187/* \brief The abstract interface of database. */
188class DatabaseNode : public ffi::Object {
189 public:
202 explicit DatabaseNode(ffi::String mod_eq_name = "structural");
203
205 virtual ~DatabaseNode();
211 virtual bool HasWorkload(const IRModule& mod) = 0;
217 virtual Workload CommitWorkload(const IRModule& mod) = 0;
222 virtual void CommitTuningRecord(const TuningRecord& record) = 0;
229 virtual ffi::Array<TuningRecord> GetTopK(const Workload& workload, int top_k) = 0;
234 virtual ffi::Array<TuningRecord> GetAllTuningRecords() = 0;
239 virtual int64_t Size() = 0;
247 virtual ffi::Optional<TuningRecord> QueryTuningRecord(const IRModule& mod, const Target& target,
248 const ffi::String& workload_name);
256 virtual ffi::Optional<s_tir::Schedule> QuerySchedule(const IRModule& mod, const Target& target,
257 const ffi::String& workload_name);
265 virtual ffi::Optional<IRModule> QueryIRModule(const IRModule& mod, const Target& target,
266 const ffi::String& workload_name);
274 TVM_FFI_ICHECK(mod_eq_);
275 return *mod_eq_;
276 }
277
278 static constexpr const bool _type_mutable = true;
279 TVM_FFI_DECLARE_OBJECT_INFO("s_tir.meta_schedule.Database", DatabaseNode, ffi::Object);
280
281 private:
283 std::unique_ptr<ModuleEquality> mod_eq_;
284};
285
288 public:
301 explicit PyDatabaseNode(ffi::String mod_eq_name = "structural");
302
308 using FHasWorkload = ffi::TypedFunction<bool(const IRModule&)>;
314 using FCommitWorkload = ffi::TypedFunction<Workload(const IRModule&)>;
319 using FCommitTuningRecord = ffi::TypedFunction<void(const TuningRecord&)>;
326 using FGetTopK = ffi::TypedFunction<ffi::Array<TuningRecord>(const Workload&, int)>;
331 using FGetAllTuningRecords = ffi::TypedFunction<ffi::Array<TuningRecord>()>;
339 using FQueryTuningRecord = ffi::TypedFunction<ffi::Optional<TuningRecord>(
340 const IRModule&, const Target&, const ffi::String&)>;
348 using FQuerySchedule = ffi::TypedFunction<ffi::Optional<s_tir::Schedule>(
349 const IRModule&, const Target&, const ffi::String&)>;
357 using FQueryIRModule = ffi::TypedFunction<ffi::Optional<IRModule>(const IRModule&, const Target&,
358 const ffi::String&)>;
363 using FSize = ffi::TypedFunction<int64_t()>;
364
383
384 static void RegisterReflection() {
385 // ffi::Functions are all not registered, because the reflection system doesn't take care of
386 // them, so it cannot be accessible on the python side. If there is such need from the future,
387 // we can then add corresponding accessor methods to help access on python.
388 // `f_has_workload` is not registered
389 // `f_commit_workload` is not registered
390 // `f_commit_tuning_record` is not registered
391 // `f_get_top_k` is not registered
392 // `f_get_all_tuning_records` is not registered
393 // `f_query_tuning_record` is not registered
394 // `f_query_schedule` is not registered
395 // `f_query_ir_module` is not registered
396 // `f_size` is not registered
397 namespace refl = tvm::ffi::reflection;
398 refl::ObjectDef<PyDatabaseNode>();
399 }
400
401 bool HasWorkload(const IRModule& mod) final {
402 TVM_FFI_ICHECK(f_has_workload != nullptr) << "PyDatabase's HasWorkload method not implemented!";
403 return f_has_workload(mod);
404 }
405
406 Workload CommitWorkload(const IRModule& mod) final {
408 << "PyDatabase's CommitWorkload method not implemented!";
409 return f_commit_workload(mod);
410 }
411
414 << "PyDatabase's CommitTuningRecord method not implemented!";
416 }
417
418 ffi::Array<TuningRecord> GetTopK(const Workload& workload, int top_k) final {
419 TVM_FFI_ICHECK(f_get_top_k != nullptr) << "PyDatabase's GetTopK method not implemented!";
420 return f_get_top_k(workload, top_k);
421 }
422
423 ffi::Array<TuningRecord> GetAllTuningRecords() final {
425 << "PyDatabase's GetAllTuningRecords method not implemented!";
427 }
428
429 ffi::Optional<TuningRecord> QueryTuningRecord(const IRModule& mod, const Target& target,
430 const ffi::String& workload_name) final {
431 if (f_query_tuning_record == nullptr) {
433 } else {
434 return f_query_tuning_record(mod, target, workload_name);
435 }
436 }
437
438 ffi::Optional<s_tir::Schedule> QuerySchedule(const IRModule& mod, const Target& target,
439 const ffi::String& workload_name) final {
440 if (f_query_schedule == nullptr) {
441 return DatabaseNode::QuerySchedule(mod, target, workload_name);
442 } else {
443 return f_query_schedule(mod, target, workload_name);
444 }
445 }
446
447 ffi::Optional<IRModule> QueryIRModule(const IRModule& mod, const Target& target,
448 const ffi::String& workload_name) final {
449 if (f_query_ir_module == nullptr) {
450 return DatabaseNode::QueryIRModule(mod, target, workload_name);
451 } else {
452 return f_query_ir_module(mod, target, workload_name);
453 }
454 }
455
457 TVM_FFI_ICHECK(f_size != nullptr) << "PyDatabase's Size method not implemented!";
458 return f_size();
459 }
460
461 static constexpr const bool _type_mutable = true;
463};
464
469class Database : public ffi::ObjectRef {
470 public:
475 explicit Database(ffi::ObjectPtr<DatabaseNode> data) : ffi::ObjectRef(data) {
476 TVM_FFI_ICHECK(data != nullptr);
477 }
482 TVM_DLL static Database MemoryDatabase(ffi::String mod_eq_name = "structural");
490 ffi::String mod_eq_name = "structural");
499 bool allow_missing, ffi::String mod_eq_name = "structural");
507 TVM_DLL static Database UnionDatabase(ffi::Array<Database, void> databases);
515 TVM_DLL static Database OrderedUnionDatabase(ffi::Array<Database, void> databases);
531 PyDatabaseNode::FCommitWorkload f_commit_workload,
532 PyDatabaseNode::FCommitTuningRecord f_commit_tuning_record,
533 PyDatabaseNode::FGetTopK f_get_top_k,
534 PyDatabaseNode::FGetAllTuningRecords f_get_all_tuning_records,
535 PyDatabaseNode::FQueryTuningRecord f_query_tuning_record,
536 PyDatabaseNode::FQuerySchedule f_query_schedule,
537 PyDatabaseNode::FQueryIRModule f_query_ir_module,
539 ffi::String mod_eq_name = "structural");
541 static ffi::Optional<Database> Current();
546
548};
549
550} // namespace meta_schedule
551} // namespace s_tir
552} // namespace tvm
553
554#endif // TVM_S_TIR_META_SCHEDULE_DATABASE_H_
Managed reference class to IRModuleNode.
Definition module.h:255
Managed reference class to TargetNode.
Definition target.h:134
RAII wrapper function to enter and exit a context object similar to python's with syntax.
Definition with_context.h:59
Managed reference to ScheduleNode.
Definition schedule.h:887
Managed reference to TraceNode.
Definition trace.h:146
virtual void CommitTuningRecord(const TuningRecord &record)=0
Add a tuning record to the database.
static constexpr const bool _type_mutable
Definition database.h:278
TVM_FFI_DECLARE_OBJECT_INFO("s_tir.meta_schedule.Database", DatabaseNode, ffi::Object)
void DumpPruned(Database destination)
Prune the database and dump it a given database.
virtual bool HasWorkload(const IRModule &mod)=0
Check if the database has the given workload.
const ModuleEquality & GetModuleEquality() const
Return a reference to the owned module equality method instance.
Definition database.h:273
virtual ffi::Array< TuningRecord > GetTopK(const Workload &workload, int top_k)=0
Get the top K valid tuning records of given workload from the database.
virtual ffi::Optional< TuningRecord > QueryTuningRecord(const IRModule &mod, const Target &target, const ffi::String &workload_name)
Query the best record of the given workload from the database.
virtual Workload CommitWorkload(const IRModule &mod)=0
Look up or add workload to the database if missing.
virtual ffi::Optional< IRModule > QueryIRModule(const IRModule &mod, const Target &target, const ffi::String &workload_name)
Query the best IRModule of the given workload from the database.
virtual int64_t Size()=0
Get the size of the database.
virtual ffi::Optional< s_tir::Schedule > QuerySchedule(const IRModule &mod, const Target &target, const ffi::String &workload_name)
Query the best schedule of the given workload from the database.
virtual ~DatabaseNode()
Default destructor.
DatabaseNode(ffi::String mod_eq_name="structural")
Constructor.
virtual ffi::Array< TuningRecord > GetAllTuningRecords()=0
Get all tuning records from the database.
Managed reference to DatabaseNode.
Definition database.h:469
void EnterWithScope()
Entering the scope of the context manager.
static Database ScheduleFnDatabase(ffi::TypedFunction< bool(s_tir::Schedule)> schedule_fn, ffi::String mod_eq_name="structural")
A database for injecting handcrafted schedule functions.
static Database JSONDatabase(ffi::String path_workload, ffi::String path_tuning_record, bool allow_missing, ffi::String mod_eq_name="structural")
Create a default database that uses JSON file for tuning records.
static Database PyDatabase(PyDatabaseNode::FHasWorkload f_has_workload, PyDatabaseNode::FCommitWorkload f_commit_workload, PyDatabaseNode::FCommitTuningRecord f_commit_tuning_record, PyDatabaseNode::FGetTopK f_get_top_k, PyDatabaseNode::FGetAllTuningRecords f_get_all_tuning_records, PyDatabaseNode::FQueryTuningRecord f_query_tuning_record, PyDatabaseNode::FQuerySchedule f_query_schedule, PyDatabaseNode::FQueryIRModule f_query_ir_module, PyDatabaseNode::FSize f_size, ffi::String mod_eq_name="structural")
Create a database with customized methods on the python-side.
static Database MemoryDatabase(ffi::String mod_eq_name="structural")
An in-memory database.
static Database OrderedUnionDatabase(ffi::Array< Database, void > databases)
A database composed of multiple databases, allowing users to guide IR rewriting using combined knowle...
static Database UnionDatabase(ffi::Array< Database, void > databases)
A database composed of multiple databases, allowing users to guide IR rewriting using combined knowle...
void ExitWithScope()
Exiting the scope of the context manager.
Database(ffi::ObjectPtr< DatabaseNode > data)
Constructor from ffi::ObjectPtr<DatabaseNode>.
Definition database.h:475
static ffi::Optional< Database > Current()
TVM_FFI_DEFINE_OBJECT_REF_METHODS_NOTNULLABLE(Database, ffi::ObjectRef, DatabaseNode)
Managed reference to MeasureCandidateNode.
Definition measure_candidate.h:57
The database with customized methods on the python-side.
Definition database.h:287
ffi::TypedFunction< ffi::Array< TuningRecord >(const Workload &, int)> FGetTopK
The function type of GetTopK method.
Definition database.h:326
ffi::TypedFunction< ffi::Optional< s_tir::Schedule >(const IRModule &, const Target &, const ffi::String &)> FQuerySchedule
The function type of QuerySchedule method.
Definition database.h:349
Workload CommitWorkload(const IRModule &mod) final
Look up or add workload to the database if missing.
Definition database.h:406
bool HasWorkload(const IRModule &mod) final
Check if the database has the given workload.
Definition database.h:401
ffi::TypedFunction< void(const TuningRecord &)> FCommitTuningRecord
The function type of CommitTuningRecord method.
Definition database.h:319
void CommitTuningRecord(const TuningRecord &record) final
Add a tuning record to the database.
Definition database.h:412
ffi::TypedFunction< bool(const IRModule &)> FHasWorkload
The function type of HasWorkload method.
Definition database.h:308
FQuerySchedule f_query_schedule
The packed function to the QuerySchedule function.
Definition database.h:378
PyDatabaseNode(ffi::String mod_eq_name="structural")
Constructor.
ffi::Array< TuningRecord > GetTopK(const Workload &workload, int top_k) final
Get the top K valid tuning records of given workload from the database.
Definition database.h:418
ffi::TypedFunction< int64_t()> FSize
The function type of Size method.
Definition database.h:363
ffi::TypedFunction< ffi::Optional< IRModule >(const IRModule &, const Target &, const ffi::String &)> FQueryIRModule
The function type of QueryIRModule method.
Definition database.h:358
FCommitTuningRecord f_commit_tuning_record
The packed function to the CommitTuningRecord function.
Definition database.h:370
int64_t Size() final
Get the size of the database.
Definition database.h:456
FCommitWorkload f_commit_workload
The packed function to the CommitWorkload function.
Definition database.h:368
ffi::Optional< TuningRecord > QueryTuningRecord(const IRModule &mod, const Target &target, const ffi::String &workload_name) final
Query the best record of the given workload from the database.
Definition database.h:429
ffi::TypedFunction< Workload(const IRModule &)> FCommitWorkload
The function type of CommitWorkload method.
Definition database.h:314
FQueryIRModule f_query_ir_module
The packed function to the QueryIRModule function.
Definition database.h:380
static void RegisterReflection()
Definition database.h:384
FGetAllTuningRecords f_get_all_tuning_records
The packed function to the GetAllTuningRecords function.
Definition database.h:374
FHasWorkload f_has_workload
The packed function to the HasWorkload function.
Definition database.h:366
FSize f_size
The packed function to the Size function.
Definition database.h:382
ffi::Optional< s_tir::Schedule > QuerySchedule(const IRModule &mod, const Target &target, const ffi::String &workload_name) final
Query the best schedule of the given workload from the database.
Definition database.h:438
ffi::Optional< IRModule > QueryIRModule(const IRModule &mod, const Target &target, const ffi::String &workload_name) final
Query the best IRModule of the given workload from the database.
Definition database.h:447
FGetTopK f_get_top_k
The packed function to the GetTopK function.
Definition database.h:372
static constexpr const bool _type_mutable
Definition database.h:461
ffi::TypedFunction< ffi::Optional< TuningRecord >(const IRModule &, const Target &, const ffi::String &)> FQueryTuningRecord
The function type of QueryTuningRecord method.
Definition database.h:340
TVM_FFI_DECLARE_OBJECT_INFO_FINAL("s_tir.meta_schedule.PyDatabase", PyDatabaseNode, DatabaseNode)
ffi::TypedFunction< ffi::Array< TuningRecord >()> FGetAllTuningRecords
The function type of GetAllTuningRecords method.
Definition database.h:331
FQueryTuningRecord f_query_tuning_record
The packed function to the QueryTuningRecord function.
Definition database.h:376
ffi::Array< TuningRecord > GetAllTuningRecords() final
Get all tuning records from the database.
Definition database.h:423
The class of tuning records.
Definition database.h:116
ffi::ObjectRef AsJSON() const
Export the tuning record to a JSON string.
bool IsValid() const
Check if this tuning record has valid trace instructions and successful run results.
s_tir::Trace trace
The trace tuned.
Definition database.h:119
ffi::Optional< ffi::Array< ArgInfo > > args_info
The argument information.
Definition database.h:127
Workload workload
The workload.
Definition database.h:121
MeasureCandidate AsMeasureCandidate() const
Construct the measure candidate given the initial IR module and trace stored in the tuning record.
ffi::Optional< ffi::Array< FloatImm > > run_secs
The profiling result in seconds.
Definition database.h:123
static void RegisterReflection()
Definition database.h:129
TVM_FFI_DECLARE_OBJECT_INFO_FINAL("s_tir.meta_schedule.TuningRecord", TuningRecordNode, ffi::Object)
ffi::Optional< Target > target
The target for tuning.
Definition database.h:125
The managed reference of TuningRecordNode.
Definition database.h:161
TVM_FFI_DEFINE_OBJECT_REF_METHODS_NOTNULLABLE(TuningRecord, ffi::ObjectRef, TuningRecordNode)
static TuningRecord FromJSON(const ffi::ObjectRef &json_obj, const Workload &workload)
Create a tuning record from a json object.
TuningRecord(s_tir::Trace trace, Workload workload, ffi::Optional< ffi::Array< FloatImm > > run_secs, ffi::Optional< Target > target, ffi::Optional< ffi::Array< ArgInfo > > args_info)
Constructor of a tuning record.
A workload, i.e. an IRModule and its structural hash.
Definition database.h:45
size_t THashCode
The type of structural hash.
Definition database.h:48
THashCode shash
The workload's structural hash.
Definition database.h:52
IRModule mod
The workload's IRModule.
Definition database.h:50
static void RegisterReflection()
Definition database.h:54
TVM_FFI_DECLARE_OBJECT_INFO_FINAL("s_tir.meta_schedule.Workload", WorkloadNode, ffi::Object)
ffi::ObjectRef AsJSON() const
Export the workload to a JSON string.
Managed reference to WorkloadNode.
Definition database.h:71
Workload(IRModule mod, THashCode shash)
Constructor of Workload.
TVM_FFI_DEFINE_OBJECT_REF_METHODS_NOTNULLABLE(Workload, ffi::ObjectRef, WorkloadNode)
static Workload FromJSON(const ffi::ObjectRef &json_obj)
Create a workload from a json object.
Workload(IRModule mod)
Constructor of Workload.
Workload(ffi::ObjectPtr< WorkloadNode > data)
Definition database.h:74
Base expr nodes in TVM.
TIR expressions.
IRModule that holds the functions and type definitions.
Definition builtin.h:25
An object that builds and maintains block scope and StmtSref mapping for Dependence analysis.
Definition analyzer.h:40
The equality check for Workload.
Definition database.h:102
bool operator()(const Workload &a, const Workload &b) const
WorkloadEqual(const ModuleEquality &mod_eq)
Definition database.h:103
The hash method for Workload.
Definition database.h:97
size_t operator()(const Workload &a) const
Definition database.h:98
Compilation target object.