mirror of
https://github.com/huggingface/lerobot.git
synced 2026-07-24 10:16:09 +00:00
[pre-commit.ci] auto fixes from pre-commit.com hooks
for more information, see https://pre-commit.ci
This commit is contained in:
+2085
-1695
File diff suppressed because one or more lines are too long
@@ -1,8 +1,9 @@
|
||||
from typing import Dict, List
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
from torch.utils.data.dataloader import default_collate
|
||||
import numpy as np
|
||||
|
||||
|
||||
def is_batch_need_padding(values: list[torch.Tensor], pad_dim: int = -1) -> int:
|
||||
return len(values[0].shape) > 0 # and len(set([v.shape[pad_dim] for v in values])) > 1
|
||||
@@ -51,7 +52,8 @@ def multidataset_collate_fn(
|
||||
if (
|
||||
key in keys_to_max_dim
|
||||
and isinstance(values[0], torch.Tensor)
|
||||
and is_batch_need_padding(values, pad_dim=pad_dim) and keys_to_max_dim[key] is not None
|
||||
and is_batch_need_padding(values, pad_dim=pad_dim)
|
||||
and keys_to_max_dim[key] is not None
|
||||
):
|
||||
max_size = keys_to_max_dim[key]
|
||||
for i in range(len(batch)):
|
||||
|
||||
@@ -32,7 +32,8 @@ IMAGENET_STATS = {
|
||||
"std": [[[0.229]], [[0.224]], [[0.225]]], # (c,1,1)
|
||||
}
|
||||
|
||||
from lerobot.common.datasets.utils_must import (EPISODES_DATASET_MAPPING, TRAINING_FEATURES, FEATURE_KEYS_MAPPING)
|
||||
from lerobot.common.datasets.utils_must import EPISODES_DATASET_MAPPING, FEATURE_KEYS_MAPPING
|
||||
|
||||
|
||||
def resolve_delta_timestamps(
|
||||
cfg: PreTrainedConfig, ds_meta: LeRobotDatasetMetadata
|
||||
@@ -66,6 +67,7 @@ def resolve_delta_timestamps(
|
||||
|
||||
return delta_timestamps
|
||||
|
||||
|
||||
def make_dataset(cfg: TrainPipelineConfig) -> LeRobotDataset | MultiLeRobotDataset:
|
||||
"""Handles the logic of setting up delta timestamps and image transforms before creating a dataset.
|
||||
|
||||
@@ -119,9 +121,9 @@ def make_dataset(cfg: TrainPipelineConfig) -> LeRobotDataset | MultiLeRobotDatas
|
||||
feature_keys_mapping=feature_keys_mapping,
|
||||
) # FIXME(mshukor): ?
|
||||
delta_timestamps[repo_id[i]] = resolve_delta_timestamps(cfg.policy, ds_meta)
|
||||
episodes[repo_id[i]] = EPISODES_DATASET_MAPPING.get(repo_id[i], cfg.dataset.episodes)
|
||||
episodes[repo_id[i]] = EPISODES_DATASET_MAPPING.get(repo_id[i], cfg.dataset.episodes)
|
||||
# training_features = TRAINING_FEATURES.get(cfg.dataset.features_version, None)
|
||||
#FIXME: (jadechoghari): check support for training features
|
||||
# FIXME: (jadechoghari): check support for training features
|
||||
training_features = None
|
||||
dataset = MultiLeRobotDataset(
|
||||
repo_id,
|
||||
@@ -148,7 +150,7 @@ def make_dataset(cfg: TrainPipelineConfig) -> LeRobotDataset | MultiLeRobotDatas
|
||||
)
|
||||
logging.info(
|
||||
"Multiple datasets were provided. Applied the following index mapping to the provided datasets: "
|
||||
f"{pformat(dataset.repo_id_to_index , indent=2)}"
|
||||
f"{pformat(dataset.repo_id_to_index, indent=2)}"
|
||||
)
|
||||
if cfg.dataset.use_imagenet_stats:
|
||||
for key in dataset.meta.camera_keys:
|
||||
|
||||
@@ -14,10 +14,9 @@
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
import contextlib
|
||||
import copy
|
||||
import logging
|
||||
import shutil
|
||||
import os
|
||||
import shutil
|
||||
from pathlib import Path
|
||||
from typing import Callable
|
||||
|
||||
@@ -32,8 +31,16 @@ from huggingface_hub import HfApi, snapshot_download
|
||||
from huggingface_hub.constants import REPOCARD_NAME
|
||||
from huggingface_hub.errors import RevisionNotFoundError
|
||||
|
||||
from lerobot.constants import HF_LEROBOT_HOME
|
||||
from lerobot.datasets.compute_stats import aggregate_stats, compute_episode_stats #aggregate_stats_per_robot_type,
|
||||
from lerobot.constants import (
|
||||
ACTION,
|
||||
HF_LEROBOT_HOME,
|
||||
OBS_ENV_STATE,
|
||||
OBS_STATE,
|
||||
)
|
||||
from lerobot.datasets.compute_stats import ( # aggregate_stats_per_robot_type,
|
||||
aggregate_stats,
|
||||
compute_episode_stats,
|
||||
)
|
||||
from lerobot.datasets.image_writer import AsyncImageWriter, write_image
|
||||
from lerobot.datasets.utils import (
|
||||
DEFAULT_FEATURES,
|
||||
@@ -43,7 +50,6 @@ from lerobot.datasets.utils import (
|
||||
_validate_feature_names,
|
||||
append_jsonlines,
|
||||
backward_compatible_episodes_stats,
|
||||
check_delta_timestamps,
|
||||
check_timestamps_sync,
|
||||
check_version_compatibility,
|
||||
create_empty_dataset_info,
|
||||
@@ -51,7 +57,6 @@ from lerobot.datasets.utils import (
|
||||
embed_images,
|
||||
get_delta_indices,
|
||||
get_episode_data_index,
|
||||
get_features_from_robot,
|
||||
get_hf_features_from_features,
|
||||
get_safe_version,
|
||||
hf_transform_to_torch,
|
||||
@@ -68,10 +73,27 @@ from lerobot.datasets.utils import (
|
||||
write_episode_stats,
|
||||
write_info,
|
||||
write_json,
|
||||
#keep_datasets_with_the_same_features_per_robot_type,
|
||||
#map_dict_pad_keys,
|
||||
#keep_datasets_with_valid_fps,
|
||||
#find_start_of_motion,
|
||||
# keep_datasets_with_the_same_features_per_robot_type,
|
||||
# map_dict_pad_keys,
|
||||
# keep_datasets_with_valid_fps,
|
||||
# find_start_of_motion,
|
||||
)
|
||||
|
||||
# mustafa stuff here
|
||||
from lerobot.datasets.utils_must import (
|
||||
OBS_IMAGE,
|
||||
OBS_IMAGE_2,
|
||||
OBS_IMAGE_3,
|
||||
ROBOT_TYPE_KEYS_MAPPING,
|
||||
TASKS_KEYS_MAPPING,
|
||||
aggregate_stats_per_robot_type,
|
||||
create_padded_features,
|
||||
find_start_of_motion,
|
||||
keep_datasets_with_the_same_features_per_robot_type,
|
||||
keep_datasets_with_valid_fps,
|
||||
map_dict_keys,
|
||||
pad_tensor,
|
||||
reshape_features_to_max_dim,
|
||||
)
|
||||
from lerobot.datasets.video_utils import (
|
||||
VideoFrame,
|
||||
@@ -81,40 +103,18 @@ from lerobot.datasets.video_utils import (
|
||||
get_video_info,
|
||||
)
|
||||
|
||||
# mustafa stuff here
|
||||
from lerobot.datasets.utils_must import (
|
||||
reshape_features_to_max_dim,
|
||||
keep_datasets_with_valid_fps,
|
||||
keep_datasets_with_the_same_features_per_robot_type,
|
||||
aggregate_stats_per_robot_type,
|
||||
create_padded_features,
|
||||
pad_tensor,
|
||||
map_dict_keys,
|
||||
find_start_of_motion,
|
||||
ROBOT_TYPE_KEYS_MAPPING,
|
||||
OBS_IMAGE,
|
||||
OBS_IMAGE_2,
|
||||
OBS_IMAGE_3,
|
||||
TASKS_KEYS_MAPPING,
|
||||
)
|
||||
from lerobot.constants import (
|
||||
ACTION,
|
||||
OBS_ENV_STATE,
|
||||
OBS_STATE,
|
||||
|
||||
)
|
||||
|
||||
CODEBASE_VERSION = "v2.1"
|
||||
LEROBOT_HOME = Path(os.getenv("LEROBOT_HOME", "~/.cache/huggingface/lerobot")).expanduser()
|
||||
|
||||
|
||||
def find_start_of_motion(velocities, window_size, threshold, motion_buffer):
|
||||
for t in range(len(velocities) - window_size):
|
||||
window_mean = velocities[t:t+window_size].mean()
|
||||
window_mean = velocities[t : t + window_size].mean()
|
||||
if window_mean > threshold:
|
||||
return max(0, t - motion_buffer) # include slight context before motion
|
||||
return 0
|
||||
|
||||
|
||||
class LeRobotDatasetMetadata:
|
||||
def __init__(
|
||||
self,
|
||||
@@ -400,7 +400,6 @@ class LeRobotDataset(torch.utils.data.Dataset):
|
||||
force_cache_sync: bool = False,
|
||||
download_videos: bool = True,
|
||||
video_backend: str | None = None,
|
||||
|
||||
# new thing by M
|
||||
feature_keys_mapping: dict[str, str] | None = None,
|
||||
max_action_dim: int = None,
|
||||
@@ -549,7 +548,10 @@ class LeRobotDataset(torch.utils.data.Dataset):
|
||||
# Load metadata
|
||||
# TODO: change
|
||||
self.meta = LeRobotDatasetMetadata(
|
||||
self.repo_id, self.root, self.revision, force_cache_sync=force_cache_sync,
|
||||
self.repo_id,
|
||||
self.root,
|
||||
self.revision,
|
||||
force_cache_sync=force_cache_sync,
|
||||
feature_keys_mapping=feature_keys_mapping,
|
||||
)
|
||||
if self.episodes is not None and self.meta._version >= packaging.version.parse("v2.1"):
|
||||
@@ -577,9 +579,13 @@ class LeRobotDataset(torch.utils.data.Dataset):
|
||||
from_ = self.episode_data_index["from"][ep_idx]
|
||||
to_ = self.episode_data_index["to"][ep_idx]
|
||||
# TODO implement advanced strategy
|
||||
self.subset_frame_ids += [frame_idx for frame_idx in range(from_ + int(self.fps*self.discard_first_n_frames), to_)]
|
||||
self.subset_frame_ids += [
|
||||
frame_idx for frame_idx in range(from_ + int(self.fps * self.discard_first_n_frames), to_)
|
||||
]
|
||||
elif self.discard_first_idle_frames:
|
||||
print(f"Discarding first idle frames: motion_threshold={self.motion_threshold}, motion_window_size={self.motion_window_size}, motion_buffer={self.motion_buffer}")
|
||||
print(
|
||||
f"Discarding first idle frames: motion_threshold={self.motion_threshold}, motion_window_size={self.motion_window_size}, motion_buffer={self.motion_buffer}"
|
||||
)
|
||||
self.robot_states = torch.stack(self.hf_dataset[OBS_STATE]).numpy() # shape: [T, D]
|
||||
self.subset_frame_ids = []
|
||||
for ep_idx in range(self.num_episodes):
|
||||
@@ -587,8 +593,10 @@ class LeRobotDataset(torch.utils.data.Dataset):
|
||||
to_ = self.episode_data_index["to"][ep_idx]
|
||||
ep_states = self.robot_states[from_:to_]
|
||||
velocities = np.linalg.norm(np.diff(ep_states, axis=0), axis=1)
|
||||
velocities = np.concatenate([[0.0], velocities])
|
||||
start_idx = find_start_of_motion(velocities, self.motion_window_size, self.motion_threshold, self.motion_buffer)
|
||||
velocities = np.concatenate([[0.0], velocities])
|
||||
start_idx = find_start_of_motion(
|
||||
velocities, self.motion_window_size, self.motion_threshold, self.motion_buffer
|
||||
)
|
||||
self.subset_frame_ids += list(range(from_ + start_idx, to_))
|
||||
|
||||
# Check timestamps
|
||||
@@ -606,7 +614,9 @@ class LeRobotDataset(torch.utils.data.Dataset):
|
||||
|
||||
# Mustafa
|
||||
self.meta.info["features"] = map_dict_keys(
|
||||
self.meta.info["features"], feature_keys_mapping=self.feature_keys_mapping, training_features=self.training_features
|
||||
self.meta.info["features"],
|
||||
feature_keys_mapping=self.feature_keys_mapping,
|
||||
training_features=self.training_features,
|
||||
)
|
||||
self.keys_to_max_dim = {
|
||||
ACTION: max_action_dim,
|
||||
@@ -619,7 +629,11 @@ class LeRobotDataset(torch.utils.data.Dataset):
|
||||
self.meta.info["features"] = reshape_features_to_max_dim(
|
||||
self.meta.info["features"], reshape_dim=-1, keys_to_max_dim=self.keys_to_max_dim
|
||||
)
|
||||
self.meta.stats = map_dict_keys(self.meta.stats, feature_keys_mapping=self.feature_keys_mapping, training_features=self.training_features)
|
||||
self.meta.stats = map_dict_keys(
|
||||
self.meta.stats,
|
||||
feature_keys_mapping=self.feature_keys_mapping,
|
||||
training_features=self.training_features,
|
||||
)
|
||||
self.robot_type = self.meta.info.get("robot_type", "")
|
||||
# Override tasks
|
||||
print(TASKS_KEYS_MAPPING.get(self.repo_id, self.meta.tasks), "previous", self.meta.tasks)
|
||||
@@ -807,7 +821,10 @@ class LeRobotDataset(torch.utils.data.Dataset):
|
||||
def _query_hf_dataset(self, query_indices: dict[str, list[int]]) -> dict:
|
||||
queries = {}
|
||||
for key, q_idx in query_indices.items():
|
||||
if key not in self.meta.video_keys and self.inverse_feature_keys_mapping.get(key, key) not in self.meta.video_keys:
|
||||
if (
|
||||
key not in self.meta.video_keys
|
||||
and self.inverse_feature_keys_mapping.get(key, key) not in self.meta.video_keys
|
||||
):
|
||||
key_ = (
|
||||
self.inverse_feature_keys_mapping.get(key, key)
|
||||
if self.inverse_feature_keys_mapping
|
||||
@@ -868,7 +885,9 @@ class LeRobotDataset(torch.utils.data.Dataset):
|
||||
print(self.meta.tasks, task_idx, self.repo_id)
|
||||
if "robot_type" not in item:
|
||||
item["robot_type"] = self.robot_type
|
||||
item = map_dict_keys(item, feature_keys_mapping=self.feature_keys_mapping, training_features=self.training_features)
|
||||
item = map_dict_keys(
|
||||
item, feature_keys_mapping=self.feature_keys_mapping, training_features=self.training_features
|
||||
)
|
||||
# Add padded features
|
||||
# item = self._add_padded_features(item, self.training_features)
|
||||
if self.image_transforms is not None:
|
||||
@@ -942,7 +961,6 @@ class LeRobotDataset(torch.utils.data.Dataset):
|
||||
|
||||
# Add frame features to episode_buffer
|
||||
for key in frame:
|
||||
|
||||
if key not in self.features:
|
||||
raise ValueError(
|
||||
f"An element of the frame is not in the features. '{key}' not in '{self.features.keys()}'."
|
||||
@@ -1159,6 +1177,7 @@ class LeRobotDataset(torch.utils.data.Dataset):
|
||||
obj.video_backend = video_backend if video_backend is not None else get_safe_default_codec()
|
||||
return obj
|
||||
|
||||
|
||||
class MultiLeRobotDatasetMeta:
|
||||
def __init__(
|
||||
self,
|
||||
@@ -1185,7 +1204,7 @@ class MultiLeRobotDatasetMeta:
|
||||
intersection.intersection_update(ds.features)
|
||||
if not intersection:
|
||||
raise RuntimeError("No common features across datasets.")
|
||||
for repo_id, ds in zip(repo_ids, datasets):
|
||||
for repo_id, ds in zip(repo_ids, datasets, strict=False):
|
||||
extra = set(ds.features) - intersection
|
||||
logging.warning(f"Disabling {extra} for repo {repo_id}")
|
||||
self.disabled_features.update(extra)
|
||||
@@ -1210,19 +1229,17 @@ class MultiLeRobotDatasetMeta:
|
||||
for k, v in feat_stats.items():
|
||||
pad_value = 0 if k in ["min", "mean"] else 1
|
||||
self.stats[robot_type_][feat_key][k] = pad_tensor(
|
||||
v, max_size=self.keys_to_max_dim.get(feat_key, -1), pad_dim=-1, pad_value=pad_value
|
||||
v,
|
||||
max_size=self.keys_to_max_dim.get(feat_key, -1),
|
||||
pad_dim=-1,
|
||||
pad_value=pad_value,
|
||||
)
|
||||
|
||||
# step 5: episodes & tasks
|
||||
self.episodes = {
|
||||
repo_id: ds.meta.episodes for repo_id, ds in zip(repo_ids, datasets)
|
||||
}
|
||||
self.tasks = {
|
||||
repo_id: ds.meta.tasks for repo_id, ds in zip(repo_ids, datasets)
|
||||
}
|
||||
self.info = {
|
||||
repo_id: ds.meta.info for repo_id, ds in zip(repo_ids, datasets)
|
||||
}
|
||||
self.episodes = {repo_id: ds.meta.episodes for repo_id, ds in zip(repo_ids, datasets, strict=False)}
|
||||
self.tasks = {repo_id: ds.meta.tasks for repo_id, ds in zip(repo_ids, datasets, strict=False)}
|
||||
self.info = {repo_id: ds.meta.info for repo_id, ds in zip(repo_ids, datasets, strict=False)}
|
||||
|
||||
|
||||
class MultiLeRobotDatasetCleaner:
|
||||
def __init__(
|
||||
@@ -1243,7 +1260,9 @@ class MultiLeRobotDatasetCleaner:
|
||||
valid_fps_datasets = keep_datasets_with_valid_fps(datasets, min_fps=min_fps, max_fps=max_fps)
|
||||
|
||||
# step 2: keep datasets with same features per robot type
|
||||
consistent_datasets, keep_mask = keep_datasets_with_the_same_features_per_robot_type(valid_fps_datasets)
|
||||
consistent_datasets, keep_mask = keep_datasets_with_the_same_features_per_robot_type(
|
||||
valid_fps_datasets
|
||||
)
|
||||
|
||||
self.cleaned_datasets = consistent_datasets
|
||||
self.keep_mask = keep_mask
|
||||
@@ -1257,7 +1276,7 @@ class MultiLeRobotDatasetCleaner:
|
||||
[0] + list(torch.cumsum(torch.tensor([len(d) for d in consistent_datasets]), dim=0))
|
||||
)
|
||||
self.cleaned_weights = np.array(self.cleaned_weights, dtype=np.float32)
|
||||
|
||||
|
||||
|
||||
class MultiLeRobotDataset(torch.utils.data.Dataset):
|
||||
"""A dataset consisting of multiple underlying `LeRobotDataset`s.
|
||||
@@ -1277,7 +1296,6 @@ class MultiLeRobotDataset(torch.utils.data.Dataset):
|
||||
download_videos: bool = True,
|
||||
local_files_only: bool = False,
|
||||
video_backend: str | None = None,
|
||||
|
||||
# add
|
||||
sampling_weights: list[float] | None = None,
|
||||
feature_keys_mapping: dict[str, dict[str, str]] | None = None,
|
||||
@@ -1298,7 +1316,7 @@ class MultiLeRobotDataset(torch.utils.data.Dataset):
|
||||
super().__init__()
|
||||
self.repo_ids = repo_ids
|
||||
self.root = Path(root) if root else HF_LEROBOT_HOME
|
||||
self.tolerances_s = tolerances_s if tolerances_s else {repo_id: 1e-4 for repo_id in repo_ids}
|
||||
self.tolerances_s = tolerances_s if tolerances_s else dict.fromkeys(repo_ids, 0.0001)
|
||||
# Construct the underlying datasets passing everything but `transform` and `delta_timestamps` which
|
||||
# are handled by this class.
|
||||
_datasets = []
|
||||
@@ -1320,7 +1338,7 @@ class MultiLeRobotDataset(torch.utils.data.Dataset):
|
||||
root=self.root / repo_id,
|
||||
episodes=episodes.get(repo_id, None) if episodes else None,
|
||||
image_transforms=image_transforms,
|
||||
delta_timestamps = delta_timestamps.get(repo_id, None) if delta_timestamps else None,
|
||||
delta_timestamps=delta_timestamps.get(repo_id, None) if delta_timestamps else None,
|
||||
tolerance_s=self.tolerances_s[repo_id],
|
||||
download_videos=download_videos,
|
||||
video_backend=video_backend,
|
||||
@@ -1385,7 +1403,6 @@ class MultiLeRobotDataset(torch.utils.data.Dataset):
|
||||
self.disabled_features = self.meta.disabled_features
|
||||
self.stats = self.meta.stats
|
||||
|
||||
|
||||
@property
|
||||
def repo_id_to_index(self):
|
||||
"""Return a mapping from dataset repo_id to a dataset index automatically created by this class.
|
||||
|
||||
@@ -860,7 +860,9 @@ def validate_episode_buffer(episode_buffer: dict, total_episodes: int, features:
|
||||
)
|
||||
|
||||
|
||||
def map_dict_keys(item: dict, feature_keys_mapping: dict, training_features: list = None, pad_key: str = "is_pad") -> dict:
|
||||
def map_dict_keys(
|
||||
item: dict, feature_keys_mapping: dict, training_features: list = None, pad_key: str = "is_pad"
|
||||
) -> dict:
|
||||
"""Maps feature keys from the dataset to the keys used in the model."""
|
||||
if feature_keys_mapping is None:
|
||||
return item
|
||||
|
||||
@@ -1,25 +1,22 @@
|
||||
"""
|
||||
Utils function by Mustafa to refactor
|
||||
"""
|
||||
import torch
|
||||
import numpy as np
|
||||
from lerobot.datasets.compute_stats import (
|
||||
aggregate_stats
|
||||
)
|
||||
|
||||
from collections import defaultdict
|
||||
from torch.utils.data.dataloader import default_collate
|
||||
import torch.nn.functional as F
|
||||
from typing import Dict, List
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
from typing import Dict, List
|
||||
import torch.nn.functional as F
|
||||
from torch.utils.data.dataloader import default_collate
|
||||
|
||||
from lerobot.datasets.compute_stats import aggregate_stats
|
||||
|
||||
|
||||
from typing import Dict, List
|
||||
OBS_IMAGE = "observation.image"
|
||||
OBS_IMAGE_2 = "observation.image2"
|
||||
OBS_IMAGE_3 = "observation.image3"
|
||||
|
||||
|
||||
def reshape_features_to_max_dim(features: dict, reshape_dim: int = -1, keys_to_max_dim: dict = {}) -> dict:
|
||||
"""Reshape features to have a maximum dimension of `max_dim`."""
|
||||
reshaped_features = {}
|
||||
@@ -37,10 +34,11 @@ def reshape_features_to_max_dim(features: dict, reshape_dim: int = -1, keys_to_m
|
||||
reshaped_features[key] = features[key]
|
||||
return reshaped_features
|
||||
|
||||
def keep_datasets_with_valid_fps(
|
||||
ls_datasets: list, min_fps: int = 1, max_fps: int = 100
|
||||
) -> list:
|
||||
print(f"Keeping datasets with fps between {min_fps} and {max_fps}. Considering {len(ls_datasets)} datasets.")
|
||||
|
||||
def keep_datasets_with_valid_fps(ls_datasets: list, min_fps: int = 1, max_fps: int = 100) -> list:
|
||||
print(
|
||||
f"Keeping datasets with fps between {min_fps} and {max_fps}. Considering {len(ls_datasets)} datasets."
|
||||
)
|
||||
for ds in ls_datasets:
|
||||
if ds.fps < min_fps or ds.fps > max_fps:
|
||||
print(f"Dataset {ds} has invalid fps: {ds.fps}. Removing it.")
|
||||
@@ -48,9 +46,8 @@ def keep_datasets_with_valid_fps(
|
||||
print(f"Keeping {len(ls_datasets)} datasets with valid fps.")
|
||||
return ls_datasets
|
||||
|
||||
def keep_datasets_with_the_same_features_per_robot_type(
|
||||
ls_datasets: list
|
||||
) -> list:
|
||||
|
||||
def keep_datasets_with_the_same_features_per_robot_type(ls_datasets: list) -> list:
|
||||
"""
|
||||
Filters datasets to only keep those with consistent feature shapes per robot type.
|
||||
|
||||
@@ -68,7 +65,8 @@ def keep_datasets_with_the_same_features_per_robot_type(
|
||||
# Collect all stats dicts for this robot type
|
||||
stats_list = [
|
||||
ep_stats
|
||||
for ds in ls_datasets if ds.meta.info["robot_type"] == robot_type
|
||||
for ds in ls_datasets
|
||||
if ds.meta.info["robot_type"] == robot_type
|
||||
for ep_stats in ds.meta.episodes_stats.values()
|
||||
]
|
||||
if not stats_list:
|
||||
@@ -84,7 +82,9 @@ def keep_datasets_with_the_same_features_per_robot_type(
|
||||
|
||||
for stats in stats_list:
|
||||
value = stats.get(key)
|
||||
if value and "mean" in value and isinstance(value["mean"], (torch.Tensor, np.ndarray)): # FIXME(mshukor): check all stats; min, mean, max
|
||||
if (
|
||||
value and "mean" in value and isinstance(value["mean"], (torch.Tensor, np.ndarray))
|
||||
): # FIXME(mshukor): check all stats; min, mean, max
|
||||
shape_counter[value["mean"].shape] += 1
|
||||
if not shape_counter:
|
||||
continue
|
||||
@@ -97,18 +97,24 @@ def keep_datasets_with_the_same_features_per_robot_type(
|
||||
if not first_ep_stats:
|
||||
continue
|
||||
value = first_ep_stats.get(key)
|
||||
if value and "mean" in value and isinstance(value["mean"], (torch.Tensor, np.ndarray)) and value["mean"].shape != main_shape:
|
||||
if (
|
||||
value
|
||||
and "mean" in value
|
||||
and isinstance(value["mean"], (torch.Tensor, np.ndarray))
|
||||
and value["mean"].shape != main_shape
|
||||
):
|
||||
datasets_to_remove.add(ds)
|
||||
break
|
||||
|
||||
# Filter out inconsistent datasets
|
||||
datasets_maks = [ds not in datasets_to_remove for ds in ls_datasets]
|
||||
filtered_datasets = [ds for ds in ls_datasets if ds not in datasets_to_remove]
|
||||
print(f"Keeping {len(filtered_datasets)} datasets. Removed {len(datasets_to_remove)} inconsistent ones. Inconsistent datasets:\n{datasets_to_remove}")
|
||||
print(
|
||||
f"Keeping {len(filtered_datasets)} datasets. Removed {len(datasets_to_remove)} inconsistent ones. Inconsistent datasets:\n{datasets_to_remove}"
|
||||
)
|
||||
return filtered_datasets, datasets_maks
|
||||
|
||||
|
||||
|
||||
def aggregate_stats_per_robot_type(ls_datasets) -> dict[str, dict[str, torch.Tensor]]:
|
||||
"""Aggregate stats of multiple LeRobot datasets into multiple set of stats per robot type.
|
||||
|
||||
@@ -133,6 +139,7 @@ def aggregate_stats_per_robot_type(ls_datasets) -> dict[str, dict[str, torch.Ten
|
||||
stats[robot_type] = stat
|
||||
return stats
|
||||
|
||||
|
||||
def str_to_torch_dtype(dtype_str):
|
||||
"""Convert a dtype string to a torch dtype."""
|
||||
mapping = {
|
||||
@@ -144,9 +151,10 @@ def str_to_torch_dtype(dtype_str):
|
||||
}
|
||||
return mapping.get(dtype_str, torch.float32) # Default to float32
|
||||
|
||||
|
||||
def create_padded_features(item: dict, features: dict = {}):
|
||||
for key, ft in features.items():
|
||||
if any([k in key for k in ["cam", "effort", "absolute"]]): # FIXME(mshukor): temporary hack
|
||||
if any([k in key for k in ["cam", "effort", "absolute"]]): # FIXME(mshukor): temporary hack
|
||||
continue
|
||||
shape = ft["shape"]
|
||||
if len(shape) == 3: # images to torch format (C, H, W)
|
||||
@@ -157,12 +165,13 @@ def create_padded_features(item: dict, features: dict = {}):
|
||||
dtype = str_to_torch_dtype(ft["dtype"])
|
||||
item[key] = torch.zeros(shape, dtype=dtype)
|
||||
item[f"{key}_padding_mask"] = torch.tensor(0, dtype=torch.int64)
|
||||
if "image" in key: # FIXME(mshukor): support other observations
|
||||
if "image" in key: # FIXME(mshukor): support other observations
|
||||
item[f"{key}_is_pad"] = torch.BoolTensor([False])
|
||||
else:
|
||||
item[f"{key}_padding_mask"] = torch.tensor(1, dtype=torch.int64)
|
||||
return item
|
||||
|
||||
|
||||
ROBOT_TYPE_KEYS_MAPPING = {
|
||||
"lerobot/stanford_hydra_dataset": "static_single_arm",
|
||||
"lerobot/iamlab_cmu_pickup_insert": "static_single_arm",
|
||||
@@ -173,6 +182,7 @@ ROBOT_TYPE_KEYS_MAPPING = {
|
||||
"lerobot/taco_play": "static_single_arm_7statedim",
|
||||
}
|
||||
|
||||
|
||||
def pad_tensor(
|
||||
tensor: torch.Tensor, max_size: int, pad_dim: int = -1, pad_value: float = 0.0
|
||||
) -> torch.Tensor:
|
||||
@@ -188,7 +198,10 @@ def pad_tensor(
|
||||
tensor = torch.nn.functional.pad(tensor, pad_sizes, value=pad_value)
|
||||
return tensor.numpy() if is_numpy else tensor
|
||||
|
||||
def map_dict_keys(item: dict, feature_keys_mapping: dict, training_features: list = None, pad_key: str = "is_pad") -> dict:
|
||||
|
||||
def map_dict_keys(
|
||||
item: dict, feature_keys_mapping: dict, training_features: list = None, pad_key: str = "is_pad"
|
||||
) -> dict:
|
||||
"""Maps feature keys from the dataset to the keys used in the model."""
|
||||
if feature_keys_mapping is None:
|
||||
return item
|
||||
@@ -205,15 +218,19 @@ def map_dict_keys(item: dict, feature_keys_mapping: dict, training_features: lis
|
||||
# breakpoint()
|
||||
return features
|
||||
|
||||
|
||||
def find_start_of_motion(velocities, window_size, threshold, motion_buffer):
|
||||
for t in range(len(velocities) - window_size):
|
||||
window_mean = velocities[t:t+window_size].mean()
|
||||
window_mean = velocities[t : t + window_size].mean()
|
||||
if window_mean > threshold:
|
||||
return max(0, t - motion_buffer) # include slight context before motion
|
||||
return 0
|
||||
|
||||
import yaml
|
||||
|
||||
import requests
|
||||
import yaml
|
||||
|
||||
|
||||
def load_yaml_mapping(name: str) -> dict:
|
||||
"""
|
||||
Loads a YAML mapping from a Hugging Face repo.
|
||||
@@ -225,13 +242,115 @@ def load_yaml_mapping(name: str) -> dict:
|
||||
|
||||
return yaml.safe_load(response.text)
|
||||
|
||||
|
||||
# Example usage
|
||||
TASKS_KEYS_MAPPING = load_yaml_mapping("tasks")
|
||||
FEATURE_KEYS_MAPPING = load_yaml_mapping("features")
|
||||
EPISODES_DATASET_MAPPING = {
|
||||
"cadene/droid_1.0.1": list(range(50)),
|
||||
"danaaubakirova/svla_so100_task5_v3": [0, 1, 2, 3, 4, 5, 6, 7, 9, 10, 11, 12, 13, 14, 15, 16, 17, 18, 19, 20, 21, 22, 24, 25, 26, 27, 28, 29, 30, 31, 32, 33, 34, 35, 36, 37, 38, 39, 40, 41, 42, 43, 44, 45, 46, 47, 48, 49, 50, 51],
|
||||
"danaaubakirova/svla_so100_task4_v3": [0, 1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 12, 13, 14, 15, 16, 17, 18, 19, 21, 22, 23, 24, 25, 26, 27, 28, 29, 30, 31, 32, 33, 34, 35, 40, 41, 42, 43, 44, 45, 46, 47, 48, 49, 50, 51, 52, 53],
|
||||
"danaaubakirova/svla_so100_task5_v3": [
|
||||
0,
|
||||
1,
|
||||
2,
|
||||
3,
|
||||
4,
|
||||
5,
|
||||
6,
|
||||
7,
|
||||
9,
|
||||
10,
|
||||
11,
|
||||
12,
|
||||
13,
|
||||
14,
|
||||
15,
|
||||
16,
|
||||
17,
|
||||
18,
|
||||
19,
|
||||
20,
|
||||
21,
|
||||
22,
|
||||
24,
|
||||
25,
|
||||
26,
|
||||
27,
|
||||
28,
|
||||
29,
|
||||
30,
|
||||
31,
|
||||
32,
|
||||
33,
|
||||
34,
|
||||
35,
|
||||
36,
|
||||
37,
|
||||
38,
|
||||
39,
|
||||
40,
|
||||
41,
|
||||
42,
|
||||
43,
|
||||
44,
|
||||
45,
|
||||
46,
|
||||
47,
|
||||
48,
|
||||
49,
|
||||
50,
|
||||
51,
|
||||
],
|
||||
"danaaubakirova/svla_so100_task4_v3": [
|
||||
0,
|
||||
1,
|
||||
2,
|
||||
3,
|
||||
4,
|
||||
5,
|
||||
6,
|
||||
7,
|
||||
8,
|
||||
9,
|
||||
10,
|
||||
11,
|
||||
12,
|
||||
13,
|
||||
14,
|
||||
15,
|
||||
16,
|
||||
17,
|
||||
18,
|
||||
19,
|
||||
21,
|
||||
22,
|
||||
23,
|
||||
24,
|
||||
25,
|
||||
26,
|
||||
27,
|
||||
28,
|
||||
29,
|
||||
30,
|
||||
31,
|
||||
32,
|
||||
33,
|
||||
34,
|
||||
35,
|
||||
40,
|
||||
41,
|
||||
42,
|
||||
43,
|
||||
44,
|
||||
45,
|
||||
46,
|
||||
47,
|
||||
48,
|
||||
49,
|
||||
50,
|
||||
51,
|
||||
52,
|
||||
53,
|
||||
],
|
||||
}
|
||||
ACTION = "action"
|
||||
OBS_STATE = "observation.state"
|
||||
@@ -243,6 +362,7 @@ TRAINING_FEATURES = {
|
||||
2: [ACTION, OBS_STATE, TASK, ROBOT, OBS_IMAGE, OBS_IMAGE_2, OBS_IMAGE_3],
|
||||
}
|
||||
|
||||
|
||||
def is_batch_need_padding(values: list[torch.Tensor], pad_dim: int = -1) -> int:
|
||||
return len(values[0].shape) > 0 # and len(set([v.shape[pad_dim] for v in values])) > 1
|
||||
|
||||
@@ -250,7 +370,7 @@ def is_batch_need_padding(values: list[torch.Tensor], pad_dim: int = -1) -> int:
|
||||
def pad_tensor_to_shape(tensor: torch.Tensor, target_shape: tuple, pad_value: float = 0.0) -> torch.Tensor:
|
||||
"""Pads a tensor to the target shape (right/bottom only)."""
|
||||
pad = []
|
||||
for actual, target in zip(reversed(tensor.shape), reversed(target_shape)):
|
||||
for actual, target in zip(reversed(tensor.shape), reversed(target_shape), strict=False):
|
||||
pad.extend([0, max(target - actual, 0)])
|
||||
return F.pad(tensor, pad, value=pad_value)
|
||||
|
||||
|
||||
@@ -14,12 +14,13 @@
|
||||
|
||||
from dataclasses import dataclass, field
|
||||
|
||||
from lerobot.configs.policies import PreTrainedConfig
|
||||
from lerobot.configs.types import FeatureType, NormalizationMode, PolicyFeature
|
||||
from lerobot.optim.optimizers import AdamWConfig
|
||||
from lerobot.optim.schedulers import (
|
||||
CosineDecayWithWarmupSchedulerConfig,
|
||||
)
|
||||
from lerobot.configs.policies import PreTrainedConfig
|
||||
from lerobot.configs.types import FeatureType, NormalizationMode, PolicyFeature
|
||||
|
||||
|
||||
@dataclass
|
||||
class PEFTConfig:
|
||||
@@ -28,6 +29,7 @@ class PEFTConfig:
|
||||
lora_dropout: float = 0.1
|
||||
target_modules: str = "q_proj,v_proj"
|
||||
|
||||
|
||||
@PreTrainedConfig.register_subclass("smolvla2")
|
||||
@dataclass
|
||||
class SmolVLA2Config(PreTrainedConfig):
|
||||
@@ -78,7 +80,7 @@ class SmolVLA2Config(PreTrainedConfig):
|
||||
train_state_proj: bool = True
|
||||
|
||||
# Training presets
|
||||
optimizer_lr: float = 2.5e-5 #1e-4
|
||||
optimizer_lr: float = 2.5e-5 # 1e-4
|
||||
optimizer_betas: tuple[float, float] = (0.9, 0.95)
|
||||
optimizer_eps: float = 1e-8
|
||||
optimizer_weight_decay: float = 1e-10
|
||||
@@ -104,19 +106,19 @@ class SmolVLA2Config(PreTrainedConfig):
|
||||
pad_language_to: str = "longest" # "max_length"
|
||||
|
||||
num_expert_layers: int = -1 # Less or equal to 0 is the default where the action expert has the same number of layers of VLM. Otherwise the expert have less layers.
|
||||
num_vlm_layers: int = 16
|
||||
past_obs_keys: str = f"image"
|
||||
num_vlm_layers: int = 16
|
||||
past_obs_keys: str = "image"
|
||||
add_local_special_image_tokens: bool = False
|
||||
|
||||
reverse_images_order: bool = False
|
||||
|
||||
state_to_prefix: bool = False
|
||||
|
||||
pad_language_to: str = "longest" # "max_length"
|
||||
pad_language_to: str = "longest" # "max_length"
|
||||
causal_action_attention_mask: bool = False
|
||||
|
||||
self_attn_every_n_layers: int = -1# Number of layers used in the VLM (first num_vlm_layers layers)
|
||||
#self_attn_every_n_layers: int = 2 # Interleave SA layers each self_attn_every_n_layers
|
||||
self_attn_every_n_layers: int = -1 # Number of layers used in the VLM (first num_vlm_layers layers)
|
||||
# self_attn_every_n_layers: int = 2 # Interleave SA layers each self_attn_every_n_layers
|
||||
expert_width_multiplier: float = 0.75 # The action expert hidden size (wrt to the VLM)
|
||||
|
||||
min_period: float = 4e-3 # sensitivity range for the timestep used in sine-cosine positional encoding
|
||||
@@ -133,7 +135,7 @@ class SmolVLA2Config(PreTrainedConfig):
|
||||
|
||||
shuffle_camera_positions: bool = False
|
||||
vlm_img_size: int = -1
|
||||
|
||||
|
||||
regression_loss: bool = False
|
||||
|
||||
def __post_init__(self):
|
||||
|
||||
@@ -54,8 +54,8 @@ policy = SmolVLAPolicy.from_pretrained("lerobot/smolvla_base")
|
||||
|
||||
import math
|
||||
import os
|
||||
import re
|
||||
import random
|
||||
import re
|
||||
from collections import deque
|
||||
|
||||
import safetensors
|
||||
@@ -65,18 +65,18 @@ from torch import Tensor, nn
|
||||
from transformers import AutoProcessor
|
||||
|
||||
from lerobot.constants import ACTION, OBS_STATE
|
||||
from lerobot.datasets import IMAGES_ORDER
|
||||
from lerobot.policies.normalize import (
|
||||
Normalize,
|
||||
Unnormalize,
|
||||
)
|
||||
from lerobot.policies.pretrained import PreTrainedPolicy
|
||||
from lerobot.policies.smolvla2.configuration_smolvla2 import SmolVLA2Config
|
||||
from lerobot.policies.smolvla.smolvlm_with_expert import SmolVLMWithExpertModel
|
||||
from lerobot.policies.smolvla2.configuration_smolvla2 import SmolVLA2Config
|
||||
from lerobot.policies.utils import (
|
||||
populate_queues,
|
||||
)
|
||||
from lerobot.utils.utils import get_safe_dtype
|
||||
from lerobot.datasets import IMAGES_ORDER
|
||||
|
||||
# Matches ".soNNN", optionally followed by "-something", up to the "_buffer_" marker
|
||||
_VARIANT_RE = re.compile(r"\.so\d+(?:-[\w]+)?_buffer_")
|
||||
@@ -368,7 +368,6 @@ class SmolVLA2Policy(PreTrainedPolicy):
|
||||
for k in self.config.input_features:
|
||||
if any([past_obs_key in k for past_obs_key in self.config.past_obs_keys.split(",")]):
|
||||
self._queues[k] = deque(maxlen=self.config.n_obs_steps)
|
||||
|
||||
|
||||
# HACK(aliberts, danaaubakirova): we overwrite this classmethod here to fix smolVLA-specific issues
|
||||
@classmethod
|
||||
@@ -389,7 +388,7 @@ class SmolVLA2Policy(PreTrainedPolicy):
|
||||
|
||||
def get_optim_params(self) -> dict:
|
||||
return self.parameters()
|
||||
|
||||
|
||||
def merge_peft_model_weights(self) -> None:
|
||||
if "lora" in self.config.peft_method:
|
||||
self.model.vlm_with_expert.merge_lora_weights()
|
||||
@@ -413,9 +412,7 @@ class SmolVLA2Policy(PreTrainedPolicy):
|
||||
state = self.prepare_state(batch)
|
||||
lang_tokens, lang_masks = self.prepare_language(batch)
|
||||
|
||||
actions = self.model.sample_actions(
|
||||
images, img_masks, lang_tokens, lang_masks, state, noise=noise
|
||||
)
|
||||
actions = self.model.sample_actions(images, img_masks, lang_tokens, lang_masks, state, noise=noise)
|
||||
# Unpad actions
|
||||
original_action_dim = self.config.action_feature.shape[0]
|
||||
actions = actions[:, :, :original_action_dim]
|
||||
@@ -466,7 +463,7 @@ class SmolVLA2Policy(PreTrainedPolicy):
|
||||
else:
|
||||
actions = torch.cat((actions[:, :1], actions[:, 1:] + actions[:, :-1]), dim=1)
|
||||
# Unpad actions
|
||||
|
||||
|
||||
original_action_dim = self.config.action_feature.shape[0]
|
||||
actions = actions[:, :, :original_action_dim]
|
||||
|
||||
@@ -520,8 +517,12 @@ class SmolVLA2Policy(PreTrainedPolicy):
|
||||
present_img_keys = [key for key in self.config.image_features if key in batch]
|
||||
missing_img_keys = [key for key in self.config.image_features if key not in batch]
|
||||
|
||||
present_img_keys = sorted(present_img_keys, key=lambda k: IMAGES_ORDER.get(k, float("inf")), reverse=self.config.reverse_images_order)
|
||||
if self.config.shuffle_camera_positions and ACTION in batch: # only during training
|
||||
present_img_keys = sorted(
|
||||
present_img_keys,
|
||||
key=lambda k: IMAGES_ORDER.get(k, float("inf")),
|
||||
reverse=self.config.reverse_images_order,
|
||||
)
|
||||
if self.config.shuffle_camera_positions and ACTION in batch: # only during training
|
||||
present_img_keys = random.sample(present_img_keys, len(present_img_keys))
|
||||
if len(present_img_keys) == 0:
|
||||
raise ValueError(
|
||||
@@ -575,7 +576,7 @@ class SmolVLA2Policy(PreTrainedPolicy):
|
||||
padding_side="right",
|
||||
max_length=self.config.tokenizer_max_length,
|
||||
return_tensors="pt",
|
||||
truncation=True, # FIXME(mshukor)
|
||||
truncation=True, # FIXME(mshukor)
|
||||
)
|
||||
lang_tokens = tokenized_prompt["input_ids"].to(device=device)
|
||||
lang_masks = tokenized_prompt["attention_mask"].to(device=device, dtype=torch.bool)
|
||||
@@ -622,7 +623,9 @@ class SmolVLA2Policy(PreTrainedPolicy):
|
||||
if self.config.relative_actions_mode == "first":
|
||||
actions = torch.cat((actions[:, :1], actions[:, 1:] - actions[:, :1]), dim=1)
|
||||
elif self.config.relative_actions_mode == "state":
|
||||
assert batch[ACTION].shape[-1] == batch[OBS_STATE].shape[-1], "Relative action mode 'state' requires the action and state to have the same dimension."
|
||||
assert batch[ACTION].shape[-1] == batch[OBS_STATE].shape[-1], (
|
||||
"Relative action mode 'state' requires the action and state to have the same dimension."
|
||||
)
|
||||
if state.ndim == 2:
|
||||
state = state.unsqueeze(1)
|
||||
actions = actions - state
|
||||
@@ -700,7 +703,9 @@ class VLAFlowMatching(nn.Module):
|
||||
# Projections are float32
|
||||
self.state_to_prefix = self.config.state_to_prefix
|
||||
if self.state_to_prefix:
|
||||
self.state_proj = nn.Linear(self.config.max_state_dim, self.vlm_with_expert.config.text_config.hidden_size)
|
||||
self.state_proj = nn.Linear(
|
||||
self.config.max_state_dim, self.vlm_with_expert.config.text_config.hidden_size
|
||||
)
|
||||
else:
|
||||
self.state_proj = nn.Linear(self.config.max_state_dim, self.vlm_with_expert.expert_hidden_size)
|
||||
self.action_in_proj = nn.Linear(self.config.max_action_dim, self.vlm_with_expert.expert_hidden_size)
|
||||
@@ -714,7 +719,7 @@ class VLAFlowMatching(nn.Module):
|
||||
)
|
||||
|
||||
self.set_requires_grad()
|
||||
# SmolVLM2 has: [fake_tok + crop_tok + crop + fake_tok + crop_tok ... + fake_tok + global_tok + global + fake_tok] + [second image] + ...
|
||||
# SmolVLM2 has: [fake_tok + crop_tok + crop + fake_tok + crop_tok ... + fake_tok + global_tok + global + fake_tok] + [second image] + ...
|
||||
self.fake_image_token = self.vlm_with_expert.processor.tokenizer.fake_image_token_id
|
||||
self.global_image_token = self.vlm_with_expert.processor.tokenizer.global_image_token_id
|
||||
self.global_image_start_token = torch.tensor(
|
||||
@@ -724,11 +729,12 @@ class VLAFlowMatching(nn.Module):
|
||||
self.add_image_special_tokens = self.config.add_image_special_tokens
|
||||
self.image_end_token = torch.tensor([self.fake_image_token], dtype=torch.long)
|
||||
self.prefix_length = self.config.prefix_length
|
||||
self.include_past_images = self.config.n_obs_steps > 1 and "image" in self.config.past_obs_keys.split(",")
|
||||
self.include_past_images = self.config.n_obs_steps > 1 and "image" in self.config.past_obs_keys.split(
|
||||
","
|
||||
)
|
||||
self.num_past_images = self.config.n_obs_steps if self.include_past_images else 1
|
||||
self.causal_attention_on_history = self.config.causal_attention_on_history
|
||||
|
||||
|
||||
def set_requires_grad(self):
|
||||
for params in self.state_proj.parameters():
|
||||
params.requires_grad = self.config.train_state_proj
|
||||
@@ -749,14 +755,21 @@ class VLAFlowMatching(nn.Module):
|
||||
return time.to(dtype=torch.float32, device=device)
|
||||
|
||||
def embed_prefix(
|
||||
self, images, img_masks, lang_tokens, lang_masks, state: torch.Tensor = None,
|
||||
pointtrackers=None, pt_masks=None, **kwargs
|
||||
self,
|
||||
images,
|
||||
img_masks,
|
||||
lang_tokens,
|
||||
lang_masks,
|
||||
state: torch.Tensor = None,
|
||||
pointtrackers=None,
|
||||
pt_masks=None,
|
||||
**kwargs,
|
||||
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
|
||||
"""Embed multiple modalities for vlm processing.
|
||||
|
||||
|
||||
Simple, extensible approach using list + torch.cat.
|
||||
Easy to add new information/modalities like point trackers, audio, etc.
|
||||
|
||||
|
||||
Args:
|
||||
images: List of image tensors
|
||||
img_masks: List of image masks
|
||||
@@ -770,104 +783,110 @@ class VLAFlowMatching(nn.Module):
|
||||
embs = []
|
||||
pad_masks = []
|
||||
att_masks = []
|
||||
|
||||
|
||||
# Process each modality type
|
||||
self._add_image_embeddings(images, img_masks, embs, pad_masks, att_masks)
|
||||
self._add_language_embeddings(lang_tokens, lang_masks, embs, pad_masks, att_masks)
|
||||
|
||||
|
||||
if state is not None and self.state_to_prefix:
|
||||
self._add_state_embeddings(state, embs, pad_masks, att_masks)
|
||||
|
||||
|
||||
# Future extensions - easy to add new modalities
|
||||
if pointtrackers is not None:
|
||||
self._add_pointtracker_embeddings(pointtrackers, pt_masks, embs, pad_masks, att_masks)
|
||||
|
||||
|
||||
# Add more modalities here as needed:
|
||||
# if audio is not None:
|
||||
# self._add_audio_embeddings(audio, audio_masks, embs, pad_masks, att_masks)
|
||||
|
||||
|
||||
# Concatenate all embeddings
|
||||
embs = torch.cat(embs, dim=1)
|
||||
pad_masks = torch.cat(pad_masks, dim=1)
|
||||
att_masks = torch.tensor(att_masks, dtype=torch.bool, device=pad_masks.device)
|
||||
|
||||
|
||||
# Handle prefix length padding
|
||||
seq_len = pad_masks.shape[1]
|
||||
if seq_len < self.prefix_length:
|
||||
embs = pad_tensor(embs, self.prefix_length, pad_value=0)
|
||||
pad_masks = pad_tensor(pad_masks, self.prefix_length, pad_value=0)
|
||||
att_masks = pad_tensor(att_masks, self.prefix_length, pad_value=0)
|
||||
|
||||
|
||||
# Expand attention masks to batch size
|
||||
bsize = pad_masks.shape[0]
|
||||
att_masks = att_masks[None, :].expand(bsize, -1)
|
||||
|
||||
|
||||
return embs, pad_masks, att_masks
|
||||
|
||||
|
||||
def _add_image_embeddings(self, images, img_masks, embs, pad_masks, att_masks):
|
||||
"""Add image embeddings with special tokens to the lists."""
|
||||
for img, img_mask in zip(images, img_masks):
|
||||
for img, img_mask in zip(images, img_masks, strict=False):
|
||||
# Add image start tokens if enabled
|
||||
if self.add_image_special_tokens:
|
||||
start_emb = self.vlm_with_expert.embed_language_tokens(
|
||||
self.global_image_start_token.to(device=img.device)
|
||||
).unsqueeze(0).expand(img.shape[0], -1, -1)
|
||||
|
||||
start_emb = (
|
||||
self.vlm_with_expert.embed_language_tokens(
|
||||
self.global_image_start_token.to(device=img.device)
|
||||
)
|
||||
.unsqueeze(0)
|
||||
.expand(img.shape[0], -1, -1)
|
||||
)
|
||||
|
||||
start_mask = torch.ones_like(start_emb[:, :, 0], dtype=torch.bool)
|
||||
embs.append(start_emb)
|
||||
pad_masks.append(start_mask)
|
||||
att_masks += [0] * start_emb.shape[1]
|
||||
|
||||
|
||||
# Process image embedding
|
||||
img_emb = self.vlm_with_expert.embed_image(img)
|
||||
|
||||
|
||||
# Normalize image embeddings
|
||||
img_emb_dim = img_emb.shape[-1]
|
||||
img_emb = img_emb * torch.tensor(img_emb_dim**0.5, dtype=img_emb.dtype, device=img_emb.device)
|
||||
|
||||
|
||||
# Expand mask to match image embedding sequence length
|
||||
bsize, num_img_embs = img_emb.shape[:2]
|
||||
expanded_mask = img_mask[:, None].expand(bsize, num_img_embs)
|
||||
|
||||
|
||||
embs.append(img_emb)
|
||||
pad_masks.append(expanded_mask)
|
||||
att_masks += [0] * num_img_embs
|
||||
|
||||
|
||||
# Add image end tokens if enabled
|
||||
if self.add_image_special_tokens:
|
||||
end_emb = self.vlm_with_expert.embed_language_tokens(
|
||||
self.image_end_token.to(device=img.device)
|
||||
).unsqueeze(0).expand(img.shape[0], -1, -1)
|
||||
|
||||
end_emb = (
|
||||
self.vlm_with_expert.embed_language_tokens(self.image_end_token.to(device=img.device))
|
||||
.unsqueeze(0)
|
||||
.expand(img.shape[0], -1, -1)
|
||||
)
|
||||
|
||||
end_mask = torch.ones_like(end_emb[:, :, 0], dtype=torch.bool)
|
||||
embs.append(end_emb)
|
||||
pad_masks.append(end_mask)
|
||||
att_masks += [0] * end_emb.shape[1]
|
||||
|
||||
|
||||
def _add_language_embeddings(self, lang_tokens, lang_masks, embs, pad_masks, att_masks):
|
||||
"""Add language embeddings to the lists."""
|
||||
lang_emb = self.vlm_with_expert.embed_language_tokens(lang_tokens)
|
||||
|
||||
|
||||
# Normalize language embeddings
|
||||
lang_emb_dim = lang_emb.shape[-1]
|
||||
lang_emb = lang_emb * math.sqrt(lang_emb_dim)
|
||||
|
||||
|
||||
embs.append(lang_emb)
|
||||
pad_masks.append(lang_masks)
|
||||
att_masks += [0] * lang_emb.shape[1]
|
||||
|
||||
|
||||
def _add_state_embeddings(self, state, embs, pad_masks, att_masks):
|
||||
"""Add state embeddings to the lists."""
|
||||
state_emb = self.state_proj(state)
|
||||
state_emb = state_emb[:, None, :] if state_emb.ndim == 2 else state_emb
|
||||
|
||||
|
||||
bsize, states_seq_len = state_emb.shape[:2]
|
||||
state_mask = torch.ones(bsize, states_seq_len, dtype=torch.bool, device=state_emb.device)
|
||||
|
||||
|
||||
embs.append(state_emb)
|
||||
pad_masks.append(state_mask)
|
||||
att_masks += [1] * states_seq_len # State tokens get causal attention
|
||||
|
||||
|
||||
def _add_pointtracker_embeddings(self, pointtrackers, pt_masks, embs, pad_masks, att_masks):
|
||||
"""Add point tracker embeddings to the lists (future extension)."""
|
||||
# TODO: Implement point tracker processing
|
||||
@@ -884,10 +903,12 @@ class VLAFlowMatching(nn.Module):
|
||||
embs = []
|
||||
pad_masks = []
|
||||
att_masks = []
|
||||
# Embed state
|
||||
# Embed state
|
||||
if not self.state_to_prefix:
|
||||
state_emb = self.state_proj(state)
|
||||
state_emb = state_emb[:, None, :] if state_emb.ndim == 2 else state_emb #.to(dtype=self.vlm_with_expert.type)
|
||||
state_emb = (
|
||||
state_emb[:, None, :] if state_emb.ndim == 2 else state_emb
|
||||
) # .to(dtype=self.vlm_with_expert.type)
|
||||
embs.append(state_emb)
|
||||
bsize = state_emb.shape[0]
|
||||
dtype = state_emb.dtype
|
||||
@@ -898,7 +919,7 @@ class VLAFlowMatching(nn.Module):
|
||||
pad_masks.append(state_mask)
|
||||
|
||||
# Set attention masks so that image and language inputs do not attend to state or actions
|
||||
att_masks += [1] + [0]*(states_seq_len - 1)
|
||||
att_masks += [1] + [0] * (states_seq_len - 1)
|
||||
# Fuse timestep + action information using an MLP
|
||||
action_emb = self.action_in_proj(noisy_actions)
|
||||
device = action_emb.device
|
||||
@@ -1001,12 +1022,12 @@ class VLAFlowMatching(nn.Module):
|
||||
x_t = torch.zeros_like(noise, dtype=torch.float32, device=device)
|
||||
expanded_time = torch.zeros(bsize, dtype=torch.float32, device=device)
|
||||
x_t = self.denoise_step(
|
||||
state,
|
||||
prefix_pad_masks,
|
||||
past_key_values,
|
||||
x_t,
|
||||
expanded_time,
|
||||
)
|
||||
state,
|
||||
prefix_pad_masks,
|
||||
past_key_values,
|
||||
x_t,
|
||||
expanded_time,
|
||||
)
|
||||
else:
|
||||
dt = -1.0 / self.config.num_steps
|
||||
dt = torch.tensor(dt, dtype=torch.float32, device=device)
|
||||
|
||||
@@ -177,8 +177,10 @@ class SmolVLMWithExpertModel(nn.Module):
|
||||
else:
|
||||
self.vlm = self.vlm.merge_and_unload()
|
||||
|
||||
def get_vlm_model(self,):
|
||||
if hasattr(self.vlm.model, "model"): # When using peft
|
||||
def get_vlm_model(
|
||||
self,
|
||||
):
|
||||
if hasattr(self.vlm.model, "model"): # When using peft
|
||||
return self.vlm.model.model
|
||||
else:
|
||||
return self.vlm.model
|
||||
|
||||
@@ -16,6 +16,7 @@
|
||||
import logging
|
||||
import time
|
||||
from contextlib import nullcontext
|
||||
from functools import partial
|
||||
from pprint import pformat
|
||||
from typing import Any
|
||||
|
||||
@@ -29,6 +30,7 @@ from lerobot.configs.train import TrainPipelineConfig
|
||||
from lerobot.datasets.factory import make_dataset
|
||||
from lerobot.datasets.sampler import EpisodeAwareSampler
|
||||
from lerobot.datasets.utils import cycle
|
||||
from lerobot.datasets.utils_must import multidataset_collate_fn
|
||||
from lerobot.envs.factory import make_env
|
||||
from lerobot.optim.factory import make_optimizer_and_scheduler
|
||||
from lerobot.policies.factory import make_policy
|
||||
@@ -51,8 +53,7 @@ from lerobot.utils.utils import (
|
||||
init_logging,
|
||||
)
|
||||
from lerobot.utils.wandb_utils import WandBLogger
|
||||
from lerobot.datasets.utils_must import multidataset_collate_fn
|
||||
from functools import partial
|
||||
|
||||
|
||||
def update_policy(
|
||||
train_metrics: MetricsTracker,
|
||||
@@ -174,7 +175,7 @@ def train(cfg: TrainPipelineConfig):
|
||||
else:
|
||||
shuffle = True
|
||||
sampler = None
|
||||
|
||||
|
||||
keys_to_max_dim = getattr(dataset.meta, "keys_to_max_dim", {})
|
||||
collate_fn = partial(multidataset_collate_fn, keys_to_max_dim=keys_to_max_dim)
|
||||
dataloader = torch.utils.data.DataLoader(
|
||||
|
||||
Reference in New Issue
Block a user