From 049e29b16caa1a5a88587d5e875686a2d21f3ff5 Mon Sep 17 00:00:00 2001 From: Maxime Ellerbach Date: Thu, 30 Jul 2026 15:19:48 +0000 Subject: [PATCH] fix(policies): vla jepa prepare model input to take index 0 and not index -1 --- .../policies/vla_jepa/configuration_vla_jepa.py | 11 ++++++++++- src/lerobot/policies/vla_jepa/modeling_vla_jepa.py | 3 ++- 2 files changed, 12 insertions(+), 2 deletions(-) diff --git a/src/lerobot/policies/vla_jepa/configuration_vla_jepa.py b/src/lerobot/policies/vla_jepa/configuration_vla_jepa.py index 8a30ee374..2496e8688 100644 --- a/src/lerobot/policies/vla_jepa/configuration_vla_jepa.py +++ b/src/lerobot/policies/vla_jepa/configuration_vla_jepa.py @@ -15,11 +15,13 @@ from __future__ import annotations from dataclasses import dataclass, field +from typing import Any from lerobot.configs.policies import PreTrainedConfig -from lerobot.configs.types import NormalizationMode +from lerobot.configs.types import FeatureType, NormalizationMode, PolicyFeature from lerobot.optim.optimizers import AdamWConfig from lerobot.optim.schedulers import CosineDecayWithWarmupSchedulerConfig +from lerobot.utils.constants import OBS_STATE @PreTrainedConfig.register_subclass("vla_jepa") @@ -122,6 +124,13 @@ class VLAJEPAConfig(PreTrainedConfig): if self.robot_state_feature is not None: self.state_dim = self.robot_state_feature.shape[0] + def set_dataset_feature_metadata(self, dataset_features: dict[str, Any]) -> None: + """Add `observation.state` to `input_features` if missing, so it gets normalized.""" + if OBS_STATE in self.input_features or OBS_STATE not in dataset_features: + return + shape = tuple(dataset_features[OBS_STATE]["shape"]) + self.input_features[OBS_STATE] = PolicyFeature(type=FeatureType.STATE, shape=shape) + def get_optimizer_preset(self) -> AdamWConfig: return AdamWConfig( lr=self.optimizer_lr, diff --git a/src/lerobot/policies/vla_jepa/modeling_vla_jepa.py b/src/lerobot/policies/vla_jepa/modeling_vla_jepa.py index 9c689d3c7..f9a1ba447 100644 --- a/src/lerobot/policies/vla_jepa/modeling_vla_jepa.py +++ b/src/lerobot/policies/vla_jepa/modeling_vla_jepa.py @@ -399,7 +399,8 @@ class VLAJEPAPolicy(PreTrainedPolicy): state = batch.get(OBS_STATE) if state is not None: if state.ndim > 2: - state = state[:, -1, :] + # deltas are forward-looking here, so index 0 is the current observation, not -1. + state = state[:, 0, :] inputs["state"] = (state.unsqueeze(1) if state.ndim == 2 else state).float() # [B, 1, dim] return inputs