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."""