Refactor simulation state history tracking into a SimHistory class.
PiperOrigin-RevId: 836578736 Change-Id: I94a852b2d2a36cdd081371f25385dbb466dd8d4b
This commit is contained in:
committed by
Copybara-Service
parent
6d9d59cb98
commit
67b09f9afb
@@ -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");
|
||||
|
||||
|
||||
@@ -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>;
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
@@ -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_
|
||||
Reference in New Issue
Block a user