From fdd9d92349c81b0e176fa53874aa30167cc0211c Mon Sep 17 00:00:00 2001 From: Steven Palma Date: Fri, 31 Jul 2026 16:53:37 +0200 Subject: [PATCH] feat(policies): early stopping + termination --- .../policies/pi0_fast/modeling_pi0_fast.py | 27 +++++++++++++------ 1 file changed, 19 insertions(+), 8 deletions(-) diff --git a/src/lerobot/policies/pi0_fast/modeling_pi0_fast.py b/src/lerobot/policies/pi0_fast/modeling_pi0_fast.py index 475563d7c..b84ba2ce3 100644 --- a/src/lerobot/policies/pi0_fast/modeling_pi0_fast.py +++ b/src/lerobot/policies/pi0_fast/modeling_pi0_fast.py @@ -604,6 +604,12 @@ class PI0FastPytorch(nn.Module): # see openpi `PI0Pytorch` ) -> torch.Tensor: """ Optimized autoregressive decoding for FAST tokens using KV Caching. + + Greedy decoding stops once every sequence emits the end-of-action marker. The + returned tensor keeps its fixed shape, with positions not generated after the + batch-wide stop left zero-filled. Stochastic decoding always runs to + ``max_decoding_steps`` so early stopping does not change the RNG state used by + subsequent calls. """ if max_decoding_steps is None: max_decoding_steps = self.config.max_action_tokens @@ -612,10 +618,11 @@ class PI0FastPytorch(nn.Module): # see openpi `PI0Pytorch` device = tokens.device lm_head = self.paligemma_with_expert.paligemma.lm_head - # detokenize_actions() cuts at the first "|", so anything generated past it is - # discarded anyway. - pipe_id = self._paligemma_tokenizer.convert_tokens_to_ids("|") - finished = torch.zeros(bsize, dtype=torch.bool, device=device) + # detokenize_actions() cuts at the first "|", so greedy decoding can stop once + # every sequence has emitted it. Keep stochastic decoding unchanged because + # skipping multinomial calls would shift the RNG state for subsequent calls. + end_of_action_token_id = self._paligemma_tokenizer.convert_tokens_to_ids("|") + finished = torch.zeros(bsize, dtype=torch.bool, device=device) if temperature == 0 else None # --- 1. PREFILL PHASE --- # Process Images + Text Prompt + BOS token once to populate the KV cache. @@ -668,7 +675,10 @@ class PI0FastPytorch(nn.Module): # see openpi `PI0Pytorch` # Initialize storage for generated tokens generated_action_tokens = torch.zeros((bsize, max_decoding_steps), dtype=torch.long, device=device) generated_action_tokens[:, 0] = next_token.squeeze(-1) - finished |= next_token.squeeze(-1) == pipe_id + if finished is not None: + finished |= next_token.squeeze(-1) == end_of_action_token_id + if bool(finished.all()): + return generated_action_tokens # Track valid tokens mask (0 for pad, 1 for valid) # We need this to tell the new token what it can attend to (images + text + past actions) @@ -719,9 +729,10 @@ class PI0FastPytorch(nn.Module): # see openpi `PI0Pytorch` generated_action_tokens[:, t] = next_token.squeeze(-1) - finished |= next_token.squeeze(-1) == pipe_id - if bool(finished.all()): - break + if finished is not None: + finished |= next_token.squeeze(-1) == end_of_action_token_id + if bool(finished.all()): + break return generated_action_tokens