From e4f1d584ad5e9d5b9929beb45ade5578dccb4127 Mon Sep 17 00:00:00 2001 From: Saran Tunyasuvunakool Date: Tue, 6 Dec 2022 08:23:07 -0800 Subject: [PATCH] Add control plugin API. PiperOrigin-RevId: 493311715 Change-Id: Iccd99fb26004f39832818e378f01f9149a0e3e3b --- doc/includes/references.h | 12 ++++++---- include/mujoco/mjplugin.h | 9 +++---- introspect/enums.py | 7 +++--- plugin/elasticity/cable.cc | 11 +++++---- plugin/elasticity/cable.h | 3 +-- plugin/elasticity/solid.cc | 11 +++++---- plugin/elasticity/solid.h | 3 +-- src/engine/engine_core_smooth.c | 2 +- src/engine/engine_forward.c | 30 ++++++++++++++++++++---- src/engine/engine_sensor.c | 6 ++--- src/user/user_model.cc | 2 +- src/user/user_objects.cc | 6 ++--- test/engine/engine_plugin_test.cc | 6 ++--- unity/Runtime/Bindings/MujocoBindings.cs | 3 ++- 14 files changed, 70 insertions(+), 41 deletions(-) diff --git a/doc/includes/references.h b/doc/includes/references.h index 14908ad4..4c34fba6 100644 --- a/doc/includes/references.h +++ b/doc/includes/references.h @@ -1148,18 +1148,19 @@ struct mjModel_ { char* names; // names of all objects, 0-terminated (nnames x 1) }; typedef struct mjModel_ mjModel; -typedef enum mjtPluginTypeBit_ { +typedef enum mjtPluginCapabilityBit_ { mjPLUGIN_ACTUATOR = 1<<0, mjPLUGIN_SENSOR = 1<<1, mjPLUGIN_PASSIVE = 1<<2, -} mjtPluginTypeBit; + mjPLUGIN_CONTROL = 1<<3, +} mjtPluginCapabilityBit; struct mjpPlugin_ { const char* name; // globally unique name identifying the plugin int nattribute; // number of configuration attributes const char* const* attributes; // name of configuration attributes - int type; // bitfield of mjtPluginTypeBits specifying the plugin type + int capabilities; // bitfield of mjtPluginCapabilityBit specifying optional plugin capabilities int needstage; // an mjtStage enum value specifying the sensor computation stage // number of mjtNums needed to store the state of a plugin instance (required) @@ -1181,10 +1182,13 @@ struct mjpPlugin_ { void (*reset)(const mjModel* m, mjData* d, int instance); // called when the plugin needs to update its outputs (required) - void (*compute)(const mjModel* m, mjData* d, int instance, int type); + void (*compute)(const mjModel* m, mjData* d, int instance, int capability_bit); // called when time integration occurs (optional) void (*advance)(const mjModel* m, mjData* d, int instance); + + // called by mjv_updateScene (optional) + void (*visualize)(const mjModel*m, mjData* d, mjvScene* scn, int instance); }; typedef struct mjpPlugin_ mjpPlugin; typedef enum mjtGridPos_ { // grid position for overlay diff --git a/include/mujoco/mjplugin.h b/include/mujoco/mjplugin.h index 466f230e..f6c8a3fd 100644 --- a/include/mujoco/mjplugin.h +++ b/include/mujoco/mjplugin.h @@ -20,11 +20,12 @@ #include -typedef enum mjtPluginTypeBit_ { +typedef enum mjtPluginCapabilityBit_ { mjPLUGIN_ACTUATOR = 1<<0, mjPLUGIN_SENSOR = 1<<1, mjPLUGIN_PASSIVE = 1<<2, -} mjtPluginTypeBit; + mjPLUGIN_CONTROL = 1<<3, +} mjtPluginCapabilityBit; struct mjpPlugin_ { const char* name; // globally unique name identifying the plugin @@ -32,7 +33,7 @@ struct mjpPlugin_ { int nattribute; // number of configuration attributes const char* const* attributes; // name of configuration attributes - int type; // bitfield of mjtPluginTypeBits specifying the plugin type + int capabilities; // bitfield of mjtPluginCapabilityBit specifying optional plugin capabilities int needstage; // an mjtStage enum value specifying the sensor computation stage // number of mjtNums needed to store the state of a plugin instance (required) @@ -54,7 +55,7 @@ struct mjpPlugin_ { void (*reset)(const mjModel* m, mjData* d, int instance); // called when the plugin needs to update its outputs (required) - void (*compute)(const mjModel* m, mjData* d, int instance, int type); + void (*compute)(const mjModel* m, mjData* d, int instance, int capability_bit); // called when time integration occurs (optional) void (*advance)(const mjModel* m, mjData* d, int instance); diff --git a/introspect/enums.py b/introspect/enums.py index 89bb479b..5d72b063 100755 --- a/introspect/enums.py +++ b/introspect/enums.py @@ -550,14 +550,15 @@ ENUMS: Mapping[str, EnumDecl] = dict([ ('mjSTEREO_SIDEBYSIDE', 2), ]), )), - ('mjtPluginTypeBit', + ('mjtPluginCapabilityBit', EnumDecl( - name='mjtPluginTypeBit', - declname='enum mjtPluginTypeBit_', + name='mjtPluginCapabilityBit', + declname='enum mjtPluginCapabilityBit_', values=dict([ ('mjPLUGIN_ACTUATOR', 1), ('mjPLUGIN_SENSOR', 2), ('mjPLUGIN_PASSIVE', 4), + ('mjPLUGIN_CONTROL', 8), ]), )), ('mjtGridPos', diff --git a/plugin/elasticity/cable.cc b/plugin/elasticity/cable.cc index 08bd6884..ff50ae31 100644 --- a/plugin/elasticity/cable.cc +++ b/plugin/elasticity/cable.cc @@ -277,7 +277,7 @@ void Cable::RegisterPlugin() { mjp_defaultPlugin(&plugin); plugin.name = "mujoco.elasticity.cable"; - plugin.type |= mjPLUGIN_PASSIVE; + plugin.capabilities |= mjPLUGIN_PASSIVE; const char* attributes[] = {"twist", "bend", "flat", "vmax"}; plugin.nattribute = sizeof(attributes) / sizeof(attributes[0]); @@ -297,10 +297,11 @@ return 0; delete reinterpret_cast(d->plugin_data[instance]); d->plugin_data[instance] = 0; }; - plugin.compute = +[](const mjModel* m, mjData* d, int instance, int type) { - auto* elasticity = reinterpret_cast(d->plugin_data[instance]); - elasticity->Compute(m, d, instance); - }; + plugin.compute = + +[](const mjModel* m, mjData* d, int instance, int capability_bit) { + auto* elasticity = reinterpret_cast(d->plugin_data[instance]); + elasticity->Compute(m, d, instance); + }; plugin.visualize = +[](const mjModel* m, mjData* d, mjvScene* scn, int instance) { auto* elasticity = reinterpret_cast(d->plugin_data[instance]); diff --git a/plugin/elasticity/cable.h b/plugin/elasticity/cable.h index f6f7542e..bf2bf75d 100644 --- a/plugin/elasticity/cable.h +++ b/plugin/elasticity/cable.h @@ -30,8 +30,7 @@ class Cable { public: // Creates a new Cable instance (allocated with `new`) or // returns null on failure. - static std::optional Create(const mjModel* m, mjData* d, - int instance); + static std::optional Create(const mjModel* m, mjData* d, int instance); Cable(Cable&&) = default; ~Cable() = default; diff --git a/plugin/elasticity/solid.cc b/plugin/elasticity/solid.cc index c79df296..fad8f73f 100644 --- a/plugin/elasticity/solid.cc +++ b/plugin/elasticity/solid.cc @@ -347,7 +347,7 @@ void Solid::RegisterPlugin() { mjp_defaultPlugin(&plugin); plugin.name = "mujoco.elasticity.solid"; - plugin.type |= mjPLUGIN_PASSIVE; + plugin.capabilities |= mjPLUGIN_PASSIVE; const char* attributes[] = {"nx", "ny", "nz", "young", "poisson", "damping"}; plugin.nattribute = sizeof(attributes) / sizeof(attributes[0]); @@ -367,10 +367,11 @@ void Solid::RegisterPlugin() { delete reinterpret_cast(d->plugin_data[instance]); d->plugin_data[instance] = 0; }; - plugin.compute = +[](const mjModel* m, mjData* d, int instance, int type) { - auto* elasticity = reinterpret_cast(d->plugin_data[instance]); - elasticity->Compute(m, d, instance); - }; + plugin.compute = + +[](const mjModel* m, mjData* d, int instance, int capability_bit) { + auto* elasticity = reinterpret_cast(d->plugin_data[instance]); + elasticity->Compute(m, d, instance); + }; mjp_registerPlugin(&plugin); } diff --git a/plugin/elasticity/solid.h b/plugin/elasticity/solid.h index 25593f56..48b4c526 100644 --- a/plugin/elasticity/solid.h +++ b/plugin/elasticity/solid.h @@ -33,8 +33,7 @@ struct Stencil { class Solid { public: // Returns a new Solid instance or nullopt on failure. - static std::optional Create(const mjModel* m, mjData* d, - int instance); + static std::optional Create(const mjModel* m, mjData* d, int instance); Solid(Solid&&) = default; Solid& operator=(Solid&& other) = default; diff --git a/src/engine/engine_core_smooth.c b/src/engine/engine_core_smooth.c index a23540f2..f2967b49 100644 --- a/src/engine/engine_core_smooth.c +++ b/src/engine/engine_core_smooth.c @@ -1454,7 +1454,7 @@ void mj_passive(const mjModel* m, mjData* d) { if (!plugin) { mju_error_i("invalid plugin slot: %d", slot); } - if (plugin->type & mjPLUGIN_PASSIVE) { + if (plugin->capabilities & mjPLUGIN_PASSIVE) { if (!plugin->compute) { mju_error_i("`compute` is a null function pointer for plugin at slot %d", slot); } diff --git a/src/engine/engine_forward.c b/src/engine/engine_forward.c index 40288d40..618e43fb 100644 --- a/src/engine/engine_forward.c +++ b/src/engine/engine_forward.c @@ -279,7 +279,7 @@ void mj_fwdActuation(const mjModel* m, mjData* d) { if (!plugin) { mju_error_i("invalid plugin slot: %d", slot); } - if (plugin->type & mjPLUGIN_ACTUATOR) { + if (plugin->capabilities & mjPLUGIN_ACTUATOR) { if (!plugin->compute) { mju_error_i("`compute` is a null function pointer for plugin at slot %d", slot); } @@ -774,9 +774,31 @@ void mj_forwardSkip(const mjModel* m, mjData* d, int skipstage, int skipsensor) } // acceleration-dependent - if (mjcb_control && !mjDISABLED(mjDSBL_ACTUATION)) { - mjcb_control(m, d); - } + if (!mjDISABLED(mjDSBL_ACTUATION)) { + // call legacy control callback if specified + if (mjcb_control) { + mjcb_control(m, d); + } + + // handle control plugins + if (m->nplugin) { + const int nslot = mjp_pluginCount(); + for (int i=0; inplugin; 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->capabilities & mjPLUGIN_CONTROL) { + if (!plugin->compute) { + mju_error_i("`compute` is a null function pointer for plugin at slot %d", slot); + } + plugin->compute(m, d, i, mjPLUGIN_CONTROL); + } + } + } +} + mj_fwdActuation(m, d); mj_fwdAcceleration(m, d); mj_fwdConstraint(m, d); diff --git a/src/engine/engine_sensor.c b/src/engine/engine_sensor.c index c2b29734..378b962a 100644 --- a/src/engine/engine_sensor.c +++ b/src/engine/engine_sensor.c @@ -357,7 +357,7 @@ void mj_sensorPos(const mjModel* m, mjData* d) { if (!plugin) { mju_error_i("invalid plugin slot: %d", slot); } - if ((plugin->type & mjPLUGIN_SENSOR) && + if ((plugin->capabilities & 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); @@ -538,7 +538,7 @@ void mj_sensorVel(const mjModel* m, mjData* d) { if (!plugin) { mju_error_i("invalid plugin slot: %d", slot); } - if ((plugin->type & mjPLUGIN_SENSOR) && plugin->needstage==mjSTAGE_VEL) { + if ((plugin->capabilities & mjPLUGIN_SENSOR) && plugin->needstage==mjSTAGE_VEL) { if (!plugin->compute) { mju_error_i("`compute` is null for plugin at slot %d", slot); } @@ -747,7 +747,7 @@ void mj_sensorAcc(const mjModel* m, mjData* d) { if (!plugin) { mju_error_i("invalid plugin slot: %d", slot); } - if ((plugin->type & mjPLUGIN_SENSOR) && plugin->needstage==mjSTAGE_ACC) { + if ((plugin->capabilities & mjPLUGIN_SENSOR) && plugin->needstage==mjSTAGE_ACC) { if (!plugin->compute) { mju_error_i("`compute` is null for plugin at slot %d", slot); } diff --git a/src/user/user_model.cc b/src/user/user_model.cc index 231db8fb..63e68a46 100644 --- a/src/user/user_model.cc +++ b/src/user/user_model.cc @@ -2633,7 +2633,7 @@ void mjCModel::TryCompile(mjModel*& m, mjData*& d, const mjVFS* vfs) { m->plugin_stateadr[i] = stateadr; m->plugin_statenum[i] = nstate; stateadr += nstate; - if (plugin->type & mjPLUGIN_SENSOR) { + if (plugin->capabilities & mjPLUGIN_SENSOR) { for (int sensor_id : plugin_to_sensors[i]) { if (!plugin->nsensordata) { mju_error_i("`reset` is null for plugin at slot %d", m->plugin[i]); diff --git a/src/user/user_objects.cc b/src/user/user_objects.cc index 3f190a95..d9a5bb62 100644 --- a/src/user/user_objects.cc +++ b/src/user/user_objects.cc @@ -819,7 +819,7 @@ void mjCBody::Compile(void) { model->ResolvePlugin(this, plugin_name, plugin_instance_name, &plugin_instance); const mjpPlugin* plugin = mjp_getPluginAtSlot(plugin_instance->plugin_slot); - if (!(plugin->type & mjPLUGIN_PASSIVE)) { + if (!(plugin->capabilities & mjPLUGIN_PASSIVE)) { throw mjCError(this, "plugin '%s' does not support passive forces", plugin->name); } } @@ -3618,7 +3618,7 @@ void mjCActuator::Compile(void) { model->ResolvePlugin(this, plugin_name, plugin_instance_name, &plugin_instance); const mjpPlugin* plugin = mjp_getPluginAtSlot(plugin_instance->plugin_slot); - if (!(plugin->type & mjPLUGIN_ACTUATOR)) { + if (!(plugin->capabilities & mjPLUGIN_ACTUATOR)) { throw mjCError(this, "plugin '%s' does not support actuators", plugin->name); } } @@ -4006,7 +4006,7 @@ void mjCSensor::Compile(void) { { model->ResolvePlugin(this, plugin_name, plugin_instance_name, &plugin_instance); const mjpPlugin* plugin = mjp_getPluginAtSlot(plugin_instance->plugin_slot); - if (!(plugin->type & mjPLUGIN_SENSOR)) { + if (!(plugin->capabilities & mjPLUGIN_SENSOR)) { throw mjCError(this, "plugin '%s' does not support sensors", plugin->name); } needstage = static_cast(plugin->needstage); diff --git a/test/engine/engine_plugin_test.cc b/test/engine/engine_plugin_test.cc index c6b46d64..0f48da76 100644 --- a/test/engine/engine_plugin_test.cc +++ b/test/engine/engine_plugin_test.cc @@ -206,7 +206,7 @@ int RegisterSensorPlugin() { plugin.nattribute = sizeof(attributes) / sizeof(*attributes); plugin.attributes = attributes; - plugin.type |= mjPLUGIN_SENSOR; + plugin.capabilities |= mjPLUGIN_SENSOR; plugin.nstate = +[](const mjModel* m, int instance) { return 3; }; plugin.nsensordata = @@ -250,7 +250,7 @@ int RegisterActuatorPlugin() { plugin.nattribute = sizeof(attributes) / sizeof(*attributes); plugin.attributes = attributes; - plugin.type |= mjPLUGIN_ACTUATOR; + plugin.capabilities |= mjPLUGIN_ACTUATOR; plugin.nstate = +[](const mjModel* m, int instance) { return 3; }; @@ -292,7 +292,7 @@ int RegisterPassivePlugin() { plugin.nattribute = sizeof(attributes) / sizeof(*attributes); plugin.attributes = attributes; - plugin.type |= mjPLUGIN_PASSIVE; + plugin.capabilities |= mjPLUGIN_PASSIVE; plugin.nstate = +[](const mjModel* m, int instance) { return 0; }; diff --git a/unity/Runtime/Bindings/MujocoBindings.cs b/unity/Runtime/Bindings/MujocoBindings.cs index 661c1913..40be7f75 100644 --- a/unity/Runtime/Bindings/MujocoBindings.cs +++ b/unity/Runtime/Bindings/MujocoBindings.cs @@ -361,10 +361,11 @@ public enum mjtLRMode : int{ mjLRMODE_MUSCLEUSER = 2, mjLRMODE_ALL = 3, } -public enum mjtPluginTypeBit : int{ +public enum mjtPluginCapabilityBit : int{ mjPLUGIN_ACTUATOR = 1, mjPLUGIN_SENSOR = 2, mjPLUGIN_PASSIVE = 4, + mjPLUGIN_CONTROL = 8, } public enum mjtGridPos : int{ mjGRID_TOPLEFT = 0,