Refactor CopyPlugins() out of TryCompile()

PiperOrigin-RevId: 698072233
Change-Id: I2138d9d8383838328611eef463448fe2f15bd173
This commit is contained in:
Yuval Tassa
2024-11-19 10:46:03 -08:00
committed by Copybara-Service
parent 821f1d3f98
commit 58d68d08d3
2 changed files with 85 additions and 82 deletions
+84 -82
View File
@@ -2493,7 +2493,90 @@ void mjCModel::CopyTree(mjModel* m) {
m->nC = nC = 2 * nOD + nv;
}
// copy plugin data
void mjCModel::CopyPlugins(mjModel* m) {
// assign plugin slots and copy plugin config attributes
{
int adr = 0;
for (int i = 0; i < nplugin; ++i) {
m->plugin[i] = plugins_[i]->plugin_slot;
const int size = plugins_[i]->flattened_attributes.size();
std::memcpy(m->plugin_attr + adr,
plugins_[i]->flattened_attributes.data(), size);
m->plugin_attradr[i] = adr;
adr += size;
}
}
// query and set plugin-related information
{
// set actuator_plugin to the plugin instance ID
std::vector<std::vector<int>> plugin_to_actuators(nplugin);
for (int i = 0; i < nu; ++i) {
if (actuators_[i]->plugin.active) {
int actuator_plugin = static_cast<mjCPlugin*>(actuators_[i]->plugin.element)->id;
m->actuator_plugin[i] = actuator_plugin;
plugin_to_actuators[actuator_plugin].push_back(i);
} else {
m->actuator_plugin[i] = -1;
}
}
for (int i = 0; i < nbody; ++i) {
if (bodies_[i]->plugin.active) {
m->body_plugin[i] = static_cast<mjCPlugin*>(bodies_[i]->plugin.element)->id;
} else {
m->body_plugin[i] = -1;
}
}
for (int i = 0; i < ngeom; ++i) {
if (geoms_[i]->plugin.active) {
m->geom_plugin[i] = static_cast<mjCPlugin*>(geoms_[i]->plugin.element)->id;
} else {
m->geom_plugin[i] = -1;
}
}
std::vector<std::vector<int>> plugin_to_sensors(nplugin);
for (int i = 0; i < nsensor; ++i) {
if (sensors_[i]->type == mjSENS_PLUGIN) {
int sensor_plugin = static_cast<mjCPlugin*>(sensors_[i]->plugin.element)->id;
m->sensor_plugin[i] = sensor_plugin;
plugin_to_sensors[sensor_plugin].push_back(i);
} else {
m->sensor_plugin[i] = -1;
}
}
// query plugin->nstate, compute and set plugin_state and plugin_stateadr
// for sensor plugins, also query plugin->nsensordata and set nsensordata
int stateadr = 0;
for (int i = 0; i < nplugin; ++i) {
const mjpPlugin* plugin = mjp_getPluginAtSlot(m->plugin[i]);
if (!plugin->nstate) {
mju_error("`nstate` is null for plugin at slot %d", m->plugin[i]);
}
int nstate = plugin->nstate(m, i);
m->plugin_stateadr[i] = stateadr;
m->plugin_statenum[i] = nstate;
stateadr += nstate;
if (plugin->capabilityflags & mjPLUGIN_SENSOR) {
for (int sensor_id : plugin_to_sensors[i]) {
if (!plugin->nsensordata) {
mju_error("`nsensordata` is null for plugin at slot %d", m->plugin[i]);
}
int nsensordata = plugin->nsensordata(m, i, sensor_id);
sensors_[sensor_id]->dim = nsensordata;
sensors_[sensor_id]->needstage =
static_cast<mjtStage>(plugin->needstage);
this->nsensordata += nsensordata;
}
}
}
m->npluginstate = stateadr;
}
}
// copy objects outside kinematic tree
void mjCModel::CopyObjects(mjModel* m) {
@@ -4061,88 +4144,7 @@ void mjCModel::TryCompile(mjModel*& m, mjData*& d, const mjVFS* vfs) {
CopyNames(m);
CopyPaths(m);
CopyTree(m);
// assign plugin slots and copy plugin config attributes
{
int adr = 0;
for (int i = 0; i < nplugin; ++i) {
m->plugin[i] = plugins_[i]->plugin_slot;
const int size = plugins_[i]->flattened_attributes.size();
std::memcpy(m->plugin_attr + adr,
plugins_[i]->flattened_attributes.data(), size);
m->plugin_attradr[i] = adr;
adr += size;
}
}
// query and set plugin-related information
{
// set actuator_plugin to the plugin instance ID
std::vector<std::vector<int>> plugin_to_actuators(nplugin);
for (int i = 0; i < nu; ++i) {
if (actuators_[i]->plugin.active) {
int actuator_plugin = static_cast<mjCPlugin*>(actuators_[i]->plugin.element)->id;
m->actuator_plugin[i] = actuator_plugin;
plugin_to_actuators[actuator_plugin].push_back(i);
} else {
m->actuator_plugin[i] = -1;
}
}
for (int i = 0; i < nbody; ++i) {
if (bodies_[i]->plugin.active) {
m->body_plugin[i] = static_cast<mjCPlugin*>(bodies_[i]->plugin.element)->id;
} else {
m->body_plugin[i] = -1;
}
}
for (int i = 0; i < ngeom; ++i) {
if (geoms_[i]->plugin.active) {
m->geom_plugin[i] = static_cast<mjCPlugin*>(geoms_[i]->plugin.element)->id;
} else {
m->geom_plugin[i] = -1;
}
}
std::vector<std::vector<int>> plugin_to_sensors(nplugin);
for (int i = 0; i < nsensor; ++i) {
if (sensors_[i]->type == mjSENS_PLUGIN) {
int sensor_plugin = static_cast<mjCPlugin*>(sensors_[i]->plugin.element)->id;
m->sensor_plugin[i] = sensor_plugin;
plugin_to_sensors[sensor_plugin].push_back(i);
} else {
m->sensor_plugin[i] = -1;
}
}
// query plugin->nstate, compute and set plugin_state and plugin_stateadr
// for sensor plugins, also query plugin->nsensordata and set nsensordata
int stateadr = 0;
for (int i = 0; i < nplugin; ++i) {
const mjpPlugin* plugin = mjp_getPluginAtSlot(m->plugin[i]);
if (!plugin->nstate) {
mju_error("`nstate` is null for plugin at slot %d", m->plugin[i]);
}
int nstate = plugin->nstate(m, i);
m->plugin_stateadr[i] = stateadr;
m->plugin_statenum[i] = nstate;
stateadr += nstate;
if (plugin->capabilityflags & mjPLUGIN_SENSOR) {
for (int sensor_id : plugin_to_sensors[i]) {
if (!plugin->nsensordata) {
mju_error("`nsensordata` is null for plugin at slot %d", m->plugin[i]);
}
int nsensordata = plugin->nsensordata(m, i, sensor_id);
sensors_[sensor_id]->dim = nsensordata;
sensors_[sensor_id]->needstage =
static_cast<mjtStage>(plugin->needstage);
this->nsensordata += nsensordata;
}
}
}
m->npluginstate = stateadr;
}
CopyPlugins(m);
// keyframe compilation needs access to nq, nv, na, nmocap, qpos0
ResolveKeyframes(m);
+1
View File
@@ -322,6 +322,7 @@ class mjCModel : public mjCModel_, private mjSpec {
void CopyPaths(mjModel*); // copy paths, compute path addresses
void CopyObjects(mjModel*); // copy objects outside kinematic tree
void CopyTree(mjModel*); // copy objects inside kinematic tree
void CopyPlugins(mjModel*); // copy plugin data
// objects created here
std::vector<mjCFlex*> flexes_; // list of flexes