Add control plugin API.
PiperOrigin-RevId: 493311715 Change-Id: Iccd99fb26004f39832818e378f01f9149a0e3e3b
This commit is contained in:
committed by
Copybara-Service
parent
60fbb77b2a
commit
e4f1d584ad
@@ -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
|
||||
|
||||
@@ -20,11 +20,12 @@
|
||||
#include <mujoco/mjvisualize.h>
|
||||
|
||||
|
||||
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);
|
||||
|
||||
+4
-3
@@ -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',
|
||||
|
||||
@@ -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<Cable*>(d->plugin_data[instance]);
|
||||
d->plugin_data[instance] = 0;
|
||||
};
|
||||
plugin.compute = +[](const mjModel* m, mjData* d, int instance, int type) {
|
||||
auto* elasticity = reinterpret_cast<Cable*>(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<Cable*>(d->plugin_data[instance]);
|
||||
elasticity->Compute(m, d, instance);
|
||||
};
|
||||
plugin.visualize =
|
||||
+[](const mjModel* m, mjData* d, mjvScene* scn, int instance) {
|
||||
auto* elasticity = reinterpret_cast<Cable*>(d->plugin_data[instance]);
|
||||
|
||||
@@ -30,8 +30,7 @@ class Cable {
|
||||
public:
|
||||
// Creates a new Cable instance (allocated with `new`) or
|
||||
// returns null on failure.
|
||||
static std::optional<Cable> Create(const mjModel* m, mjData* d,
|
||||
int instance);
|
||||
static std::optional<Cable> Create(const mjModel* m, mjData* d, int instance);
|
||||
Cable(Cable&&) = default;
|
||||
~Cable() = default;
|
||||
|
||||
|
||||
@@ -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<Solid*>(d->plugin_data[instance]);
|
||||
d->plugin_data[instance] = 0;
|
||||
};
|
||||
plugin.compute = +[](const mjModel* m, mjData* d, int instance, int type) {
|
||||
auto* elasticity = reinterpret_cast<Solid*>(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<Solid*>(d->plugin_data[instance]);
|
||||
elasticity->Compute(m, d, instance);
|
||||
};
|
||||
|
||||
mjp_registerPlugin(&plugin);
|
||||
}
|
||||
|
||||
@@ -33,8 +33,7 @@ struct Stencil {
|
||||
class Solid {
|
||||
public:
|
||||
// Returns a new Solid instance or nullopt on failure.
|
||||
static std::optional<Solid> Create(const mjModel* m, mjData* d,
|
||||
int instance);
|
||||
static std::optional<Solid> Create(const mjModel* m, mjData* d, int instance);
|
||||
Solid(Solid&&) = default;
|
||||
|
||||
Solid& operator=(Solid&& other) = default;
|
||||
|
||||
@@ -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);
|
||||
}
|
||||
|
||||
@@ -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; 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->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);
|
||||
|
||||
@@ -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);
|
||||
}
|
||||
|
||||
@@ -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]);
|
||||
|
||||
@@ -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<mjtStage>(plugin->needstage);
|
||||
|
||||
@@ -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; };
|
||||
|
||||
|
||||
@@ -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,
|
||||
|
||||
Reference in New Issue
Block a user