mirror of
https://github.com/huggingface/lerobot.git
synced 2026-07-31 21:49:45 +00:00
fix(rollout): apply torch compile mode correctly (#4268)
* fix(rollout): apply torch compile mode correctly * chore(rollout): clarify intent --------- Co-authored-by: Patrick Ribbsaeter <patrickswedish@gmail.com>
This commit is contained in:
@@ -22,6 +22,7 @@ and :class:`DatasetContext` — assembled into :class:`RolloutContext`.
|
|||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
import logging
|
import logging
|
||||||
|
from copy import copy
|
||||||
from dataclasses import dataclass, field
|
from dataclasses import dataclass, field
|
||||||
from threading import Event
|
from threading import Event
|
||||||
from typing import TYPE_CHECKING
|
from typing import TYPE_CHECKING
|
||||||
@@ -69,6 +70,35 @@ else:
|
|||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
|
||||||
|
def _wrap_predict_action_chunk_with_torch_compile(
|
||||||
|
policy: PreTrainedPolicy,
|
||||||
|
*,
|
||||||
|
backend: str,
|
||||||
|
mode: str,
|
||||||
|
) -> bool:
|
||||||
|
"""Install the JIT wrapper and report whether it was configured successfully.
|
||||||
|
|
||||||
|
``torch.compile`` compiles lazily on the first invocation, so success here
|
||||||
|
does not guarantee that backend compilation will succeed during warm-up.
|
||||||
|
"""
|
||||||
|
if not hasattr(torch, "compile"):
|
||||||
|
logger.warning("torch.compile is not available in this PyTorch build")
|
||||||
|
return False
|
||||||
|
|
||||||
|
try:
|
||||||
|
policy.predict_action_chunk = torch.compile(
|
||||||
|
policy.predict_action_chunk,
|
||||||
|
backend=backend,
|
||||||
|
mode=mode,
|
||||||
|
)
|
||||||
|
except Exception as exc:
|
||||||
|
logger.warning("Failed to configure torch.compile: %s", exc)
|
||||||
|
return False
|
||||||
|
|
||||||
|
logger.info("torch.compile configured for predict_action_chunk")
|
||||||
|
return True
|
||||||
|
|
||||||
|
|
||||||
def _resolve_action_key_order(
|
def _resolve_action_key_order(
|
||||||
policy_action_names: list[str] | None, dataset_action_names: list[str]
|
policy_action_names: list[str] | None, dataset_action_names: list[str]
|
||||||
) -> list[str]:
|
) -> list[str]:
|
||||||
@@ -241,18 +271,19 @@ def build_rollout_context(
|
|||||||
policy.eval()
|
policy.eval()
|
||||||
logger.info("Policy loaded: type=%s, device=%s", policy_config.type, cfg.device)
|
logger.info("Policy loaded: type=%s, device=%s", policy_config.type, cfg.device)
|
||||||
|
|
||||||
|
torch_compile_active = cfg.use_torch_compile
|
||||||
if cfg.use_torch_compile and policy.type not in ("pi0", "pi05"):
|
if cfg.use_torch_compile and policy.type not in ("pi0", "pi05"):
|
||||||
try:
|
torch_compile_active = _wrap_predict_action_chunk_with_torch_compile(
|
||||||
if hasattr(torch, "compile"):
|
policy,
|
||||||
compile_kwargs = {
|
backend=cfg.torch_compile_backend,
|
||||||
"backend": cfg.torch_compile_backend,
|
mode=cfg.torch_compile_mode,
|
||||||
"mode": cfg.torch_compile_mode,
|
)
|
||||||
"options": {"triton.cudagraphs": False},
|
|
||||||
}
|
if cfg.use_torch_compile and not torch_compile_active:
|
||||||
policy.predict_action_chunk = torch.compile(policy.predict_action_chunk, **compile_kwargs)
|
# RolloutConfig.__post_init__ reloads the policy configuration, so avoid
|
||||||
logger.info("torch.compile applied to predict_action_chunk")
|
# dataclasses.replace when carrying the effective state downstream.
|
||||||
except Exception as e:
|
cfg = copy(cfg)
|
||||||
logger.warning("Failed to apply torch.compile: %s", e)
|
cfg.use_torch_compile = False
|
||||||
|
|
||||||
# --- 2. Robot-side processors (user-supplied or defaults) --------
|
# --- 2. Robot-side processors (user-supplied or defaults) --------
|
||||||
if (
|
if (
|
||||||
@@ -470,7 +501,7 @@ def build_rollout_context(
|
|||||||
task=task_str,
|
task=task_str,
|
||||||
fps=cfg.fps,
|
fps=cfg.fps,
|
||||||
device=cfg.device,
|
device=cfg.device,
|
||||||
use_torch_compile=cfg.use_torch_compile,
|
use_torch_compile=torch_compile_active,
|
||||||
compile_warmup_inferences=cfg.compile_warmup_inferences,
|
compile_warmup_inferences=cfg.compile_warmup_inferences,
|
||||||
shutdown_event=shutdown_event,
|
shutdown_event=shutdown_event,
|
||||||
)
|
)
|
||||||
|
|||||||
Reference in New Issue
Block a user