338 lines
12 KiB
Python
338 lines
12 KiB
Python
"""LeRobot data loading bridge for typed training configs."""
|
|
|
|
from __future__ import annotations
|
|
|
|
import logging
|
|
import multiprocessing as mp
|
|
from typing import Any, Dict, Tuple
|
|
|
|
import torch
|
|
import torch.distributed as dist
|
|
|
|
from wall_x.data.backends.lerobot.config import LerobotConfig
|
|
from wall_x.data.backends.lerobot.utils import load_norm_stats
|
|
from wall_x.model.core.action.normalizer import (
|
|
create_normalizers_from_lerobot_norm_stats,
|
|
)
|
|
|
|
logger = logging.getLogger(__name__)
|
|
|
|
|
|
def load_lerobot_normalizers(cfg: Any):
|
|
"""Create model normalizers from the LeRobot norm stats configured for a run."""
|
|
data = cfg.data
|
|
|
|
norm_stats_path = getattr(data, "norm_stats_path", None)
|
|
if not norm_stats_path:
|
|
return None
|
|
|
|
key_mappings = getattr(data, "key_mappings", None)
|
|
if not key_mappings:
|
|
raise ValueError(
|
|
"LeRobot normalizer from norm_stats_path requires data.key_mappings"
|
|
)
|
|
|
|
lerobot_config = getattr(data, "lerobot_config", None)
|
|
if not isinstance(lerobot_config, dict) or not lerobot_config.get("repo_id"):
|
|
raise ValueError(
|
|
"LeRobot normalizer from norm_stats_path requires "
|
|
"data.lerobot_config.repo_id"
|
|
)
|
|
|
|
dataset_name = str(lerobot_config["repo_id"])
|
|
norm_stats = load_norm_stats(
|
|
norm_stats_path,
|
|
key_mappings,
|
|
dof_config=dict(cfg.task.dof_config or {}),
|
|
agent_pos_config=dict(cfg.task.agent_pos_config or {}),
|
|
)
|
|
normalizer_action, normalizer_propri = create_normalizers_from_lerobot_norm_stats(
|
|
norm_stats,
|
|
dataset_name,
|
|
cfg.action_dim,
|
|
cfg.propri_dim,
|
|
)
|
|
return normalizer_action, normalizer_propri, norm_stats_path, dataset_name
|
|
|
|
|
|
class _LerobotDatasetWrapper:
|
|
"""Trainer-facing wrapper aligning PreprocessedDataset with v1 API.
|
|
|
|
PreprocessedDataset internally switches ``self._dataset`` between
|
|
its train/val splits via ``_train()`` / ``_eval()``. Its
|
|
``get_train_dataloader`` / ``get_val_dataloader`` return
|
|
``(dataloader, sampler)`` tuples and no-argument calls are supported
|
|
(they read rank/world_size/seed from the inner object itself).
|
|
|
|
This wrapper:
|
|
- Caches the rebuilt train dataloader / sampler so
|
|
``set_epoch(epoch)`` can reset shuffling per-epoch.
|
|
- Owns the val dataloader so the trainer's ``val_loop`` can do
|
|
``self.dataset.get_val_dataloader()`` and iterate directly (matching
|
|
what the v1/v2 wrappers return).
|
|
"""
|
|
|
|
def __init__(
|
|
self,
|
|
inner,
|
|
train_dataloader: torch.utils.data.DataLoader,
|
|
train_sampler,
|
|
train_num: int,
|
|
val_dataloader: torch.utils.data.DataLoader = None,
|
|
val_num: int = 0,
|
|
):
|
|
self._inner = inner
|
|
self._train_dataloader = train_dataloader
|
|
self._train_sampler = train_sampler
|
|
self._train_num = train_num
|
|
self._val_dataloader = val_dataloader
|
|
self.global_train_iters = mp.Value("i", train_num)
|
|
self.global_val_iters = mp.Value("i", val_num)
|
|
|
|
def __len__(self) -> int:
|
|
return self._train_num
|
|
|
|
def _activate_train_split(self) -> None:
|
|
if hasattr(self._inner, "_train"):
|
|
self._inner._train()
|
|
|
|
def get_train_dataloader(self):
|
|
self._activate_train_split()
|
|
return self._train_dataloader
|
|
|
|
def get_val_dataloader(self):
|
|
# PreprocessedDataset shares one ``_dataset`` pointer between its
|
|
# train and val splits (flipped by ``_train()`` / ``_eval()``).
|
|
# The val DataLoader's DistributedSampler caches total_size sized
|
|
# to the val split but ``__iter__`` reads ``len(self.dataset)``
|
|
# live - if a preceding train_loop left the pointer at train, that
|
|
# live len is ~20x total_size and DistributedSampler asserts.
|
|
# Rebuild each time so ``_eval()`` runs and a fresh sampler is
|
|
# snapped to the current (val) split length. Mirrors the train-side
|
|
# rebuild-on-every-epoch pattern.
|
|
if self._val_dataloader is None:
|
|
return None
|
|
self._val_dataloader, _ = self._inner.get_val_dataloader()
|
|
return self._val_dataloader
|
|
|
|
def set_epoch(self, epoch: int) -> None:
|
|
"""Seed the per-epoch shuffle in the train DistributedSampler."""
|
|
self._activate_train_split()
|
|
if self._train_sampler is not None and hasattr(
|
|
self._train_sampler, "set_epoch"
|
|
):
|
|
self._train_sampler.set_epoch(epoch)
|
|
|
|
|
|
def load_trainer_data_config(cfg: Any) -> LerobotConfig:
|
|
"""Build the inference/trainer data config from a typed TrainConfig."""
|
|
raw_yaml = dict(getattr(cfg, "_raw_yaml", {}) or {})
|
|
raw_data = dict(getattr(cfg, "_raw_data", {}) or {})
|
|
data = getattr(cfg, "data", None)
|
|
|
|
data_section = dict(raw_yaml.get("data", {}) or {})
|
|
data_section.update(raw_data)
|
|
|
|
if data is not None:
|
|
for key in (
|
|
"resolution",
|
|
"train_test_split",
|
|
"priority_order",
|
|
"camera_name_mapping",
|
|
):
|
|
value = getattr(data, key, None)
|
|
if value is not None:
|
|
data_section.setdefault(key, value)
|
|
|
|
raw_yaml["data"] = data_section
|
|
raw_yaml.setdefault("model_type", getattr(cfg, "model_type", "qwen2_5"))
|
|
return load_trainer_data_config_from_yaml_dict(raw_yaml)
|
|
|
|
|
|
def load_trainer_data_config_from_yaml_dict(yaml_dict: Dict[str, Any]) -> LerobotConfig:
|
|
"""Build the LeRobot runtime config from a raw training YAML dict."""
|
|
return LerobotConfig.from_yaml_dict(yaml_dict)
|
|
|
|
|
|
def _build_flat_config(cfg: Any) -> Tuple[Dict[str, Any], Dict[str, Any]]:
|
|
"""Map typed TrainConfig -> (flat_config, lerobot_config) for legacy entry.
|
|
|
|
``load_lerobot_data`` expects a 2509-style flat dict plus a separate
|
|
``lerobot_config`` carrying ``repo_id`` / ``root``. This function is
|
|
the one place that translation lives; keep it surgical so future
|
|
field additions on ``LeRobotDataConfig`` do not require touching the
|
|
legacy loader.
|
|
"""
|
|
model = cfg.model
|
|
data = cfg.data
|
|
hp = cfg.hyperparams
|
|
raw = getattr(cfg, "_raw_yaml", {}) or {}
|
|
raw_data = dict(getattr(cfg, "_raw_data", {}) or {})
|
|
|
|
lerobot_cfg = dict(data.lerobot_config or {})
|
|
if "repo_id" not in lerobot_cfg:
|
|
raise ValueError(
|
|
"lerobot requires data.lerobot_config.repo_id to be set "
|
|
"(HuggingFace LeRobot dataset id)."
|
|
)
|
|
|
|
data_section: Dict[str, Any] = {
|
|
"key_mappings": data.key_mappings,
|
|
"action_horizon": cfg.task.action_horizon,
|
|
"train_test_split": data.train_test_split,
|
|
"seed": hp.seed,
|
|
"resolution": data.resolution,
|
|
}
|
|
if raw_data.get("max_length") is not None:
|
|
data_section["max_length"] = raw_data["max_length"]
|
|
if data.priority_order is not None:
|
|
data_section["priority_order"] = data.priority_order
|
|
if data.camera_name_mapping is not None:
|
|
data_section["camera_name_mapping"] = data.camera_name_mapping
|
|
data_section.setdefault(
|
|
"use_state_string_representation",
|
|
cfg.task.use_state_string_representation,
|
|
)
|
|
data_section.setdefault(
|
|
"state_bins",
|
|
raw_data.get("state_bins", raw.get("state_bins", 256)),
|
|
)
|
|
|
|
# Dof/agent_pos totals for the collator's zero-pad step. When resuming
|
|
# from a checkpoint trained on a larger action space, task.dof_config
|
|
# should include an ``action_padding`` key that absorbs the diff; the
|
|
# collator right-pads action/agent_pos tensors to these totals with
|
|
# dof_mask/agent_pos_mask zeroed on padded dims so loss doesn't flow
|
|
# through them.
|
|
dof_total = int(sum((cfg.task.dof_config or {}).values()))
|
|
agent_pos_total = int(sum((cfg.task.agent_pos_config or {}).values()))
|
|
|
|
flat: Dict[str, Any] = {
|
|
"model_type": cfg.model_type,
|
|
"processor_path": getattr(model, "processor_path", "") or "",
|
|
"norm_stats_path": data.norm_stats_path or raw.get("norm_stats_path"),
|
|
"batch_size_per_gpu": hp.batch_size_per_gpu,
|
|
"eval_batch_size_per_gpu": raw.get(
|
|
"eval_batch_size_per_gpu", hp.batch_size_per_gpu
|
|
),
|
|
"num_workers": data.num_workers,
|
|
"padding_side": data.padding_side,
|
|
"use_fast_tokenizer": data.use_fast_tokenizer,
|
|
"action_tokenizer_path": data.action_tokenizer_path,
|
|
"noise_scheduler": data.noise_scheduler or {},
|
|
"dof_total_dim": dof_total,
|
|
"agent_pos_total_dim": agent_pos_total,
|
|
"dof_config": dict(cfg.task.dof_config or {}),
|
|
"agent_pos_config": dict(cfg.task.agent_pos_config or {}),
|
|
"use_state_string_representation": cfg.task.use_state_string_representation,
|
|
"state_bins": int(
|
|
raw_data.get("state_bins")
|
|
or data_section.get("state_bins")
|
|
or raw.get("state_bins")
|
|
or 256
|
|
),
|
|
"data": data_section,
|
|
}
|
|
return flat, lerobot_cfg
|
|
|
|
|
|
def load_lerobot_v2(
|
|
cfg: Any,
|
|
) -> Tuple[_LerobotDatasetWrapper, torch.utils.data.DataLoader, int]:
|
|
"""Build lerobot (wrapper, dataloader, train_num) from TrainConfig.
|
|
|
|
The third return value ``train_num`` is a snapshot of
|
|
``len(train_dataloader)`` at construction time. It matches
|
|
``wrapper.global_train_iters.value`` initially but does not track
|
|
subsequent rebuilds inside ``set_epoch`` - callers doing dynamic
|
|
resampling should read from the mp.Value, not from this snapshot.
|
|
"""
|
|
from wall_x.data.backends.lerobot.loader import load_lerobot_data
|
|
|
|
flat_cfg, lerobot_cfg = _build_flat_config(cfg)
|
|
|
|
if dist.is_initialized():
|
|
rank = dist.get_rank()
|
|
world_size = dist.get_world_size()
|
|
else:
|
|
rank = 0
|
|
world_size = 1
|
|
|
|
seed = cfg.hyperparams.seed
|
|
inner, _ = load_lerobot_data(
|
|
flat_cfg,
|
|
lerobot_cfg,
|
|
rank=rank,
|
|
world_size=world_size,
|
|
seed=seed,
|
|
)
|
|
|
|
# PreprocessedDataset.get_*_dataloader returns (dataloader, sampler).
|
|
# Build val first, train second, so the inner ``_dataset`` pointer is
|
|
# left at the train split when we finish - workers fork from that
|
|
# state on first iteration.
|
|
val_dataloader, _ = inner.get_val_dataloader()
|
|
val_num = len(val_dataloader) if val_dataloader is not None else 0
|
|
|
|
train_dataloader, train_sampler = inner.get_train_dataloader()
|
|
train_num = len(train_dataloader)
|
|
|
|
if rank == 0:
|
|
logger.info(
|
|
"\n%s\nLeRobot Data Loading Configuration:\n"
|
|
" RANK: %d\n WORLD SIZE: %d\n"
|
|
" BATCH SIZE PER DEVICE: %d\n GLOBAL BATCH SIZE: %d\n"
|
|
" TRAIN BATCHES: %d\n VAL BATCHES: %d\n"
|
|
" NUM WORKERS: %d\n REPO ID: %s\n%s",
|
|
"=" * 50,
|
|
rank,
|
|
world_size,
|
|
flat_cfg["batch_size_per_gpu"],
|
|
flat_cfg["batch_size_per_gpu"] * world_size,
|
|
train_num,
|
|
val_num,
|
|
flat_cfg["num_workers"],
|
|
lerobot_cfg.get("repo_id"),
|
|
"=" * 50,
|
|
)
|
|
|
|
wrapper = _LerobotDatasetWrapper(
|
|
inner,
|
|
train_dataloader,
|
|
train_sampler,
|
|
train_num,
|
|
val_dataloader=val_dataloader,
|
|
val_num=val_num,
|
|
)
|
|
return wrapper, train_dataloader, train_num
|
|
|
|
|
|
def build(cfg, ctx):
|
|
"""Backend Protocol entry - returns a ``DataBundle``.
|
|
|
|
Wraps ``load_lerobot_v2`` (which returns the trainer-facing triple)
|
|
into the unified ``DataBundle`` shape every backend exposes.
|
|
"""
|
|
from wall_x.data._bundle import DataBundle
|
|
|
|
wrapper, train_dataloader, train_num = load_lerobot_v2(cfg)
|
|
|
|
# PreprocessedDataset shares one ``self._dataset`` pointer between
|
|
# train and val splits (flipped by ``_train()`` / ``_eval()``).
|
|
# ``wrapper.get_val_dataloader()`` flips the pointer to val. Flip back once
|
|
# here so the initial train loop starts from the right split even if callers
|
|
# inspect the raw ``train_dataloader`` before invoking ``set_epoch``.
|
|
val_loader = wrapper.get_val_dataloader()
|
|
inner = wrapper._inner
|
|
if hasattr(inner, "_train"):
|
|
inner._train()
|
|
|
|
return DataBundle(
|
|
dataset=wrapper,
|
|
train_loader=train_dataloader,
|
|
val_loader=val_loader,
|
|
train_iters=train_num,
|
|
val_iters=wrapper.global_val_iters.value,
|
|
set_epoch=wrapper.set_epoch,
|
|
)
|