From 39cdc4dea3d3e7a64090ca4e0d28c531ee307b19 Mon Sep 17 00:00:00 2001 From: Khalil Meftah Date: Fri, 24 Jul 2026 09:24:39 +0200 Subject: [PATCH] fix: nan in eval --- ...modeling_temporal_siglip_value_function.py | 35 +++++++++++++++---- .../test_temporal_siglip_value_function.py | 6 ++++ 2 files changed, 34 insertions(+), 7 deletions(-) diff --git a/src/lerobot/rewards/temporal_siglip_value_function/modeling_temporal_siglip_value_function.py b/src/lerobot/rewards/temporal_siglip_value_function/modeling_temporal_siglip_value_function.py index a605658e7..986c46efa 100644 --- a/src/lerobot/rewards/temporal_siglip_value_function/modeling_temporal_siglip_value_function.py +++ b/src/lerobot/rewards/temporal_siglip_value_function/modeling_temporal_siglip_value_function.py @@ -131,21 +131,42 @@ class TemporalSiglipVFRewardModel(DistributionalValueMixin, PreTrainedRewardMode frame_tokens + self.time_embedding(torch.arange(history_steps, device=frame_tokens.device))[None] ) - causal_mask = torch.triu( - torch.ones(history_steps, history_steps, dtype=torch.bool, device=frame_tokens.device), - diagonal=1, - ) frame_valid = torch.stack(masks).any(0) + attention_mask = self._make_temporal_attention_mask(frame_valid) hidden = self.temporal_transformer( frame_tokens, - mask=causal_mask, - src_key_padding_mask=~frame_valid, - is_causal=True, + mask=attention_mask, ) # The history window is ordered oldest→current and the current frame is # always the final, non-padding element. return hidden[:, -1] + def _make_temporal_attention_mask(self, frame_valid: Tensor) -> Tensor: + """Combine causal and padding masks without fully masked padded queries. + + A left-padded causal query has no valid past keys. PyTorch's optimized + eval path returns NaNs for such rows, which then contaminate later valid + tokens. Padded queries attend only to themselves; valid queries retain + causal attention and cannot attend to padded keys. + """ + batch_size, history_steps = frame_valid.shape + causal_mask = torch.triu( + torch.ones(history_steps, history_steps, dtype=torch.bool, device=frame_valid.device), + diagonal=1, + ) + attention_mask = causal_mask[None].expand(batch_size, -1, -1) | (~frame_valid)[:, None, :] + padded_queries = (~frame_valid).nonzero(as_tuple=False) + attention_mask[ + padded_queries[:, 0], + padded_queries[:, 1], + padded_queries[:, 1], + ] = False + return ( + attention_mask[:, None] + .expand(-1, self.config.num_heads, -1, -1) + .reshape(batch_size * self.config.num_heads, history_steps, history_steps) + ) + def _fit_state_dim(self, state: Tensor) -> Tensor: if state.shape[-1] > self.config.state_dim: return state[..., : self.config.state_dim] diff --git a/tests/rewards/test_temporal_siglip_value_function.py b/tests/rewards/test_temporal_siglip_value_function.py index c183129c6..beaec5f7a 100644 --- a/tests/rewards/test_temporal_siglip_value_function.py +++ b/tests/rewards/test_temporal_siglip_value_function.py @@ -127,3 +127,9 @@ def test_temporal_model_forward(monkeypatch): loss, metrics = model(batch) assert torch.isfinite(loss) assert -1.0 <= metrics["predicted_value_mean"] <= 0.0 + + model.eval() + with torch.no_grad(): + eval_loss, eval_metrics = model(batch) + assert torch.isfinite(eval_loss) + assert -1.0 <= eval_metrics["predicted_value_mean"] <= 0.0