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:
Steven Palma
2026-07-31 14:59:17 +02:00
committed by GitHub
parent e4152a2481
commit 228ecd480f
+43 -12
View File
@@ -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,
) )