Refactor simulation state history tracking into a SimHistory class.

PiperOrigin-RevId: 836578736
Change-Id: I94a852b2d2a36cdd081371f25385dbb466dd8d4b
This commit is contained in:
Haroon Qureshi
2025-11-25 02:37:06 -08:00
committed by Copybara-Service
parent 6d9d59cb98
commit 67b09f9afb
6 changed files with 211 additions and 97 deletions
+11 -14
View File
@@ -446,14 +446,14 @@ void App::HandleKeyboardEvents() {
SetSpeedIndex(tmp_.speed_index - 1);
} else if (ImGui_IsChordJustPressed(ImGuiKey_LeftArrow)) {
if (physics_->GetStepControl().IsPaused()) {
physics_->LoadHistory(ui_.scrub_idx - 1);
physics_->LoadHistory(physics_->GetHistoryIndex() - 1);
}
} else if (ImGui_IsChordJustPressed(ImGuiKey_RightArrow)) {
if (physics_->GetStepControl().IsPaused()) {
if (ui_.scrub_idx == 0) {
if (physics_->GetHistoryIndex() == 0) {
physics_->GetStepControl().RequestSingleStep();
} else {
physics_->LoadHistory(ui_.scrub_idx + 1);
physics_->LoadHistory(physics_->GetHistoryIndex() + 1);
}
}
} else if (ImGui_IsChordJustPressed(ImGuiKey_Space)) {
@@ -1258,43 +1258,40 @@ void App::StatusBarGui() {
ImGui::Text("%s", " |");
// Frame scrubber.
const int max_history =
std::min<int>(physics_->GetStepCount(), physics_->GetHistorySize());
toolbox::ScopedStyle style;
style.Var(ImGuiStyleVar_FrameBorderSize, 0);
style.Color(ImGuiCol_Button, ImGui::GetStyle().Colors[ImGuiCol_WindowBg]);
ImGui::SameLine();
if (ImGui::Button(ICON_PREV_FRAME)) {
ui_.scrub_idx = std::max(-max_history, ui_.scrub_idx - 1);
physics_->LoadHistory(ui_.scrub_idx);
physics_->LoadHistory(physics_->GetHistoryIndex() - 1);
}
ImGui::SetItemTooltip("%s", "Previous Frame");
style.Reset();
ImGui::SameLine();
ImGui::SetNextItemWidth(450);
if (ImGui::SliderInt("##ScrubIndex", &ui_.scrub_idx, -max_history, 0)) {
physics_->LoadHistory(ui_.scrub_idx);
int index = physics_->GetHistoryIndex();
int history_size = physics_->GetSimHistory().Size();
if (ImGui::SliderInt("##ScrubIndex", &index, 1 - history_size, 0)) {
physics_->LoadHistory(index);
}
style.Var(ImGuiStyleVar_FrameBorderSize, 0);
style.Color(ImGuiCol_Button, ImGui::GetStyle().Colors[ImGuiCol_WindowBg]);
ImGui::SameLine();
if (ImGui::Button(ICON_NEXT_FRAME)) {
if (ui_.scrub_idx == 0) {
if (physics_->GetHistoryIndex() == 0) {
physics_->GetStepControl().RequestSingleStep();
} else {
ui_.scrub_idx = std::min(0, ui_.scrub_idx + 1);
physics_->LoadHistory(ui_.scrub_idx);
physics_->LoadHistory(physics_->GetHistoryIndex() + 1);
}
}
ImGui::SetItemTooltip("%s", "Next Frame");
ImGui::SameLine();
if (ImGui::Button(ICON_CURR_FRAME)) {
ui_.scrub_idx = 0;
physics_->LoadHistory(ui_.scrub_idx);
physics_->LoadHistory(0);
}
ImGui::SetItemTooltip("%s", "Current Frame");
-1
View File
@@ -74,7 +74,6 @@ class App {
int watch_index = 0;
int camera_idx = toolbox::kTumbleCameraIdx;
int key_idx = 0;
int scrub_idx = 0;
Style style = kLight;
using Dict = std::unordered_map<std::string, std::string>;
+21 -61
View File
@@ -14,14 +14,12 @@
#include "experimental/toolbox/physics.h"
#include <algorithm>
#include <chrono>
#include <climits>
#include <functional>
#include <ratio>
#include <span>
#include <string>
#include <utility>
#include <vector>
#include "experimental/toolbox/helpers.h"
#include "experimental/toolbox/step_control.h"
@@ -55,6 +53,8 @@ bool Physics::ProcessPendingLoad() {
Clear();
step_control_.SetSpeed(100.f);
std::string model_file = std::move(pending_load_.value());
pending_load_.reset();
@@ -73,7 +73,8 @@ bool Physics::ProcessPendingLoad() {
on_model_loaded_(model_file);
InitHistory();
const int state_size = mj_stateSize(model_, mjSTATE_INTEGRATION);
sim_history_.Init(state_size);
return model_ && data_;
}
@@ -84,12 +85,6 @@ void Physics::Clear() {
data_ = nullptr;
mj_deleteModel(model_);
model_ = nullptr;
history_.clear();
history_cursor_ = 0;
steps_ = 0;
step_control_.SetSpeed(100.f);
error_ = "";
}
}
@@ -98,7 +93,6 @@ void Physics::Reset() {
mj_resetData(model_, data_);
mj_forward(model_, data_);
error_ = "";
history_cursor_ = 0;
}
bool Physics::Update() {
@@ -117,7 +111,12 @@ bool Physics::Update() {
StepControl::Status status = step_control_.Advance(model_, data_);
if (status == StepControl::Status::kOk) {
AddToHistory();
std::span<mjtNum> state = sim_history_.AddToHistory();
if (!state.empty()) {
mj_getState(model_, data_, state.data(), mjSTATE_INTEGRATION);
}
// If we are adding to the history we didn't have a divergence error
error_ = "";
} else if (status == StepControl::Status::kPaused) {
// do nothing
} else if (status == StepControl::Status::kAutoReset) {
@@ -143,59 +142,20 @@ bool Physics::UpdateState(mjtNum* state, unsigned int state_sig) {
return true;
}
void Physics::InitHistory() {
const int state_size = mj_stateSize(model_, mjSTATE_INTEGRATION);
void Physics::LoadHistory(int offset) {
std::span<mjtNum> state = sim_history_.SetIndex(offset);
if (!state.empty()) {
// Pause simulation when entering history mode.
step_control_.Pause();
// History buffer will be smaller of 2000 states or 100 MB.
constexpr int kMaxBytes = 1e8;
constexpr int kMaxHistory = 2000;
const int state_bytes = state_size * sizeof(mjtNum);
const int history_length = std::min(INT_MAX / state_bytes, kMaxHistory);
const int history_bytes = std::min(state_bytes * history_length, kMaxBytes);
const int num_history = history_bytes / state_bytes;
history_.resize(num_history);
for (std::vector<mjtNum>& state : history_) {
state.resize(state_size, 0);
}
history_cursor_ = 0;
}
void Physics::AddToHistory() {
if (!history_.empty()) {
mjtNum* state = history_[history_cursor_].data();
mj_getState(model_, data_, state, mjSTATE_INTEGRATION);
history_cursor_ = (history_cursor_ + 1) % history_.size();
steps_++;
// If we are adding to the history we didn't have a divergence error
error_ = "";
// Load the state into the data buffer.
mj_setState(model_, data_, state.data(), mjSTATE_INTEGRATION);
mj_forward(model_, data_);
}
}
int Physics::LoadHistory(int offset) {
// No history to load.
if (steps_ == 0) {
return 0;
}
// Pause simulation when entering history mode.
step_control_.Pause();
// Ensure the offset is within a valid range. It's a negative value since
// we will be going backwards from the "latest" frame.
const int max_history = std::min<int>(steps_, history_.size());
offset = std::clamp(offset, -max_history + 1, 0);
// Determine the index in the history buffer that corresponds to the frame
// index.
const int idx = (history_cursor_ + offset - 1) % history_.size();
const mjtNum* state = history_[idx].data();
// Load the state into the data buffer.
mj_setState(model_, data_, state, mjSTATE_INTEGRATION);
mj_forward(model_, data_);
return offset;
int Physics::GetHistoryIndex() const {
return sim_history_.GetIndex();
}
} // namespace mujoco::toolbox
+8 -21
View File
@@ -19,9 +19,9 @@
#include <optional>
#include <string>
#include <string_view>
#include <vector>
#include <mujoco/mujoco.h>
#include "experimental/toolbox/sim_history.h"
#include "experimental/toolbox/step_control.h"
namespace mujoco::toolbox {
@@ -38,16 +38,13 @@ class Physics {
Physics(const Physics&) = delete;
Physics& operator=(const Physics&) = delete;
// Access the step controller
StepControl& GetStepControl() { return step_control_; }
SimHistory& GetSimHistory() { return sim_history_; }
// Loads a model from the given path. An empty string will load an empty
// scene.
void LoadModel(std::string model_file, const mjVFS* vfs = nullptr);
// Clears the simulation, clearing all loaded state.
void Clear();
// Resets the simulation using mj_resetData
void Reset();
@@ -57,17 +54,9 @@ class Physics {
// Sets the state of the simulation.
bool UpdateState(mjtNum* state, unsigned int state_sig);
// Returns the number of steps the simulation has taken.
int GetStepCount() const { return steps_; }
// Returns the number of states in the history buffer.
int GetHistorySize() const { return history_.size(); }
// Loads a state from the history buffer at the given offset into the current
// physics state.
//
// Calling this function will automatically pause the simulation.
int LoadHistory(int offset);
// Loads a state from the history buffer at the given offset in the past.
void LoadHistory(int offset);
int GetHistoryIndex() const;
// Returns the MuJoCo data structures owned by this Simulation object.
mjModel* GetModel() { return model_; }
@@ -79,23 +68,21 @@ class Physics {
bool ProcessPendingLoad();
private:
void InitHistory();
void AddToHistory();
// Clears the simulation, clearing all loaded state.
void Clear();
mjModel* model_ = nullptr;
mjData* data_ = nullptr;
std::vector<std::vector<mjtNum>> history_;
int history_cursor_ = 0;
OnModelLoadedFn on_model_loaded_;
int steps_ = 0;
std::optional<std::string> pending_load_;
const mjVFS* vfs_;
std::string error_;
StepControl step_control_;
SimHistory sim_history_;
};
} // namespace mujoco::toolbox
+77
View File
@@ -0,0 +1,77 @@
// Copyright 2025 DeepMind Technologies Limited
//
// Licensed under the Apache License, Version 2.0 (the "License");
// you may not use this file except in compliance with the License.
// You may obtain a copy of the License at
//
// http://www.apache.org/licenses/LICENSE-2.0
//
// Unless required by applicable law or agreed to in writing, software
// distributed under the License is distributed on an "AS IS" BASIS,
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
// See the License for the specific language governing permissions and
// limitations under the License.
#include "experimental/toolbox/sim_history.h"
#include <algorithm>
#include <climits>
#include <span>
#include <mujoco/mujoco.h>
namespace mujoco::toolbox {
void SimHistory::Init(int state_size, int max_history, int max_bytes) {
// History buffer will be smaller of number of states and total memory.
const int state_bytes = state_size * sizeof(mjtNum);
const int history_length = std::min(INT_MAX / state_bytes, max_history);
const int history_bytes = std::min(state_bytes * history_length, max_bytes);
const int max_states = std::max(1, history_bytes / state_bytes);
history_.resize(max_states);
for (State& state : history_) {
state.resize(state_size, 0);
}
cursor_ = 0;
offset_ = 0;
size_ = 0;
}
std::span<mjtNum> SimHistory::AddToHistory() {
const int max_size = history_.size();
if (offset_ != 0) {
// offset will be a negative number between 1 - history_.size() and 0.
size_ += offset_;
cursor_ += offset_;
if (cursor_ < 0) {
cursor_ += max_size;
}
offset_ = 0;
}
std::span<mjtNum> state;
if (max_size > 0 && cursor_ < max_size) {
state = history_[cursor_];
cursor_ = (cursor_ + 1) % max_size;
size_ = std::min(size_ + 1, max_size);
}
return state;
}
std::span<mjtNum> SimHistory::SetIndex(int offset) {
const int size = history_.size();
if (size > 0) {
offset_ = std::clamp<int>(offset, 1 - size, 0);
}
if (history_.empty()) {
return {};
}
const int actual_index =
(cursor_ - 1 + offset_ + history_.size()) % history_.size();
return history_[actual_index];
}
} // namespace mujoco::toolbox
+94
View File
@@ -0,0 +1,94 @@
// Copyright 2025 DeepMind Technologies Limited
//
// Licensed under the Apache License, Version 2.0 (the "License");
// you may not use this file except in compliance with the License.
// You may obtain a copy of the License at
//
// http://www.apache.org/licenses/LICENSE-2.0
//
// Unless required by applicable law or agreed to in writing, software
// distributed under the License is distributed on an "AS IS" BASIS,
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
// See the License for the specific language governing permissions and
// limitations under the License.
#ifndef MUJOCO_SRC_EXPERIMENTAL_TOOLBOX_SIM_HISTORY_H_
#define MUJOCO_SRC_EXPERIMENTAL_TOOLBOX_SIM_HISTORY_H_
#include <span>
#include <vector>
#include <mujoco/mujoco.h>
namespace mujoco::toolbox {
// A historical buffer of simulation data.
class SimHistory {
public:
SimHistory() = default;
// Simulation data is stored a an array of mjtNum; see mj_getState().
using State = std::vector<mjtNum>;
// Clears and initializes the history buffer to store state.
void Init(int state_size, int max_history = 2000, int max_bytes = 1e8);
// Adds an uninitialized state to the history and returns a reference to it
// so that the caller can populate the data. Also resets the current index to
// 0; see SetIndex() for details.
std::span<mjtNum> AddToHistory();
// Returns the history at the given index (i.e. the number of steps) in the
// past. The `offset` will be clamped internally to the range [0, Size() - 1].
// This function returns the valid, clamped value.
//
// For example, calling SetIndex(0) will return the most recently recorded
// state. Calling SetIndex(-N) will return the state from N steps ago.
//
// Note that future calls to `AddToHistory` will begin recording from the
// newly set index, effectively creating a new "branch" of the history buffer.
// If you do not want to lose any states, you must call `SetIndex(0)` before
// before resuming playback.
//
// For example, consider you have recorded 6 states:
// N(-5) N(-4) N(-3) N(-2) N(-1) N(0)
// -----------------------------------------^
//
// You then call `SetIndex(-3)`:
// N(-5) N(-4) N(-3) N(-2) N(-1) N(0)
// ------------------^
//
// And then call `AddToHistory()
// N(-5) N(-4) N(-3) N(-2) N(-1) N(0)
// | | | x x x
// v v v
// N'(-3) N'(-2) N'(-1) N'(0)
// -------------------------^
//
// In this case, we effectively erase states 0 to -3 of the previous history,
// "copy" the older states into the new branch, and add the most recent state
// at the "head" of the history buffer.
std::span<mjtNum> SetIndex(int offset);
// Returns the currently set index.
int GetIndex() const { return offset_; }
// Returns the number of states in the history buffer.
int Size() const { return size_; }
private:
// The history of states.
std::vector<State> history_;
// The index at which the next AddToHistory() call will write.
int cursor_ = 0;
// The most recently requested offset from SetIndex().
int offset_ = 0;
// The total number of states available in the history buffer.
int size_ = 0;
};
} // namespace mujoco::toolbox
#endif // MUJOCO_SRC_EXPERIMENTAL_TOOLBOX_SIM_HISTORY_H_