mirror of
https://github.com/huggingface/lerobot.git
synced 2026-07-25 18:56:09 +00:00
fix temporal siglip value function
This commit is contained in:
+3
-2
@@ -142,8 +142,9 @@ class TemporalSiglipVFRewardModel(DistributionalValueMixin, PreTrainedRewardMode
|
|||||||
src_key_padding_mask=~frame_valid,
|
src_key_padding_mask=~frame_valid,
|
||||||
is_causal=True,
|
is_causal=True,
|
||||||
)
|
)
|
||||||
last_valid = frame_valid.long().sum(-1).sub(1).clamp_min(0)
|
# The history window is ordered oldest→current and the current frame is
|
||||||
return hidden[torch.arange(batch_size, device=hidden.device), last_valid]
|
# always the final, non-padding element.
|
||||||
|
return hidden[:, -1]
|
||||||
|
|
||||||
def _fit_state_dim(self, state: Tensor) -> Tensor:
|
def _fit_state_dim(self, state: Tensor) -> Tensor:
|
||||||
if state.shape[-1] > self.config.state_dim:
|
if state.shape[-1] > self.config.state_dim:
|
||||||
|
|||||||
+6
-2
@@ -52,17 +52,21 @@ class TemporalSiglipImageProcessorStep(ProcessorStep):
|
|||||||
)
|
)
|
||||||
continue
|
continue
|
||||||
|
|
||||||
image = observation[key].float()
|
image = observation[key]
|
||||||
if image.ndim == 4:
|
if image.ndim == 4:
|
||||||
image = image[:, None]
|
image = image[:, None]
|
||||||
if image.ndim != 5 or image.shape[2] != 3:
|
if image.ndim != 5 or image.shape[2] != 3:
|
||||||
raise ValueError(f"Expected {key} as [B,T,3,H,W], got {tuple(image.shape)}")
|
raise ValueError(f"Expected {key} as [B,T,3,H,W], got {tuple(image.shape)}")
|
||||||
batch_size, history = image.shape[:2]
|
batch_size, history = image.shape[:2]
|
||||||
|
image = image.float() / 127.5 - 1.0 if image.dtype == torch.uint8 else image.float() * 2.0 - 1.0
|
||||||
image = image.flatten(0, 1).permute(0, 2, 3, 1)
|
image = image.flatten(0, 1).permute(0, 2, 3, 1)
|
||||||
image = image * 2.0 - 1.0
|
|
||||||
if image.shape[1:3] != self.image_resolution:
|
if image.shape[1:3] != self.image_resolution:
|
||||||
image = resize_with_pad_torch(image, *self.image_resolution)
|
image = resize_with_pad_torch(image, *self.image_resolution)
|
||||||
observation[key] = image.permute(0, 3, 1, 2).unflatten(0, (batch_size, history))
|
observation[key] = image.permute(0, 3, 1, 2).unflatten(0, (batch_size, history))
|
||||||
|
padding_key = f"{key}_is_pad"
|
||||||
|
if padding_key in observation:
|
||||||
|
observation[key + IMAGE_MASK_SUFFIX] = ~observation[padding_key].bool()
|
||||||
|
else:
|
||||||
observation[key + IMAGE_MASK_SUFFIX] = torch.ones(
|
observation[key + IMAGE_MASK_SUFFIX] = torch.ones(
|
||||||
batch_size, history, dtype=torch.bool, device=image.device
|
batch_size, history, dtype=torch.bool, device=image.device
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -53,12 +53,15 @@ def test_temporal_image_processor():
|
|||||||
)
|
)
|
||||||
transition = {
|
transition = {
|
||||||
TransitionKey.OBSERVATION: {
|
TransitionKey.OBSERVATION: {
|
||||||
CAMERAS[0]: torch.full((1, 2, 3, 20, 16), 0.5),
|
CAMERAS[0]: torch.full((1, 2, 3, 20, 16), 128, dtype=torch.uint8),
|
||||||
|
CAMERAS[0] + "_is_pad": torch.tensor([[True, False]]),
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
observation = step(transition)[TransitionKey.OBSERVATION]
|
observation = step(transition)[TransitionKey.OBSERVATION]
|
||||||
assert observation[CAMERAS[0]].shape == (1, 2, 3, 32, 32)
|
assert observation[CAMERAS[0]].shape == (1, 2, 3, 32, 32)
|
||||||
assert observation[CAMERAS[0] + ".mask"].shape == (1, 2)
|
assert observation[CAMERAS[0] + ".mask"].shape == (1, 2)
|
||||||
|
assert observation[CAMERAS[0] + ".mask"].tolist() == [[False, True]]
|
||||||
|
assert -1.0 <= observation[CAMERAS[0]].min() <= observation[CAMERAS[0]].max() <= 1.0
|
||||||
|
|
||||||
|
|
||||||
def test_temporal_model_forward(monkeypatch):
|
def test_temporal_model_forward(monkeypatch):
|
||||||
@@ -89,7 +92,7 @@ def test_temporal_model_forward(monkeypatch):
|
|||||||
model = modeling.TemporalSiglipVFRewardModel(_config())
|
model = modeling.TemporalSiglipVFRewardModel(_config())
|
||||||
batch = {
|
batch = {
|
||||||
**{key: torch.rand(1, 2, 3, 16, 16) for key in CAMERAS},
|
**{key: torch.rand(1, 2, 3, 16, 16) for key in CAMERAS},
|
||||||
**{key + ".mask": torch.ones(1, 2, dtype=torch.bool) for key in CAMERAS},
|
**{key + ".mask": torch.tensor([[False, True]]) for key in CAMERAS},
|
||||||
OBS_STATE: torch.rand(1, 2, 4),
|
OBS_STATE: torch.rand(1, 2, 4),
|
||||||
OBS_LANGUAGE_TOKENS: torch.ones(1, 4, dtype=torch.long),
|
OBS_LANGUAGE_TOKENS: torch.ones(1, 4, dtype=torch.long),
|
||||||
OBS_LANGUAGE_ATTENTION_MASK: torch.ones(1, 4, dtype=torch.bool),
|
OBS_LANGUAGE_ATTENTION_MASK: torch.ones(1, 4, dtype=torch.bool),
|
||||||
|
|||||||
Reference in New Issue
Block a user