Implement plugin mechanism for actuators and sensors.

PiperOrigin-RevId: 474874088
Change-Id: I65a8ffdf845f4fa0f8266c165883a747ab5812d8
This commit is contained in:
Saran Tunyasuvunakool
2022-09-16 12:20:56 -07:00
committed by Copybara-Service
parent f556d4d94f
commit 1e2a9a53bc
29 changed files with 2121 additions and 51 deletions
+2
View File
@@ -39,6 +39,8 @@ set(MUJOCO_ENGINE_SRCS
engine_io.c
engine_io.h
engine_macro.h
engine_plugin.cc
engine_plugin.h
engine_print.c
engine_print.h
engine_ray.c
+44
View File
@@ -19,6 +19,7 @@
#include <mujoco/mjdata.h>
#include <mujoco/mjmodel.h>
#include <mujoco/mjplugin.h>
#include "engine/engine_callback.h"
#include "engine/engine_collision_driver.h"
#include "engine/engine_core_constraint.h"
@@ -27,6 +28,7 @@
#include "engine/engine_inverse.h"
#include "engine/engine_io.h"
#include "engine/engine_macro.h"
#include "engine/engine_plugin.h"
#include "engine/engine_sensor.h"
#include "engine/engine_solver.h"
#include "engine/engine_support.h"
@@ -194,6 +196,11 @@ void mj_fwdActuation(const mjModel* m, mjData* d) {
// force = gain .* [ctrl/act] + bias
for (int i=0; i<nu; i++) {
// skip actuator plugins -- these are handled after builtin actuator types
if (m->actuator_plugin[i] >= 0) {
continue;
}
// extract gain info
prm = m->actuator_gainprm + mjNGAIN*i;
@@ -262,6 +269,24 @@ void mj_fwdActuation(const mjModel* m, mjData* d) {
force[i] += bias;
}
// handle actuator plugins
if (m->nplugin) {
const int nslot = mjp_pluginCount();
for (int i=0; i<m->nplugin; i++) {
const int slot = m->plugin[i];
const mjpPlugin* plugin = mjp_getPluginAtSlotUnsafe(slot, nslot);
if (!plugin) {
mju_error_i("invalid plugin slot: %d", slot);
}
if (plugin->type & mjPLUGIN_ACTUATOR) {
if (!plugin->compute) {
mju_error_i("`compute` is a null function pointer for plugin at slot %d", slot);
}
plugin->compute(m, d, i, mjPLUGIN_ACTUATOR);
}
}
}
// clamp actuator_force
for (int i=0; i<nu; i++) {
if (m->actuator_forcelimited[i]) {
@@ -275,6 +300,10 @@ void mj_fwdActuation(const mjModel* m, mjData* d) {
// act_dot for stateful actuators
for (int i=nu-na; i<nu; i++) {
if (m->actuator_plugin[i] >= 0) {
continue;
}
// extract info
prm = m->actuator_dynprm + i*mjNDYN;
int j = i-(nu-na);
@@ -481,6 +510,21 @@ static void mj_advance(const mjModel* m, mjData* d,
// advance time
d->time += m->opt.timestep;
// advance plugin states
if (m->nplugin) {
const int nslot = mjp_pluginCount();
for (int i = 0; i < m->nplugin; ++i) {
const int slot = m->plugin[i];
const mjpPlugin* plugin = mjp_getPluginAtSlotUnsafe(slot, nslot);
if (!plugin) {
mju_error_i("invalid plugin slot: %d", slot);
}
if (plugin->advance) {
plugin->advance(m, d, i);
}
}
}
}
// Euler integrator, semi-implicit in velocity, possibly skipping factorisation
+101 -9
View File
@@ -22,8 +22,10 @@
#include <string.h>
#include <mujoco/mjmodel.h>
#include <mujoco/mjplugin.h>
#include <mujoco/mjxmacro.h>
#include "engine/engine_macro.h"
#include "engine/engine_plugin.h"
#include "engine/engine_util_blas.h"
#include "engine/engine_util_errmem.h"
#include "engine/engine_vfs.h"
@@ -396,9 +398,10 @@ mjModel* mj_makeModel(int nq, int nv, int nu, int na, int nbody, int njnt,
int ntex, int ntexdata, int nmat, int npair, int nexclude,
int neq, int ntendon, int nwrap, int nsensor,
int nnumeric, int nnumericdata, int ntext, int ntextdata,
int ntuple, int ntupledata, int nkey, int nmocap,
int nuser_body, int nuser_jnt, int nuser_geom, int nuser_site, int nuser_cam,
int nuser_tendon, int nuser_actuator, int nuser_sensor, int nnames) {
int ntuple, int ntupledata, int nkey, int nmocap, int nplugin,
int npluginattr, int nuser_body, int nuser_jnt, int nuser_geom,
int nuser_site, int nuser_cam, int nuser_tendon, int nuser_actuator,
int nuser_sensor, int nnames) {
intptr_t offset = 0;
// allocate mjModel
@@ -449,6 +452,8 @@ mjModel* mj_makeModel(int nq, int nv, int nu, int na, int nbody, int njnt,
m->ntupledata = ntupledata;
m->nkey = nkey;
m->nmocap = nmocap;
m->nplugin = nplugin;
m->npluginattr = npluginattr;
m->nuser_body = nuser_body;
m->nuser_jnt = nuser_jnt;
m->nuser_geom = nuser_geom;
@@ -533,10 +538,10 @@ mjModel* mj_copyModel(mjModel* dest, const mjModel* src) {
src->ntex, src->ntexdata, src->nmat, src->npair, src->nexclude,
src->neq, src->ntendon, src->nwrap, src->nsensor,
src->nnumeric, src->nnumericdata, src->ntext, src->ntextdata,
src->ntuple, src->ntupledata, src->nkey, src->nmocap,
src->nuser_body, src->nuser_jnt, src->nuser_geom, src->nuser_site,
src->nuser_cam, src->nuser_tendon, src->nuser_actuator, src->nuser_sensor,
src->nnames);
src->ntuple, src->ntupledata, src->nkey, src->nmocap, src->nplugin,
src->npluginattr, src->nuser_body, src->nuser_jnt, src->nuser_geom,
src->nuser_site, src->nuser_cam, src->nuser_tendon, src->nuser_actuator,
src->nuser_sensor, src->nnames);
}
if (!dest) {
mju_error("Failed to make mjModel. Invalid sizes.");
@@ -705,7 +710,8 @@ mjModel* mj_loadModel(const char* filename, const mjVFS* vfs) {
info[21], info[22], info[23], info[24], info[25], info[26], info[27],
info[28], info[29], info[30], info[31], info[32], info[33], info[34],
info[35], info[36], info[37], info[38], info[39], info[40], info[41],
info[42], info[43], info[44], info[45], info[46], info[47], info[48]);
info[42], info[43], info[44], info[45], info[46], info[47], info[48],
info[49], info[50]);
if (!m || m->nbuffer!=info[getnint()-1]) {
if (fp) {
fclose(fp);
@@ -884,6 +890,17 @@ static mjData* _makeData(const mjModel* m) {
// set pointers into buffer, reset data
mj_setPtrData(m, d);
// copy plugins into d, required for deletion
d->nplugin = m->nplugin;
for (int i = 0; i < m->nplugin; ++i) {
d->plugin[i] = m->plugin[i];
const mjpPlugin* plugin = mjp_getPluginAtSlot(m->plugin[i]);
if (!plugin->init) {
mju_error_i("`init` is a null function pointer for plugin at slot %d", m->plugin[i]);
}
plugin->init(m, d, i);
}
return d;
}
@@ -927,6 +944,17 @@ mjData* mj_copyData(mjData* dest, const mjModel* m, const mjData* src) {
dest->stack = save_stack;
mj_setPtrData(m, dest);
// save plugin_data, since the X macro copying block below will override it
const size_t plugin_data_size = sizeof(*dest->plugin_data) * dest->nplugin;
uintptr_t* save_plugin_data = NULL;
if (plugin_data_size) {
save_plugin_data = (uintptr_t*)mju_malloc(plugin_data_size);
if (!save_plugin_data) {
mju_error("failed to allocate temporary memory for plugin_data");
}
memcpy(save_plugin_data, dest->plugin_data, plugin_data_size);
}
// copy buffer
{
MJDATA_POINTERS_PREAMBLE(m)
@@ -936,6 +964,22 @@ mjData* mj_copyData(mjData* dest, const mjModel* m, const mjData* src) {
#undef X
}
// restore plugin_data
if (plugin_data_size) {
memcpy(dest->plugin_data, save_plugin_data, plugin_data_size);
free(save_plugin_data);
save_plugin_data = NULL;
}
// copy plugin instances
dest->nplugin = m->nplugin;
for (int i = 0; i < m->nplugin; ++i) {
const mjpPlugin* plugin = mjp_getPluginAtSlot(m->plugin[i]);
if (plugin->copy) {
plugin->copy(dest, m, src, i);
}
}
return dest;
}
@@ -966,6 +1010,12 @@ mjtNum* mj_stackAlloc(mjData* d, int size) {
// clear data, set defaults
static void _resetData(const mjModel* m, mjData* d, unsigned char debug_value) {
//------------------------------ save plugin state and data
mjtNum* plugin_state = mju_malloc(sizeof(mjtNum) * m->npluginstate);
memcpy(plugin_state, d->plugin_state, sizeof(mjtNum) * m->npluginstate);
uintptr_t* plugindata = mju_malloc(sizeof(uintptr_t) * m->nplugin);
memcpy(plugindata, d->plugin_data, sizeof(uintptr_t) * m->nplugin);
//------------------------------ clear header
// clear stack pointer
@@ -1048,6 +1098,22 @@ static void _resetData(const mjModel* m, mjData* d, unsigned char debug_value) {
d->mocap_quat[4*i] = 1.0;
}
}
// restore pluginstate and plugindata
memcpy(d->plugin_state, plugin_state, sizeof(mjtNum) * m->npluginstate);
mju_free(plugin_state);
memcpy(d->plugin_data, plugindata, sizeof(uintptr_t) * m->nplugin);
mju_free(plugindata);
// restore the plugin array back into d and reset the instances
for (int i = 0; i < m->nplugin; ++i) {
d->plugin[i] = m->plugin[i];
const mjpPlugin* plugin = mjp_getPluginAtSlot(m->plugin[i]);
if (!plugin->reset) {
mju_error_i("`reset` is a null function pointer for plugin at slot %d", m->plugin[i]);
}
plugin->reset(m, d, i);
}
}
@@ -1087,6 +1153,13 @@ void mj_resetDataKeyframe(const mjModel* m, mjData* d, int key) {
// de-allocate mjData
void mj_deleteData(mjData* d) {
if (d) {
// destroy plugin instances
for (int i = 0; i < d->nplugin; ++i) {
const mjpPlugin* plugin = mjp_getPluginAtSlot(d->plugin[i]);
if (plugin->destroy) {
plugin->destroy(d, i);
}
}
mju_free(d->buffer);
mju_free(d->stack);
mju_free(d);
@@ -1145,6 +1218,9 @@ static int sensorSize(mjtSensor sensor_type, int nuser_sensor) {
case mjSENS_USER:
return nuser_sensor;
case mjSENS_PLUGIN:
return -1;
// don't use a 'default' case, so compiler warns about missing values
}
return -1;
@@ -1202,6 +1278,8 @@ static int numObjects(const mjModel* m, mjtObj objtype) {
return m->ntuple;
case mjOBJ_KEY:
return m->nkey;
case mjOBJ_PLUGIN:
return m->nplugin;
}
return -2;
}
@@ -1251,6 +1329,10 @@ const char* mj_validateReferences(const mjModel* m) {
X(skin_bonevertid, nskinbonevert, nskinvert , 0 ) \
X(pair_geom1, npair, ngeom , 0 ) \
X(pair_geom2, npair, ngeom , 0 ) \
X(actuator_plugin, nu, nplugin , 0 ) \
X(sensor_plugin, nsensor, nplugin , 0 ) \
X(plugin_stateadr, nplugin, npluginstate , 0 ) \
X(plugin_attradr, nplugin, npluginattr , 0 ) \
X(tendon_adr, ntendon, nwrap , m->tendon_num ) \
X(tendon_matid, ntendon, nmat , 0 ) \
X(numeric_adr, nnumeric, nnumericdata , m->numeric_size ) \
@@ -1459,7 +1541,17 @@ const char* mj_validateReferences(const mjModel* m) {
}
for (int i=0; i<m->nsensor; i++) {
mjtSensor sensor_type = m->sensor_type[i];
int sensor_size = sensorSize(sensor_type, m->nuser_sensor);
int sensor_size;
if (sensor_type == mjSENS_PLUGIN) {
const mjpPlugin* plugin = mjp_getPluginAtSlot(m->plugin[m->sensor_plugin[i]]);
if (!plugin->nsensordata) {
mju_error_i("`nsensordata` is a null function pointer for plugin at slot %d",
m->plugin[m->sensor_plugin[i]]);
}
sensor_size = plugin->nsensordata(m, m->sensor_plugin[i], i);
} else {
sensor_size = sensorSize(sensor_type, m->nuser_sensor);
}
if (sensor_size < 0) {
return "Invalid model: Bad sensor_type.";
}
+4 -3
View File
@@ -52,9 +52,10 @@ mjModel* mj_makeModel(int nq, int nv, int nu, int na, int nbody, int njnt,
int ntex, int ntexdata, int nmat, int npair, int nexclude,
int neq, int ntendon, int nwrap, int nsensor,
int nnumeric, int nnumericdata, int ntext, int ntextdata,
int ntuple, int ntupledata, int nkey, int nmocap,
int nuser_body, int nuser_jnt, int nuser_geom, int nuser_site, int nuser_cam,
int nuser_tendon, int nuser_actuator, int nuser_sensor, int nnames);
int ntuple, int ntupledata, int nkey, int nmocap, int nplugin,
int npluginattr, int nuser_body, int nuser_jnt, int nuser_geom,
int nuser_site, int nuser_cam, int nuser_tendon, int nuser_actuator,
int nuser_sensor, int nnames);
// copy mjModel; allocate new if dest is NULL
MJAPI mjModel* mj_copyModel(mjModel* dest, const mjModel* src);
+431
View File
@@ -0,0 +1,431 @@
// Copyright 2022 DeepMind Technologies Limited
//
// Licensed under the Apache License, Version 2.0 (the "License");
// you may not use this file except in compliance with the License.
// You may obtain a copy of the License at
//
// http://www.apache.org/licenses/LICENSE-2.0
//
// Unless required by applicable law or agreed to in writing, software
// distributed under the License is distributed on an "AS IS" BASIS,
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
// See the License for the specific language governing permissions and
// limitations under the License.
// Plugin registration is implemented in C++, unlike the rest of the engine code which is in C.
// This is because C++ provides a standard cross-platform mutex, which we use to guard the global
// plugin table in order to make the API thread-safe. We only expose a C API externally, and this
// entire file can in principle be re-implemented in C if necessary, without breaking any external
// or internal MuJoCo code elsewhere.
#include "engine/engine_plugin.h"
#include <atomic>
#include <cstddef>
#include <cstdlib>
#include <cstring>
#include <memory>
#include <mutex>
#include <new>
#include <shared_mutex>
#include <type_traits>
#include <utility>
#include <vector>
#ifdef __APPLE__
#include <Availability.h>
#if !defined(MAC_OS_X_VERSION_MIN_REQUIRED) && defined(__MAC_OS_X_VERSION_MIN_REQUIRED)
#define MAC_OS_X_VERSION_MIN_REQUIRED __MAC_OS_X_VERSION_MIN_REQUIRED
#endif
#endif
#include <mujoco/mjplugin.h>
#include "engine/engine_util_errmem.h"
// set default plugin definition
void mjp_defaultPlugin(mjpPlugin* plugin) {
std::memset(plugin, 0, sizeof(*plugin));
}
namespace {
constexpr int kMaxNameLength = 1024;
constexpr int kMaxAttributes = 255;
constexpr int kCacheLine = 64;
// A table of registered plugins, implemented as a linked list of array "blocks".
// This is a compromise that maintains a good degree of memory locality while not invalidating
// existing pointers when growing the table. It is expected that for most users, the number of
// plugins loaded into a program will be small enough to fit in the initial block, and so the global
// table will behave like an array. Since pointers are never invalidated, we do not need to apply a
// read lock on the global table when resolving a plugin.
struct alignas(kCacheLine) PluginTable {
static constexpr int kBlockSize = 15;
PluginTable() {
for (int i = 0; i < kBlockSize; ++i) {
mjp_defaultPlugin(&plugins[i]);
}
}
mjpPlugin plugins[kBlockSize];
PluginTable* next = nullptr;
};
static_assert(
sizeof(PluginTable) / kCacheLine ==
sizeof(PluginTable::plugins) / kCacheLine + (sizeof(PluginTable::plugins) % kCacheLine > 0),
"PluginTable::next doesn't fit in the same cache line as the end of PluginTable::plugins");
using Mutex = std::shared_mutex;
class Global {
public:
Global() {
new(mutex_) Mutex;
}
PluginTable& table() {
return table_;
}
std::atomic_int& count() {
return count_;
}
Mutex& mutex() {
return *std::launder(reinterpret_cast<Mutex*>(&mutex_));
}
private:
PluginTable table_;
std::atomic_int count_;
// A mutex whose destructor is never run.
// When a C++ program terminates, the destructors for function static objects and globals will be
// executed by whichever thread started that termination but there is no guarantee that other
// threads have terminated. In other words, a static object may be accessed by another thread
// after it is deleted. We avoid destruction issues by never running the destructor.
alignas(Mutex) unsigned char mutex_[sizeof(Mutex)];
};
Global& GetGlobal() {
static Global global;
static_assert(std::is_trivially_destructible_v<decltype(global)>);
return global;
}
// return the length of a null-terminated string, or -1 if it is not terminated after kMaxNameLength
int strnlen(const char* s) {
for (int i = 0; i < kMaxNameLength; ++i) {
if (!s[i]) {
return i;
}
}
return -1;
}
// copy a null-terminated string into a new heap-allocated char array managed by a unique_ptr
std::unique_ptr<char[]> CopyName(const char* s) {
int len = strnlen(s);
if (len == -1) {
return nullptr;
}
std::unique_ptr<char[]> out(new(std::nothrow) char[len + 1]);
if (!out) {
return nullptr;
}
std::strncpy(out.get(), s, len);
out.get()[len] = '\0';
return out;
}
// check if two plugins are identical
bool PluginsAreIdentical(const mjpPlugin& plugin1, const mjpPlugin& plugin2) {
if (plugin1.name && !plugin2.name) {
return false;
}
if (plugin2.name && !plugin1.name) {
return false;
}
if (plugin1.name && plugin2.name &&
std::strncmp(plugin1.name, plugin2.name, kMaxNameLength)) {
return false;
}
if (plugin1.nattribute != plugin2.nattribute) {
return false;
}
for (int i = 0; i < plugin1.nattribute; ++i) {
if (plugin1.attributes[i] && !plugin2.attributes[i]) {
return false;
}
if (plugin2.attributes[i] && !plugin1.attributes[i]) {
return false;
}
if (plugin1.attributes[i] && plugin2.attributes[i] &&
std::strncmp(plugin1.attributes[i], plugin2.attributes[i],
kMaxNameLength)) {
return false;
}
}
const char* ptr1 = reinterpret_cast<const char*>(&plugin1.attributes) +
sizeof(plugin1.attributes);
const char* ptr2 = reinterpret_cast<const char*>(&plugin2.attributes) +
sizeof(plugin2.attributes);
std::size_t remaining_size =
sizeof(mjpPlugin) - (ptr1 - reinterpret_cast<const char*>(&plugin1));
return !std::memcmp(ptr1, ptr2, remaining_size);
}
} // namespace
// globally register a plugin (thread-safe), return new slot id
int mjp_registerPlugin(const mjpPlugin* plugin) {
if (!plugin->name) {
mju_error("plugin->name is a null pointer");
} else if (plugin->name[0] == '\0') {
mju_error("plugin->name is an empty string");
} else if (plugin->nattribute < 0) {
mju_error("plugin->nattribute is negative");
} else if (plugin->nattribute > kMaxAttributes) {
mju_error_i("plugin->nattribute exceeds the maximum limit of ",
kMaxAttributes);
}
char err[512];
err[0] = '\0';
// ========= ATTENTION! ==========================================================================
// Do not handle objects with nontrivial destructors outside of this lambda.
// Do not call mju_error inside this lambda.
int slot = [&]() -> int {
// check and copy the plugin name
std::unique_ptr<char[]> name = CopyName(plugin->name);
if (!name) {
if (strnlen(plugin->name) == -1) {
std::snprintf(err, sizeof(err),
"plugin->name length exceeds the maximum limit of %d", kMaxNameLength);
} else {
std::snprintf(err, sizeof(err), "failed to allocate memory for plugin name");
}
return -1;
}
// check and copy plugin attributes
std::vector<std::unique_ptr<char[]>> attributes_vec;
if (plugin->nattribute) {
attributes_vec.reserve(plugin->nattribute);
for (int i = 0; i < plugin->nattribute; ++i) {
std::unique_ptr<char[]> attr = CopyName(plugin->attributes[i]);
if (!attr) {
if (strnlen(plugin->attributes[i]) == -1) {
std::snprintf(
err, sizeof(err),
"plugin->attributes[%d] exceeds the maximum limit of %d", i, kMaxAttributes);
} else {
std::snprintf(err, sizeof(err), "failed to allocate memory for plugin attribute");
}
return -1;
}
attributes_vec.emplace_back(std::move(attr));
}
}
// exclusively lock the global plugin table
Global& global = GetGlobal();
std::unique_lock lock(global.mutex());
int count = global.count().load(std::memory_order_acquire);
int local_idx = 0;
PluginTable* table = &global.table();
// check if a non-identical plugin with the same name has already been registered
for (int i = 0; i < count; ++i, ++local_idx) {
if (local_idx == PluginTable::kBlockSize) {
local_idx = 0;
table = table->next;
}
mjpPlugin& existing = table->plugins[local_idx];
if (std::strcmp(plugin->name, existing.name) == 0) {
if (PluginsAreIdentical(*plugin, existing)) {
return i;
} else {
std::snprintf(err, sizeof(err), "plugin '%s' is already registered", plugin->name);
return -1;
}
}
}
// allocate a new block of PluginTable if the last allocated block is full
if (local_idx == PluginTable::kBlockSize) {
local_idx = 0;
#if defined(MAC_OS_X_VERSION_MIN_REQUIRED) && MAC_OS_X_VERSION_MIN_REQUIRED < MAC_OS_X_VERSION_10_14
// aligned nothrow new is not available until macOS 10.14
posix_memalign(reinterpret_cast<void**>(&table->next),
alignof(PluginTable), sizeof(PluginTable));
if (table->next) new(table->next) PluginTable;
#else
table->next = new(std::nothrow) PluginTable;
#endif
if (!table->next) {
std::snprintf(err, sizeof(err), "failed to allocate memory for the global plugin table");
return -1;
}
table = table->next;
}
// release the attribute names from unique_ptr into a plain array
const char** attributes = nullptr;
if (plugin->nattribute) {
attributes = new(std::nothrow) const char*[plugin->nattribute];
if (!attributes) {
std::snprintf(err, sizeof(err), "failed to allocate memory for plugin attribute array");
return -1;
}
for (int i = 0; i < plugin->nattribute; ++i) {
attributes[i] = attributes_vec[i].release();
}
}
// all checked passed, actually register the plugin into the global table
mjpPlugin& registered_plugin = table->plugins[local_idx];
registered_plugin = *plugin;
registered_plugin.name = name.release();
registered_plugin.attributes = attributes;
// increment the global plugin count with a release memory barrier
global.count().store(count + 1, std::memory_order_release);
return count;
}();
// ========= ATTENTION! ==========================================================================
// End of safe lambda, do not handle objects with non-trivial destructors beyond this point.
// plugin registration failed, throw an mju_error
if (slot < 0) {
err[sizeof(err) - 1] = '\0';
mju_error(err);
}
return slot;
}
// look up plugin by slot number, assuming that mjp_pluginCount has already been called
const mjpPlugin* mjp_getPluginAtSlotUnsafe(int slot, int nslot) {
if (slot < 0 || slot >= nslot) {
return nullptr;
}
Global& global = GetGlobal();
PluginTable* table = &global.table();
// iterate over blocks in the global table until the local index is less the block size
int local_idx = slot;
while (local_idx >= PluginTable::kBlockSize) {
local_idx -= PluginTable::kBlockSize;
table = table->next;
if (!table) {
return nullptr;
}
}
// local_idx is now a valid index into the current block
const mjpPlugin& plugin = table->plugins[local_idx];
if (!plugin.name) {
return nullptr;
}
return &plugin;
}
// look up plugin by name, assuming that mjp_pluginCount has already been called
const mjpPlugin* mjp_getPluginUnsafe(const char* name, int* slot, int nslot) {
if (slot) *slot = -1;
if (!name || !name[0]) {
return nullptr;
}
Global& plugin = GetGlobal();
PluginTable* table = &plugin.table();
int found_slot = 0;
while (table) {
for (int i = 0;
i < PluginTable::kBlockSize && found_slot < nslot;
++i, ++found_slot) {
const mjpPlugin& plugin = table->plugins[i];
// reached an uninitialized plugin, which means that iterated beyond the plugin count
// this should never happen if `count` was actually returned by mjp_pluginCount
if (!plugin.name) {
return nullptr;
}
if (std::strcmp(plugin.name, name) == 0) {
if (slot) *slot = found_slot;
return &plugin;
}
}
table = table->next;
}
return nullptr;
}
// return the number of globally registered plugins
int mjp_pluginCount() {
return GetGlobal().count().load(std::memory_order_acquire);
}
// look up a plugin by slot number
const mjpPlugin* mjp_getPluginAtSlot(int slot) {
const int count = mjp_pluginCount();
// mjp_pluginCount uses memory_order_acquire which acts as a barrier that guarantees that all
// plugins up to `count` have been completely inserted
return mjp_getPluginAtSlotUnsafe(slot, count);
}
// look up a plugin by name, optionally also get its registered slot number
const mjpPlugin* mjp_getPlugin(const char* name, int* slot) {
const int count = mjp_pluginCount();
int found_slot = -1;
const mjpPlugin* plugin = mjp_getPluginUnsafe(name, &found_slot, count);
if (slot) *slot = found_slot;
return plugin;
}
namespace {
// seek the nth config attrib of a plugin instance by counting null terminators
const char* PluginAttrSeek(const mjModel* m, int plugin_id, int attrib_id) {
const char* ptr = m->plugin_attr + m->plugin_attradr[plugin_id];
for (int i = 0; i < attrib_id; ++i) {
while (*ptr) {
++ptr;
}
++ptr;
}
return ptr;
}
} // namespace
// return a config attribute of a plugin instance
// NULL: invalid plugin instance ID or attribute name
const char* mj_getPluginConfig(const mjModel* m, int plugin_id, const char* attrib) {
if (plugin_id < 0 || plugin_id >= m->nplugin || attrib == nullptr) {
return nullptr;
}
const mjpPlugin* plugin = mjp_getPluginAtSlot(m->plugin[plugin_id]);
if (!plugin) {
return nullptr;
}
for (int i = 0; i < plugin->nattribute; ++i) {
if (std::strcmp(plugin->attributes[i], attrib) == 0) {
return PluginAttrSeek(m, plugin_id, i);
}
}
return nullptr;
}
+62
View File
@@ -0,0 +1,62 @@
// Copyright 2022 DeepMind Technologies Limited
//
// Licensed under the Apache License, Version 2.0 (the "License");
// you may not use this file except in compliance with the License.
// You may obtain a copy of the License at
//
// http://www.apache.org/licenses/LICENSE-2.0
//
// Unless required by applicable law or agreed to in writing, software
// distributed under the License is distributed on an "AS IS" BASIS,
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
// See the License for the specific language governing permissions and
// limitations under the License.
#ifndef MUJOCO_SRC_ENGINE_ENGINE_PLUGIN_H_
#define MUJOCO_SRC_ENGINE_ENGINE_PLUGIN_H_
#include <mujoco/mjexport.h>
#include <mujoco/mjplugin.h>
#ifdef __cplusplus
extern "C" {
#endif
// set default plugin definition
MJAPI void mjp_defaultPlugin(mjpPlugin* plugin);
// globally register a plugin (thread-safe), return new slot id
MJAPI int mjp_registerPlugin(const mjpPlugin* plugin);
// return the number of globally registered plugins
MJAPI int mjp_pluginCount();
// look up a plugin by name, optionally also get its registered slot number
MJAPI const mjpPlugin* mjp_getPlugin(const char* name, int* slot);
// look up a plugin by slot number
MJAPI const mjpPlugin* mjp_getPluginAtSlot(int slot);
// return a config attribute of a plugin instance
// NULL: invalid plugin instance ID or attribute name
MJAPI const char* mj_getPluginConfig(const mjModel* m, int plugin_id, const char* attrib);
// =================================================================================================
// MuJoCo-internal functions beyond this point.
// "Unsafe" suffix indicates that improper use of these functions may result in data races.
//
// The unsafe functions assume that called mjp_pluginCount has already been called, and that it is
// safe to assume that all plugins up to `count` have been completely written into the global table.
// =================================================================================================
// look up a plugin by name, assuming that mjp_pluginCount has already been called
const mjpPlugin* mjp_getPluginUnsafe(const char* name, int* slot, int nslot);
// look up a plugin by slot number, assuming that mjp_pluginCount has already been called
const mjpPlugin* mjp_getPluginAtSlotUnsafe(int slot, int nslot);
#ifdef __cplusplus
}
#endif
#endif // MUJOCO_SRC_ENGINE_ENGINE_PLUGIN_H_
+88
View File
@@ -18,10 +18,12 @@
#include <mujoco/mjdata.h>
#include <mujoco/mjmodel.h>
#include <mujoco/mjplugin.h>
#include "engine/engine_callback.h"
#include "engine/engine_core_smooth.h"
#include "engine/engine_io.h"
#include "engine/engine_macro.h"
#include "engine/engine_plugin.h"
#include "engine/engine_ray.h"
#include "engine/engine_support.h"
#include "engine/engine_util_blas.h"
@@ -200,6 +202,11 @@ void mj_sensorPos(const mjModel* m, mjData* d) {
// process sensors matching stage
for (int i=0; i<m->nsensor; i++) {
// skip sensor plugins -- these are handled after builtin sensor types
if (m->sensor_type[i] == mjSENS_PLUGIN) {
continue;
}
if (m->sensor_needstage[i]==mjSTAGE_POS) {
// get sensor info
objtype = m->sensor_objtype[i];
@@ -342,6 +349,25 @@ void mj_sensorPos(const mjModel* m, mjData* d) {
add_noise(m, d, mjSTAGE_POS);
}
// compute plugin sensor values
if (m->nplugin) {
const int nslot = mjp_pluginCount();
for (int i=0; i<m->nplugin; i++) {
const int slot = m->plugin[i];
const mjpPlugin* plugin = mjp_getPluginAtSlotUnsafe(slot, nslot);
if (!plugin) {
mju_error_i("invalid plugin slot: %d", slot);
}
if ((plugin->type & mjPLUGIN_SENSOR) &&
(plugin->needstage==mjSTAGE_POS || plugin->needstage==mjSTAGE_NONE)) {
if (!plugin->compute) {
mju_error_i("`compute` is a null function pointer for plugin at slot %d", slot);
}
plugin->compute(m, d, i, mjPLUGIN_SENSOR);
}
}
}
// cutoff
apply_cutoff(m, d, mjSTAGE_POS);
}
@@ -362,6 +388,11 @@ void mj_sensorVel(const mjModel* m, mjData* d) {
// process sensors matching stage
int subtreeVel = 0;
for (int i=0; i<m->nsensor; i++) {
// skip sensor plugins -- these are handled after builtin sensor types
if (m->sensor_type[i] == mjSENS_PLUGIN) {
continue;
}
if (m->sensor_needstage[i]==mjSTAGE_VEL) {
// get sensor info
type = m->sensor_type[i];
@@ -499,6 +530,32 @@ void mj_sensorVel(const mjModel* m, mjData* d) {
add_noise(m, d, mjSTAGE_VEL);
}
// trigger computation of plugins
if (m->nplugin) {
const int nslot = mjp_pluginCount();
for (int i=0; i<m->nplugin; i++) {
const int slot = m->plugin[i];
const mjpPlugin* plugin = mjp_getPluginAtSlotUnsafe(slot, nslot);
if (!plugin) {
mju_error_i("invalid plugin slot: %d", slot);
}
if ((plugin->type & mjPLUGIN_SENSOR) && plugin->needstage==mjSTAGE_VEL) {
if (!plugin->compute) {
mju_error_i("`compute` is null for plugin at slot %d", slot);
}
if (subtreeVel == 0) {
// compute subtree_linvel, subtree_angmom
// TODO(b/247107630): add a flag to allow plugin to specify whether it actually needs this
mj_subtreeVel(m, d);
// mark computed
subtreeVel = 1;
}
plugin->compute(m, d, i, mjPLUGIN_SENSOR);
}
}
}
// cutoff
apply_cutoff(m, d, mjSTAGE_VEL);
}
@@ -520,6 +577,11 @@ void mj_sensorAcc(const mjModel* m, mjData* d) {
// process sensors matching stage
int rnePost = 0;
for (int i=0; i<m->nsensor; i++) {
// skip sensor plugins -- these are handled after builtin sensor types
if (m->sensor_type[i] == mjSENS_PLUGIN) {
continue;
}
if (m->sensor_needstage[i]==mjSTAGE_ACC) {
// get sensor info
type = m->sensor_type[i];
@@ -677,6 +739,32 @@ void mj_sensorAcc(const mjModel* m, mjData* d) {
add_noise(m, d, mjSTAGE_ACC);
}
// trigger computation of plugins
if (m->nplugin) {
const int nslot = mjp_pluginCount();
for (int i=0; i<m->nplugin; i++) {
const int slot = m->plugin[i];
const mjpPlugin* plugin = mjp_getPluginAtSlotUnsafe(slot, nslot);
if (!plugin) {
mju_error_i("invalid plugin slot: %d", slot);
}
if ((plugin->type & mjPLUGIN_SENSOR) && plugin->needstage==mjSTAGE_ACC) {
if (!plugin->compute) {
mju_error_i("`compute` is null for plugin at slot %d", slot);
}
if (rnePost == 0) {
// compute cacc, cfrc_int, cfrc_ext
// TODO(b/247107630): add a flag to allow plugin to specify whether it actually needs this
mj_rnePostConstraint(m, d);
// mark computed
rnePost = 1;
}
plugin->compute(m, d, i, mjPLUGIN_SENSOR);
}
}
}
// cutoff
apply_cutoff(m, d, mjSTAGE_ACC);
}
+4 -21
View File
@@ -440,107 +440,90 @@ static int _getnumadr(const mjModel* m, mjtObj type, int** padr) {
case mjOBJ_XBODY:
*padr = m->name_bodyadr;
return m->nbody;
break;
case mjOBJ_JOINT:
*padr = m->name_jntadr;
return m->njnt;
break;
case mjOBJ_GEOM:
*padr = m->name_geomadr;
return m->ngeom;
break;
case mjOBJ_SITE:
*padr = m->name_siteadr;
return m->nsite;
break;
case mjOBJ_CAMERA:
*padr = m->name_camadr;
return m->ncam;
break;
case mjOBJ_LIGHT:
*padr = m->name_lightadr;
return m->nlight;
break;
case mjOBJ_MESH:
*padr = m->name_meshadr;
return m->nmesh;
break;
case mjOBJ_SKIN:
*padr = m->name_skinadr;
return m->nskin;
break;
case mjOBJ_HFIELD:
*padr = m->name_hfieldadr;
return m->nhfield;
break;
case mjOBJ_TEXTURE:
*padr = m->name_texadr;
return m->ntex;
break;
case mjOBJ_MATERIAL:
*padr = m->name_matadr;
return m->nmat;
break;
case mjOBJ_PAIR:
*padr = m->name_pairadr;
return m->npair;
break;
case mjOBJ_EXCLUDE:
*padr = m->name_excludeadr;
return m->nexclude;
break;
case mjOBJ_EQUALITY:
*padr = m->name_eqadr;
return m->neq;
break;
case mjOBJ_TENDON:
*padr = m->name_tendonadr;
return m->ntendon;
break;
case mjOBJ_ACTUATOR:
*padr = m->name_actuatoradr;
return m->nu;
break;
case mjOBJ_SENSOR:
*padr = m->name_sensoradr;
return m->nsensor;
break;
case mjOBJ_NUMERIC:
*padr = m->name_numericadr;
return m->nnumeric;
break;
case mjOBJ_TEXT:
*padr = m->name_textadr;
return m->ntext;
break;
case mjOBJ_TUPLE:
*padr = m->name_tupleadr;
return m->ntuple;
break;
case mjOBJ_KEY:
*padr = m->name_keyadr;
return m->nkey;
break;
case mjOBJ_PLUGIN:
*padr = m->name_pluginadr;
return m->nplugin;
default:
*padr = 0;
+7
View File
@@ -831,6 +831,9 @@ const char* mju_type2Str(int type) {
case mjOBJ_KEY:
return "key";
case mjOBJ_PLUGIN:
return "plugin";
default:
return 0;
}
@@ -932,6 +935,10 @@ int mju_str2Type(const char* str) {
return mjOBJ_KEY;
}
else if (!strcmp(str, "plugin")) {
return mjOBJ_PLUGIN;
}
else {
return mjOBJ_UNKNOWN;
}