From d91beb4b8128663e6e9e2bb99d8d24297326e28b Mon Sep 17 00:00:00 2001 From: xuxin <747302550@qq.com> Date: Mon, 23 Jun 2025 11:12:16 +0800 Subject: [PATCH 1/3] fix the copy error in obs buffer --- legged_gym/envs/base/observation_buffer.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/legged_gym/envs/base/observation_buffer.py b/legged_gym/envs/base/observation_buffer.py index de062478..a7d2856d 100644 --- a/legged_gym/envs/base/observation_buffer.py +++ b/legged_gym/envs/base/observation_buffer.py @@ -17,7 +17,7 @@ def reset(self, reset_idxs, new_obs): def insert(self, new_obs): # Shift observations back. - self.obs_buf[:, : self.num_obs * (self.include_history_steps - 1)] = self.obs_buf[:,self.num_obs : self.num_obs * self.include_history_steps] + self.obs_buf[:, : self.num_obs * (self.include_history_steps - 1)] = self.obs_buf[:,self.num_obs : self.num_obs * self.include_history_steps].clone() # Add new observation. self.obs_buf[:, -self.num_obs:] = new_obs From b0e119e7dae4f9dfa5ca10fe530fde527b4dc373 Mon Sep 17 00:00:00 2001 From: xuxin <747302550@qq.com> Date: Mon, 23 Jun 2025 11:12:51 +0800 Subject: [PATCH 2/3] add 3d and circular implementations for obs buffer --- legged_gym/envs/base/observation_buffer_3d.py | 37 ++++++++++++++++ .../base/observation_buffer_3d_circular.py | 43 +++++++++++++++++++ 2 files changed, 80 insertions(+) create mode 100644 legged_gym/envs/base/observation_buffer_3d.py create mode 100644 legged_gym/envs/base/observation_buffer_3d_circular.py diff --git a/legged_gym/envs/base/observation_buffer_3d.py b/legged_gym/envs/base/observation_buffer_3d.py new file mode 100644 index 00000000..28da0753 --- /dev/null +++ b/legged_gym/envs/base/observation_buffer_3d.py @@ -0,0 +1,37 @@ +# created by Xu Xin, 2025 +import torch + +class ObservationBuffer3D: + """using a 3d tensor for storage (num_envs, num_obs, history_steps) \\ + 0 = old obs -> end = new obs + """ + def __init__(self, num_envs, num_obs, history_length, device): + self.num_envs = num_envs + self.num_obs = num_obs + self.history_length = history_length + self.device = device + + self.obs_history = torch.zeros((self.num_envs, self.num_obs, self.history_length), + device=self.device, dtype=torch.float) + + def reset(self, reset_ids, new_obs): + self.obs_history[reset_ids,:,:] = new_obs.unsqueeze(-1).repeat(1, 1, self.history_length) + + def insert(self, new_obs): + self.obs_history[:, :, :-1] = self.obs_history[:, :, 1:].clone() + self.obs_history[:, :, -1].copy_(new_obs) + + def get_obs_vec(self, obs_ids): + """Gets history of observations indexed by obs_ids. + + Arguments: + obs_ids: An array of integers with which to index the desired + observations, where 0 is the latest observation and + include_history_steps - 1 is the oldest observation. + """ + obs = [] + for obs_id in reversed(sorted(obs_ids)): + slice_idx = self.history_length - obs_id - 1 + obs.append(self.obs_history[:, :, slice_idx]) + return torch.cat(obs, dim=-1) + diff --git a/legged_gym/envs/base/observation_buffer_3d_circular.py b/legged_gym/envs/base/observation_buffer_3d_circular.py new file mode 100644 index 00000000..a7f2494e --- /dev/null +++ b/legged_gym/envs/base/observation_buffer_3d_circular.py @@ -0,0 +1,43 @@ +# created by Xu Xin, 2025 +import isaacgym +import torch +import numpy as np +from observation_buffer_3d import ObservationBuffer3D +from observation_buffer import ObservationBuffer + +class ObservationBuffer3D_circular: + """using a 3d tensor for storage (num_envs, num_obs, history_steps) \\ + 0 = old obs -> end = new obs + """ + def __init__(self, num_envs, num_obs, history_length, device): + self.num_envs = num_envs + self.num_obs = num_obs + self.history_length = history_length + self.device = device + + self.obs_history = torch.zeros((self.num_envs, self.num_obs, self.history_length), + device=self.device, dtype=torch.float) + self.current_index = 0 # a common pointer for all envs + + def reset(self, reset_ids, new_obs): + self.obs_history[reset_ids, :, :] = new_obs.unsqueeze(-1).repeat(1, 1, self.history_length) + + def insert(self, new_obs): + # write the data to the pointer + self.obs_history[:, :, self.current_index] = new_obs + # update the pointer to a unwrittend place + self.current_index = (self.current_index + 1) % self.history_length + + def get_obs_vec(self, obs_ids): + sorted_obs_ids = sorted(obs_ids,reverse=True) + + # get the real index according to the pointer + indices = (self.current_index - 1 - torch.tensor(sorted_obs_ids, device=self.device)) % self.history_length + + expanded_indices = indices.expand(self.num_envs, self.num_obs, -1) + selected = torch.gather(self.obs_history, 2, expanded_indices) + + # (env, obs, time) -> (env, time, obs) + selected = selected.permute(0, 2, 1) + # reshape to (num_envs, num_obs * len(obs_ids)) + return selected.reshape(self.num_envs, -1) From 59d777d244cb44452dd2df7143970703dd07a69d Mon Sep 17 00:00:00 2001 From: xuxin <747302550@qq.com> Date: Mon, 23 Jun 2025 11:13:10 +0800 Subject: [PATCH 3/3] add a test script for obs buffer --- .../base/observation_buffer_3d_circular.py | 95 +++++++++++++++++++ 1 file changed, 95 insertions(+) diff --git a/legged_gym/envs/base/observation_buffer_3d_circular.py b/legged_gym/envs/base/observation_buffer_3d_circular.py index a7f2494e..f0c96490 100644 --- a/legged_gym/envs/base/observation_buffer_3d_circular.py +++ b/legged_gym/envs/base/observation_buffer_3d_circular.py @@ -41,3 +41,98 @@ def get_obs_vec(self, obs_ids): selected = selected.permute(0, 2, 1) # reshape to (num_envs, num_obs * len(obs_ids)) return selected.reshape(self.num_envs, -1) + +def test2(): + # speed test across three implementations + import time + num_envs = 4096 + num_obs = 12 + include_history_steps = 40 + device = torch.device("cuda") + large_buffer1 = ObservationBuffer(num_envs, num_obs, include_history_steps, device) + large_buffer2 = ObservationBuffer3D(num_envs, num_obs, include_history_steps, device) + large_buffer3 = ObservationBuffer3D_circular(num_envs, num_obs, include_history_steps, device) + + test_iter = 10000 + obs = torch.rand(num_envs, num_obs, test_iter, device=device) + + large_buffer1.insert(obs[...,0]) + start = time.time() + for it in range(test_iter): + large_buffer1.insert(obs[...,it]) + end = time.time() + print(f"{large_buffer1.__class__} insert {test_iter} using: {end-start:.4f} seconds") + + large_buffer2.insert(obs[...,0]) + start = time.time() + for it in range(test_iter): + large_buffer2.insert(obs[...,it]) + end = time.time() + print(f"{large_buffer2.__class__} insert {test_iter} using: {end-start:.4f} seconds") + + large_buffer3.insert(obs[...,0]) + start = time.time() + for it in range(test_iter): + large_buffer3.insert(obs[...,it]) + end = time.time() + print(f"{large_buffer3.__class__} insert {test_iter} using: {end-start:.4f} seconds") + + + r1 = large_buffer1.get_obs_vec(np.arange(include_history_steps)) + r2 = large_buffer2.get_obs_vec(np.arange(include_history_steps)) + r3 = large_buffer3.get_obs_vec(np.arange(include_history_steps)) + diff2 = r2 - r1 + diff3 = r3 - r1 + if (diff2!=0).sum() >0: + print(it, (diff2!=0).sum()) + print((diff2!=0).nonzero()) + return + if (diff3!=0).sum() >0: + print(it, (diff3!=0).sum()) + print((diff3!=0).nonzero()) + return + print(diff2.mean().item(),diff3.mean().item()) + +def test3(): + # larget batch error test + num_envs = 400 + num_obs = 12 + include_history_steps = 40 + device = torch.device("cuda") + large_buffer1 = ObservationBuffer(num_envs, num_obs, include_history_steps, device) + large_buffer2 = ObservationBuffer3D(num_envs, num_obs, include_history_steps, device) + large_buffer3 = ObservationBuffer3D_circular(num_envs, num_obs, include_history_steps, device) + + test_iter = 10000 + obs = torch.rand(num_envs, num_obs, test_iter, device=device) + obs_r = torch.zeros(num_envs, num_obs, device=device) + + for it in range(test_iter): + large_buffer1.insert(obs[...,it]) + large_buffer2.insert(obs[...,it]) + large_buffer3.insert(obs[...,it]) + r1 = large_buffer1.get_obs_vec(np.arange(include_history_steps)) + r2 = large_buffer2.get_obs_vec(np.arange(include_history_steps)) + r3 = large_buffer3.get_obs_vec(np.arange(include_history_steps)) + diff2 = r2 - r1 + diff3 = r3 - r1 + + if it % 100 == 0: + large_buffer1.reset(torch.arange(num_envs,device=device),obs_r) + large_buffer2.reset(torch.arange(num_envs,device=device),obs_r) + large_buffer3.reset(torch.arange(num_envs,device=device),obs_r) + + if (diff2!=0).sum() >0: + print(it, (diff2!=0).sum()) + print((diff2!=0).nonzero()) + return + + if (diff3!=0).sum() >0: + print(it, (diff3!=0).sum()) + print((diff3!=0).nonzero()) + return + print(f"larget batch error test: ok") + +if __name__ == "__main__": + test2() + test3() \ No newline at end of file