Add control plugin API.

PiperOrigin-RevId: 493311715
Change-Id: Iccd99fb26004f39832818e378f01f9149a0e3e3b
This commit is contained in:
Saran Tunyasuvunakool
2022-12-06 08:23:07 -08:00
committed by Copybara-Service
parent 60fbb77b2a
commit e4f1d584ad
14 changed files with 70 additions and 41 deletions
+1 -1
View File
@@ -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);
}
+26 -4
View File
@@ -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);
+3 -3
View File
@@ -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);
}
+1 -1
View File
@@ -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]);
+3 -3
View File
@@ -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);