Refactor CopyPlugins() out of TryCompile()
PiperOrigin-RevId: 698072233 Change-Id: I2138d9d8383838328611eef463448fe2f15bd173
This commit is contained in:
committed by
Copybara-Service
parent
821f1d3f98
commit
58d68d08d3
+84
-82
@@ -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);
|
||||
|
||||
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user