Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
545 changes: 311 additions & 234 deletions python/tvm/auto_scheduler/measure.py

Large diffs are not rendered by default.

39 changes: 1 addition & 38 deletions python/tvm/auto_scheduler/measure_record.py
Original file line number Diff line number Diff line change
Expand Up @@ -21,9 +21,7 @@

import tvm._ffi
from tvm.runtime import Object
from .compute_dag import ComputeDAG
from .measure import MeasureErrorNo, MeasureInput, MeasureCallback
from .search_task import SearchTask
from .measure import MeasureErrorNo, MeasureCallback
from . import _ffi_api


Expand Down Expand Up @@ -175,38 +173,3 @@ def load_best(filename, workload_key=None, target=None):
best_res = res

return best_inp, best_res


def recover_measure_input(inp, rebuild_state=False):
"""
Recover a deserialized MeasureInput by rebuilding the missing fields.
1. Rebuid the compute_dag in inp.task
2. (Optional) Rebuild the stages in inp.state

Parameters
----------
inp: MeasureInput
The deserialized MeasureInput
rebuild_state: bool = False
Whether rebuild the stages in MeasureInput.State

Returns
-------
new_input: MeasureInput
The fully recovered MeasureInput with all fields rebuilt.
"""
task = inp.task
new_task = SearchTask(
ComputeDAG(task.workload_key),
task.workload_key,
task.target,
task.target_host,
task.hardware_params,
)

if rebuild_state:
new_state = new_task.compute_dag.infer_bound_from_state(inp.state)
else:
new_state = inp.state

return MeasureInput(new_task, new_state)
42 changes: 9 additions & 33 deletions python/tvm/auto_scheduler/utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -129,32 +129,6 @@ def deserialize_args(args):
return ret


class NoDaemonProcess(multiprocessing.Process):
@property
def daemon(self):
return False

@daemon.setter
def daemon(self, value):
pass


class NoDaemonContext(type(multiprocessing.get_context())):
Process = NoDaemonProcess


class NoDaemonPool(multiprocessing.pool.Pool):
"""A no daemon pool version of multiprocessing.Pool.
This allows us to start new processes inside the worker function"""

def __init__(self, *args, **kwargs):
kwargs["context"] = NoDaemonContext()
super().__init__(*args, **kwargs)

def __reduce__(self):
pass


def kill_child_processes(parent_pid, sig=signal.SIGTERM):
"""kill all child processes recursively"""
try:
Expand All @@ -169,17 +143,19 @@ def kill_child_processes(parent_pid, sig=signal.SIGTERM):
return


def _func_wrapper(que, func, args, kwargs):
Comment thread
tkonolige marked this conversation as resolved.
"""Call function and return the result over the queue."""
if kwargs:
que.put(func(*args, **kwargs))
else:
que.put(func(*args))


def call_func_with_timeout(timeout, func, args=(), kwargs=None):
"""Call a function with timeout"""

def func_wrapper(que):
if kwargs:
que.put(func(*args, **kwargs))
else:
que.put(func(*args))

que = multiprocessing.Queue(2)
process = multiprocessing.Process(target=func_wrapper, args=(que,))
process = multiprocessing.Process(target=_func_wrapper, args=(que, func, args, kwargs))
process.start()
process.join(timeout)

Expand Down
35 changes: 35 additions & 0 deletions python/tvm/auto_scheduler/workload_registry.py
Original file line number Diff line number Diff line change
Expand Up @@ -175,6 +175,41 @@ def workload_key_to_tensors(workload_key):
return lookup(*args)


def get_workload_func(task):
"""Get the workload function for a given task

Parameters
----------
task : SearchTask
Task to get workload of.

Returns
-------
workload : callable
The registered workload function.
"""
name = workload_func_name(task.workload_key)
lookup = WORKLOAD_FUNC_REGISTRY[name]
assert callable(lookup)
return lookup


def workload_func_name(workload_key):
"""Decode a workload key to the registered function name.

Parameters
----------
workload_key : str
The input workload key.

Returns
-------
name : str
The function name of this workload key.
"""
return decode_workload_key_to_func_args(workload_key)[0]


def save_workload_func_registry(filename):
"""Dump workload function registry to a pickle binary file.

Expand Down
17 changes: 17 additions & 0 deletions python/tvm/testing.py
Original file line number Diff line number Diff line change
Expand Up @@ -634,6 +634,23 @@ def requires_micro(*args):
return _compose(args, _requires_micro)


def requires_rpc(*args):
"""Mark a test as requiring rpc to run.

Parameters
----------
f : function
Function to mark
"""
_requires_rpc = [
pytest.mark.skipif(
tvm.support.libinfo().get("USE_RPC", "OFF") != "ON",
reason="RPC support not enabled. Set USE_RPC=ON in config.cmake to enable.",
)
]
return _compose(args, _requires_rpc)


def _target_to_requirement(target):
# mapping from target to decorator
if target.startswith("cuda"):
Expand Down
75 changes: 73 additions & 2 deletions src/auto_scheduler/measure_record.cc
100755 → 100644
Original file line number Diff line number Diff line change
Expand Up @@ -107,18 +107,68 @@ struct Handler<::tvm::auto_scheduler::StateNode> {
}
};

template <>
struct Handler<::tvm::auto_scheduler::HardwareParamsNode> {
inline static void Write(dmlc::JSONWriter* writer,
const ::tvm::auto_scheduler::HardwareParamsNode& data) {
writer->BeginArray(false);
writer->WriteArrayItem(data.num_cores);
writer->WriteArrayItem(data.vector_unit_bytes);
writer->WriteArrayItem(data.cache_line_bytes);
writer->WriteArrayItem(data.max_shared_memory_per_block);
writer->WriteArrayItem(data.max_registers_per_block);
writer->WriteArrayItem(data.max_threads_per_block);
writer->WriteArrayItem(data.max_vthread_extent);
writer->WriteArrayItem(data.warp_size);
writer->EndArray();
}
inline static void Read(dmlc::JSONReader* reader,
::tvm::auto_scheduler::HardwareParamsNode* data) {
bool s;
reader->BeginArray();
s = reader->NextArrayItem();
CHECK(s);
reader->Read(&data->num_cores);
s = reader->NextArrayItem();
CHECK(s);
reader->Read(&data->vector_unit_bytes);
s = reader->NextArrayItem();
CHECK(s);
reader->Read(&data->cache_line_bytes);
s = reader->NextArrayItem();
CHECK(s);
reader->Read(&data->max_shared_memory_per_block);
s = reader->NextArrayItem();
CHECK(s);
reader->Read(&data->max_registers_per_block);
s = reader->NextArrayItem();
CHECK(s);
reader->Read(&data->max_threads_per_block);
s = reader->NextArrayItem();
CHECK(s);
reader->Read(&data->max_vthread_extent);
s = reader->NextArrayItem();
CHECK(s);
reader->Read(&data->warp_size);
s = reader->NextArrayItem();
CHECK(!s);
}
};

template <>
struct Handler<::tvm::auto_scheduler::SearchTaskNode> {
inline static void Write(dmlc::JSONWriter* writer,
const ::tvm::auto_scheduler::SearchTaskNode& data) {
writer->BeginArray(false);
writer->WriteArrayItem(std::string(data.workload_key));
writer->WriteArrayItem(data.target->str());
writer->WriteArrayItem(*data.hardware_params.get());
Comment thread
tkonolige marked this conversation as resolved.
writer->EndArray();
}
inline static void Read(dmlc::JSONReader* reader, ::tvm::auto_scheduler::SearchTaskNode* data) {
bool s;
std::string str_value;
auto hardware_params_node = ::tvm::make_object<::tvm::auto_scheduler::HardwareParamsNode>();
reader->BeginArray();
s = reader->NextArrayItem();
ICHECK(s);
Expand All @@ -129,7 +179,12 @@ struct Handler<::tvm::auto_scheduler::SearchTaskNode> {
reader->Read(&str_value);
data->target = ::tvm::Target(str_value);
s = reader->NextArrayItem();
ICHECK(!s);
if (s) {
reader->Read(hardware_params_node.get());
s = reader->NextArrayItem();
data->hardware_params = ::tvm::auto_scheduler::HardwareParams(hardware_params_node);
ICHECK(!s);
}
}
};

Expand Down Expand Up @@ -216,7 +271,7 @@ namespace auto_scheduler {
TVM_REGISTER_OBJECT_TYPE(RecordToFileNode);
TVM_REGISTER_OBJECT_TYPE(RecordReaderNode);

const std::string AUTO_SCHEDULER_LOG_VERSION = "v0.2"; // NOLINT(*)
const std::string AUTO_SCHEDULER_LOG_VERSION = "v0.3"; // NOLINT(*)

RecordToFile::RecordToFile(String filename) {
auto node = make_object<RecordToFileNode>();
Expand Down Expand Up @@ -340,5 +395,21 @@ TVM_REGISTER_GLOBAL("auto_scheduler.SaveRecords")
std::ofstream ofs(filename, std::ofstream::app);
WriteMeasureRecords(&ofs, in, res);
});

TVM_REGISTER_GLOBAL("auto_scheduler.SerializeMeasureInput")
.set_body_typed([](const MeasureInput& input) {
std::ostringstream os;
dmlc::JSONWriter writer(&os);
writer.Write(*input.get());
return os.str();
});

TVM_REGISTER_GLOBAL("auto_scheduler.DeserializeMeasureInput").set_body_typed([](String json) {
std::istringstream ss(json);
dmlc::JSONReader reader(&ss);
auto inp = make_object<MeasureInputNode>();
reader.Read(inp.get());
return ObjectRef(inp);
});
} // namespace auto_scheduler
} // namespace tvm
21 changes: 17 additions & 4 deletions tests/python/unittest/test_auto_scheduler_measure.py
Original file line number Diff line number Diff line change
Expand Up @@ -17,6 +17,7 @@

""" Test measurement and log serialization. """

import multiprocessing
import tvm
from tvm import topi
from tvm import te, auto_scheduler
Expand Down Expand Up @@ -182,12 +183,10 @@ def test_recover_measure_input():

raw_inp = inputs[0]

correct_inp = auto_scheduler.measure_record.recover_measure_input(raw_inp)
correct_inp = auto_scheduler.measure.recover_measure_input(raw_inp)
assert str(correct_inp.task.compute_dag) == str(inp.task.compute_dag)

correct_inp = auto_scheduler.measure_record.recover_measure_input(
raw_inp, rebuild_state=True
)
correct_inp = auto_scheduler.measure.recover_measure_input(raw_inp, rebuild_state=True)
assert str(correct_inp.state) == str(inp.state)


Expand Down Expand Up @@ -232,11 +231,25 @@ def test_measure_local_builder_rpc_runner():
del measure_ctx


def measure_local_builder_rpc_runner_spawn():
assert multiprocessing.get_start_method(False) == "spawn"
test_measure_local_builder_rpc_runner()


@tvm.testing.requires_llvm
def test_measure_local_builder_rpc_runner_spawn():
ctx = multiprocessing.get_context("spawn")
p = ctx.Process(target=measure_local_builder_rpc_runner_spawn)
p.start()
p.join()


if __name__ == "__main__":
test_record_split_reorder_fuse_annotation()
test_record_compute_at_root_inline_cache_read_write()
test_record_follow_split_follow_fused_split()
test_record_pragma_storage_align_rfactor()
test_recover_measure_input()
test_measure_local_builder_runner()
test_measure_local_builder_runner_spawn()
test_measure_local_builder_rpc_runner()
19 changes: 17 additions & 2 deletions tests/python/unittest/test_auto_scheduler_search_policy.py
Original file line number Diff line number Diff line change
Expand Up @@ -18,6 +18,7 @@
"""Test search policy"""

import random
import multiprocessing
import numpy as np
import tempfile

Expand All @@ -26,6 +27,7 @@
from tvm import auto_scheduler

from test_auto_scheduler_common import matmul_auto_scheduler_test, PropagatingThread
import multiprocessing


def search_common(
Expand Down Expand Up @@ -122,6 +124,19 @@ def test_sketch_search_policy_basic():
t.join()


def sketch_search_policy_basic_spawn():
assert multiprocessing.get_start_method(False) == "spawn"
test_sketch_search_policy_basic()


@tvm.testing.requires_llvm
def test_sketch_search_policy_basic_spawn():
ctx = multiprocessing.get_context("spawn")
p = ctx.Process(target=sketch_search_policy_basic_spawn)
p.start()
p.join()


@tvm.testing.requires_llvm
def test_sketch_search_policy_xgbmodel():
# wrap the search in a new thread to avoid the conflict
Expand Down Expand Up @@ -156,9 +171,8 @@ def test_sketch_search_policy_cuda_rpc_runner():
t.join()


@tvm.testing.requires_cuda
def test_sketch_search_policy_cuda_xgbmodel_rpc_runner():
if not tvm.runtime.enabled("cuda"):
return
measure_ctx = auto_scheduler.LocalRPCMeasureContext()
# wrap the search in a new thread to avoid the conflict
# between python's multiprocessing and tvm's thread pool
Expand All @@ -179,6 +193,7 @@ def test_sketch_search_policy_cuda_xgbmodel_rpc_runner():
if __name__ == "__main__":
test_workload_registry_search_basic()
test_sketch_search_policy_basic()
test_sketch_search_policy_basic_spawn()
test_sketch_search_policy_xgbmodel()
test_sketch_search_policy_cuda_rpc_runner()
test_sketch_search_policy_cuda_xgbmodel_rpc_runner()
Loading