Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
1 change: 1 addition & 0 deletions conf/offpolicy/config.yaml
Original file line number Diff line number Diff line change
Expand Up @@ -6,6 +6,7 @@ defaults:
training:
task_name: G1WalkFlat
device: null
collector_infer_device: cpu
logger: tensorboard
wandb_project: unilab
wandb_entity: null
Expand Down
4 changes: 4 additions & 0 deletions scripts/train_offpolicy.py
Original file line number Diff line number Diff line change
Expand Up @@ -178,6 +178,7 @@ def build_runner(algo_name: str, cfg: DictConfig):
num_gpus = int(getattr(cfg.training, "num_gpus", 1))
multi_gpu_sync_mode = str(getattr(cfg.training, "multi_gpu_sync_mode", "local_sgd"))
multi_gpu_sync_interval = int(getattr(cfg.training, "multi_gpu_sync_interval", 1))
collector_infer_device = str(getattr(cfg.training, "collector_infer_device", "cpu") or "cpu")

sync_collection = not bool(cfg.training.no_sync_collection)

Expand Down Expand Up @@ -318,6 +319,7 @@ def build_runner(algo_name: str, cfg: DictConfig):
trace_cuda_events=cfg.training.trace_cuda_events,
seed=cfg.algo.seed,
nan_guard_cfg=_nan_guard_cfg,
collector_infer_device=collector_infer_device,
)

_learner = _learner_cls(device=_device, **_learner_kwargs)
Expand Down Expand Up @@ -348,6 +350,7 @@ def build_runner(algo_name: str, cfg: DictConfig):
verbose_metrics=verbose_metrics,
seed=cfg.algo.seed,
nan_guard_cfg=_nan_guard_cfg,
collector_infer_device=collector_infer_device,
)

if algo_name == "td3":
Expand Down Expand Up @@ -433,6 +436,7 @@ def build_runner(algo_name: str, cfg: DictConfig):
verbose_metrics=verbose_metrics,
actor_kwargs=_actor_kwargs,
nan_guard_cfg=_nan_guard_cfg,
collector_infer_device=collector_infer_device,
)

if algo_name == "flashsac":
Expand Down
3 changes: 3 additions & 0 deletions src/unilab/algos/torch/flash_sac/double_buffer.py
Original file line number Diff line number Diff line change
Expand Up @@ -55,6 +55,7 @@ def build_flashsac_double_buffer_runner(
ensure_registries()
apply_training_seed(cfg.algo.seed, torch_runtime=True, cuda=True)
device = cfg.training.device or get_default_device()
collector_infer_device = str(getattr(cfg.training, "collector_infer_device", "cpu") or "cpu")
_validate_flashsac_double_buffer_runtime(
cfg,
device=device,
Expand Down Expand Up @@ -152,6 +153,7 @@ def build_flashsac_double_buffer_runner(
trace_thread_time=cfg.training.trace_thread_time,
trace_cuda_events=cfg.training.trace_cuda_events,
nan_guard_cfg=nan_guard_cfg,
collector_infer_device=collector_infer_device,
)

learner = FlashSACLearner(device=device, **learner_kwargs)
Expand Down Expand Up @@ -183,4 +185,5 @@ def build_flashsac_double_buffer_runner(
replay_prefetch_mode=replay_prefetch_mode,
verbose_metrics=verbose_metrics,
nan_guard_cfg=nan_guard_cfg,
collector_infer_device=collector_infer_device,
)
6 changes: 6 additions & 0 deletions src/unilab/algos/torch/offpolicy/double_buffer_runner.py
Original file line number Diff line number Diff line change
Expand Up @@ -282,6 +282,10 @@ def learn(
f"{self.replay_transfer_backend.get('backend')} "
f"({self.replay_transfer_backend.get('device_family')})"
)
logger.log_status(
"Collector infer device: "
f"{self.collector_infer_device_raw} -> {self.collector_infer_device}"
)
logger.log_status("Replay learner lightweight: fixed (log_interval=1)")
if self.verbose_metrics:
logger.log_status("Verbose metrics: enabled (field-level pack CSV)")
Expand Down Expand Up @@ -332,6 +336,8 @@ def learn(
"obs_dim": self.obs_dim,
"action_dim": self.action_dim,
"actor_kwargs": self.actor_kwargs,
"collector_infer_device": self.collector_infer_device,
"collector_infer_device_raw": self.collector_infer_device_raw,
"seed": derive_worker_seed(self.seed, worker_index=0),
"trace_enabled": self.trace_enabled,
"trace_thread_time": self.trace_thread_time,
Expand Down
18 changes: 14 additions & 4 deletions src/unilab/algos/torch/offpolicy/multi_gpu_runner.py
Original file line number Diff line number Diff line change
Expand Up @@ -2,7 +2,7 @@

Architecture:
Main process → creates ReplayBuffer (host-only), WeightSync, queues
→ spawns Collector subprocess (CPU, env simulation)
→ spawns Collector subprocess (CPU env I/O, configurable inference device)
→ spawns N Learner workers via mp.spawn (one per GPU)
Learner rank i → samples packed CPU replay rows to its rank device through
a rank-local H2D pipeline, then either averages gradients
Expand Down Expand Up @@ -276,6 +276,11 @@ def _learner_worker(
f"{sync_mode} (interval={sync_interval} iteration"
f"{'s' if sync_interval != 1 else ''})"
)
logger.log_status(
"Collector infer device: "
f"{runner_kwargs.get('collector_infer_device_raw', 'cpu')} -> "
f"{runner_kwargs.get('collector_infer_device', 'cpu')}"
)
if sync_mode == "local_sgd":
logger.log_status(
"Local-SGD optimizer state: rank-local; parameters averaged at sync boundary"
Expand Down Expand Up @@ -571,8 +576,9 @@ def _learner_worker(
class MultiGPUOffPolicyRunner(OffPolicyRunner):
"""Multi-GPU off-policy runner.

Keeps a single Collector on CPU and spawns *num_gpus* Learner workers via
``torch.multiprocessing.spawn``. Each worker processes an independent
Keeps a single Collector process and spawns *num_gpus* Learner workers via
``torch.multiprocessing.spawn``. Env I/O remains CPU/numpy while collector
actor inference can use a configured device. Each worker processes an independent
mini-batch from the same shared ReplayBuffer through a rank-local H2D
pipeline. SAC defaults to local-SGD: ranks apply local updates and average
parameters at runner-controlled synchronization boundaries. Strict per-update
Expand Down Expand Up @@ -733,7 +739,7 @@ def _learn_multi_gpu(
for _ in range(self.num_gpus)
]

# --- Start Collector (CPU, single process, unchanged) ---
# --- Start Collector (single process, device-configurable inference) ---
weight_param_shapes = {k: v.shape for k, v in self.learner.actor.state_dict().items()}
collector_kwargs = {
"env_name": self.env_name,
Expand All @@ -758,6 +764,8 @@ def _learn_multi_gpu(
"obs_dim": self.obs_dim,
"action_dim": self.action_dim,
"actor_kwargs": self.actor_kwargs,
"collector_infer_device": self.collector_infer_device,
"collector_infer_device_raw": self.collector_infer_device_raw,
"seed": derive_worker_seed(self.seed, worker_index=0),
"collector_pack_request_queue": collector_pack_request_queues,
"collector_pack_ready_queue": collector_pack_ready_queues,
Expand Down Expand Up @@ -798,6 +806,8 @@ def _learn_multi_gpu(
"algo_type": self.algo_type,
"obs_normalization": self.obs_normalization,
"shared_obs_normalizer_stats": shared_obs_normalizer_stats,
"collector_infer_device": self.collector_infer_device,
"collector_infer_device_raw": self.collector_infer_device_raw,
}

try:
Expand Down
16 changes: 14 additions & 2 deletions src/unilab/algos/torch/offpolicy/runner.py
Original file line number Diff line number Diff line change
Expand Up @@ -17,7 +17,7 @@
from unilab.ipc.replay_buffer import ReplayBuffer
from unilab.logging import OffPolicyLogger, TraceRecorder
from unilab.training.seed import apply_training_seed, derive_worker_seed
from unilab.utils.device import get_default_device
from unilab.utils.device import get_default_device, resolve_torch_device_alias
from unilab.utils.nan_guard import NanGuardCfg


Expand Down Expand Up @@ -176,13 +176,19 @@ def __init__(
trace_thread_time: bool = False,
trace_cuda_events: bool = True,
nan_guard_cfg: NanGuardCfg | None = None,
collector_infer_device: str | None = "cpu",
):
self.collector_infer_device_raw = str(collector_infer_device or "cpu")
self.collector_infer_device = resolve_torch_device_alias(
self.collector_infer_device_raw,
default="cpu",
)
super().__init__(
env_name=env_name,
env_cfg_overrides={},
rl_cfg={},
device=device,
collector_device="cpu",
collector_device=self.collector_infer_device,
num_envs=num_envs,
sim_backend=sim_backend,
)
Expand Down Expand Up @@ -332,6 +338,8 @@ def learn(
"obs_dim": self.obs_dim,
"action_dim": self.action_dim,
"actor_kwargs": self.actor_kwargs,
"collector_infer_device": self.collector_infer_device,
"collector_infer_device_raw": self.collector_infer_device_raw,
"seed": derive_worker_seed(self.seed, worker_index=0),
"trace_enabled": self.trace_enabled,
"trace_thread_time": self.trace_thread_time,
Expand Down Expand Up @@ -362,6 +370,10 @@ def learn(
logger.set_collection_sync(self.sync_collection, self.env_steps_per_sync)
if hasattr(self.learner, "use_symmetry") and self.learner.use_symmetry:
logger.log_status("Symmetry augmentation: enabled")
logger.log_status(
"Collector infer device: "
f"{self.collector_infer_device_raw} -> {self.collector_infer_device}"
)
self._active_logger = logger
logger.start()

Expand Down
33 changes: 24 additions & 9 deletions src/unilab/algos/torch/offpolicy/worker.py
Original file line number Diff line number Diff line change
Expand Up @@ -490,6 +490,8 @@ def off_policy_collector_fn(
collector_pack_ready_queue=None,
collector_pack_shared_slots=None,
nan_guard_cfg=None,
collector_infer_device: str = "cpu",
collector_infer_device_raw: str | None = None,
**kwargs,
):
"""Entry point for the off-policy collector subprocess.
Expand Down Expand Up @@ -529,6 +531,8 @@ def off_policy_collector_fn(
collector_pack_ready_queue=collector_pack_ready_queue,
collector_pack_shared_slots=collector_pack_shared_slots,
nan_guard_cfg=nan_guard_cfg,
collector_infer_device=collector_infer_device,
collector_infer_device_raw=collector_infer_device_raw,
)


Expand Down Expand Up @@ -563,6 +567,8 @@ def _run_collector(
collector_pack_ready_queue,
collector_pack_shared_slots,
nan_guard_cfg=None,
collector_infer_device: str = "cpu",
collector_infer_device_raw: str | None = None,
):
del learning_starts
from unilab.base import registry
Expand Down Expand Up @@ -601,7 +607,10 @@ def _run_collector(
weight_sync.trace_recorder = trace_recorder
weight_sync.trace_thread_time = trace_thread_time

# Build actor (always on CPU for env interaction)
collector_infer_device = str(collector_infer_device or "cpu")
collector_infer_device_raw = str(collector_infer_device_raw or collector_infer_device)

# Build actor on the resolved collector inference device. Env I/O remains numpy.
obs_dim, action_dim = resolve_collector_actor_dims(
env,
obs_dim=obs_dim,
Expand All @@ -613,7 +622,7 @@ def _run_collector(
action_dim,
actor_hidden_dim,
use_layer_norm,
"cpu",
collector_infer_device,
num_envs,
**(actor_kwargs or {}),
)
Expand All @@ -635,8 +644,8 @@ def _run_collector(
from collections import defaultdict

ep_reward_components = defaultdict(list)
timing_accum_ms = defaultdict(float)
timing_counts = defaultdict(int)
timing_accum_ms: defaultdict[str, float] = defaultdict(float)
timing_counts: defaultdict[str, int] = defaultdict(int)
done_count_window = 0
timeout_count_window = 0
terminated_count_window = 0
Expand Down Expand Up @@ -711,16 +720,18 @@ def _run_collector(
# Select action
with torch.no_grad():
_t_infer_ns = _time.perf_counter_ns()
obs_torch = torch.from_numpy(obs_np_input)
dones_torch = torch.from_numpy(prev_dones_np)
obs_torch = torch.from_numpy(obs_np_input).to(collector_infer_device)
dones_torch = torch.from_numpy(prev_dones_np).to(collector_infer_device)
priv_info_np = resolve_offpolicy_actor_priv_info(
algo_type=algo_type,
obs_np=obs_np,
critic_np=critic_np,
info=info_dict,
)
priv_info_torch = (
torch.from_numpy(priv_info_np) if priv_info_np is not None else None
torch.from_numpy(priv_info_np).to(collector_infer_device)
if priv_info_np is not None
else None
)
actions_torch = sample_offpolicy_actions(
actor=actor,
Expand All @@ -729,13 +740,17 @@ def _run_collector(
prev_dones_torch=dones_torch,
priv_info_torch=priv_info_torch,
)
actions_np = actions_torch.numpy()
actions_np = actions_torch.detach().cpu().numpy()
if trace_recorder:
trace_recorder.add_slice(
"collector/actor_infer_cpu",
"collector/actor_infer",
category="collector",
start_ns=_t_infer_ns,
end_ns=_time.perf_counter_ns(),
args={
"collector_infer_device_raw": collector_infer_device_raw,
"collector_infer_device": collector_infer_device,
},
)
phase_start_ns = _record_phase_ms(cycle_timing_ms, "action_select_ms", phase_start_ns)

Expand Down
Loading
Loading