mirror of
https://github.com/huggingface/lerobot.git
synced 2026-07-25 02:36:11 +00:00
fix: nan in eval
This commit is contained in:
+28
-7
@@ -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]
|
||||
|
||||
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user