Multithreaded mesh compilation.

PiperOrigin-RevId: 651358120
Change-Id: Id6272aac84ec6ac1b709d2f4a44b7593307bd158
This commit is contained in:
Yuval Tassa
2024-07-11 04:20:59 -07:00
committed by Copybara-Service
parent e7301edd20
commit 3d1d1d0718
7 changed files with 173 additions and 38 deletions
+1 -1
View File
@@ -1376,7 +1376,7 @@ void mjCMesh::Process() {
// perform computation with convex mesh if volume is negative
if (GetVolumeRef(type) <= 0 && exactmeshinertia) {
mju_warning("Malformed mesh %s, computing mesh inertia from convex hull", name.c_str());
mju_warning("Malformed mesh '%s', computing mesh inertia from convex hull", name.c_str());
exactmeshinertia = false;
ComputeVolume(CoM, type, facecen, exactmeshinertia);
}
+111 -12
View File
@@ -15,6 +15,7 @@
#include "user/user_model.h"
#include <algorithm>
#include <atomic>
#include <cmath>
#include <array>
#include <csetjmp>
@@ -22,7 +23,9 @@
#include <cstdint>
#include <cstdlib>
#include <cstring>
#include <exception>
#include <map>
#include <mutex>
#include <string>
#include <string_view>
#include <thread>
@@ -51,6 +54,7 @@ namespace {
namespace mju = ::mujoco::util;
using std::string;
using std::vector;
constexpr int kMaxCompilerThreads = 16;
} // namespace
//---------------------------------- CONSTRUCTOR AND DESTRUCTOR ------------------------------------
@@ -133,7 +137,9 @@ mjCModel& mjCModel::operator=(const mjCModel& other) {
mjCDef* subtree = new mjCDef(*other.defaults_[0]);
*this += *subtree;
for (const auto& [name, def] : other.def_map) {
std::size_t index = std::find(other.defaults_.begin(), other.defaults_.end(), def) - other.defaults_.begin();
std::size_t index =
std::find(other.defaults_.begin(), other.defaults_.end(), def) -
other.defaults_.begin();
def_map[name] = defaults_[index];
}
@@ -1552,8 +1558,9 @@ void mjCModel::LengthRange(mjModel* m, mjData* data) {
m->opt.timestep = LRopt.timestep;
}
// number of threads available, max 16
const int nthread = mjMIN(16, std::thread::hardware_concurrency()/2);
// number of threads available
int hardware_threads = std::thread::hardware_concurrency();
const int nthread = mjMAX(1, mjMIN(kMaxCompilerThreads, hardware_threads/2));
// count actuators that need computation
int cnt = 0;
@@ -1592,8 +1599,8 @@ void mjCModel::LengthRange(mjModel* m, mjData* data) {
// multiple threads
else {
// allocate mjData for each thread
char err[16][200];
mjData* pdata[16] = {data};
char err[kMaxCompilerThreads][200];
mjData* pdata[kMaxCompilerThreads] = {data};
for (int i=1; i<nthread; i++) {
pdata[i] = mj_makeData(m);
}
@@ -1605,7 +1612,7 @@ void mjCModel::LengthRange(mjModel* m, mjData* data) {
}
// prepare thread function arguments, clear errors
LRThreadArg arg[16];
LRThreadArg arg[kMaxCompilerThreads];
for (int i=0; i<nthread; i++) {
LRThreadArg temp = {m, pdata[i], i*num, num, &LRopt, err[i], 200};
arg[i] = temp;
@@ -1613,7 +1620,7 @@ void mjCModel::LengthRange(mjModel* m, mjData* data) {
}
// launch threads
std::thread th[16];
std::thread th[kMaxCompilerThreads];
for (int i=0; i<nthread; i++) {
th[i] = std::thread(LRfunc, arg+i);
}
@@ -3184,8 +3191,9 @@ void mjCModel::ProcessLists(bool checkrepeat) {
// error handler for low-level engine
constexpr int kErrorBufferSize = 500;
static thread_local std::jmp_buf error_jmp_buf;
static thread_local char errortext[500] = "";
static thread_local char errortext[kErrorBufferSize] = "";
static void errorhandler(const char* msg) {
mju::strcpy_arr(errortext, msg);
std::longjmp(error_jmp_buf, 1);
@@ -3193,9 +3201,14 @@ static void errorhandler(const char* msg) {
// warning handler for low-level engine
static thread_local char warningtext[500] = "";
static thread_local char warningtext[kErrorBufferSize] = ""; // top-level warning buffer
static thread_local std::string* local_warningtext_ptr = nullptr; // sub-thread warning buffer
static void warninghandler(const char* msg) {
mju::strcpy_arr(warningtext, msg);
if (local_warningtext_ptr) {
*local_warningtext_ptr = msg;
} else {
mju::strcpy_arr(warningtext, msg);
}
}
@@ -3279,6 +3292,86 @@ mjModel* mjCModel::Compile(const mjVFS* vfs, mjModel** m) {
}
// mesh compilation function to be used in threads
void CompileMesh(mjCMesh* mesh, const mjVFS* vfs, std::exception_ptr& exception,
std::mutex& exception_mutex, std::string* warningtext) {
// set warning text buffer to this local thread
local_warningtext_ptr = warningtext;
auto previous_handler = _mjPRIVATE__get_tls_warning_fn();
_mjPRIVATE__set_tls_warning_fn(warninghandler);
// compile the mesh, catch exception to be rethrown later
try {
mesh->Compile(vfs);
} catch (...) {
std::lock_guard<std::mutex> lock(exception_mutex);
if (!exception) {
exception = std::current_exception();
}
}
// restore warning handler to top-level
_mjPRIVATE__set_tls_warning_fn(previous_handler);
local_warningtext_ptr = nullptr;
}
// multi-threaded mesh compilation
void mjCModel::CompileMeshes(const mjVFS* vfs) {
std::vector<std::thread> threads;
int nmesh = meshes_.size();
int hardware_threads = std::thread::hardware_concurrency();
int nthread = std::max(1, std::min(kMaxCompilerThreads, hardware_threads / 2));
threads.reserve(nthread);
// holds an exception thrown by a worker thread
std::exception_ptr exception;
std::mutex except_mutex;
std::atomic_int next_mesh = 0;
std::vector<std::string> mesh_warningtext(nmesh);
for (int i = 0; i < nthread; ++i) {
threads.emplace_back([&] {
for (int meshid = next_mesh++; meshid < nmesh; meshid = next_mesh++) {
auto& mesh = meshes_[meshid];
CompileMesh(mesh, vfs, exception, except_mutex,
&mesh_warningtext[meshid]);
}
});
}
// join threads
for (auto& thread : threads) {
if (thread.joinable()) {
thread.join();
}
}
// concatenate all warnings from threads, copy into warningtext
std::string concatenated_warnings;
bool has_warning = false;
for (int i = 0; i < nmesh; i++) {
if (!mesh_warningtext[i].empty()) {
if (has_warning) {
concatenated_warnings += "\n";
}
concatenated_warnings += mesh_warningtext[i];
has_warning = true;
}
}
mju::strcpy_arr(warningtext, concatenated_warnings.c_str());
// if exception was caught, rethrow it
if (exception) {
std::rethrow_exception(exception);
}
}
void mjCModel::TryCompile(mjModel*& m, mjData*& d, const mjVFS* vfs) {
// check if nan test works
double test = mjNAN;
@@ -3352,8 +3445,14 @@ void mjCModel::TryCompile(mjModel*& m, mjData*& d, const mjVFS* vfs) {
SetNuser();
// compile meshes (needed for geom compilation)
for (int i=0; i<meshes_.size(); i++) {
meshes_[i]->Compile(vfs);
if (usethread) {
// multi-threaded mesh compile
CompileMeshes(vfs);
} else {
// single-threaded mesh compile
for (int i=0; i < meshes_.size(); i++) {
meshes_[i]->Compile(vfs);
}
}
// compile objects in kinematic tree
+3
View File
@@ -261,6 +261,9 @@ class mjCModel : public mjCModel_, private mjSpec {
// clear objects allocated by Compile
void Clear();
// multi-threaded mesh compilation
void CompileMeshes(const mjVFS* vfs);
// if asset name is missing, set to filename
template<class T> void SetDefaultNames(std::vector<T*>& assets);