From 67b09f9afb459e732a099c076a5c602fb89472e2 Mon Sep 17 00:00:00 2001 From: Haroon Qureshi Date: Tue, 25 Nov 2025 02:37:06 -0800 Subject: [PATCH] Refactor simulation state history tracking into a SimHistory class. PiperOrigin-RevId: 836578736 Change-Id: I94a852b2d2a36cdd081371f25385dbb466dd8d4b --- src/experimental/studio/app.cc | 25 +++---- src/experimental/studio/app.h | 1 - src/experimental/toolbox/physics.cc | 82 ++++++--------------- src/experimental/toolbox/physics.h | 29 +++----- src/experimental/toolbox/sim_history.cc | 77 ++++++++++++++++++++ src/experimental/toolbox/sim_history.h | 94 +++++++++++++++++++++++++ 6 files changed, 211 insertions(+), 97 deletions(-) create mode 100644 src/experimental/toolbox/sim_history.cc create mode 100644 src/experimental/toolbox/sim_history.h diff --git a/src/experimental/studio/app.cc b/src/experimental/studio/app.cc index 5023f7da..e32476b7 100644 --- a/src/experimental/studio/app.cc +++ b/src/experimental/studio/app.cc @@ -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(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"); diff --git a/src/experimental/studio/app.h b/src/experimental/studio/app.h index 88eb0ff4..fbdd9231 100644 --- a/src/experimental/studio/app.h +++ b/src/experimental/studio/app.h @@ -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; diff --git a/src/experimental/toolbox/physics.cc b/src/experimental/toolbox/physics.cc index abe5f803..8210f9af 100644 --- a/src/experimental/toolbox/physics.cc +++ b/src/experimental/toolbox/physics.cc @@ -14,14 +14,12 @@ #include "experimental/toolbox/physics.h" -#include #include -#include #include #include +#include #include #include -#include #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 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 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& 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(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 diff --git a/src/experimental/toolbox/physics.h b/src/experimental/toolbox/physics.h index 9675f804..948df74b 100644 --- a/src/experimental/toolbox/physics.h +++ b/src/experimental/toolbox/physics.h @@ -19,9 +19,9 @@ #include #include #include -#include #include +#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> history_; - int history_cursor_ = 0; OnModelLoadedFn on_model_loaded_; - int steps_ = 0; std::optional pending_load_; const mjVFS* vfs_; std::string error_; StepControl step_control_; + SimHistory sim_history_; }; } // namespace mujoco::toolbox diff --git a/src/experimental/toolbox/sim_history.cc b/src/experimental/toolbox/sim_history.cc new file mode 100644 index 00000000..09d613bc --- /dev/null +++ b/src/experimental/toolbox/sim_history.cc @@ -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 +#include +#include +#include + +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 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 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 SimHistory::SetIndex(int offset) { + const int size = history_.size(); + if (size > 0) { + offset_ = std::clamp(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 diff --git a/src/experimental/toolbox/sim_history.h b/src/experimental/toolbox/sim_history.h new file mode 100644 index 00000000..a9ea8dd7 --- /dev/null +++ b/src/experimental/toolbox/sim_history.h @@ -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 +#include +#include + +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; + + // 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 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 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 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_