mirror of
https://github.com/huggingface/lerobot.git
synced 2026-07-30 21:19:40 +00:00
fixing vla jepa prepare model input to take index 0 and not index -1
This commit is contained in:
@@ -15,11 +15,13 @@
|
|||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
from dataclasses import dataclass, field
|
from dataclasses import dataclass, field
|
||||||
|
from typing import Any
|
||||||
|
|
||||||
from lerobot.configs.policies import PreTrainedConfig
|
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.optimizers import AdamWConfig
|
||||||
from lerobot.optim.schedulers import CosineDecayWithWarmupSchedulerConfig
|
from lerobot.optim.schedulers import CosineDecayWithWarmupSchedulerConfig
|
||||||
|
from lerobot.utils.constants import OBS_STATE
|
||||||
|
|
||||||
|
|
||||||
@PreTrainedConfig.register_subclass("vla_jepa")
|
@PreTrainedConfig.register_subclass("vla_jepa")
|
||||||
@@ -130,6 +132,26 @@ class VLAJEPAConfig(PreTrainedConfig):
|
|||||||
if self.robot_state_feature is not None:
|
if self.robot_state_feature is not None:
|
||||||
self.state_dim = self.robot_state_feature.shape[0]
|
self.state_dim = self.robot_state_feature.shape[0]
|
||||||
|
|
||||||
|
def set_dataset_feature_metadata(self, dataset_features: dict[str, Any]) -> None:
|
||||||
|
"""Declare `observation.state` as an input feature so the normalizer picks it up.
|
||||||
|
|
||||||
|
`make_policy` only fills `input_features` from the dataset when it is empty, so a
|
||||||
|
fine-tune inheriting `input_features` from a pretrained checkpoint keeps that
|
||||||
|
checkpoint's keys. The starVLA bases were vision-only, so `observation.state` was
|
||||||
|
absent — and `NormalizerProcessorStep` iterates its `features` dict, meaning the
|
||||||
|
state reached `VLAJEPAActionHead.state_encoder` as raw joint values (tens of
|
||||||
|
degrees) while every other token in the DiT sequence is O(1).
|
||||||
|
|
||||||
|
This hook (invoked by `make_policy` once the dataset metadata is known, before the
|
||||||
|
processors are built) adds the key back, so `STATE: MEAN_STD` applies. The
|
||||||
|
relative-action step runs *before* the normalizer and caches the raw state, so the
|
||||||
|
relative/absolute round-trip is unaffected.
|
||||||
|
"""
|
||||||
|
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:
|
def get_optimizer_preset(self) -> AdamWConfig:
|
||||||
return AdamWConfig(
|
return AdamWConfig(
|
||||||
lr=self.optimizer_lr,
|
lr=self.optimizer_lr,
|
||||||
|
|||||||
@@ -452,7 +452,13 @@ class VLAJEPAPolicy(PreTrainedPolicy):
|
|||||||
state = batch.get(OBS_STATE)
|
state = batch.get(OBS_STATE)
|
||||||
if state is not None:
|
if state is not None:
|
||||||
if state.ndim > 2:
|
if state.ndim > 2:
|
||||||
state = state[:, -1, :]
|
# `observation_delta_indices` is FORWARD-looking here (it feeds future frames to the
|
||||||
|
# world model), so a stacked state is [B, T_obs, dim] at deltas 0, +stride, ..., and
|
||||||
|
# index 0 — not -1 — is the CURRENT observation. Taking [:, -1] would condition the
|
||||||
|
# action head on the state at the END of the chunk during training while inference
|
||||||
|
# (single frame, ndim == 2) passes the current state, and the relative-action anchor
|
||||||
|
# in RelativeActionsProcessorStep also uses index 0. Matches the image slice above.
|
||||||
|
state = state[:, 0, :]
|
||||||
inputs["state"] = (state.unsqueeze(1) if state.ndim == 2 else state).float() # [B, 1, dim]
|
inputs["state"] = (state.unsqueeze(1) if state.ndim == 2 else state).float() # [B, 1, dim]
|
||||||
|
|
||||||
return inputs
|
return inputs
|
||||||
|
|||||||
Reference in New Issue
Block a user