chore(scripts): add multiprocessing_context safeguards

This commit is contained in:
Steven Palma
2026-07-24 13:14:00 +02:00
parent acbab791f1
commit 0f174bd0cc
2 changed files with 29 additions and 9 deletions
+17 -4
View File
@@ -14,6 +14,7 @@
import builtins import builtins
import datetime as dt import datetime as dt
import json import json
import multiprocessing
import os import os
import tempfile import tempfile
from dataclasses import dataclass, field from dataclasses import dataclass, field
@@ -102,10 +103,11 @@ class TrainPipelineConfig(HubMixin):
prefetch_factor: int = 4 prefetch_factor: int = 4
persistent_workers: bool = True persistent_workers: bool = True
# DataLoader worker start method. "spawn" is safer than "fork" with # DataLoader worker start method. "spawn" is safer than "fork" with
# non-fork-safe libs (PyAV / torchcodec / ffmpeg — see #2488), but # non-fork-safe libs, but adds some worker-startup time per run
# adds some worker-startup time per run since workers re-import # since workers re-import modules instead of inheriting parent state.
# modules instead of inheriting parent state. # Override with `--dataloader_multiprocessing_context=fork` when appropriate,
dataloader_multiprocessing_context: str = "spawn" # or set it to `null` to use Python's platform default.
dataloader_multiprocessing_context: str | None = "spawn"
steps: int = 100_000 steps: int = 100_000
# Run policy in the simulation environment every N steps to measure reward/success (0 = disabled). # Run policy in the simulation environment every N steps to measure reward/success (0 = disabled).
env_eval_freq: int = 20_000 env_eval_freq: int = 20_000
@@ -217,6 +219,17 @@ class TrainPipelineConfig(HubMixin):
self.reward_model.pretrained_path = str(policy_dir) self.reward_model.pretrained_path = str(policy_dir)
def validate(self) -> None: 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() self._resolve_pretrained_from_cli()
if self.policy is None and self.reward_model is None: if self.policy is None and self.reward_model is None:
+12 -5
View File
@@ -71,6 +71,16 @@ from lerobot.utils.utils import (
from .lerobot_eval import eval_policy_all 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( def update_policy(
train_metrics: MetricsTracker, train_metrics: MetricsTracker,
policy: PreTrainedPolicy, policy: PreTrainedPolicy,
@@ -473,9 +483,7 @@ def train(cfg: TrainPipelineConfig, accelerator: "Accelerator | None" = None):
pin_memory=device.type == "cuda", pin_memory=device.type == "cuda",
drop_last=False, drop_last=False,
collate_fn=collate_fn, collate_fn=collate_fn,
prefetch_factor=cfg.prefetch_factor if cfg.num_workers > 0 else None, **_dataloader_worker_kwargs(cfg),
persistent_workers=cfg.persistent_workers and cfg.num_workers > 0,
multiprocessing_context=cfg.dataloader_multiprocessing_context if cfg.num_workers > 0 else None,
) )
# Build eval dataloader if a held-out split exists # 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", pin_memory=device.type == "cuda",
drop_last=False, drop_last=False,
collate_fn=eval_collate_fn, collate_fn=eval_collate_fn,
prefetch_factor=cfg.prefetch_factor if cfg.num_workers > 0 else None, **_dataloader_worker_kwargs(cfg),
persistent_workers=cfg.persistent_workers and cfg.num_workers > 0,
) )
# Prepare everything with accelerator # Prepare everything with accelerator