Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
36 changes: 30 additions & 6 deletions src/inference.rs
Original file line number Diff line number Diff line change
Expand Up @@ -613,6 +613,8 @@ impl State for Reset {
shared_state,
creation_time: std::time::Instant::now(),
session_input_vec: self.session_input_vec,
lpf_prev_qpos: enum_map::EnumMap::default(),
lpf_last_update: std::time::Instant::now(),
}),
result: Ok(()),
}
Expand All @@ -624,6 +626,9 @@ pub struct Operate {
shared_state: Pin<Box<Store>>,
creation_time: std::time::Instant,
session_input_vec: Vec<ort::session::SessionInputValue<'static>>,
// Per-actuator low-pass filter state for commanded qpos
lpf_prev_qpos: enum_map::EnumMap<crate::robot_description::ActuatorId, f64>,
lpf_last_update: std::time::Instant,
}

impl std::fmt::Debug for Operate {
Expand Down Expand Up @@ -744,19 +749,36 @@ impl Operate {
.map_err(std::io::Error::other)?;

let actuator_states = &mut robot_description.actuators.actuator_states;

let cutoff_hz = robot_description.lpf_cutoff_hz;
let prev_map = &mut self.lpf_prev_qpos;
let now = std::time::Instant::now();
let dt = now.duration_since(self.lpf_last_update).as_secs_f64();
self.lpf_last_update = now;

for (i, command) in commands.iter().enumerate() {
let actuator_id = cmd_idx_to_actuator_id[i];
let act_state = &mut actuator_states[actuator_id];
// get the normalized qpso
// get the normalized qpos
let normalized_qpos =
robot_description::normalize_actuator_qpos(act_state.feedback.qpos);
let err = *command as f64 - normalized_qpos;
let final_command = act_state.feedback.qpos + err * robot_description.policy_scale;
// TODO: action scale
// act_state.command.qpos = *command as f64 * robot_description.policy_scale;
let unfiltered = act_state.feedback.qpos + err * robot_description.policy_scale;

// One-pole LPF: y = y_prev + alpha * (x - y_prev); alpha = 1 - exp(-2*pi*fc*dt)
let filtered = if cutoff_hz <= 0.0 || dt <= 0.0 {
unfiltered
} else {
let alpha = 1.0 - (-2.0 * std::f64::consts::PI * cutoff_hz * dt).exp();
let y_prev = prev_map[actuator_id];
let y = y_prev + alpha * (unfiltered - y_prev);
prev_map[actuator_id] = y;
y
};
let final_command = filtered;
act_state.command.qpos = final_command;
act_state.command.qvel = 0.0; // no velocity
act_state.command.qfrc = 0.0; // no force
act_state.command.qvel = 0.0; // no velocity command
act_state.command.qfrc = 0.0; // no torque command
act_state.command.kp =
robot_description.policy_position[actuator_id].kp * robot_description.kp_scale;
act_state.command.kd =
Expand Down Expand Up @@ -876,6 +898,8 @@ impl State for Operate {
shared_state,
creation_time: self.creation_time,
session_input_vec: self.session_input_vec,
lpf_prev_qpos: self.lpf_prev_qpos,
lpf_last_update: self.lpf_last_update,
}),
result: Ok(()),
}
Expand Down
4 changes: 4 additions & 0 deletions src/main.rs
Original file line number Diff line number Diff line change
Expand Up @@ -61,6 +61,10 @@ pub struct Args {
/// derivative gain scale
#[arg(long, value_name = "FLOAT", default_value_t = 1.0)]
kd_scale: f64,

/// low-pass filter cutoff in Hz for policy outputs (0 disables filtering)
#[arg(long, value_name = "FLOAT", default_value_t = 6.0)]
lpf_cutoff_hz: f64,
}

async fn driver() -> std::io::Result<()> {
Expand Down
6 changes: 6 additions & 0 deletions src/robot_description.rs
Original file line number Diff line number Diff line change
Expand Up @@ -385,6 +385,10 @@ pub struct Args {
/// derivative gain scale
#[arg(long, value_name = "FLOAT", default_value_t = 1.0)]
kd_scale: f64,

/// low-pass filter cutoff in Hz for policy outputs (0 disables filtering)
#[arg(long, value_name = "FLOAT", default_value_t = 6.0)]
lpf_cutoff_hz: f64,
}

pub struct RobotDescription {
Expand All @@ -398,6 +402,7 @@ pub struct RobotDescription {
pub policy_scale: f64,
pub kp_scale: f64,
pub kd_scale: f64,
pub lpf_cutoff_hz: f64,
}

impl Default for RobotDescription {
Expand All @@ -419,6 +424,7 @@ impl RobotDescription {
kp_scale: args.kp_scale,
kd_scale: args.kd_scale,
policy_scale: args.policy_scale,
lpf_cutoff_hz: args.lpf_cutoff_hz,
home_position: enum_map! {
ActuatorId::Lsp => ActuatorCommand { qpos: 0.0, kp: 100.0, kd: 8.284, ..Default::default() },
ActuatorId::Lsr => ActuatorCommand { qpos: (10.0_f64).to_radians(), kp: 100.0, kd: 8.257, ..Default::default() },
Expand Down