Implement plugin mechanism for actuators and sensors.
PiperOrigin-RevId: 474874088 Change-Id: I65a8ffdf845f4fa0f8266c165883a747ab5812d8
This commit is contained in:
committed by
Copybara-Service
parent
f556d4d94f
commit
1e2a9a53bc
@@ -17,32 +17,65 @@
|
||||
#include <cfloat>
|
||||
#include <cstdio>
|
||||
#include <cstring>
|
||||
#include <functional>
|
||||
#include <iostream>
|
||||
#include <map>
|
||||
#include <sstream>
|
||||
#include <string>
|
||||
#include <string_view>
|
||||
#include <utility>
|
||||
#include <vector>
|
||||
|
||||
#include <mujoco/mjmodel.h>
|
||||
#include <mujoco/mjvisualize.h>
|
||||
#include "engine/engine_macro.h"
|
||||
#include "engine/engine_plugin.h"
|
||||
#include "engine/engine_util_errmem.h"
|
||||
#include "engine/engine_util_misc.h"
|
||||
#include "user/user_composite.h"
|
||||
#include "user/user_model.h"
|
||||
#include "user/user_objects.h"
|
||||
#include "user/user_util.h"
|
||||
#include "xml/xml_util.h"
|
||||
#include "tinyxml2.h"
|
||||
|
||||
namespace {
|
||||
|
||||
using std::string;
|
||||
using std::vector;
|
||||
using tinyxml2::XMLElement;
|
||||
|
||||
void ReadPluginConfigs(tinyxml2::XMLElement* elem, mjCPlugin* pp) {
|
||||
std::map<std::string, std::string, std::less<>> config_attribs;
|
||||
XMLElement* child = elem->FirstChildElement();
|
||||
while (child) {
|
||||
std::string_view name = child->Value();
|
||||
if (name == "config") {
|
||||
std::string key, value;
|
||||
mjXUtil::ReadAttrTxt(child, "key", key, /* required = */ true);
|
||||
if (config_attribs.find(key) != config_attribs.end()) {
|
||||
std::string err = "duplicate config key: " + key;
|
||||
throw mjXError(child, err.c_str());
|
||||
}
|
||||
mjXUtil::ReadAttrTxt(child, "value", value, /* required = */ true);
|
||||
config_attribs[key] = value;
|
||||
}
|
||||
child = child->NextSiblingElement();
|
||||
}
|
||||
|
||||
if (!pp && !config_attribs.empty()) {
|
||||
throw mjXError(elem,
|
||||
"plugin configuration attributes cannot be used in an "
|
||||
"element that references a predefined plugin instance");
|
||||
} else if (pp) {
|
||||
pp->config_attribs = std::move(config_attribs);
|
||||
}
|
||||
}
|
||||
} // namespace
|
||||
|
||||
|
||||
//---------------------------------- MJCF schema ---------------------------------------------------
|
||||
|
||||
static const int nMJCF = 165;
|
||||
static const int nMJCF = 183;
|
||||
static const char* MJCF[nMJCF][mjXATTRNUM] = {
|
||||
{"mujoco", "!", "1", "model"},
|
||||
{"<"},
|
||||
@@ -150,6 +183,17 @@ static const char* MJCF[nMJCF][mjXATTRNUM] = {
|
||||
"gain", "user", "group"},
|
||||
{">"},
|
||||
|
||||
{"extension", "*", "0"},
|
||||
{"<"},
|
||||
{"required", "*", "1", "plugin"},
|
||||
{"<"},
|
||||
{"instance", "*", "1", "name"},
|
||||
{"<"},
|
||||
{"config", "*", "2", "key", "value"},
|
||||
{">"},
|
||||
{">"},
|
||||
{">"},
|
||||
|
||||
{"custom", "*", "0"},
|
||||
{"<"},
|
||||
{"numeric", "*", "3", "name", "size", "data"},
|
||||
@@ -308,6 +352,13 @@ static const char* MJCF[nMJCF][mjXATTRNUM] = {
|
||||
"lmin", "lmax", "vmax", "fpmax", "fvmax"},
|
||||
{"adhesion", "*", "9", "name", "class", "group",
|
||||
"forcelimited", "ctrlrange", "forcerange", "user", "body", "gain"},
|
||||
{"plugin", "*", "19", "name", "class", "plugin", "instance", "group",
|
||||
"ctrllimited", "forcelimited", "ctrlrange", "forcerange",
|
||||
"lengthrange", "gear", "cranklength", "joint", "jointinparent",
|
||||
"site", "tendon", "cranksite", "slidersite", "user"},
|
||||
{"<"},
|
||||
{"config", "*", "2", "key", "value"},
|
||||
{">"},
|
||||
{">"},
|
||||
|
||||
{"sensor", "*", "0"},
|
||||
@@ -350,6 +401,11 @@ static const char* MJCF[nMJCF][mjXATTRNUM] = {
|
||||
{"clock", "*", "4", "name", "cutoff", "noise", "user"},
|
||||
{"user", "*", "9", "name", "objtype", "objname", "datatype", "needstage",
|
||||
"dim", "cutoff", "noise", "user"},
|
||||
{"plugin", "*", "9", "name", "plugin", "instance", "cutoff", "objtype", "objname", "reftype", "refname",
|
||||
"user"},
|
||||
{"<"},
|
||||
{"config", "*", "2", "key", "value"},
|
||||
{">"},
|
||||
{">"},
|
||||
|
||||
{"keyframe", "*", "0"},
|
||||
@@ -711,6 +767,11 @@ void mjXReader::Parse(XMLElement* root) {
|
||||
}
|
||||
readingdefaults = false;
|
||||
|
||||
for (section = root->FirstChildElement("extension"); section;
|
||||
section = section->NextSiblingElement("extension")) {
|
||||
Extension(section);
|
||||
}
|
||||
|
||||
for (section = root->FirstChildElement("custom"); section;
|
||||
section = section->NextSiblingElement("custom")) {
|
||||
Custom(section);
|
||||
@@ -1627,6 +1688,18 @@ void mjXReader::OneActuator(XMLElement* elem, mjCActuator* pact) {
|
||||
pact->biastype = mjBIAS_NONE;
|
||||
}
|
||||
|
||||
else if (type == "plugin") {
|
||||
pact->is_plugin = true;
|
||||
ReadAttrTxt(elem, "plugin", pact->plugin_name);
|
||||
ReadAttrTxt(elem, "instance", pact->plugin_instance_name);
|
||||
if (pact->plugin_instance_name.empty()) {
|
||||
pact->plugin_instance = model->AddPlugin();
|
||||
} else {
|
||||
model->hasImplicitPluginElem = true;
|
||||
}
|
||||
ReadPluginConfigs(elem, pact->plugin_instance);
|
||||
}
|
||||
|
||||
else { // SHOULD NOT OCCUR
|
||||
throw mjXError(elem, "unrecognized actuator type: %s", type.c_str());
|
||||
}
|
||||
@@ -1908,6 +1981,61 @@ void mjXReader::Default(XMLElement* section, int parentid) {
|
||||
|
||||
|
||||
|
||||
// extension section parser
|
||||
void mjXReader::Extension(XMLElement* section) {
|
||||
XMLElement* elem = section->FirstChildElement();
|
||||
while (elem) {
|
||||
// get sub-element name
|
||||
std::string_view name = elem->Value();
|
||||
|
||||
if (name == "required") {
|
||||
std::string plugin_name;
|
||||
int plugin_slot = -1;
|
||||
ReadAttrTxt(elem, "plugin", plugin_name, /* required = */ true);
|
||||
const mjpPlugin* plugin = mjp_getPlugin(plugin_name.c_str(), &plugin_slot);
|
||||
if (!plugin) {
|
||||
throw mjXError(elem, "unknown plugin '%s'", plugin_name.c_str());
|
||||
}
|
||||
|
||||
bool already_declared = false;
|
||||
for (const auto& [existing_plugin, existing_slot] : model->active_plugins) {
|
||||
if (plugin == existing_plugin) {
|
||||
already_declared = true;
|
||||
break;
|
||||
}
|
||||
}
|
||||
if (!already_declared) {
|
||||
model->active_plugins.emplace_back(std::make_pair(plugin, plugin_slot));
|
||||
}
|
||||
|
||||
XMLElement* child = elem->FirstChildElement();
|
||||
while (child) {
|
||||
if (std::string(child->Value())=="instance") {
|
||||
if (model->hasImplicitPluginElem) {
|
||||
throw mjXError(
|
||||
child, "explicit plugin instance must appear before implicit plugin elements");
|
||||
}
|
||||
mjCPlugin* pp = model->AddPlugin();
|
||||
GetXMLPos(child, pp);
|
||||
ReadAttrTxt(child, "name", pp->name, /* required = */ true);
|
||||
if (pp->name.empty()) {
|
||||
throw mjXError(child, "plugin instance must have a name");
|
||||
}
|
||||
ReadPluginConfigs(child, pp);
|
||||
pp->plugin_slot = plugin_slot;
|
||||
pp->nstate = -1; // actual value to be filled in by the plugin later
|
||||
}
|
||||
child = child->NextSiblingElement();
|
||||
}
|
||||
}
|
||||
|
||||
// advance to next element
|
||||
elem = elem->NextSiblingElement();
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
|
||||
// custom section parser
|
||||
void mjXReader::Custom(XMLElement* section) {
|
||||
string text, name;
|
||||
@@ -2800,6 +2928,38 @@ void mjXReader::Sensor(XMLElement* section) {
|
||||
psen->datatype = (mjtDataType)n;
|
||||
}
|
||||
|
||||
else if (type=="plugin") {
|
||||
psen->type = mjSENS_PLUGIN;
|
||||
ReadAttrTxt(elem, "plugin", psen->plugin_name);
|
||||
ReadAttrTxt(elem, "instance", psen->plugin_instance_name);
|
||||
if (psen->plugin_instance_name.empty()) {
|
||||
psen->plugin_instance = model->AddPlugin();
|
||||
} else {
|
||||
model->hasImplicitPluginElem = true;
|
||||
}
|
||||
ReadPluginConfigs(elem, psen->plugin_instance);
|
||||
ReadAttrTxt(elem, "objtype", text);
|
||||
psen->objtype = (mjtObj)mju_str2Type(text.c_str());
|
||||
ReadAttrTxt(elem, "objname", psen->objname);
|
||||
if (psen->objtype != mjOBJ_UNKNOWN && psen->objname.empty()) {
|
||||
throw mjXError(elem, "objtype is specified but objname is not");
|
||||
}
|
||||
if (psen->objtype == mjOBJ_UNKNOWN && !psen->objname.empty()) {
|
||||
throw mjXError(elem, "objname is specified but objtype is not");
|
||||
}
|
||||
ReadAttrTxt(elem, "reftype", text);
|
||||
psen->reftype = (mjtObj)mju_str2Type(text.c_str());
|
||||
ReadAttrTxt(elem, "refname", psen->refname);
|
||||
if (psen->reftype != mjOBJ_UNKNOWN && psen->refname.empty()) {
|
||||
throw mjXError(elem, "reftype is specified but refname is not");
|
||||
}
|
||||
if (psen->reftype == mjOBJ_UNKNOWN && !psen->refname.empty()) {
|
||||
throw mjXError(elem, "refname is specified but reftype is not");
|
||||
}
|
||||
}
|
||||
|
||||
GetXMLPos(elem, psen);
|
||||
|
||||
// advance to next element
|
||||
elem = elem->NextSiblingElement();
|
||||
}
|
||||
|
||||
@@ -37,6 +37,7 @@ class mjXReader : public mjXBase {
|
||||
private:
|
||||
// XML section specific to MJCF
|
||||
void Default(tinyxml2::XMLElement* section, int parentid); // default section
|
||||
void Extension(tinyxml2::XMLElement* section); // extension section
|
||||
void Custom(tinyxml2::XMLElement* section); // custom section
|
||||
void Visual(tinyxml2::XMLElement* section); // visual section
|
||||
void Statistic(tinyxml2::XMLElement* section); // statistic section
|
||||
|
||||
@@ -18,10 +18,15 @@
|
||||
#include <cstddef>
|
||||
#include <cstdio>
|
||||
#include <string>
|
||||
#include <unordered_set>
|
||||
|
||||
#include <mujoco/mjmodel.h>
|
||||
#include <mujoco/mjplugin.h>
|
||||
#include "engine/engine_io.h"
|
||||
#include "engine/engine_plugin.h"
|
||||
#include "engine/engine_util_errmem.h"
|
||||
#include "engine/engine_util_misc.h"
|
||||
#include "user/user_objects.h"
|
||||
#include "user/user_util.h"
|
||||
#include "xml/xml_util.h"
|
||||
#include "tinyxml2.h"
|
||||
@@ -604,12 +609,38 @@ void mjXWriter::OneActuator(XMLElement* elem, mjCActuator* pact, mjCDef* def) {
|
||||
WriteAttr(elem, "lengthrange", 2, pact->lengthrange, def->actuator.lengthrange);
|
||||
WriteAttr(elem, "gear", 6, pact->gear, def->actuator.gear);
|
||||
WriteAttr(elem, "cranklength", 1, &pact->cranklength, &def->actuator.cranklength);
|
||||
WriteAttrKey(elem, "dyntype", dyn_map, dyn_sz, pact->dyntype, def->actuator.dyntype);
|
||||
WriteAttrKey(elem, "gaintype", gain_map, gain_sz, pact->gaintype, def->actuator.gaintype);
|
||||
WriteAttrKey(elem, "biastype", bias_map, bias_sz, pact->biastype, def->actuator.biastype);
|
||||
WriteAttr(elem, "dynprm", mjNDYN, pact->dynprm, def->actuator.dynprm);
|
||||
WriteAttr(elem, "gainprm", mjNGAIN, pact->gainprm, def->actuator.gainprm);
|
||||
WriteAttr(elem, "biasprm", mjNBIAS, pact->biasprm, def->actuator.biasprm);
|
||||
|
||||
// plugins: write config attributes
|
||||
if (pact->is_plugin) {
|
||||
if (!pact->plugin_instance_name.empty()) {
|
||||
WriteAttrTxt(elem, "instance", pact->plugin_instance_name);
|
||||
} else {
|
||||
WriteAttrTxt(elem, "plugin", pact->plugin_name);
|
||||
const mjpPlugin* plugin = mjp_getPluginAtSlot(
|
||||
pact->plugin_instance->plugin_slot);
|
||||
const char* c = &pact->plugin_instance->flattened_attributes[0];
|
||||
for (int i = 0; i < plugin->nattribute; ++i) {
|
||||
std::string value(c);
|
||||
if (!value.empty()) {
|
||||
XMLElement* config_elem = InsertEnd(elem, "config");
|
||||
WriteAttrTxt(config_elem, "key", plugin->attributes[i]);
|
||||
WriteAttrTxt(config_elem, "value", value);
|
||||
c += value.size();
|
||||
}
|
||||
++c;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// non-plugins: write actuator parameters
|
||||
else {
|
||||
WriteAttrKey(elem, "dyntype", dyn_map, dyn_sz, pact->dyntype, def->actuator.dyntype);
|
||||
WriteAttrKey(elem, "gaintype", gain_map, gain_sz, pact->gaintype, def->actuator.gaintype);
|
||||
WriteAttrKey(elem, "biastype", bias_map, bias_sz, pact->biastype, def->actuator.biastype);
|
||||
WriteAttr(elem, "dynprm", mjNDYN, pact->dynprm, def->actuator.dynprm);
|
||||
WriteAttr(elem, "gainprm", mjNGAIN, pact->gainprm, def->actuator.gainprm);
|
||||
WriteAttr(elem, "biasprm", mjNBIAS, pact->biasprm, def->actuator.biasprm);
|
||||
}
|
||||
|
||||
// userdata
|
||||
if (writingdefaults) {
|
||||
@@ -659,6 +690,7 @@ void mjXWriter::Write(FILE* fp) {
|
||||
writingdefaults = true;
|
||||
Default(root, model->defaults[0]);
|
||||
writingdefaults = false;
|
||||
Extension(root);
|
||||
Custom(root);
|
||||
Asset(root);
|
||||
Body(InsertEnd(root, "worldbody"), model->GetWorld());
|
||||
@@ -1020,6 +1052,69 @@ void mjXWriter::Default(XMLElement* root, mjCDef* def) {
|
||||
|
||||
|
||||
|
||||
// extension section
|
||||
void mjXWriter::Extension(XMLElement* root) {
|
||||
// skip section if there is no required plugin
|
||||
if (model->active_plugins.empty()) {
|
||||
return;
|
||||
}
|
||||
|
||||
// create section
|
||||
XMLElement* section = InsertEnd(root, "extension");
|
||||
|
||||
// keep track of plugins whose <required> section have been created
|
||||
std::unordered_set<const mjpPlugin*> seen_plugins;
|
||||
|
||||
// write all plugins
|
||||
const mjpPlugin* last_plugin = nullptr;
|
||||
XMLElement* required_elem = nullptr;
|
||||
for (int i = 0; i < model->plugins.size(); ++i) {
|
||||
mjCPlugin* pp = static_cast<mjCPlugin*>(model->GetObject(mjOBJ_PLUGIN, i));
|
||||
|
||||
if (pp->name.empty()) {
|
||||
// reached the first unnamed plugin instance, meaning that it was created through an
|
||||
// "implicit" plugin element, e.g. sensor or actuator
|
||||
break;
|
||||
}
|
||||
|
||||
// check if we need to open a new <required> section
|
||||
const mjpPlugin* plugin = mjp_getPluginAtSlot(pp->plugin_slot);
|
||||
if (plugin != last_plugin) {
|
||||
required_elem = InsertEnd(section, "required");
|
||||
WriteAttrTxt(required_elem, "plugin", plugin->name);
|
||||
seen_plugins.insert(plugin);
|
||||
last_plugin = plugin;
|
||||
}
|
||||
|
||||
// write instance element
|
||||
XMLElement* elem = InsertEnd(required_elem, "instance");
|
||||
WriteAttrTxt(elem, "name", pp->name);
|
||||
|
||||
// write plugin config attributes
|
||||
const char* c = &pp->flattened_attributes[0];
|
||||
for (int i = 0; i < plugin->nattribute; ++i) {
|
||||
std::string value(c);
|
||||
if (!value.empty()) {
|
||||
XMLElement* config_elem = InsertEnd(elem, "config");
|
||||
WriteAttrTxt(config_elem, "key", plugin->attributes[i]);
|
||||
WriteAttrTxt(config_elem, "value", value);
|
||||
c += value.size();
|
||||
}
|
||||
++c;
|
||||
}
|
||||
}
|
||||
|
||||
// write <required> elements for plugins without explicit instances
|
||||
for (const auto& [plugin, slot] : model->active_plugins) {
|
||||
if (seen_plugins.find(plugin) == seen_plugins.end()) {
|
||||
required_elem = InsertEnd(section, "required");
|
||||
WriteAttrTxt(required_elem, "plugin", plugin->name);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
|
||||
// custom section
|
||||
void mjXWriter::Custom(XMLElement* root) {
|
||||
XMLElement* elem;
|
||||
@@ -1400,7 +1495,12 @@ void mjXWriter::Actuator(XMLElement* root) {
|
||||
// write all actuators
|
||||
for (int i=0; i<num; i++) {
|
||||
mjCActuator* pact = (mjCActuator*)model->GetObject(mjOBJ_ACTUATOR, i);
|
||||
XMLElement* elem = InsertEnd(section, "general");
|
||||
XMLElement* elem;
|
||||
if (pact->is_plugin) {
|
||||
elem = InsertEnd(section, "plugin");
|
||||
} else {
|
||||
elem = InsertEnd(section, "general");
|
||||
}
|
||||
OneActuator(elem, pact, pact->def);
|
||||
}
|
||||
}
|
||||
@@ -1593,6 +1693,34 @@ void mjXWriter::Sensor(XMLElement* root) {
|
||||
elem = InsertEnd(section, "clock");
|
||||
break;
|
||||
|
||||
|
||||
// plugin-controlled sensor
|
||||
case mjSENS_PLUGIN:
|
||||
elem = InsertEnd(section, "plugin");
|
||||
if (psen->objtype != mjOBJ_UNKNOWN) {
|
||||
WriteAttrTxt(elem, "objtype", mju_type2Str(psen->objtype));
|
||||
WriteAttrTxt(elem, "objname", psen->objname);
|
||||
}
|
||||
if (!psen->plugin_instance_name.empty()) {
|
||||
WriteAttrTxt(elem, "instance", psen->plugin_instance_name);
|
||||
} else {
|
||||
WriteAttrTxt(elem, "plugin", psen->plugin_name);
|
||||
const mjpPlugin* plugin = mjp_getPluginAtSlot(
|
||||
psen->plugin_instance->plugin_slot);
|
||||
const char* c = &psen->plugin_instance->flattened_attributes[0];
|
||||
for (int i = 0; i < plugin->nattribute; ++i) {
|
||||
std::string value(c);
|
||||
if (!value.empty()) {
|
||||
XMLElement* config_elem = InsertEnd(elem, "config");
|
||||
WriteAttrTxt(config_elem, "key", plugin->attributes[i]);
|
||||
WriteAttrTxt(config_elem, "value", value);
|
||||
c += value.size();
|
||||
}
|
||||
++c;
|
||||
}
|
||||
}
|
||||
break;
|
||||
|
||||
// user-defined sensor
|
||||
case mjSENS_USER:
|
||||
elem = InsertEnd(section, "user");
|
||||
@@ -1610,11 +1738,13 @@ void mjXWriter::Sensor(XMLElement* root) {
|
||||
// write name, noise, userdata
|
||||
WriteAttrTxt(elem, "name", psen->name);
|
||||
WriteAttr(elem, "cutoff", 1, &psen->cutoff, &zero);
|
||||
WriteAttr(elem, "noise", 1, &psen->noise, &zero);
|
||||
if (psen->type != mjSENS_PLUGIN) {
|
||||
WriteAttr(elem, "noise", 1, &psen->noise, &zero);
|
||||
}
|
||||
WriteVector(elem, "user", psen->userdata);
|
||||
|
||||
// add reference if present
|
||||
if (psen->reftype > 0) {
|
||||
if (psen->reftype != mjOBJ_UNKNOWN) {
|
||||
WriteAttrTxt(elem, "reftype", mju_type2Str(psen->reftype));
|
||||
WriteAttrTxt(elem, "refname", psen->refname);
|
||||
}
|
||||
|
||||
@@ -37,6 +37,7 @@ class mjXWriter : public mjXBase {
|
||||
void Visual(tinyxml2::XMLElement* root); // visual section
|
||||
void Statistic(tinyxml2::XMLElement* root); // statistic section
|
||||
void Default(tinyxml2::XMLElement* root, mjCDef* def); // default section
|
||||
void Extension(tinyxml2::XMLElement* root); // extension section
|
||||
void Custom(tinyxml2::XMLElement* root); // custom section
|
||||
void Asset(tinyxml2::XMLElement* root); // asset section
|
||||
void Body(tinyxml2::XMLElement* elem, mjCBody* body); // body/world section
|
||||
|
||||
Reference in New Issue
Block a user