Allow multiple actuators to use the same PID plugin.

Maintain separate state for each actuator in the Plugin object.

PiperOrigin-RevId: 587760569
Change-Id: I633a1b94f0dde9377120145066a26aef6a4033b8
This commit is contained in:
Nimrod Gileadi
2023-12-04 10:07:20 -08:00
committed by Copybara-Service
parent 236dcb4ffd
commit c062a8db19
3 changed files with 122 additions and 62 deletions
+55 -55
View File
@@ -107,46 +107,39 @@ std::unique_ptr<Pid> Pid::Create(const mjModel* m, int instance) {
return nullptr;
}
int actuator_idx = -1;
std::vector<int> actuators;
for (int i = 0; i < m->nu; i++) {
if (m->actuator_plugin[i] == instance) {
if (actuator_idx != -1) {
mju_warning("multiple actuators found for plugin instance %d",
instance);
return nullptr;
}
actuator_idx = i;
actuators.push_back(i);
}
}
if (actuator_idx == -1) {
if (actuators.empty()) {
mju_warning("actuator not found for plugin instance %d", instance);
return nullptr;
}
return std::unique_ptr<Pid>(new Pid(config, actuator_idx));
return std::unique_ptr<Pid>(new Pid(config, std::move(actuators)));
}
void Pid::Reset(mjtNum* plugin_state) {
integral_ = 0.0;
previous_ctrl_ = 0.0;
}
void Pid::Reset(mjtNum* plugin_state) {}
mjtNum Pid::GetCtrl(const mjModel* m, const mjData* d, const State& state,
mjtNum Pid::GetCtrl(const mjModel* m, const mjData* d, int actuator_idx,
const State& state,
bool actearly) const {
mjtNum ctrl = 0;
if (m->actuator_dyntype[actuator_idx_] == mjDYN_NONE) {
ctrl = d->ctrl[actuator_idx_];
if (m->actuator_dyntype[actuator_idx] == mjDYN_NONE) {
ctrl = d->ctrl[actuator_idx];
// clamp ctrl
if (m->actuator_ctrllimited[actuator_idx_]) {
ctrl = mju_clip(ctrl, m->actuator_ctrlrange[2 * actuator_idx_],
m->actuator_ctrlrange[2 * actuator_idx_ + 1]);
if (m->actuator_ctrllimited[actuator_idx]) {
ctrl = mju_clip(ctrl, m->actuator_ctrlrange[2 * actuator_idx],
m->actuator_ctrlrange[2 * actuator_idx + 1]);
}
} else {
// Use of act instead of ctrl, to create integrated-velocity controllers or
// to filter the controls.
int actadr = m->actuator_actadr[actuator_idx_] +
m->actuator_actnum[actuator_idx_] - 1;
int actadr = m->actuator_actadr[actuator_idx] +
m->actuator_actnum[actuator_idx] - 1;
if (actearly) {
ctrl = NextActivation(m, d, actuator_idx_, actadr, d->act_dot[actadr]);
ctrl = NextActivation(m, d, actuator_idx, actadr, d->act_dot[actadr]);
} else {
ctrl = d->act[actadr];
}
@@ -160,48 +153,55 @@ mjtNum Pid::GetCtrl(const mjModel* m, const mjData* d, const State& state,
}
void Pid::ActDot(const mjModel* m, mjData* d, int instance) const {
State state = GetState(m, d, actuator_idx_);
mjtNum ctrl = GetCtrl(m, d, state, /*actearly=*/false);
mjtNum error = ctrl - d->actuator_length[actuator_idx_];
for (int actuator_idx : actuators_) {
State state = GetState(m, d, actuator_idx);
mjtNum ctrl = GetCtrl(m, d, actuator_idx, state, /*actearly=*/false);
mjtNum error = ctrl - d->actuator_length[actuator_idx];
int state_idx = m->actuator_actadr[actuator_idx_];
if (config_.i_gain) {
mjtNum integral = state.integral + error * m->opt.timestep;
if (config_.i_max.has_value()) {
integral = mju_clip(integral, -*config_.i_max, *config_.i_max);
int state_idx = m->actuator_actadr[actuator_idx];
if (config_.i_gain) {
mjtNum integral = state.integral + error * m->opt.timestep;
if (config_.i_max.has_value()) {
integral = mju_clip(integral, -*config_.i_max, *config_.i_max);
}
d->act_dot[state_idx] = (integral - d->act[state_idx]) / m->opt.timestep;
++state_idx;
}
if (config_.slew_max.has_value()) {
d->act_dot[state_idx] = (ctrl - d->act[state_idx]) / m->opt.timestep;
++state_idx;
}
d->act_dot[state_idx] = (integral - d->act[state_idx]) / m->opt.timestep;
++state_idx;
}
if (config_.slew_max.has_value()) {
d->act_dot[state_idx] = (ctrl - d->act[state_idx]) / m->opt.timestep;
++state_idx;
}
}
void Pid::Compute(const mjModel* m, mjData* d, int instance) {
State state = GetState(m, d, actuator_idx_);
mjtNum ctrl = GetCtrl(m, d, state, m->actuator_actearly[actuator_idx_]);
for (int i = 0; i < actuators_.size(); i++) {
int actuator_idx = actuators_[i];
State state = GetState(m, d, actuator_idx);
mjtNum ctrl =
GetCtrl(m, d, actuator_idx, state, m->actuator_actearly[actuator_idx]);
mjtNum error = ctrl - d->actuator_length[actuator_idx_];
mjtNum error = ctrl - d->actuator_length[actuator_idx];
mjtNum ctrl_dot = m->actuator_dyntype[actuator_idx_] == mjDYN_NONE
? 0
: d->act_dot[m->actuator_actadr[actuator_idx_] +
m->actuator_actnum[actuator_idx_] - 1];
mjtNum error_dot = ctrl_dot - d->actuator_velocity[actuator_idx_];
mjtNum ctrl_dot = m->actuator_dyntype[actuator_idx] == mjDYN_NONE
? 0
: d->act_dot[m->actuator_actadr[actuator_idx] +
m->actuator_actnum[actuator_idx] - 1];
mjtNum error_dot = ctrl_dot - d->actuator_velocity[actuator_idx];
if (config_.i_gain) {
integral_ = state.integral + error * m->opt.timestep;
if (config_.i_max.has_value()) {
integral_ = mju_clip(integral_, -*config_.i_max, *config_.i_max);
mjtNum integral = 0;
if (config_.i_gain) {
integral = state.integral + error * m->opt.timestep;
if (config_.i_max.has_value()) {
integral =
mju_clip(integral, -*config_.i_max, *config_.i_max);
}
}
}
d->actuator_force[actuator_idx_] = config_.p_gain * error +
config_.d_gain * error_dot +
config_.i_gain * integral_;
previous_ctrl_ = ctrl;
d->actuator_force[actuator_idx] = config_.p_gain * error +
config_.d_gain * error_dot +
config_.i_gain * integral;
}
}
void Pid::Advance(const mjModel* m, mjData* d, int instance) const {
@@ -279,7 +279,7 @@ void Pid::RegisterPlugin() {
mjp_registerPlugin(&plugin);
}
Pid::Pid(PidConfig config, int actuator_idx)
: config_(std::move(config)), actuator_idx_(actuator_idx) {}
Pid::Pid(PidConfig config, std::vector<int> actuators)
: config_(std::move(config)), actuators_(std::move(actuators)) {}
} // namespace mujoco::plugin::actuator
+6 -7
View File
@@ -17,6 +17,7 @@
#include <memory>
#include <optional>
#include <vector>
#include <mujoco/mjdata.h>
#include <mujoco/mjmodel.h>
@@ -73,7 +74,7 @@ class Pid {
static void RegisterPlugin();
private:
Pid(PidConfig config, int actuator_idx);
Pid(PidConfig config, std::vector<int> actuators);
struct State {
mjtNum previous_ctrl = 0;
@@ -88,14 +89,12 @@ class Pid {
// Returns the PID setpoint, which is normally d->ctrl, but can be d->act for
// actuators with dyntype != none.
mjtNum GetCtrl(const mjModel* m, const mjData* d, const State& state,
bool actearly) const;
mjtNum GetCtrl(const mjModel* m, const mjData* d, int actuator_idx,
const State& state, bool actearly) const;
PidConfig config_;
int actuator_idx_ = 0;
mjtNum integral_ = 0.0;
mjtNum previous_ctrl_ = 0.0;
// set of actuator IDs controlled by this plugin instance.
std::vector<int> actuators_;
};
} // namespace mujoco::plugin::actuator
+61
View File
@@ -562,6 +562,67 @@ TEST_F(PidTest, CopyData) {
EXPECT_EQ(d1->qpos[1], d2->qpos[1]);
}
TEST_F(PidTest, MultipleActuatorsSamePlugin) {
constexpr absl::string_view kModelXml = R"(
<mujoco>
<extension>
<plugin plugin="mujoco.pid">
<instance name="pid">
<config key="kp" value="4.0"/>
<config key="slewmax" value="0.75"/>
</instance>
</plugin>
</extension>
<option gravity="0 0 0" timestep="0.001"/>
<worldbody>
<body>
<joint name="j1" type="slide"/>
<geom size="0.01"/>
</body>
<body pos="0.04 0 0" >
<joint name="j2" type="slide"/>
<geom size="0.01"/>
</body>
</worldbody>
<actuator>
<plugin joint="j1" plugin="mujoco.pid" instance="pid" />
<plugin joint="j2" plugin="mujoco.pid" instance="pid" />
</actuator>
</mujoco>
)";
char error[1024] = {0};
mjModel* m = LoadModelFromString(kModelXml, error, sizeof(error));
ASSERT_THAT(m, NotNull()) << error;
absl::Cleanup m_deleter = [m] { mj_deleteModel(m); };
// having a slew rate means that there should be one extra state variable
// for the plugin.
EXPECT_EQ(m->actuator_actnum[0], 1);
EXPECT_EQ(m->actuator_actnum[1], 1);
mjData* d = mj_makeData(m);
absl::Cleanup d_deleter = [d] { mj_deleteData(d); };
// Set different ctrls for the two actuators, and check that they're
// independent.
d->ctrl[0] = 1.0;
d->ctrl[1] = -1.0;
for (int i = 0; i < 2; i++) {
mj_step(m, d);
EXPECT_EQ(d->actuator_force[0], -d->actuator_force[1])
<< "actuator_force mismatch at step " << i;
EXPECT_EQ(d->qfrc_actuator[0], -d->qfrc_actuator[1])
<< "qfrc_actuator mismatch at step " << i;
EXPECT_EQ(d->qpos[0], -d->qpos[1]) << "qpos mismatch at step " << i;
}
}
TEST_F(PidTest, InvalidClamp) {
constexpr absl::string_view kModelXml = R"(
<mujoco>