fix crashes caused by test runner (#4475)

* Fix crashes in test runner

* plus some fixes

* changes per review
This commit is contained in:
Ashwini Khade 2020-07-10 14:04:22 -07:00 committed by GitHub
parent 26ebcfab88
commit 6a9a9a35be
No known key found for this signature in database
GPG key ID: 4AEE18F83AFDEB23
2 changed files with 61 additions and 33 deletions

View file

@ -49,13 +49,35 @@ void ORT_CALLBACK RunTestCase(ORT_CALLBACK_INSTANCE pci, void* context, ORT_WORK
}
}
void PTestRunner::Start(ORT_CALLBACK_INSTANCE, size_t concurrent_runs) {
concurrent_runs = std::min<size_t>(std::max<size_t>(1, concurrent_runs), c_.GetDataCount());
next_test_to_run = 0;
for (size_t i = 0; i != concurrent_runs; ++i) {
if (!ScheduleNew()) {
throw std::runtime_error("ScheduleNew task failed");
void PTestRunner::Start(ORT_CALLBACK_INSTANCE pci, size_t concurrent_runs) noexcept {
bool atleast_one_run_scheduled = false;
try {
concurrent_runs = std::min<size_t>(std::max<size_t>(1, concurrent_runs), c_.GetDataCount());
next_test_to_run = 0;
for (size_t i = 0; i < concurrent_runs; ++i) {
if (!ScheduleNew()) {
if (i == 0) {
// This should be a rare case.
// Log an error. While exiting Finish will be called which will take care of the cleanup
LOGF_DEFAULT(ERROR, "No task scheduled for test case: %s\n", c_.GetTestCaseName().c_str());
}
// In cases when concurrent_runs == number of test cases, before i reaches "concurrent_runs"
// an older tasks may finish and schedule a new task from OnTaskFinished method so no task is left to be scheduled
// from here.
break;
}
atleast_one_run_scheduled = true;
}
} catch (const std::exception& ex) {
LOGF_DEFAULT(ERROR, "Cannot schedule tasks for test %s. Failing with error :%s\n", c_.GetTestCaseName().c_str(), ex.what());
} catch (...) {
LOGF_DEFAULT(ERROR, "Cannot schedule tasks for test %s.\n", c_.GetTestCaseName().c_str());
}
// If no task was scheduled then call Finish to perform cleanup
// Otherwise the last task to complete will call Finish
if (!atleast_one_run_scheduled) {
Finish(pci);
}
}
@ -309,7 +331,7 @@ void DataRunner::RunTask(size_t task_id, ORT_CALLBACK_INSTANCE pci) {
EXECUTE_RESULT res = EXECUTE_RESULT::UNKNOWN_ERROR;
try {
res = RunTaskImpl(task_id);
} catch (std::exception& ex) {
} catch (const std::exception& ex) {
res = EXECUTE_RESULT::WITH_EXCEPTION;
LOGS_DEFAULT(ERROR) << c_.GetTestCaseName() << ":" << ex.what();
}
@ -461,12 +483,16 @@ EXECUTE_RESULT DataRunner::RunTaskImpl(size_t task_id) {
return res;
}
void SeqTestRunner::Start(ORT_CALLBACK_INSTANCE pci, size_t) {
const size_t data_count = c_.GetDataCount();
for (size_t idx_repeat = 0; idx_repeat != repeat_count_; ++idx_repeat) {
for (size_t idx_data = 0; idx_data != data_count; ++idx_data) {
RunTask(idx_data, nullptr);
void SeqTestRunner::Start(ORT_CALLBACK_INSTANCE pci, size_t) noexcept {
try {
const size_t data_count = c_.GetDataCount();
for (size_t idx_repeat = 0; idx_repeat != repeat_count_; ++idx_repeat) {
for (size_t idx_data = 0; idx_data != data_count; ++idx_data) {
RunTask(idx_data, nullptr);
}
}
} catch (...) {
LOGS_DEFAULT(ERROR) << "SeqTestRunner::Start - Encountred exception with running tests";
}
Finish(pci);
@ -478,7 +504,6 @@ void RunSingleTestCase(const ITestCase& info, Ort::Env& env, const Ort::SessionO
std::shared_ptr<TestCaseResult> ret;
size_t data_count = info.GetDataCount();
try {
std::unique_ptr<DataRunner> r;
std::string node_name = info.GetNodeName();
auto sf2 = sf.Clone();
sf2.SetLogId(info.GetTestCaseName().c_str());
@ -489,30 +514,30 @@ void RunSingleTestCase(const ITestCase& info, Ort::Env& env, const Ort::SessionO
concurrent_runs = 1;
}
// DataRunner owns itself. In case of success or failures DataRunner will delete itself.
// In case of SeqTestRunner SeqTestRunner::Start will call DataRunner::Finish at the end and will delete itself.
// In case of PTestRunner on success the last task will call Finish and in cases of failure Start will itself call DataRunner::Finish
// before exiting. Start does not throw so we do not need to worry about this case.
DataRunner* r;
if (concurrent_runs > 1 && data_count > 1) {
r.reset(new PTestRunner(session_object.release(), info, tpool, on_finished));
r = new PTestRunner(session_object.release(), info, tpool, on_finished);
} else {
r.reset(new SeqTestRunner(session_object.release(), info, repeat_count, on_finished));
r = new SeqTestRunner(session_object.release(), info, repeat_count, on_finished);
}
r->Start(pci, concurrent_runs);
// DataRunner::Finish will delete itself, so now that we know everything has started without any exceptions
// we can release it from the unique_ptr
r.release();
// both PTestRunner and SeqTestRunner call DataRunner::Finish which will delete itself and call on_finished callback
// at this point we know everything has started without any exceptions so simply return.
return;
} catch (const Ort::Exception& ex) {
if (ex.GetOrtErrorCode() != ORT_NOT_IMPLEMENTED) {
throw;
}
LOGF_DEFAULT(ERROR, "Test %s failed:%s", info.GetTestCaseName().c_str(), ex.what());
std::string node_name;
ret = std::make_shared<TestCaseResult>(data_count, EXECUTE_RESULT::NOT_SUPPORT, "");
} catch (onnxruntime::NotImplementedException& ex) {
} catch (const Ort::Exception& ex) {
LOGF_DEFAULT(ERROR, "Test %s failed:%s", info.GetTestCaseName().c_str(), ex.what());
std::string node_name;
ret = std::make_shared<TestCaseResult>(data_count, EXECUTE_RESULT::NOT_SUPPORT, "");
}
// we will land here in case of session creation failures.
// in all other cases DataRunner::Finish will call this.
on_finished(ret, pci);
}

View file

@ -43,6 +43,11 @@ class DataRunner {
//Time spent in Session::Run. It only make sense when SeqTestRunner was used
::onnxruntime::TIME_SPEC spent_time_;
// DataRunner destructor should not be made public. DataRunner owns itself and only it knows
// when it can destoy itself. Multiple tasks share the session object owned by DataRunner and only
// when all these tasks complete can it be destroyed.
virtual ~DataRunner();
private:
OrtSession* session;
CALL_BACK on_finished;
@ -55,13 +60,11 @@ class DataRunner {
TestCaseCallBack on_finished1);
virtual void OnTaskFinished(size_t task_id, EXECUTE_RESULT res, ORT_CALLBACK_INSTANCE pci) noexcept = 0;
void RunTask(size_t task_id, ORT_CALLBACK_INSTANCE pci);
virtual ~DataRunner();
virtual void Start(ORT_CALLBACK_INSTANCE pci, size_t concurrent_runs) = 0;
virtual void Start(ORT_CALLBACK_INSTANCE pci, size_t concurrent_runs) noexcept = 0;
void Finish(ORT_CALLBACK_INSTANCE pci) {
std::shared_ptr<TestCaseResult> res = result;
CALL_BACK callback = on_finished;
CALL_BACK on_finished_local = on_finished;
res->SetSpentTime(spent_time_);
const std::vector<EXECUTE_RESULT>& er = res->GetExcutionResult();
for (size_t i = 0; i != er.size(); ++i) {
@ -90,7 +93,7 @@ class DataRunner {
break;
}
delete this;
callback(res, pci);
on_finished_local(res, pci);
}
};
@ -101,7 +104,7 @@ class SeqTestRunner : public DataRunner {
public:
SeqTestRunner(OrtSession* session1, const ITestCase& c, size_t repeat_count, TestCaseCallBack on_finished1);
void Start(ORT_CALLBACK_INSTANCE pci, size_t concurrent_runs) override;
void Start(ORT_CALLBACK_INSTANCE pci, size_t concurrent_runs) noexcept override;
void OnTaskFinished(size_t, EXECUTE_RESULT, ORT_CALLBACK_INSTANCE) noexcept override {}
};
@ -112,7 +115,7 @@ class PTestRunner : public DataRunner {
void OnTaskFinished(size_t task_id, EXECUTE_RESULT res, ORT_CALLBACK_INSTANCE pci) noexcept override;
public:
void Start(ORT_CALLBACK_INSTANCE pci, size_t concurrent_runs) override;
void Start(ORT_CALLBACK_INSTANCE pci, size_t concurrent_runs) noexcept override;
PTestRunner(OrtSession* session1, const ITestCase& c, PThreadPool tpool, TestCaseCallBack on_finished1);