From 29972349d21f809d3dff7163abe1cf85537f262f Mon Sep 17 00:00:00 2001 From: Octi Zhang Date: Sat, 26 Sep 2026 17:41:43 -0700 Subject: [PATCH] Sample keyboard reset-buffer batches across all environments --- .../keyboard-reset-buffer-coverage.rst | 5 +++++ .../keyboard/mdp/commands/typing_commands.py | 15 ++++++++------- 2 files changed, 13 insertions(+), 7 deletions(-) create mode 100644 source/isaaclab_tasks/changelog.d/keyboard-reset-buffer-coverage.rst diff --git a/source/isaaclab_tasks/changelog.d/keyboard-reset-buffer-coverage.rst b/source/isaaclab_tasks/changelog.d/keyboard-reset-buffer-coverage.rst new file mode 100644 index 000000000000..837f911fb7ee --- /dev/null +++ b/source/isaaclab_tasks/changelog.d/keyboard-reset-buffer-coverage.rst @@ -0,0 +1,5 @@ +Fixed +^^^^^ + +* Fixed keyboard reset-buffer partial batches to sample distinct environments across the full scene + instead of favoring its first clone variants. Buffer capacity was unchanged. diff --git a/source/isaaclab_tasks/isaaclab_tasks/contrib/keyboard/mdp/commands/typing_commands.py b/source/isaaclab_tasks/isaaclab_tasks/contrib/keyboard/mdp/commands/typing_commands.py index 6f34273918da..ac61b6d9dcaa 100644 --- a/source/isaaclab_tasks/isaaclab_tasks/contrib/keyboard/mdp/commands/typing_commands.py +++ b/source/isaaclab_tasks/isaaclab_tasks/contrib/keyboard/mdp/commands/typing_commands.py @@ -531,13 +531,14 @@ def _build_buffer(self): with tqdm(total=cap, desc="[typing] building reset-curriculum buffer (IK)", unit="snap") as pbar: for start in range(0, cap, self.num_envs): n = min(self.num_envs, cap - start) - ids = all_ids[:n] + # Partial batches must not favor the first clone variants when envs are grouped. + ids = all_ids if n == self.num_envs else torch.randperm(self.num_envs, device=self.device)[:n] # Load this batch's cached command into the live state so the reset-IK aims at the right key. - self.target[:n] = tgt[start : start + n] - self.typed[:n] = typd[start : start + n] - self.target_len[:n] = tlen[start : start + n] - self.typed_len[:n] = typlen[start : start + n] - self.prefix_len[:n] = self._prefix_len()[:n] + self.target[ids] = tgt[start : start + n] + self.typed[ids] = typd[start : start + n] + self.target_len[ids] = tlen[start : start + n] + self.typed_len[ids] = typlen[start : start + n] + self.prefix_len[ids] = self._prefix_len()[ids] self._solve_reset_pose(ids) state = get_reset_state(self._env, ids, self._cur_reset_assets, is_relative=True) if self._buf_state is None: @@ -548,7 +549,7 @@ def _build_buffer(self): ee_quat = self.robot.data.body_quat_w.torch[:, self._ik_body_idx] tip = ee_pos + quat_apply(ee_quat, self._ik_offset) reach = torch.linalg.norm(tip - (self.target_key_pos_w() + self._ik_hover), dim=-1) - self._buf_reach[start : start + n] = reach[:n] + self._buf_reach[start : start + n] = reach[ids] pbar.update(n) self._buffer_built = True self._log_buffer_stats()