Update train test split (#41)

* update

* lint
This commit is contained in:
Lufang Chen
2025-09-27 12:51:25 +08:00
committed by GitHub
parent 421db17d53
commit 17335bcc3d
2 changed files with 76 additions and 20 deletions
+73 -19
View File
@@ -2,9 +2,10 @@
LeRobot Dataset Loader - Distributed Version LeRobot Dataset Loader - Distributed Version
""" """
import numpy as np
import torch import torch
from torch.utils.data import DistributedSampler from torch.utils.data import DistributedSampler, random_split
from lerobot.datasets.lerobot_dataset import LeRobotDataset from lerobot.datasets.lerobot_dataset import LeRobotDataset, LeRobotDatasetMetadata
from typing import Protocol, SupportsIndex, TypeVar from typing import Protocol, SupportsIndex, TypeVar
from qwen_vl_utils.vision_process import smart_resize from qwen_vl_utils.vision_process import smart_resize
from wall_x.data.config import X2RDataProcessingConfig from wall_x.data.config import X2RDataProcessingConfig
@@ -40,8 +41,29 @@ class Dataset(Protocol[T_co]):
class PreprocessedDataset(Dataset[T_co]): class PreprocessedDataset(Dataset[T_co]):
def __init__(self, dataset, config, dataload_config, seed=42, rank=0, world_size=1): def __init__(
self,
dataset,
config,
dataload_config,
seed=42,
rank=0,
world_size=1,
test_only=False,
):
self.hf_dataset = dataset
if test_only:
self._dataset = dataset self._dataset = dataset
else:
self._dataset = None
self.train_dataset, self.val_dataset = random_split(
dataset,
[0.95, 0.05],
torch.Generator().manual_seed(seed) if seed is not None else None,
)
self._train()
self.seed = seed self.seed = seed
self.rank = rank self.rank = rank
self.world_size = world_size self.world_size = world_size
@@ -60,11 +82,11 @@ class PreprocessedDataset(Dataset[T_co]):
priority_order=self.dataload_config.get("priority_order", None), priority_order=self.dataload_config.get("priority_order", None),
) )
self._cam_key_mapping = CAMERA_KEY_MAPPINGS[self._dataset.meta.repo_id] self._cam_key_mapping = CAMERA_KEY_MAPPINGS[self.hf_dataset.meta.repo_id]
def _vision_preprocess(self, frames): def _vision_preprocess(self, frames):
processed_frames = [] processed_frames = []
for key in self._dataset.meta.camera_keys: for key in self.hf_dataset.meta.camera_keys:
from PIL import Image from PIL import Image
current_obs = frames[key].clone().permute(1, 2, 0) current_obs = frames[key].clone().permute(1, 2, 0)
@@ -131,6 +153,12 @@ class PreprocessedDataset(Dataset[T_co]):
def __len__(self) -> int: def __len__(self) -> int:
return len(self._dataset) return len(self._dataset)
def _eval(self):
self._dataset = self.val_dataset
def _train(self):
self._dataset = self.train_dataset
def get_train_dataloader(self): def get_train_dataloader(self):
""" """
Get distributed training dataloader Get distributed training dataloader
@@ -140,6 +168,7 @@ class PreprocessedDataset(Dataset[T_co]):
world_size: Total number of processes world_size: Total number of processes
seed: Random seed for reproducibility seed: Random seed for reproducibility
""" """
self._train()
batch_size = self.config.get("batch_size_per_gpu", 8) batch_size = self.config.get("batch_size_per_gpu", 8)
num_workers = self.config.get("num_workers", 4) num_workers = self.config.get("num_workers", 4)
@@ -160,7 +189,7 @@ class PreprocessedDataset(Dataset[T_co]):
sampler=sampler, # Use distributed sampler instead of shuffle=True sampler=sampler, # Use distributed sampler instead of shuffle=True
num_workers=num_workers, num_workers=num_workers,
collate_fn=DataCollator( collate_fn=DataCollator(
self.config, self.dataload_config, self._dataset.meta.stats self.config, self.dataload_config, self.hf_dataset.meta.stats
), ),
pin_memory=True, # Enable for GPU training pin_memory=True, # Enable for GPU training
persistent_workers=num_workers > 0, # Only if num_workers > 0 persistent_workers=num_workers > 0, # Only if num_workers > 0
@@ -174,6 +203,7 @@ class PreprocessedDataset(Dataset[T_co]):
""" """
Get distributed evaluation dataloader (no shuffling for consistent evaluation) Get distributed evaluation dataloader (no shuffling for consistent evaluation)
""" """
self._eval()
batch_size = self.config.get( batch_size = self.config.get(
"eval_batch_size_per_gpu", self.config.get("batch_size_per_gpu", 8) "eval_batch_size_per_gpu", self.config.get("batch_size_per_gpu", 8)
@@ -195,7 +225,7 @@ class PreprocessedDataset(Dataset[T_co]):
sampler=sampler, sampler=sampler,
num_workers=num_workers, num_workers=num_workers,
collate_fn=DataCollator( collate_fn=DataCollator(
self.config, self.dataload_config, self._dataset.meta.stats self.config, self.dataload_config, self.hf_dataset.meta.stats
), ),
pin_memory=True, pin_memory=True,
persistent_workers=num_workers > 0, persistent_workers=num_workers > 0,
@@ -410,9 +440,15 @@ def load_lerobot_data(
# Set seed for reproducibility # Set seed for reproducibility
torch.manual_seed(seed) torch.manual_seed(seed)
dataset_fps = 50
dataload_config = get_data_configs(config["data"]) dataload_config = get_data_configs(config["data"])
# repo_id = "lerobot/aloha_mobile_cabinet"
repo_id = lerobot_config.get("repo_id", "lerobot/aloha_mobile_cabinet")
root = lerobot_config.get("root", None)
meta_info = LeRobotDatasetMetadata(repo_id)
dataset_fps = meta_info.fps
episodes_num = meta_info.total_episodes
delta_timestamps = { delta_timestamps = {
# action chunk # action chunk
"action": [ "action": [
@@ -421,21 +457,33 @@ def load_lerobot_data(
], ],
} }
batch_size = config.get("batch_size_per_gpu", 8) batch_size = config.get("batch_size_per_gpu", 8)
episodes = np.arange(episodes_num).tolist()
# repo_id = "lerobot/aloha_mobile_cabinet" train_test_split = dataload_config.get("train_test_split", 0.95)
repo_id = lerobot_config.get("repo_id", "lerobot/aloha_mobile_cabinet") train_episodes = episodes[: int(episodes_num * train_test_split)]
root = lerobot_config.get("root", None) test_episodes = episodes[int(episodes_num * train_test_split) :]
dataset = LeRobotDataset(
repo_id, root=root, delta_timestamps=delta_timestamps, video_backend="pyav" train_dataset = LeRobotDataset(
repo_id,
root=root,
episodes=train_episodes,
delta_timestamps=delta_timestamps,
video_backend="pyav",
) )
if rank == 0: if rank == 0:
print(f"Selected episodes: {dataset.episodes}") print(f"Selected train episodes: {train_dataset.episodes}")
print(f"Number of episodes selected: {dataset.num_episodes}") print(f"Number of train episodes selected: {train_dataset.num_episodes}")
print(f"Number of frames selected: {dataset.num_frames}") print(f"Number of train frames selected: {train_dataset.num_frames}")
print(f"Selected test episodes: {test_episodes}")
dataset = PreprocessedDataset( dataset = PreprocessedDataset(
dataset, config, dataload_config, seed=seed, rank=rank, world_size=world_size train_dataset,
config,
dataload_config,
seed=seed,
rank=rank,
world_size=world_size,
) )
# Calculate samples per process # Calculate samples per process
@@ -516,7 +564,13 @@ def get_data_configs(config):
class TestDataset(PreprocessedDataset): class TestDataset(PreprocessedDataset):
def __init__(self, dataset, config, dataload_config, seed=42): def __init__(self, dataset, config, dataload_config, seed=42):
super().__init__( super().__init__(
dataset, config, dataload_config, seed=seed, rank=0, world_size=1 dataset,
config,
dataload_config,
seed=seed,
rank=0,
world_size=1,
test_only=True,
) )
def get_dataloader(self): def get_dataloader(self):
@@ -528,7 +582,7 @@ class TestDataset(PreprocessedDataset):
self, self,
batch_size=1, batch_size=1,
collate_fn=DataCollator( collate_fn=DataCollator(
self.config, self.dataload_config, self._dataset.meta.stats self.config, self.dataload_config, self.hf_dataset.meta.stats
), ),
) )
+2
View File
@@ -253,6 +253,7 @@ class QwenVlAct_Trainer:
# Initialize training dataloader for current epoch # Initialize training dataloader for current epoch
if isinstance(self.dataset, PreprocessedDataset): if isinstance(self.dataset, PreprocessedDataset):
if getattr(self, "train_dataloader", None) is not None: if getattr(self, "train_dataloader", None) is not None:
self.dataset._train()
self.train_sampler.set_epoch(epoch) self.train_sampler.set_epoch(epoch)
else: else:
self.train_dataloader, self.train_sampler = ( self.train_dataloader, self.train_sampler = (
@@ -463,6 +464,7 @@ class QwenVlAct_Trainer:
""" """
# Initialize validation dataloader # Initialize validation dataloader
if getattr(self, "val_dataloader", None) is not None: if getattr(self, "val_dataloader", None) is not None:
self.dataset._eval()
self.val_sampler.set_epoch(0) self.val_sampler.set_epoch(0)
else: else:
self.val_dataloader, self.val_sampler = self.dataset.get_val_dataloader() self.val_dataloader, self.val_sampler = self.dataset.get_val_dataloader()