From d6c605e8c5761c2b3d1c57b28b0a5c3e12dc18ca Mon Sep 17 00:00:00 2001 From: "Duhyeon, Kim" <49020301+dudududukim@users.noreply.github.com> Date: Thu, 23 Jul 2026 21:37:57 +0900 Subject: [PATCH] refactor(pi05): remove unused variables in embed_suffix method (#3263) * refactor(pi05): remove unused variables in embed_suffix method * Refactor embed_suffix to streamline pad_masks handling Removed unused pad_masks list and simplified its creation. Signed-off-by: Duhyeon, Kim <49020301+dudududukim@users.noreply.github.com> --------- Signed-off-by: Duhyeon, Kim <49020301+dudududukim@users.noreply.github.com> Co-authored-by: Steven Palma --- src/lerobot/policies/pi05/modeling_pi05.py | 16 ++++------------ 1 file changed, 4 insertions(+), 12 deletions(-) diff --git a/src/lerobot/policies/pi05/modeling_pi05.py b/src/lerobot/policies/pi05/modeling_pi05.py index 33896a9fa..d45f5a5c2 100644 --- a/src/lerobot/policies/pi05/modeling_pi05.py +++ b/src/lerobot/policies/pi05/modeling_pi05.py @@ -524,8 +524,6 @@ class PI05Pytorch(nn.Module): # see openpi `PI0Pytorch` def embed_suffix(self, noisy_actions, timestep): """Embed noisy_actions, timestep to prepare for Expert Gemma processing.""" - embs = [] - pad_masks = [] att_masks = [] # Embed timestep using sine-cosine positional encoding @@ -551,23 +549,17 @@ class PI05Pytorch(nn.Module): # see openpi `PI0Pytorch` return F.silu(x) time_emb = self._apply_checkpoint(time_mlp_func, time_emb) - action_time_emb = action_emb adarms_cond = time_emb - embs.append(action_time_emb) - bsize, action_time_dim = action_time_emb.shape[:2] - action_time_mask = torch.ones(bsize, action_time_dim, dtype=torch.bool, device=timestep.device) - pad_masks.append(action_time_mask) + bsize, action_time_dim = action_emb.shape[:2] + pad_masks = torch.ones(bsize, action_time_dim, dtype=torch.bool, device=timestep.device) # Set attention masks so that image, language and state inputs do not attend to action tokens att_masks += [1] + ([0] * (self.config.chunk_size - 1)) - - embs = torch.cat(embs, dim=1) - pad_masks = torch.cat(pad_masks, dim=1) - att_masks = torch.tensor(att_masks, dtype=embs.dtype, device=embs.device) + att_masks = torch.tensor(att_masks, dtype=action_emb.dtype, device=action_emb.device) att_masks = att_masks[None, :].expand(bsize, len(att_masks)) - return embs, pad_masks, att_masks, adarms_cond + return action_emb, pad_masks, att_masks, adarms_cond def forward(self, images, img_masks, tokens, masks, actions, noise, time) -> Tensor: """Do a full training forward pass and compute the loss."""