From 0f174bd0cc7f7de3ad88da66deffcbf90186f237 Mon Sep 17 00:00:00 2001 From: Steven Palma Date: Fri, 24 Jul 2026 13:14:00 +0200 Subject: [PATCH] chore(scripts): add multiprocessing_context safeguards --- src/lerobot/configs/train.py | 21 +++++++++++++++++---- src/lerobot/scripts/lerobot_train.py | 17 ++++++++++++----- 2 files changed, 29 insertions(+), 9 deletions(-) diff --git a/src/lerobot/configs/train.py b/src/lerobot/configs/train.py index d268b88a4..2de140e08 100644 --- a/src/lerobot/configs/train.py +++ b/src/lerobot/configs/train.py @@ -14,6 +14,7 @@ import builtins import datetime as dt import json +import multiprocessing import os import tempfile from dataclasses import dataclass, field @@ -102,10 +103,11 @@ class TrainPipelineConfig(HubMixin): prefetch_factor: int = 4 persistent_workers: bool = True # DataLoader worker start method. "spawn" is safer than "fork" with - # non-fork-safe libs (PyAV / torchcodec / ffmpeg — see #2488), but - # adds some worker-startup time per run since workers re-import - # modules instead of inheriting parent state. - dataloader_multiprocessing_context: str = "spawn" + # non-fork-safe libs, but adds some worker-startup time per run + # since workers re-import modules instead of inheriting parent state. + # Override with `--dataloader_multiprocessing_context=fork` when appropriate, + # or set it to `null` to use Python's platform default. + dataloader_multiprocessing_context: str | None = "spawn" steps: int = 100_000 # Run policy in the simulation environment every N steps to measure reward/success (0 = disabled). env_eval_freq: int = 20_000 @@ -217,6 +219,17 @@ class TrainPipelineConfig(HubMixin): self.reward_model.pretrained_path = str(policy_dir) def validate(self) -> None: + available_contexts = multiprocessing.get_all_start_methods() + if ( + self.dataloader_multiprocessing_context is not None + and self.dataloader_multiprocessing_context not in available_contexts + ): + raise ValueError( + "`dataloader_multiprocessing_context` must be None or one of " + f"{available_contexts} on this platform, got " + f"{self.dataloader_multiprocessing_context!r}." + ) + self._resolve_pretrained_from_cli() if self.policy is None and self.reward_model is None: diff --git a/src/lerobot/scripts/lerobot_train.py b/src/lerobot/scripts/lerobot_train.py index bca92629a..8bfa16a98 100644 --- a/src/lerobot/scripts/lerobot_train.py +++ b/src/lerobot/scripts/lerobot_train.py @@ -71,6 +71,16 @@ from lerobot.utils.utils import ( from .lerobot_eval import eval_policy_all +def _dataloader_worker_kwargs(cfg: TrainPipelineConfig) -> dict[str, Any]: + """Return worker-only DataLoader options, disabling them for single-process loading.""" + workers_enabled = cfg.num_workers > 0 + return { + "prefetch_factor": cfg.prefetch_factor if workers_enabled else None, + "persistent_workers": cfg.persistent_workers and workers_enabled, + "multiprocessing_context": cfg.dataloader_multiprocessing_context if workers_enabled else None, + } + + def update_policy( train_metrics: MetricsTracker, policy: PreTrainedPolicy, @@ -473,9 +483,7 @@ def train(cfg: TrainPipelineConfig, accelerator: "Accelerator | None" = None): pin_memory=device.type == "cuda", drop_last=False, collate_fn=collate_fn, - prefetch_factor=cfg.prefetch_factor if cfg.num_workers > 0 else None, - persistent_workers=cfg.persistent_workers and cfg.num_workers > 0, - multiprocessing_context=cfg.dataloader_multiprocessing_context if cfg.num_workers > 0 else None, + **_dataloader_worker_kwargs(cfg), ) # Build eval dataloader if a held-out split exists @@ -501,8 +509,7 @@ def train(cfg: TrainPipelineConfig, accelerator: "Accelerator | None" = None): pin_memory=device.type == "cuda", drop_last=False, collate_fn=eval_collate_fn, - prefetch_factor=cfg.prefetch_factor if cfg.num_workers > 0 else None, - persistent_workers=cfg.persistent_workers and cfg.num_workers > 0, + **_dataloader_worker_kwargs(cfg), ) # Prepare everything with accelerator