mirror of
https://github.com/huggingface/lerobot.git
synced 2026-07-29 20:49:42 +00:00
Compare commits
8 Commits
| Author | SHA1 | Date | |
|---|---|---|---|
| 0371e99117 | |||
| 9e30807eeb | |||
| dd08d4eb53 | |||
| 6e5f6df6e7 | |||
| 265abe6c79 | |||
| b4e2d0b610 | |||
| 5594eba06a | |||
| 207183c2f8 |
+3
-2
@@ -321,10 +321,11 @@ SmolVLA ships with `freeze_vision_encoder=True`. Unfreezing usually **improves p
|
||||
|
||||
```bash
|
||||
lerobot-train ... --policy.type=smolvla \
|
||||
--policy.freeze_vision_encoder=false \
|
||||
--policy.train_expert_only=false
|
||||
--policy.fine_tune_vision_encoder=true
|
||||
```
|
||||
|
||||
This selectively trains the vision encoder and connector while leaving the language model frozen. Their learning rate defaults to `0.1 × optimizer_lr`; adjust it with `--policy.vision_encoder_lr_multiplier` if needed.
|
||||
|
||||
### 7.7 Signals to stop / keep going
|
||||
|
||||
- Train loss plateaus → stop, save a Hub checkpoint.
|
||||
|
||||
@@ -58,7 +58,7 @@ final_action = postprocessor(action)
|
||||
|
||||
## Hardware API redesign
|
||||
|
||||
PR [#777](https://github.com/huggingface/lerobot/pull/777) improves the LeRobot calibration but is **not backward-compatible**. Below is a overview of what changed and how you can continue to work with datasets created before this pull request.
|
||||
PR [#777](https://github.com/huggingface/lerobot/pull/777) improves the LeRobot calibration but is **not backward-compatible**. Below is an overview of what changed and how you can continue to work with datasets created before this pull request.
|
||||
|
||||
### What changed?
|
||||
|
||||
@@ -129,8 +129,8 @@ python examples/backward_compatibility/replay.py \
|
||||
|
||||
Policies output actions in the same format as the datasets (`torch.Tensors`). Therefore, the same transformations should be applied.
|
||||
|
||||
To find these transformations, we recommend to first try and and replay an episode of the dataset your policy was trained on using the section above.
|
||||
Then, add these same transformations on your inference script (shown here in the `record.py` script):
|
||||
To find these transformations, we recommend first replaying an episode of the dataset your policy was trained on using the section above.
|
||||
Then, add these same transformations to your inference script (shown here in the `record.py` script):
|
||||
|
||||
```diff
|
||||
action_values = predict_action(
|
||||
|
||||
@@ -40,10 +40,10 @@ This tutorial guides you through updating the firmware of Feetech motors using t
|
||||
For each motor you want to update:
|
||||
|
||||
1. **Select the motor** from the list by clicking on it
|
||||
2. **Click on Upgrade tab**:
|
||||
3. **Click on Online button**:
|
||||
- If an potential firmware update is found, it will be displayed in the box
|
||||
4. **Click on Upgrade button**:
|
||||
2. **Click the Upgrade tab**:
|
||||
3. **Click the Online button**:
|
||||
- If a potential firmware update is found, it will be displayed in the box
|
||||
4. **Click the Upgrade button**:
|
||||
- The update progress will be displayed
|
||||
|
||||
## Step 6: Verify Update
|
||||
|
||||
@@ -22,7 +22,7 @@ With processors, you choose the learning features you want to use for your polic
|
||||
## Three pipelines
|
||||
|
||||
We often compose three pipelines. Depending on your setup, some can be empty if action and observation spaces already match.
|
||||
Each of these pipelines handle different conversions between different action and observation spaces. Below is a quick explanation of each pipeline.
|
||||
Each of these pipelines handles different conversions between different action and observation spaces. Below is a quick explanation of each pipeline.
|
||||
|
||||
1. Pipeline 1: Teleop action space → dataset action space (phone pose → EE targets)
|
||||
2. Pipeline 2: Dataset action space → robot command space (EE targets → joints)
|
||||
@@ -74,15 +74,15 @@ In the phone to SO-100 follower examples we use the following adapters:
|
||||
- `robot_action_to_transition`: transforms the teleop action dict to a pipeline transition.
|
||||
- `transition_to_robot_action`: transforms the pipeline transition to a robot action dict.
|
||||
- `observation_to_transition`: transforms the robot observation dict to a pipeline transition.
|
||||
- `transition_to_observation`: transforms the pipeline transition to a observation dict.
|
||||
- `transition_to_observation`: transforms the pipeline transition to an observation dict.
|
||||
|
||||
Checkout [src/lerobot/processor/converters.py](https://github.com/huggingface/lerobot/blob/main/src/lerobot/processor/converters.py) for more details.
|
||||
Check out [src/lerobot/processor/converters.py](https://github.com/huggingface/lerobot/blob/main/src/lerobot/processor/converters.py) for more details.
|
||||
|
||||
## Dataset feature contracts
|
||||
|
||||
Dataset features are determined by the keys saved in the dataset. Each step can declare what features it modifies in a contract called `transform_features(...)`. Once you build a processor, the processor can then aggregate all of these features with `aggregate_pipeline_dataset_features()` and merge multiple feature dicts with `combine_feature_dicts(...)`.
|
||||
|
||||
Below is and example of how we declare features with the `transform_features` method in the phone to SO-100 follower examples:
|
||||
Below is an example of how we declare features with the `transform_features` method in the phone to SO-100 follower examples:
|
||||
|
||||
```python
|
||||
def transform_features(
|
||||
|
||||
+2
-2
@@ -57,7 +57,7 @@ policy_cfg.rtc_config = RTCConfig(
|
||||
policy = PI0Policy.from_pretrained("lerobot/pi0_base", policy_cfg=policy_cfg, device="cuda")
|
||||
|
||||
# Now use predict_action_chunk with RTC parameters
|
||||
inference_delay = 4 # How many steps of inference latency, this values should be calculated based on the inference latency of the policy
|
||||
inference_delay = 4 # How many steps of inference latency, this value should be calculated based on the inference latency of the policy
|
||||
|
||||
# Initialize the action queue
|
||||
action_queue = ActionQueue(policy_cfg.rtc_config)
|
||||
@@ -100,7 +100,7 @@ Typical values: 8-12 steps
|
||||
RTCConfig(execution_horizon=10)
|
||||
```
|
||||
|
||||
**`max_guidance_weight`**: How strongly to enforce consistency with the previous chunk. This is a hyperparameter that can be tuned to balance the smoothness of the transitions and the reactivity of the policy. For 10 steps flow matching (SmolVLA, Pi0, Pi0.5), a value of 10.0 is a optimal value.
|
||||
**`max_guidance_weight`**: How strongly to enforce consistency with the previous chunk. This is a hyperparameter that can be tuned to balance the smoothness of the transitions and the reactivity of the policy. For 10 steps flow matching (SmolVLA, Pi0, Pi0.5), a value of 10.0 is an optimal value.
|
||||
|
||||
**`prefix_attention_schedule`**: How to weight consistency across the overlap region.
|
||||
|
||||
|
||||
@@ -70,6 +70,19 @@ cd lerobot && lerobot-train \
|
||||
GPU allows it, as long as loading times remain short.
|
||||
</Tip>
|
||||
|
||||
For tasks that require adapting visual features, such as distinguishing new colors or shapes, selectively
|
||||
fine-tune the vision encoder and its connector:
|
||||
|
||||
```bash
|
||||
lerobot-train ... \
|
||||
--policy.path=lerobot/smolvla_base \
|
||||
--policy.fine_tune_vision_encoder=true
|
||||
```
|
||||
|
||||
This keeps the language model frozen with the default `train_expert_only=true` setting and trains the vision
|
||||
path at `0.1` times the main learning rate by default. Fine-tuning the vision encoder increases memory use and
|
||||
can reduce the model's general visual knowledge, so enable it only when the frozen encoder is insufficient.
|
||||
|
||||
Fine-tuning is an art. For a complete overview of the options for finetuning, run
|
||||
|
||||
```bash
|
||||
|
||||
@@ -50,11 +50,11 @@ lerobot-edit-dataset \
|
||||
Divide a dataset into multiple subsets.
|
||||
|
||||
```bash
|
||||
# Split by fractions (e.g. 80% train, 20% test, 20% val)
|
||||
# Split by fractions (e.g. 60% train, 20% val, 20% test)
|
||||
lerobot-edit-dataset \
|
||||
--repo_id lerobot/pusht \
|
||||
--operation.type split \
|
||||
--operation.splits '{"train": 0.8, "test": 0.2, "val": 0.2}'
|
||||
--operation.splits '{"train": 0.6, "val": 0.2, "test": 0.2}'
|
||||
|
||||
# Split by specific episode indices
|
||||
lerobot-edit-dataset \
|
||||
|
||||
Binary file not shown.
|
Before Width: | Height: | Size: 682 KiB |
@@ -19,6 +19,7 @@ import copy
|
||||
import logging
|
||||
import shutil
|
||||
from pathlib import Path
|
||||
from typing import Any, NotRequired, TypedDict
|
||||
|
||||
import datasets
|
||||
import pandas as pd
|
||||
@@ -49,8 +50,32 @@ from .utils import (
|
||||
)
|
||||
from .video_utils import concatenate_video_files, get_video_duration_in_s
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
def merge_video_feature_info_for_aggregate(all_metadata: list[LeRobotDatasetMetadata]) -> dict[str, dict]:
|
||||
type FeatureDict = dict[str, dict[str, Any]]
|
||||
type ChunkFile = tuple[int, int]
|
||||
|
||||
|
||||
class IndexState(TypedDict):
|
||||
chunk: int
|
||||
file: int
|
||||
src_to_dst: NotRequired[dict[ChunkFile, ChunkFile]]
|
||||
|
||||
|
||||
class VideoIndex(TypedDict):
|
||||
chunk: int
|
||||
file: int
|
||||
latest_duration: float
|
||||
episode_duration: float
|
||||
src_to_offset: NotRequired[dict[ChunkFile, float]]
|
||||
src_to_dst: NotRequired[dict[ChunkFile, ChunkFile]]
|
||||
dst_file_durations: NotRequired[dict[ChunkFile, float]]
|
||||
|
||||
|
||||
type VideoIndexState = dict[str, VideoIndex]
|
||||
|
||||
|
||||
def merge_video_feature_info_for_aggregate(all_metadata: list[LeRobotDatasetMetadata]) -> FeatureDict:
|
||||
"""Create a merged video feature info dictionary for aggregation. The video encoder info is merged field-by-field: each key is kept only when every source agrees; otherwise that key is set to ``null`` (or ``{}`` for ``video.extra_options``) and a warning is logged.
|
||||
|
||||
Args:
|
||||
@@ -59,14 +84,14 @@ def merge_video_feature_info_for_aggregate(all_metadata: list[LeRobotDatasetMeta
|
||||
Returns:
|
||||
dict: A dictionary of merged video feature info.
|
||||
"""
|
||||
merged_info = copy.deepcopy(all_metadata[0].features)
|
||||
merged_info: FeatureDict = copy.deepcopy(all_metadata[0].features)
|
||||
video_keys = [k for k in merged_info if merged_info[k].get("dtype") == "video"]
|
||||
|
||||
for vk in video_keys:
|
||||
video_infos = [m.features.get(vk, {}).get("info") or {} for m in all_metadata]
|
||||
base_video_info = video_infos[0]
|
||||
|
||||
merged_encoder_info: dict = {}
|
||||
merged_encoder_info: dict[str, Any] = {}
|
||||
fallback_keys: list[str] = []
|
||||
for info_key in VIDEO_ENCODER_INFO_KEYS:
|
||||
values = [info.get(info_key, None) for info in video_infos]
|
||||
@@ -80,7 +105,7 @@ def merge_video_feature_info_for_aggregate(all_metadata: list[LeRobotDatasetMeta
|
||||
merged_encoder_info[info_key] = {} if info_key == "video.extra_options" else None
|
||||
|
||||
if fallback_keys:
|
||||
logging.warning(
|
||||
logger.warning(
|
||||
f"Merging heterogeneous or incomplete video encoder metadata for feature {vk}. "
|
||||
f"Setting these keys to null: {fallback_keys}.",
|
||||
)
|
||||
@@ -92,7 +117,7 @@ def merge_video_feature_info_for_aggregate(all_metadata: list[LeRobotDatasetMeta
|
||||
return merged_info
|
||||
|
||||
|
||||
def validate_all_metadata(all_metadata: list[LeRobotDatasetMetadata]):
|
||||
def validate_all_metadata(all_metadata: list[LeRobotDatasetMetadata]) -> tuple[int, str | None, FeatureDict]:
|
||||
"""Validates that all dataset metadata have consistent properties.
|
||||
|
||||
Ensures all datasets have the same fps, robot_type, and features to guarantee
|
||||
@@ -129,7 +154,9 @@ def validate_all_metadata(all_metadata: list[LeRobotDatasetMetadata]):
|
||||
return fps, robot_type, features
|
||||
|
||||
|
||||
def update_data_df(df, src_meta, dst_meta):
|
||||
def update_data_df(
|
||||
df: pd.DataFrame, src_meta: LeRobotDatasetMetadata, dst_meta: LeRobotDatasetMetadata
|
||||
) -> pd.DataFrame:
|
||||
"""Updates a data DataFrame with new indices and task mappings for aggregation.
|
||||
|
||||
Adjusts episode indices, frame indices, and task indices to account for
|
||||
@@ -154,12 +181,12 @@ def update_data_df(df, src_meta, dst_meta):
|
||||
|
||||
|
||||
def update_meta_data(
|
||||
df,
|
||||
dst_meta,
|
||||
meta_idx,
|
||||
data_idx,
|
||||
videos_idx,
|
||||
):
|
||||
df: pd.DataFrame,
|
||||
dst_meta: LeRobotDatasetMetadata,
|
||||
meta_idx: IndexState,
|
||||
data_idx: IndexState,
|
||||
videos_idx: VideoIndexState,
|
||||
) -> pd.DataFrame:
|
||||
"""Updates metadata DataFrame with new chunk, file, and timestamp indices.
|
||||
|
||||
Adjusts all indices and timestamps to account for previously aggregated
|
||||
@@ -289,7 +316,7 @@ def aggregate_datasets(
|
||||
chunk_size: int | None = None,
|
||||
concatenate_videos: bool = True,
|
||||
concatenate_data: bool = True,
|
||||
):
|
||||
) -> None:
|
||||
"""Aggregates multiple LeRobot datasets into a single unified dataset.
|
||||
|
||||
This is the main function that orchestrates the aggregation process by:
|
||||
@@ -309,7 +336,7 @@ def aggregate_datasets(
|
||||
concatenate_videos: When False, keep one mp4 per source file instead of packing into shards.
|
||||
concatenate_data: When False, keep one parquet per source file instead of packing into shards.
|
||||
"""
|
||||
logging.info("Start aggregate_datasets")
|
||||
logger.info("Start aggregate_datasets")
|
||||
|
||||
if data_files_size_in_mb is None:
|
||||
data_files_size_in_mb = DEFAULT_DATA_FILE_SIZE_IN_MB
|
||||
@@ -341,15 +368,15 @@ def aggregate_datasets(
|
||||
video_files_size_in_mb=video_files_size_in_mb,
|
||||
)
|
||||
|
||||
logging.info("Find all tasks")
|
||||
logger.info("Find all tasks")
|
||||
unique_tasks = pd.concat([m.tasks for m in all_metadata]).index.unique()
|
||||
dst_meta.tasks = pd.DataFrame(
|
||||
{"task_index": range(len(unique_tasks))}, index=pd.Index(unique_tasks, name="task")
|
||||
)
|
||||
|
||||
meta_idx = {"chunk": 0, "file": 0}
|
||||
data_idx = {"chunk": 0, "file": 0}
|
||||
videos_idx = {
|
||||
meta_idx: IndexState = {"chunk": 0, "file": 0}
|
||||
data_idx: IndexState = {"chunk": 0, "file": 0}
|
||||
videos_idx: VideoIndexState = {
|
||||
key: {"chunk": 0, "file": 0, "latest_duration": 0, "episode_duration": 0} for key in video_keys
|
||||
}
|
||||
|
||||
@@ -373,12 +400,17 @@ def aggregate_datasets(
|
||||
dst_meta.info.total_frames += src_meta.total_frames
|
||||
|
||||
finalize_aggregation(dst_meta, all_metadata)
|
||||
logging.info("Aggregation complete.")
|
||||
logger.info("Aggregation complete.")
|
||||
|
||||
|
||||
def aggregate_videos(
|
||||
src_meta, dst_meta, videos_idx, video_files_size_in_mb, chunk_size, concatenate_videos=True
|
||||
):
|
||||
src_meta: LeRobotDatasetMetadata,
|
||||
dst_meta: LeRobotDatasetMetadata,
|
||||
videos_idx: VideoIndexState,
|
||||
video_files_size_in_mb: float,
|
||||
chunk_size: int,
|
||||
concatenate_videos: bool = True,
|
||||
) -> VideoIndexState:
|
||||
"""Aggregates video chunks from a source dataset into the destination dataset.
|
||||
|
||||
Handles video file concatenation and rotation based on file size limits.
|
||||
@@ -406,15 +438,16 @@ def aggregate_videos(
|
||||
videos_idx[key]["dst_file_durations"] = {}
|
||||
|
||||
for key, video_idx in videos_idx.items():
|
||||
unique_chunk_file_pairs = {
|
||||
(chunk, file)
|
||||
for chunk, file in zip(
|
||||
src_meta.episodes[f"videos/{key}/chunk_index"],
|
||||
src_meta.episodes[f"videos/{key}/file_index"],
|
||||
strict=False,
|
||||
)
|
||||
}
|
||||
unique_chunk_file_pairs = sorted(unique_chunk_file_pairs)
|
||||
unique_chunk_file_pairs: list[ChunkFile] = sorted(
|
||||
{
|
||||
(chunk, file)
|
||||
for chunk, file in zip(
|
||||
src_meta.episodes[f"videos/{key}/chunk_index"],
|
||||
src_meta.episodes[f"videos/{key}/file_index"],
|
||||
strict=False,
|
||||
)
|
||||
}
|
||||
)
|
||||
|
||||
chunk_idx = video_idx["chunk"]
|
||||
file_idx = video_idx["file"]
|
||||
@@ -489,7 +522,14 @@ def aggregate_videos(
|
||||
return videos_idx
|
||||
|
||||
|
||||
def aggregate_data(src_meta, dst_meta, data_idx, data_files_size_in_mb, chunk_size, concatenate_data=True):
|
||||
def aggregate_data(
|
||||
src_meta: LeRobotDatasetMetadata,
|
||||
dst_meta: LeRobotDatasetMetadata,
|
||||
data_idx: IndexState,
|
||||
data_files_size_in_mb: float,
|
||||
chunk_size: int,
|
||||
concatenate_data: bool = True,
|
||||
) -> IndexState:
|
||||
"""Aggregates data chunks from a source dataset into the destination dataset.
|
||||
|
||||
Reads source data files, updates indices to match the aggregated dataset,
|
||||
@@ -510,14 +550,16 @@ def aggregate_data(src_meta, dst_meta, data_idx, data_files_size_in_mb, chunk_si
|
||||
Returns:
|
||||
dict: Updated data_idx with current chunk and file indices.
|
||||
"""
|
||||
unique_chunk_file_ids = {
|
||||
(c, f)
|
||||
for c, f in zip(
|
||||
src_meta.episodes["data/chunk_index"], src_meta.episodes["data/file_index"], strict=False
|
||||
)
|
||||
}
|
||||
|
||||
unique_chunk_file_ids = sorted(unique_chunk_file_ids)
|
||||
unique_chunk_file_ids: list[ChunkFile] = sorted(
|
||||
{
|
||||
(c, f)
|
||||
for c, f in zip(
|
||||
src_meta.episodes["data/chunk_index"],
|
||||
src_meta.episodes["data/file_index"],
|
||||
strict=False,
|
||||
)
|
||||
}
|
||||
)
|
||||
contains_images = len(dst_meta.image_keys) > 0
|
||||
|
||||
# retrieve features schema for proper image typing in parquet
|
||||
@@ -525,7 +567,7 @@ def aggregate_data(src_meta, dst_meta, data_idx, data_files_size_in_mb, chunk_si
|
||||
|
||||
# Track source to destination file mapping for metadata update
|
||||
# This is critical for handling datasets that are already results of a merge
|
||||
src_to_dst: dict[tuple[int, int], tuple[int, int]] = {}
|
||||
src_to_dst: dict[ChunkFile, ChunkFile] = {}
|
||||
|
||||
for src_chunk_idx, src_file_idx in unique_chunk_file_ids:
|
||||
src_path = src_meta.root / DEFAULT_DATA_PATH.format(
|
||||
@@ -564,7 +606,13 @@ def aggregate_data(src_meta, dst_meta, data_idx, data_files_size_in_mb, chunk_si
|
||||
return data_idx
|
||||
|
||||
|
||||
def aggregate_metadata(src_meta, dst_meta, meta_idx, data_idx, videos_idx):
|
||||
def aggregate_metadata(
|
||||
src_meta: LeRobotDatasetMetadata,
|
||||
dst_meta: LeRobotDatasetMetadata,
|
||||
meta_idx: IndexState,
|
||||
data_idx: IndexState,
|
||||
videos_idx: VideoIndexState,
|
||||
) -> IndexState:
|
||||
"""Aggregates metadata from a source dataset into the destination dataset.
|
||||
|
||||
Reads source metadata files, updates all indices and timestamps,
|
||||
@@ -580,16 +628,16 @@ def aggregate_metadata(src_meta, dst_meta, meta_idx, data_idx, videos_idx):
|
||||
Returns:
|
||||
dict: Updated meta_idx with current chunk and file indices.
|
||||
"""
|
||||
chunk_file_ids = {
|
||||
(c, f)
|
||||
for c, f in zip(
|
||||
src_meta.episodes["meta/episodes/chunk_index"],
|
||||
src_meta.episodes["meta/episodes/file_index"],
|
||||
strict=False,
|
||||
)
|
||||
}
|
||||
|
||||
chunk_file_ids = sorted(chunk_file_ids)
|
||||
chunk_file_ids: list[ChunkFile] = sorted(
|
||||
{
|
||||
(c, f)
|
||||
for c, f in zip(
|
||||
src_meta.episodes["meta/episodes/chunk_index"],
|
||||
src_meta.episodes["meta/episodes/file_index"],
|
||||
strict=False,
|
||||
)
|
||||
}
|
||||
)
|
||||
for chunk_idx, file_idx in chunk_file_ids:
|
||||
src_path = src_meta.root / DEFAULT_EPISODES_PATH.format(chunk_index=chunk_idx, file_index=file_idx)
|
||||
df = pd.read_parquet(src_path)
|
||||
@@ -622,16 +670,16 @@ def aggregate_metadata(src_meta, dst_meta, meta_idx, data_idx, videos_idx):
|
||||
def append_or_create_parquet_file(
|
||||
df: pd.DataFrame,
|
||||
src_path: Path,
|
||||
idx: dict[str, int],
|
||||
idx: IndexState,
|
||||
max_mb: float,
|
||||
chunk_size: int,
|
||||
default_path: str,
|
||||
contains_images: bool = False,
|
||||
aggr_root: Path = None,
|
||||
aggr_root: Path | None = None,
|
||||
hf_features: datasets.Features | None = None,
|
||||
concatenate: bool = True,
|
||||
one_row_group_per_episode: bool = False,
|
||||
) -> tuple[dict[str, int], tuple[int, int]]:
|
||||
) -> tuple[IndexState, ChunkFile]:
|
||||
"""Appends data to an existing parquet file or creates a new one based on size constraints.
|
||||
|
||||
Manages file rotation when size limits are exceeded to prevent individual files
|
||||
@@ -654,7 +702,13 @@ def append_or_create_parquet_file(
|
||||
Returns:
|
||||
tuple: (updated_idx, (dst_chunk, dst_file)) where updated_idx is the index dict
|
||||
and (dst_chunk, dst_file) is the actual destination file the data was written to.
|
||||
|
||||
Raises:
|
||||
ValueError: If aggr_root is not provided.
|
||||
"""
|
||||
if aggr_root is None:
|
||||
raise ValueError("aggr_root must be provided.")
|
||||
|
||||
dst_chunk, dst_file = idx["chunk"], idx["file"]
|
||||
dst_path = aggr_root / default_path.format(chunk_index=dst_chunk, file_index=dst_file)
|
||||
|
||||
@@ -698,7 +752,9 @@ def append_or_create_parquet_file(
|
||||
return idx, (dst_chunk, dst_file)
|
||||
|
||||
|
||||
def finalize_aggregation(aggr_meta, all_metadata):
|
||||
def finalize_aggregation(
|
||||
aggr_meta: LeRobotDatasetMetadata, all_metadata: list[LeRobotDatasetMetadata]
|
||||
) -> None:
|
||||
"""Finalizes the dataset aggregation by writing summary files and statistics.
|
||||
|
||||
Writes the tasks file, info file with total counts and splits, and
|
||||
@@ -708,16 +764,16 @@ def finalize_aggregation(aggr_meta, all_metadata):
|
||||
aggr_meta: Aggregated dataset metadata.
|
||||
all_metadata: List of all source dataset metadata objects.
|
||||
"""
|
||||
logging.info("write tasks")
|
||||
logger.info("write tasks")
|
||||
write_tasks(aggr_meta.tasks, aggr_meta.root)
|
||||
|
||||
logging.info("write info")
|
||||
logger.info("write info")
|
||||
aggr_meta.info.total_tasks = len(aggr_meta.tasks)
|
||||
aggr_meta.info.total_episodes = sum(m.total_episodes for m in all_metadata)
|
||||
aggr_meta.info.total_frames = sum(m.total_frames for m in all_metadata)
|
||||
aggr_meta.info.splits = {"train": f"0:{sum(m.total_episodes for m in all_metadata)}"}
|
||||
write_info(aggr_meta.info, aggr_meta.root)
|
||||
|
||||
logging.info("write stats")
|
||||
logger.info("write stats")
|
||||
aggr_meta.stats = aggregate_stats([m.stats for m in all_metadata])
|
||||
write_stats(aggr_meta.stats, aggr_meta.root)
|
||||
|
||||
@@ -302,6 +302,33 @@ def _pad_evo1_stats(
|
||||
return padded_stats
|
||||
|
||||
|
||||
def _refresh_evo1_normalization_steps(
|
||||
config: Evo1Config,
|
||||
preprocessor: PolicyProcessorPipeline,
|
||||
postprocessor: PolicyProcessorPipeline,
|
||||
) -> None:
|
||||
"""Re-pad checkpoint-loaded (un)normalizer stats/features to EVO1's fixed widths.
|
||||
|
||||
Loading a checkpoint injects the raw dataset stats (unpadded to max_state_dim/max_action_dim)
|
||||
into the (un)normalizer via the generic override path in make_pre_post_processors. Those stats
|
||||
and their declared features must be re-padded/reshaped to EVO1's fixed widths, otherwise
|
||||
normalization fails against the padded state/action tensors (e.g. state padded to 24 vs. 8-dim
|
||||
LIBERO stats). Padding is a no-op when stats are already at the target width.
|
||||
"""
|
||||
normalization_features = _evo1_normalization_features(config)
|
||||
action_features = _evo1_action_features(config)
|
||||
for step in preprocessor.steps:
|
||||
if isinstance(step, NormalizerProcessorStep):
|
||||
step.features = normalization_features
|
||||
step.stats = _pad_evo1_stats(config, step.stats)
|
||||
step.to(device=step.device, dtype=step.dtype)
|
||||
for step in postprocessor.steps:
|
||||
if isinstance(step, UnnormalizerProcessorStep):
|
||||
step.features = action_features
|
||||
step.stats = _pad_evo1_stats(config, step.stats)
|
||||
step.to(device=step.device, dtype=step.dtype)
|
||||
|
||||
|
||||
def reconcile_evo1_processors(
|
||||
config: Evo1Config,
|
||||
preprocessor: PolicyProcessorPipeline,
|
||||
@@ -309,16 +336,19 @@ def reconcile_evo1_processors(
|
||||
) -> tuple[PolicyProcessorPipeline, PolicyProcessorPipeline]:
|
||||
"""Reconcile checkpoint-loaded pipelines with the current EVO1 config.
|
||||
|
||||
Two things cannot be restored from a serialized pipeline alone: the EVO1 batch converter
|
||||
(converters are plain functions and are never serialized), and eval-time CLI overrides of the
|
||||
action postprocessing flags (`postprocess_action_dim`, `binarize_gripper`, `gripper_*`). This
|
||||
restores the converter and rebuilds the action step from the current config so those overrides
|
||||
take effect.
|
||||
Three things cannot be restored from a serialized pipeline alone: the EVO1 batch converter
|
||||
(converters are plain functions and are never serialized), eval-time CLI overrides of the
|
||||
action postprocessing flags (`postprocess_action_dim`, `binarize_gripper`, `gripper_*`), and the
|
||||
(un)normalizer stats/features when the generic override path injects raw, unpadded dataset
|
||||
stats. This restores the converter, re-pads the normalization stats to EVO1's fixed widths, and
|
||||
rebuilds the action step from the current config so those overrides take effect.
|
||||
"""
|
||||
# Pipelines reloaded from a checkpoint come back with the default batch converter, which drops
|
||||
# non-observation extras (embodiment_id, state_mask, custom task fields) needed by EVO1.
|
||||
preprocessor.to_transition = evo1_batch_to_transition
|
||||
|
||||
_refresh_evo1_normalization_steps(config, preprocessor, postprocessor)
|
||||
|
||||
action_step = Evo1ActionProcessorStep(
|
||||
action_dim=_evo1_action_dim(config),
|
||||
binarize_gripper=config.binarize_gripper,
|
||||
|
||||
@@ -67,6 +67,8 @@ class SmolVLAConfig(PreTrainedConfig):
|
||||
|
||||
# Finetuning settings
|
||||
freeze_vision_encoder: bool = True
|
||||
fine_tune_vision_encoder: bool = False # Fine-tune vision + connector; takes priority over freezing.
|
||||
vision_encoder_lr_multiplier: float = 0.1
|
||||
train_expert_only: bool = True
|
||||
train_state_proj: bool = True
|
||||
|
||||
@@ -110,6 +112,12 @@ class SmolVLAConfig(PreTrainedConfig):
|
||||
super().__post_init__()
|
||||
|
||||
"""Input validation (not exhaustive)."""
|
||||
if self.fine_tune_vision_encoder:
|
||||
self.freeze_vision_encoder = False
|
||||
if self.vision_encoder_lr_multiplier <= 0:
|
||||
raise ValueError(
|
||||
f"`vision_encoder_lr_multiplier` must be positive, got {self.vision_encoder_lr_multiplier}."
|
||||
)
|
||||
if self.n_action_steps > self.chunk_size:
|
||||
raise ValueError(
|
||||
f"The chunk size is the upper bound for the number of action steps per model invocation. Got "
|
||||
|
||||
@@ -186,8 +186,27 @@ class SmolVLAPolicy(PreTrainedPolicy):
|
||||
if model_value is not None:
|
||||
model_value.rtc_processor = self.rtc_processor
|
||||
|
||||
def get_optim_params(self) -> dict:
|
||||
return self.parameters()
|
||||
def get_optim_params(self):
|
||||
if not self.config.fine_tune_vision_encoder:
|
||||
return self.parameters()
|
||||
|
||||
vision_params = []
|
||||
other_params = []
|
||||
for name, param in self.named_parameters():
|
||||
if not param.requires_grad:
|
||||
continue
|
||||
if ".vision_model." in name or ".connector." in name:
|
||||
vision_params.append(param)
|
||||
else:
|
||||
other_params.append(param)
|
||||
|
||||
return [
|
||||
{"params": other_params},
|
||||
{
|
||||
"params": vision_params,
|
||||
"lr": self.config.optimizer_lr * self.config.vision_encoder_lr_multiplier,
|
||||
},
|
||||
]
|
||||
|
||||
def _get_action_chunk(
|
||||
self, batch: dict[str, Tensor], noise: Tensor | None = None, **kwargs: Unpack[ActionSelectKwargs]
|
||||
@@ -493,6 +512,7 @@ class VLAFlowMatching(nn.Module):
|
||||
self.vlm_with_expert = SmolVLMWithExpertModel(
|
||||
model_id=self.config.vlm_model_name,
|
||||
freeze_vision_encoder=self.config.freeze_vision_encoder,
|
||||
fine_tune_vision_encoder=self.config.fine_tune_vision_encoder,
|
||||
train_expert_only=self.config.train_expert_only,
|
||||
load_vlm_weights=self.config.load_vlm_weights,
|
||||
attention_mode=self.config.attention_mode,
|
||||
|
||||
@@ -78,6 +78,7 @@ class SmolVLMWithExpertModel(nn.Module):
|
||||
load_vlm_weights: bool = True,
|
||||
train_expert_only: bool = True,
|
||||
freeze_vision_encoder: bool = False,
|
||||
fine_tune_vision_encoder: bool = False,
|
||||
attention_mode: str = "self_attn",
|
||||
num_expert_layers: int = -1,
|
||||
num_vlm_layers: int = -1,
|
||||
@@ -141,6 +142,7 @@ class SmolVLMWithExpertModel(nn.Module):
|
||||
self.num_key_value_heads = self.config.text_config.num_key_value_heads
|
||||
|
||||
self.freeze_vision_encoder = freeze_vision_encoder
|
||||
self.fine_tune_vision_encoder = fine_tune_vision_encoder
|
||||
self.train_expert_only = train_expert_only
|
||||
self.attention_mode = attention_mode
|
||||
self.expert_hidden_size = lm_expert_config.hidden_size
|
||||
@@ -150,10 +152,6 @@ class SmolVLMWithExpertModel(nn.Module):
|
||||
return self.vlm.model
|
||||
|
||||
def set_requires_grad(self):
|
||||
if self.freeze_vision_encoder:
|
||||
self.get_vlm_model().vision_model.eval()
|
||||
for params in self.get_vlm_model().vision_model.parameters():
|
||||
params.requires_grad = False
|
||||
if self.train_expert_only:
|
||||
self.vlm.eval()
|
||||
for params in self.vlm.parameters():
|
||||
@@ -176,6 +174,18 @@ class SmolVLMWithExpertModel(nn.Module):
|
||||
for name, params in self.vlm.named_parameters():
|
||||
if any(k in name for k in frozen_layers):
|
||||
params.requires_grad = False
|
||||
|
||||
if self.freeze_vision_encoder:
|
||||
self.get_vlm_model().vision_model.eval()
|
||||
for params in self.get_vlm_model().vision_model.parameters():
|
||||
params.requires_grad = False
|
||||
|
||||
if self.fine_tune_vision_encoder:
|
||||
for params in self.get_vlm_model().vision_model.parameters():
|
||||
params.requires_grad = True
|
||||
for params in self.get_vlm_model().connector.parameters():
|
||||
params.requires_grad = True
|
||||
|
||||
# To avoid unused params issue with distributed training
|
||||
for name, params in self.lm_expert.named_parameters():
|
||||
if "lm_head" in name:
|
||||
@@ -184,11 +194,15 @@ class SmolVLMWithExpertModel(nn.Module):
|
||||
def train(self, mode: bool = True):
|
||||
super().train(mode)
|
||||
|
||||
if self.train_expert_only:
|
||||
self.vlm.eval()
|
||||
|
||||
if self.freeze_vision_encoder:
|
||||
self.get_vlm_model().vision_model.eval()
|
||||
|
||||
if self.train_expert_only:
|
||||
self.vlm.eval()
|
||||
if self.fine_tune_vision_encoder:
|
||||
self.get_vlm_model().vision_model.train(mode)
|
||||
self.get_vlm_model().connector.train(mode)
|
||||
|
||||
def embed_image(self, image: torch.Tensor):
|
||||
patch_attention_mask = None
|
||||
|
||||
@@ -510,10 +510,10 @@ class ForwardKinematicsJointsToEEAction(RobotActionProcessorStep):
|
||||
# We only use the ee pose in the dataset, so we don't need the joint positions
|
||||
for n in self.motor_names:
|
||||
features[PipelineFeatureType.ACTION].pop(f"{n}.pos", None)
|
||||
# We specify the dataset features of this step that we want to be stored in the dataset
|
||||
# Store end-effector features as actions in the dataset schema
|
||||
for k in ["x", "y", "z", "wx", "wy", "wz", "gripper_pos"]:
|
||||
features[PipelineFeatureType.ACTION][f"ee.{k}"] = PolicyFeature(
|
||||
type=FeatureType.STATE, shape=(1,)
|
||||
type=FeatureType.ACTION, shape=(1,)
|
||||
)
|
||||
return features
|
||||
|
||||
|
||||
@@ -13,34 +13,18 @@
|
||||
# limitations under the License.
|
||||
|
||||
from .transforms import (
|
||||
CoarseDropout,
|
||||
GammaCorrection,
|
||||
GaussianNoise,
|
||||
GaussianPatchBrightness,
|
||||
ImageTransformConfig,
|
||||
ImageTransforms,
|
||||
ImageTransformsConfig,
|
||||
JPEGCompression,
|
||||
MotionBlur,
|
||||
PlanckianJitter,
|
||||
RandomShadow,
|
||||
RandomSubsetApply,
|
||||
SharpnessJitter,
|
||||
make_transform_from_config,
|
||||
)
|
||||
|
||||
__all__ = [
|
||||
"CoarseDropout",
|
||||
"GammaCorrection",
|
||||
"GaussianNoise",
|
||||
"GaussianPatchBrightness",
|
||||
"ImageTransformConfig",
|
||||
"ImageTransforms",
|
||||
"ImageTransformsConfig",
|
||||
"JPEGCompression",
|
||||
"MotionBlur",
|
||||
"PlanckianJitter",
|
||||
"RandomShadow",
|
||||
"RandomSubsetApply",
|
||||
"SharpnessJitter",
|
||||
"make_transform_from_config",
|
||||
|
||||
@@ -14,13 +14,11 @@
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
import collections
|
||||
import math
|
||||
from collections.abc import Callable, Sequence
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Any
|
||||
|
||||
import torch
|
||||
from torchvision.io import decode_image, encode_jpeg
|
||||
from torchvision.transforms import v2
|
||||
from torchvision.transforms.v2 import (
|
||||
Transform,
|
||||
@@ -146,471 +144,6 @@ class SharpnessJitter(Transform):
|
||||
return self._call_kernel(F.adjust_sharpness, inpt, sharpness_factor=sharpness_factor)
|
||||
|
||||
|
||||
class GaussianNoise(Transform):
|
||||
"""Add Gaussian noise to simulate camera sensor noise.
|
||||
|
||||
Models readout noise from ADC quantization, which increases in low-light conditions.
|
||||
Common in real-robot setups where wrist cameras operate in suboptimal lighting.
|
||||
|
||||
Args:
|
||||
std: Range (min, max) for noise standard deviation in pixel-value scale (0-255).
|
||||
"""
|
||||
|
||||
def __init__(self, std: float | Sequence[float] = (5.0, 25.0)) -> None:
|
||||
super().__init__()
|
||||
if isinstance(std, (int, float)):
|
||||
self.std = (0.0, float(std))
|
||||
elif isinstance(std, Sequence) and len(std) == 2:
|
||||
self.std = (float(std[0]), float(std[1]))
|
||||
else:
|
||||
raise TypeError("std must be a number or a sequence with length 2.")
|
||||
if not 0.0 <= self.std[0] <= self.std[1]:
|
||||
raise ValueError(f"std must satisfy 0 <= min <= max, but got {self.std}.")
|
||||
|
||||
def make_params(self, flat_inputs: list[Any]) -> dict[str, Any]:
|
||||
return {
|
||||
"std": torch.empty(1).uniform_(self.std[0], self.std[1]).item(),
|
||||
"seed": torch.randint(0, torch.iinfo(torch.int64).max, ()).item(),
|
||||
}
|
||||
|
||||
def transform(self, inpt: Any, params: dict[str, Any]) -> Any:
|
||||
if isinstance(inpt, torch.Tensor) and inpt.is_floating_point():
|
||||
generator = torch.Generator(device=inpt.device).manual_seed(params["seed"])
|
||||
noise = torch.randn(inpt.shape, device=inpt.device, dtype=inpt.dtype, generator=generator)
|
||||
return (inpt + noise * (params["std"] / 255.0)).clamp(0.0, 1.0)
|
||||
return inpt
|
||||
|
||||
|
||||
class MotionBlur(Transform):
|
||||
"""Apply directional motion blur to simulate fast robot or object movement.
|
||||
|
||||
Generates a 1D averaging kernel along a random direction, applied via depthwise convolution.
|
||||
|
||||
Args:
|
||||
kernel_size: An odd kernel size or a range containing at least one odd kernel size.
|
||||
"""
|
||||
|
||||
def __init__(self, kernel_size: int | Sequence[int] = (3, 11)) -> None:
|
||||
super().__init__()
|
||||
if isinstance(kernel_size, int):
|
||||
self.kernel_size = (kernel_size, kernel_size)
|
||||
elif isinstance(kernel_size, Sequence) and len(kernel_size) == 2:
|
||||
self.kernel_size = (int(kernel_size[0]), int(kernel_size[1]))
|
||||
else:
|
||||
raise TypeError("kernel_size must be an int or a sequence with length 2.")
|
||||
if not 1 <= self.kernel_size[0] <= self.kernel_size[1]:
|
||||
raise ValueError(f"kernel_size must satisfy 1 <= min <= max, but got {self.kernel_size}.")
|
||||
self._first_odd_kernel_size = self.kernel_size[0] + (self.kernel_size[0] + 1) % 2
|
||||
if self._first_odd_kernel_size > self.kernel_size[1]:
|
||||
raise ValueError(f"kernel_size range must contain an odd value, but got {self.kernel_size}.")
|
||||
|
||||
def make_params(self, flat_inputs: list[Any]) -> dict[str, Any]:
|
||||
num_odd_sizes = (self.kernel_size[1] - self._first_odd_kernel_size) // 2 + 1
|
||||
size_index = int(torch.randint(0, num_odd_sizes, ()).item())
|
||||
ks = self._first_odd_kernel_size + 2 * size_index
|
||||
angle = torch.empty(1).uniform_(0, 360).item()
|
||||
return {"kernel_size": ks, "angle": angle}
|
||||
|
||||
def transform(self, inpt: Any, params: dict[str, Any]) -> Any:
|
||||
if not isinstance(inpt, torch.Tensor) or not inpt.is_floating_point():
|
||||
return inpt
|
||||
if inpt.ndim < 3:
|
||||
raise ValueError(f"MotionBlur expects [..., C, H, W] input, but got shape {inpt.shape}.")
|
||||
|
||||
kernel_size = params["kernel_size"]
|
||||
radius = kernel_size // 2
|
||||
angle = math.radians(params["angle"])
|
||||
positions = torch.linspace(-radius, radius, kernel_size, device=inpt.device)
|
||||
x_coords = (positions * math.cos(angle)).round().to(torch.long) + radius
|
||||
y_coords = (positions * math.sin(angle)).round().to(torch.long) + radius
|
||||
kernel = torch.zeros((kernel_size, kernel_size), device=inpt.device, dtype=inpt.dtype)
|
||||
kernel[y_coords, x_coords] = 1
|
||||
kernel /= kernel.sum()
|
||||
|
||||
channels, height, width = inpt.shape[-3:]
|
||||
flat_input = inpt.reshape(-1, channels, height, width)
|
||||
depthwise_kernel = kernel.expand(channels, 1, kernel_size, kernel_size)
|
||||
padded = torch.nn.functional.pad(flat_input, (radius,) * 4, mode="replicate")
|
||||
output = torch.nn.functional.conv2d(padded, depthwise_kernel, groups=channels)
|
||||
return output.reshape(inpt.shape).clamp(0.0, 1.0)
|
||||
|
||||
|
||||
class JPEGCompression(Transform):
|
||||
"""Simulate JPEG compression artifacts (block artifacts, color banding).
|
||||
|
||||
Models quality degradation from video compression in network-streamed camera feeds.
|
||||
|
||||
Args:
|
||||
quality: Range (min, max) for JPEG quality factor (lower = more artifacts).
|
||||
"""
|
||||
|
||||
def __init__(self, quality: int | Sequence[int] = (15, 75)) -> None:
|
||||
super().__init__()
|
||||
if isinstance(quality, int):
|
||||
self.quality = (quality, quality)
|
||||
elif isinstance(quality, Sequence) and len(quality) == 2:
|
||||
self.quality = (int(quality[0]), int(quality[1]))
|
||||
else:
|
||||
raise TypeError("quality must be an int or a sequence with length 2.")
|
||||
if not 1 <= self.quality[0] <= self.quality[1] <= 100:
|
||||
raise ValueError(f"quality must satisfy 1 <= min <= max <= 100, but got {self.quality}.")
|
||||
|
||||
def make_params(self, flat_inputs: list[Any]) -> dict[str, Any]:
|
||||
return {"quality": int(torch.randint(self.quality[0], self.quality[1] + 1, (1,)).item())}
|
||||
|
||||
def transform(self, inpt: Any, params: dict[str, Any]) -> Any:
|
||||
if not isinstance(inpt, torch.Tensor) or not inpt.is_floating_point():
|
||||
return inpt
|
||||
if inpt.ndim < 3:
|
||||
raise ValueError(f"JPEGCompression expects [..., C, H, W] input, but got shape {inpt.shape}.")
|
||||
|
||||
channels, height, width = inpt.shape[-3:]
|
||||
if channels not in (1, 3):
|
||||
raise ValueError(f"JPEGCompression expects 1 or 3 channels, but got {channels}.")
|
||||
|
||||
flat_input = inpt.reshape(-1, channels, height, width)
|
||||
flat_uint8 = (flat_input.clamp(0.0, 1.0) * 255).round().to(torch.uint8).cpu()
|
||||
decoded_frames = [
|
||||
decode_image(encode_jpeg(frame, quality=params["quality"])) for frame in flat_uint8.unbind()
|
||||
]
|
||||
output = torch.stack(decoded_frames).to(device=inpt.device, dtype=inpt.dtype) / 255.0
|
||||
return output.reshape(inpt.shape)
|
||||
|
||||
|
||||
class GaussianPatchBrightness(Transform):
|
||||
"""Apply spatially-varying brightness with Gaussian patches.
|
||||
|
||||
Simulates uneven overhead lighting, spotlights, and shadow patches commonly
|
||||
encountered in real robot workspaces with multiple light sources.
|
||||
|
||||
Args:
|
||||
num_patches: Range (min, max) for number of brightness patches.
|
||||
sigma_range: Range for Gaussian sigma as fraction of image size.
|
||||
factor_range: Range for brightness factor (< 1 darkens, > 1 brightens).
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
num_patches: int | Sequence[int] = (1, 4),
|
||||
sigma_range: Sequence[float] = (0.05, 0.25),
|
||||
factor_range: Sequence[float] = (0.4, 1.6),
|
||||
) -> None:
|
||||
super().__init__()
|
||||
if isinstance(num_patches, int):
|
||||
self.num_patches = (num_patches, num_patches)
|
||||
elif isinstance(num_patches, Sequence) and len(num_patches) == 2:
|
||||
self.num_patches = (int(num_patches[0]), int(num_patches[1]))
|
||||
else:
|
||||
raise TypeError("num_patches must be an int or a sequence with length 2.")
|
||||
if not 1 <= self.num_patches[0] <= self.num_patches[1]:
|
||||
raise ValueError(f"num_patches must satisfy 1 <= min <= max, but got {self.num_patches}.")
|
||||
if not isinstance(sigma_range, Sequence) or len(sigma_range) != 2:
|
||||
raise TypeError("sigma_range must be a sequence with length 2.")
|
||||
self.sigma_range = (float(sigma_range[0]), float(sigma_range[1]))
|
||||
if not 0.0 < self.sigma_range[0] <= self.sigma_range[1]:
|
||||
raise ValueError(f"sigma_range must satisfy 0 < min <= max, but got {self.sigma_range}.")
|
||||
if not isinstance(factor_range, Sequence) or len(factor_range) != 2:
|
||||
raise TypeError("factor_range must be a sequence with length 2.")
|
||||
self.factor_range = (float(factor_range[0]), float(factor_range[1]))
|
||||
if not 0.0 <= self.factor_range[0] <= self.factor_range[1]:
|
||||
raise ValueError(f"factor_range must satisfy 0 <= min <= max, but got {self.factor_range}.")
|
||||
|
||||
def make_params(self, flat_inputs: list[Any]) -> dict[str, Any]:
|
||||
n = int(torch.randint(self.num_patches[0], self.num_patches[1] + 1, (1,)).item())
|
||||
return {
|
||||
"centers": torch.rand(n, 2).tolist(),
|
||||
"sigmas": torch.empty(n).uniform_(self.sigma_range[0], self.sigma_range[1]).tolist(),
|
||||
"factors": torch.empty(n).uniform_(self.factor_range[0], self.factor_range[1]).tolist(),
|
||||
}
|
||||
|
||||
def transform(self, inpt: Any, params: dict[str, Any]) -> Any:
|
||||
if not isinstance(inpt, torch.Tensor) or not inpt.is_floating_point():
|
||||
return inpt
|
||||
h, w = inpt.shape[-2:]
|
||||
mask = torch.ones(h, w, device=inpt.device, dtype=inpt.dtype)
|
||||
grid_y = torch.linspace(0, 1, h, device=inpt.device, dtype=inpt.dtype)
|
||||
grid_x = torch.linspace(0, 1, w, device=inpt.device, dtype=inpt.dtype)
|
||||
yy, xx = torch.meshgrid(grid_y, grid_x, indexing="ij")
|
||||
for (cy, cx), sigma, factor in zip(
|
||||
params["centers"], params["sigmas"], params["factors"], strict=True
|
||||
):
|
||||
gauss = torch.exp(-((yy - cy) ** 2 + (xx - cx) ** 2) / (2 * sigma**2))
|
||||
mask = mask * (1.0 + (factor - 1.0) * gauss)
|
||||
broadcast_shape = (1,) * (inpt.ndim - 2) + (h, w)
|
||||
return (inpt * mask.reshape(broadcast_shape)).clamp(0.0, 1.0)
|
||||
|
||||
|
||||
class RandomShadow(Transform):
|
||||
"""Add random vertical band shadow with smooth edges.
|
||||
|
||||
Simulates cast shadows from objects or people near the robot workspace.
|
||||
Symmetric: randomly brightens or darkens to prevent BatchNorm stats shift.
|
||||
|
||||
Args:
|
||||
opacity: Range (min, max) for shadow/highlight opacity.
|
||||
"""
|
||||
|
||||
def __init__(self, opacity: float | Sequence[float] = (0.3, 0.6)) -> None:
|
||||
super().__init__()
|
||||
if isinstance(opacity, (int, float)):
|
||||
self.opacity = (float(opacity), float(opacity))
|
||||
elif isinstance(opacity, Sequence) and len(opacity) == 2:
|
||||
self.opacity = (float(opacity[0]), float(opacity[1]))
|
||||
else:
|
||||
raise TypeError("opacity must be a number or a sequence with length 2.")
|
||||
if not 0.0 <= self.opacity[0] <= self.opacity[1] <= 1.0:
|
||||
raise ValueError(f"opacity must satisfy 0 <= min <= max <= 1, but got {self.opacity}.")
|
||||
|
||||
def make_params(self, flat_inputs: list[Any]) -> dict[str, Any]:
|
||||
return {
|
||||
"opacity": torch.empty(1).uniform_(self.opacity[0], self.opacity[1]).item(),
|
||||
"start": torch.rand(1).item(),
|
||||
"width": torch.empty(1).uniform_(1 / 3, 2 / 3).item(),
|
||||
"direction": -1.0 if torch.rand(1).item() < 0.5 else 1.0,
|
||||
}
|
||||
|
||||
def transform(self, inpt: Any, params: dict[str, Any]) -> Any:
|
||||
if not isinstance(inpt, torch.Tensor) or not inpt.is_floating_point():
|
||||
return inpt
|
||||
if inpt.ndim < 3:
|
||||
raise ValueError(f"RandomShadow expects [..., C, H, W] input, but got shape {inpt.shape}.")
|
||||
|
||||
h, w = inpt.shape[-2:]
|
||||
band_width = max(1, min(w, round(params["width"] * w)))
|
||||
x_start = round(params["start"] * (w - band_width))
|
||||
x_end = x_start + band_width
|
||||
mask = torch.ones(h, w, device=inpt.device, dtype=inpt.dtype)
|
||||
mask[:, x_start:x_end] = 1.0 + params["direction"] * params["opacity"]
|
||||
|
||||
smoothing_size = min(8, h, w)
|
||||
if smoothing_size > 1:
|
||||
batched_mask = mask[None, None]
|
||||
small = torch.nn.functional.avg_pool2d(batched_mask, smoothing_size, stride=smoothing_size)
|
||||
mask = torch.nn.functional.interpolate(small, size=(h, w), mode="bilinear", align_corners=False)[
|
||||
0, 0
|
||||
]
|
||||
|
||||
broadcast_shape = (1,) * (inpt.ndim - 2) + (h, w)
|
||||
return (inpt * mask.reshape(broadcast_shape)).clamp(0.0, 1.0)
|
||||
|
||||
|
||||
class CoarseDropout(Transform):
|
||||
"""Drop random rectangular patches to simulate partial occlusion.
|
||||
|
||||
Models objects, hands, or cables passing through the camera field of view
|
||||
during robot manipulation.
|
||||
|
||||
Args:
|
||||
max_holes: Maximum number of rectangular patches to drop.
|
||||
max_height_frac: Maximum patch height as fraction of image height.
|
||||
max_width_frac: Maximum patch width as fraction of image width.
|
||||
fill_value: Value to fill dropped regions with.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
max_holes: int = 8,
|
||||
max_height_frac: float = 0.07,
|
||||
max_width_frac: float = 0.07,
|
||||
fill_value: float = 0.0,
|
||||
) -> None:
|
||||
super().__init__()
|
||||
if not isinstance(max_holes, int):
|
||||
raise TypeError("max_holes must be an int.")
|
||||
if max_holes < 1:
|
||||
raise ValueError(f"max_holes must be at least 1, but got {max_holes}.")
|
||||
if not 0.0 < max_height_frac <= 1.0:
|
||||
raise ValueError(f"max_height_frac must be in (0, 1], but got {max_height_frac}.")
|
||||
if not 0.0 < max_width_frac <= 1.0:
|
||||
raise ValueError(f"max_width_frac must be in (0, 1], but got {max_width_frac}.")
|
||||
if not 0.0 <= fill_value <= 1.0:
|
||||
raise ValueError(f"fill_value must be in [0, 1], but got {fill_value}.")
|
||||
self.max_holes = max_holes
|
||||
self.max_height_frac = max_height_frac
|
||||
self.max_width_frac = max_width_frac
|
||||
self.fill_value = fill_value
|
||||
|
||||
def make_params(self, flat_inputs: list[Any]) -> dict[str, Any]:
|
||||
n = int(torch.randint(1, self.max_holes + 1, (1,)).item())
|
||||
sizes = torch.rand(n, 2)
|
||||
sizes[:, 0] *= self.max_height_frac
|
||||
sizes[:, 1] *= self.max_width_frac
|
||||
return {"sizes": sizes.tolist(), "positions": torch.rand(n, 2).tolist()}
|
||||
|
||||
def transform(self, inpt: Any, params: dict[str, Any]) -> Any:
|
||||
if not isinstance(inpt, torch.Tensor) or not inpt.is_floating_point():
|
||||
return inpt
|
||||
if inpt.ndim < 3:
|
||||
raise ValueError(f"CoarseDropout expects [..., C, H, W] input, but got shape {inpt.shape}.")
|
||||
|
||||
h, w = inpt.shape[-2:]
|
||||
result = inpt.clone()
|
||||
for (height_frac, width_frac), (y_frac, x_frac) in zip(
|
||||
params["sizes"], params["positions"], strict=True
|
||||
):
|
||||
hole_h = max(1, min(h, round(height_frac * h)))
|
||||
hole_w = max(1, min(w, round(width_frac * w)))
|
||||
y = round(y_frac * (h - hole_h))
|
||||
x = round(x_frac * (w - hole_w))
|
||||
result[..., y : y + hole_h, x : x + hole_w] = self.fill_value
|
||||
return result
|
||||
|
||||
|
||||
class GammaCorrection(Transform):
|
||||
"""Apply random gamma correction to simulate exposure variation.
|
||||
|
||||
Models different camera auto-exposure settings and sensor response curves.
|
||||
Uses log-symmetric sampling so brightening and darkening are equally likely,
|
||||
preventing BatchNorm statistics shift.
|
||||
|
||||
Args:
|
||||
gamma: Range (min, max) for gamma value. Values < 1 brighten, > 1 darken.
|
||||
"""
|
||||
|
||||
def __init__(self, gamma: float | Sequence[float] = (0.5, 2.0)) -> None:
|
||||
super().__init__()
|
||||
if isinstance(gamma, (int, float)):
|
||||
gamma = float(gamma)
|
||||
if gamma <= 0:
|
||||
raise ValueError(f"gamma must be positive, but got {gamma}.")
|
||||
self.gamma = (min(gamma, 1.0 / gamma), max(gamma, 1.0 / gamma))
|
||||
elif isinstance(gamma, Sequence) and len(gamma) == 2:
|
||||
self.gamma = (float(gamma[0]), float(gamma[1]))
|
||||
else:
|
||||
raise TypeError("gamma must be a number or a sequence with length 2.")
|
||||
if not 0.0 < self.gamma[0] <= self.gamma[1]:
|
||||
raise ValueError(f"gamma must satisfy 0 < min <= max, but got {self.gamma}.")
|
||||
|
||||
def make_params(self, flat_inputs: list[Any]) -> dict[str, Any]:
|
||||
log_lo = math.log(self.gamma[0])
|
||||
log_hi = math.log(self.gamma[1])
|
||||
gamma = math.exp(torch.empty(1).uniform_(log_lo, log_hi).item())
|
||||
return {"gamma": gamma}
|
||||
|
||||
def transform(self, inpt: Any, params: dict[str, Any]) -> Any:
|
||||
if isinstance(inpt, torch.Tensor) and inpt.is_floating_point():
|
||||
return inpt.pow(params["gamma"]).clamp(0.0, 1.0)
|
||||
return inpt
|
||||
|
||||
|
||||
# From the paper authors' MIT-licensed reference implementation:
|
||||
# https://github.com/TheZino/PlanckianJitter
|
||||
_PLANCKIAN_BLACKBODY_COEFFICIENTS = (
|
||||
(0.6743, 0.4029, 0.0013),
|
||||
(0.6281, 0.4241, 0.1665),
|
||||
(0.5919, 0.4372, 0.2513),
|
||||
(0.5623, 0.4457, 0.3154),
|
||||
(0.5376, 0.4515, 0.3672),
|
||||
(0.5163, 0.4555, 0.4103),
|
||||
(0.4979, 0.4584, 0.4468),
|
||||
(0.4816, 0.4604, 0.4782),
|
||||
(0.4672, 0.4619, 0.5053),
|
||||
(0.4542, 0.4630, 0.5289),
|
||||
(0.4426, 0.4638, 0.5497),
|
||||
(0.4320, 0.4644, 0.5681),
|
||||
(0.4223, 0.4648, 0.5844),
|
||||
(0.4135, 0.4651, 0.5990),
|
||||
(0.4054, 0.4653, 0.6121),
|
||||
(0.3980, 0.4654, 0.6239),
|
||||
(0.3911, 0.4655, 0.6346),
|
||||
(0.3847, 0.4656, 0.6444),
|
||||
(0.3787, 0.4656, 0.6532),
|
||||
(0.3732, 0.4656, 0.6613),
|
||||
(0.3680, 0.4655, 0.6688),
|
||||
(0.3632, 0.4655, 0.6756),
|
||||
(0.3586, 0.4655, 0.6820),
|
||||
(0.3544, 0.4654, 0.6878),
|
||||
(0.3503, 0.4653, 0.6933),
|
||||
)
|
||||
_PLANCKIAN_MIN_TEMPERATURE = 3_000
|
||||
_PLANCKIAN_MAX_TEMPERATURE = 15_000
|
||||
_PLANCKIAN_TEMPERATURE_STEP = 500
|
||||
|
||||
|
||||
class PlanckianJitter(Transform):
|
||||
"""Simulate color temperature shift along the Planckian locus.
|
||||
|
||||
Samples one black-body temperature and applies the corresponding correlated red
|
||||
and blue channel scaling while preserving the green channel. Coefficients between
|
||||
the tabulated 500 K intervals are linearly interpolated.
|
||||
|
||||
Reference: Zini et al., "Planckian Jitter", CVPR 2022 Workshop.
|
||||
|
||||
Args:
|
||||
temperature: A fixed color temperature or range in Kelvin. Supported values
|
||||
are between 3000 K and 15000 K.
|
||||
"""
|
||||
|
||||
def __init__(self, temperature: int | Sequence[int] = (3_000, 15_000)) -> None:
|
||||
super().__init__()
|
||||
if isinstance(temperature, int):
|
||||
self.temperature = (temperature, temperature)
|
||||
elif isinstance(temperature, Sequence) and len(temperature) == 2:
|
||||
self.temperature = (int(temperature[0]), int(temperature[1]))
|
||||
else:
|
||||
raise TypeError("temperature must be an int or a sequence with length 2.")
|
||||
if not (
|
||||
_PLANCKIAN_MIN_TEMPERATURE
|
||||
<= self.temperature[0]
|
||||
<= self.temperature[1]
|
||||
<= _PLANCKIAN_MAX_TEMPERATURE
|
||||
):
|
||||
raise ValueError(
|
||||
"temperature must satisfy "
|
||||
f"{_PLANCKIAN_MIN_TEMPERATURE} <= min <= max <= {_PLANCKIAN_MAX_TEMPERATURE}, "
|
||||
f"but got {self.temperature}."
|
||||
)
|
||||
|
||||
def make_params(self, flat_inputs: list[Any]) -> dict[str, Any]:
|
||||
temperature = int(torch.randint(self.temperature[0], self.temperature[1] + 1, ()).item())
|
||||
return {"temperature": temperature}
|
||||
|
||||
def transform(self, inpt: Any, params: dict[str, Any]) -> Any:
|
||||
if not isinstance(inpt, torch.Tensor) or not inpt.is_floating_point():
|
||||
return inpt
|
||||
if inpt.ndim < 3 or inpt.shape[-3] != 3:
|
||||
raise ValueError(f"PlanckianJitter expects [..., 3, H, W] input, but got shape {inpt.shape}.")
|
||||
|
||||
table_position = (params["temperature"] - _PLANCKIAN_MIN_TEMPERATURE) / _PLANCKIAN_TEMPERATURE_STEP
|
||||
left_index = math.floor(table_position)
|
||||
right_index = min(left_index + 1, len(_PLANCKIAN_BLACKBODY_COEFFICIENTS) - 1)
|
||||
interpolation_weight = table_position - left_index
|
||||
|
||||
left = torch.tensor(
|
||||
_PLANCKIAN_BLACKBODY_COEFFICIENTS[left_index],
|
||||
device=inpt.device,
|
||||
dtype=inpt.dtype,
|
||||
)
|
||||
right = torch.tensor(
|
||||
_PLANCKIAN_BLACKBODY_COEFFICIENTS[right_index],
|
||||
device=inpt.device,
|
||||
dtype=inpt.dtype,
|
||||
)
|
||||
coefficients = torch.lerp(left, right, interpolation_weight)
|
||||
scale = torch.stack(
|
||||
(
|
||||
coefficients[0] / coefficients[1],
|
||||
coefficients.new_tensor(1.0),
|
||||
coefficients[2] / coefficients[1],
|
||||
)
|
||||
)
|
||||
broadcast_shape = (1,) * (inpt.ndim - 3) + (3, 1, 1)
|
||||
return (inpt * scale.reshape(broadcast_shape)).clamp(0.0, 1.0)
|
||||
|
||||
|
||||
_CUSTOM_TRANSFORMS: dict[str, type[Transform]] = {
|
||||
"SharpnessJitter": SharpnessJitter,
|
||||
"GaussianNoise": GaussianNoise,
|
||||
"MotionBlur": MotionBlur,
|
||||
"JPEGCompression": JPEGCompression,
|
||||
"GaussianPatchBrightness": GaussianPatchBrightness,
|
||||
"RandomShadow": RandomShadow,
|
||||
"CoarseDropout": CoarseDropout,
|
||||
"GammaCorrection": GammaCorrection,
|
||||
"PlanckianJitter": PlanckianJitter,
|
||||
}
|
||||
|
||||
|
||||
@dataclass
|
||||
class ImageTransformConfig:
|
||||
"""
|
||||
@@ -683,17 +216,16 @@ class ImageTransformsConfig:
|
||||
|
||||
|
||||
def make_transform_from_config(cfg: ImageTransformConfig) -> Transform:
|
||||
if cfg.type in _CUSTOM_TRANSFORMS:
|
||||
return _CUSTOM_TRANSFORMS[cfg.type](**cfg.kwargs)
|
||||
if cfg.type == "SharpnessJitter":
|
||||
return SharpnessJitter(**cfg.kwargs)
|
||||
|
||||
transform_cls = getattr(v2, cfg.type, None)
|
||||
if isinstance(transform_cls, type) and issubclass(transform_cls, Transform):
|
||||
return transform_cls(**cfg.kwargs)
|
||||
|
||||
valid_custom = ", ".join(sorted(_CUSTOM_TRANSFORMS.keys()))
|
||||
raise ValueError(
|
||||
f"Transform '{cfg.type}' is not valid. It must be a class in "
|
||||
f"torchvision.transforms.v2 or one of: {valid_custom}."
|
||||
f"torchvision.transforms.v2 or 'SharpnessJitter'."
|
||||
)
|
||||
|
||||
|
||||
|
||||
@@ -28,17 +28,9 @@ from lerobot.scripts.lerobot_imgtransform_viz import (
|
||||
save_each_transform,
|
||||
)
|
||||
from lerobot.transforms import (
|
||||
CoarseDropout,
|
||||
GammaCorrection,
|
||||
GaussianNoise,
|
||||
GaussianPatchBrightness,
|
||||
ImageTransformConfig,
|
||||
ImageTransforms,
|
||||
ImageTransformsConfig,
|
||||
JPEGCompression,
|
||||
MotionBlur,
|
||||
PlanckianJitter,
|
||||
RandomShadow,
|
||||
RandomSubsetApply,
|
||||
SharpnessJitter,
|
||||
make_transform_from_config,
|
||||
@@ -463,153 +455,3 @@ def test_save_each_transform(img_tensor_factory, tmp_path):
|
||||
assert (transform_dir / file_name).exists(), (
|
||||
f"{file_name} was not found in {transform} directory."
|
||||
)
|
||||
|
||||
|
||||
# --- Tests for robotics-relevant augmentations ---
|
||||
|
||||
ROBOTICS_TRANSFORMS = [
|
||||
("GaussianNoise", GaussianNoise, {"std": (5.0, 25.0)}),
|
||||
("MotionBlur", MotionBlur, {"kernel_size": (3, 11)}),
|
||||
("JPEGCompression", JPEGCompression, {"quality": (15, 75)}),
|
||||
("GaussianPatchBrightness", GaussianPatchBrightness, {}),
|
||||
("RandomShadow", RandomShadow, {"opacity": (0.3, 0.6)}),
|
||||
("CoarseDropout", CoarseDropout, {"max_holes": 8}),
|
||||
("GammaCorrection", GammaCorrection, {"gamma": (0.5, 2.0)}),
|
||||
("PlanckianJitter", PlanckianJitter, {"temperature": (3_000, 15_000)}),
|
||||
]
|
||||
|
||||
|
||||
@pytest.mark.parametrize("name,cls,kwargs", ROBOTICS_TRANSFORMS, ids=[t[0] for t in ROBOTICS_TRANSFORMS])
|
||||
def test_robotics_transform_shape_preserved(name, cls, kwargs, img_tensor_factory):
|
||||
img = img_tensor_factory()
|
||||
tf = cls(**kwargs)
|
||||
out = tf(img)
|
||||
assert out.shape == img.shape, f"{name} changed shape: {img.shape} -> {out.shape}"
|
||||
|
||||
|
||||
@pytest.mark.parametrize("name,cls,kwargs", ROBOTICS_TRANSFORMS, ids=[t[0] for t in ROBOTICS_TRANSFORMS])
|
||||
def test_robotics_transform_output_range(name, cls, kwargs, img_tensor_factory):
|
||||
img = img_tensor_factory()
|
||||
tf = cls(**kwargs)
|
||||
out = tf(img)
|
||||
assert out.min() >= -0.01, f"{name} min below range: {out.min():.4f}"
|
||||
assert out.max() <= 1.01, f"{name} max above range: {out.max():.4f}"
|
||||
|
||||
|
||||
@pytest.mark.parametrize("name,cls,kwargs", ROBOTICS_TRANSFORMS, ids=[t[0] for t in ROBOTICS_TRANSFORMS])
|
||||
def test_robotics_transform_float_output(name, cls, kwargs, img_tensor_factory):
|
||||
img = img_tensor_factory()
|
||||
tf = cls(**kwargs)
|
||||
out = tf(img)
|
||||
assert out.is_floating_point(), f"{name} output dtype={out.dtype}"
|
||||
|
||||
|
||||
@pytest.mark.parametrize("name,cls,kwargs", ROBOTICS_TRANSFORMS, ids=[t[0] for t in ROBOTICS_TRANSFORMS])
|
||||
def test_robotics_transform_non_float_passthrough(name, cls, kwargs):
|
||||
int_img = torch.randint(0, 255, (3, 32, 32), dtype=torch.uint8)
|
||||
tf = cls(**kwargs)
|
||||
out = tf(int_img)
|
||||
assert torch.equal(out, int_img), f"{name} modified non-float input"
|
||||
|
||||
|
||||
@pytest.mark.parametrize("name,cls,kwargs", ROBOTICS_TRANSFORMS, ids=[t[0] for t in ROBOTICS_TRANSFORMS])
|
||||
def test_robotics_transform_via_config(name, cls, kwargs):
|
||||
cfg = ImageTransformConfig(type=name, kwargs=kwargs)
|
||||
tf = make_transform_from_config(cfg)
|
||||
assert isinstance(tf, cls), f"Config produced {type(tf)}, expected {cls}"
|
||||
|
||||
|
||||
def test_make_transform_error_message_includes_custom():
|
||||
"""Error message should list all registered custom transforms."""
|
||||
with pytest.raises(ValueError, match="GaussianNoise"):
|
||||
make_transform_from_config(ImageTransformConfig(type="NonExistent"))
|
||||
|
||||
|
||||
@pytest.mark.parametrize("name,cls,kwargs", ROBOTICS_TRANSFORMS, ids=[t[0] for t in ROBOTICS_TRANSFORMS])
|
||||
@pytest.mark.parametrize("shape", [(4, 3, 32, 32), (2, 4, 3, 16, 16)])
|
||||
def test_robotics_transform_supports_temporal_batches(name, cls, kwargs, shape):
|
||||
img = torch.rand(shape)
|
||||
out = cls(**kwargs)(img)
|
||||
assert out.shape == img.shape, f"{name} changed shape: {img.shape} -> {out.shape}"
|
||||
assert out.min() >= 0
|
||||
assert out.max() <= 1
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"cls,kwargs",
|
||||
[
|
||||
(GaussianNoise, {"std": (25.0, 25.0)}),
|
||||
(MotionBlur, {"kernel_size": 5}),
|
||||
(JPEGCompression, {"quality": 10}),
|
||||
(
|
||||
GaussianPatchBrightness,
|
||||
{"num_patches": 1, "sigma_range": (0.2, 0.2), "factor_range": (0.5, 0.5)},
|
||||
),
|
||||
(RandomShadow, {"opacity": 0.5}),
|
||||
(CoarseDropout, {"max_holes": 1, "fill_value": 0.0}),
|
||||
(GammaCorrection, {"gamma": (2.0, 2.0)}),
|
||||
(PlanckianJitter, {"temperature": 3_000}),
|
||||
],
|
||||
)
|
||||
def test_robotics_transform_is_not_silent_noop(cls, kwargs):
|
||||
img = torch.rand(3, 32, 32)
|
||||
out = cls(**kwargs)(img)
|
||||
assert not torch.equal(out, img)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"transform",
|
||||
[
|
||||
GaussianNoise(std=25),
|
||||
RandomShadow(opacity=0.5),
|
||||
CoarseDropout(max_holes=4),
|
||||
],
|
||||
)
|
||||
def test_robotics_transform_random_params_are_reused(transform):
|
||||
img = torch.rand(3, 32, 32)
|
||||
params = transform.make_params([img])
|
||||
torch.testing.assert_close(transform.transform(img, params), transform.transform(img, params))
|
||||
|
||||
|
||||
def test_motion_blur_kernel_size_stays_in_configured_range():
|
||||
transform = MotionBlur(kernel_size=(4, 10))
|
||||
sampled_sizes = {transform.make_params([])["kernel_size"] for _ in range(100)}
|
||||
assert sampled_sizes <= {5, 7, 9}
|
||||
assert sampled_sizes
|
||||
|
||||
|
||||
def test_gamma_correction_scalar_below_one_defines_symmetric_range():
|
||||
transform = GammaCorrection(gamma=0.5)
|
||||
assert transform.gamma == (0.5, 2.0)
|
||||
assert transform(torch.rand(3, 8, 8)).shape == (3, 8, 8)
|
||||
|
||||
|
||||
def test_planckian_jitter_uses_correlated_temperature_coefficients():
|
||||
img = torch.full((2, 3, 8, 8), 0.25)
|
||||
out = PlanckianJitter(temperature=3_000)(img)
|
||||
torch.testing.assert_close(out[:, 1], img[:, 1])
|
||||
assert torch.all(out[:, 0] > out[:, 1])
|
||||
assert torch.all(out[:, 2] < out[:, 1])
|
||||
|
||||
|
||||
def test_random_shadow_supports_small_images():
|
||||
img = torch.rand(3, 7, 7)
|
||||
assert RandomShadow()(img).shape == img.shape
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"cls,kwargs",
|
||||
[
|
||||
(GaussianNoise, {"std": (-1.0, 1.0)}),
|
||||
(MotionBlur, {"kernel_size": 4}),
|
||||
(JPEGCompression, {"quality": (0, 75)}),
|
||||
(GaussianPatchBrightness, {"sigma_range": (0.0, 0.25)}),
|
||||
(RandomShadow, {"opacity": (0.3, 1.1)}),
|
||||
(CoarseDropout, {"max_holes": 0}),
|
||||
(GammaCorrection, {"gamma": 0.0}),
|
||||
(PlanckianJitter, {"temperature": (2_000, 6_500)}),
|
||||
],
|
||||
)
|
||||
def test_robotics_transform_rejects_invalid_config(cls, kwargs):
|
||||
with pytest.raises(ValueError):
|
||||
cls(**kwargs)
|
||||
|
||||
@@ -496,6 +496,60 @@ def test_evo1_processor_save_load_round_trip_applies_config_overrides(tmp_path):
|
||||
assert "embodiment_id" in processed
|
||||
|
||||
|
||||
def test_reconcile_evo1_processors_repads_overridden_stats(tmp_path):
|
||||
"""Loading a checkpoint and injecting raw (unpadded) dataset stats must be re-padded.
|
||||
|
||||
Regression test: lerobot-train passes the raw dataset stats as normalizer/unnormalizer
|
||||
overrides when resuming from a checkpoint (e.g. stage2 from a stage1 checkpoint). Those stats
|
||||
are at the dataset dims (e.g. LIBERO state=8/action=7), but EVO1 pads state/action to
|
||||
max_state_dim/max_action_dim before normalization, so reconcile_evo1_processors must re-pad the
|
||||
stats or normalization crashes with a shape mismatch.
|
||||
"""
|
||||
config = make_config()
|
||||
preprocessor, postprocessor = make_evo1_pre_post_processors(config, dataset_stats=make_stats())
|
||||
preprocessor.save_pretrained(tmp_path)
|
||||
postprocessor.save_pretrained(tmp_path)
|
||||
|
||||
# Reload with the generic override path injecting raw, unpadded dataset stats.
|
||||
raw_stats = make_stats()
|
||||
loaded_pre = PolicyProcessorPipeline.from_pretrained(
|
||||
tmp_path,
|
||||
config_filename=f"{POLICY_PREPROCESSOR_DEFAULT_NAME}.json",
|
||||
overrides={"normalizer_processor": {"stats": raw_stats}},
|
||||
to_transition=batch_to_transition,
|
||||
to_output=transition_to_batch,
|
||||
)
|
||||
loaded_post = PolicyProcessorPipeline.from_pretrained(
|
||||
tmp_path,
|
||||
config_filename=f"{POLICY_POSTPROCESSOR_DEFAULT_NAME}.json",
|
||||
overrides={"unnormalizer_processor": {"stats": raw_stats}},
|
||||
to_transition=policy_action_to_transition,
|
||||
to_output=transition_to_policy_action,
|
||||
)
|
||||
|
||||
# Sanity: the override really injected unpadded stats before reconciliation.
|
||||
normalizer = next(step for step in loaded_pre.steps if isinstance(step, NormalizerProcessorStep))
|
||||
assert normalizer._tensor_stats[OBS_STATE]["min"].shape == (STATE_DIM,)
|
||||
|
||||
loaded_pre, loaded_post = reconcile_evo1_processors(config, loaded_pre, loaded_post)
|
||||
|
||||
normalizer = next(step for step in loaded_pre.steps if isinstance(step, NormalizerProcessorStep))
|
||||
unnormalizer = next(step for step in loaded_post.steps if isinstance(step, UnnormalizerProcessorStep))
|
||||
assert normalizer._tensor_stats[OBS_STATE]["min"].shape == (MAX_STATE_DIM,)
|
||||
assert normalizer._tensor_stats[ACTION]["min"].shape == (MAX_ACTION_DIM,)
|
||||
assert unnormalizer._tensor_stats[ACTION]["min"].shape == (MAX_ACTION_DIM,)
|
||||
|
||||
# Normalizing a padded state must not raise (this is the exact runtime path that crashed).
|
||||
processed = loaded_pre(
|
||||
{
|
||||
"task": "pick the block",
|
||||
OBS_STATE: torch.zeros(STATE_DIM),
|
||||
f"{OBS_IMAGES}.front": torch.rand(3, 16, 16),
|
||||
}
|
||||
)
|
||||
assert processed[OBS_STATE].shape == (1, MAX_STATE_DIM)
|
||||
|
||||
|
||||
def test_evo1_policy_forward_and_inference_use_batched_embedding(monkeypatch):
|
||||
monkeypatch.setattr(modeling_evo1, "Evo1Model", DummyEvo1Model)
|
||||
policy = modeling_evo1.Evo1Policy(make_config())
|
||||
|
||||
@@ -0,0 +1,45 @@
|
||||
#!/usr/bin/env python
|
||||
|
||||
# Copyright 2025 The HuggingFace Inc. team. All rights reserved.
|
||||
#
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
|
||||
import pytest
|
||||
|
||||
from lerobot.configs import FeatureType, PipelineFeatureType, PolicyFeature
|
||||
from lerobot.robots.so_follower.robot_kinematic_processor import (
|
||||
ForwardKinematicsJointsToEEAction,
|
||||
ForwardKinematicsJointsToEEObservation,
|
||||
)
|
||||
|
||||
MOTOR_NAMES = ["shoulder_pan", "shoulder_lift", "elbow_flex", "wrist_flex", "wrist_roll", "gripper"]
|
||||
EE_KEYS = {f"ee.{k}" for k in ["x", "y", "z", "wx", "wy", "wz", "gripper_pos"]}
|
||||
|
||||
|
||||
def _joint_bucket(feature_type: FeatureType) -> dict[str, PolicyFeature]:
|
||||
return {f"{n}.pos": PolicyFeature(type=feature_type, shape=(1,)) for n in MOTOR_NAMES}
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("step_cls", "bucket", "feature_type"),
|
||||
[
|
||||
(ForwardKinematicsJointsToEEAction, PipelineFeatureType.ACTION, FeatureType.ACTION),
|
||||
(ForwardKinematicsJointsToEEObservation, PipelineFeatureType.OBSERVATION, FeatureType.STATE),
|
||||
],
|
||||
)
|
||||
def test_fk_feature_schema(step_cls, bucket, feature_type):
|
||||
features = {PipelineFeatureType.ACTION: {}, PipelineFeatureType.OBSERVATION: {}}
|
||||
features[bucket] = _joint_bucket(feature_type)
|
||||
out = step_cls(kinematics=None, motor_names=MOTOR_NAMES).transform_features(features)[bucket]
|
||||
assert set(out) == EE_KEYS
|
||||
assert {feature.type for feature in out.values()} == {feature_type}
|
||||
Reference in New Issue
Block a user