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
60 changes: 42 additions & 18 deletions src/unilab/scripts/train_appo.py
Original file line number Diff line number Diff line change
Expand Up @@ -131,35 +131,59 @@ def run_motrix_play_loop(
play_env_num: int,
num_steps: int | None = None,
) -> None:
import numpy as np
from tensordict import TensorDict

if env.state is None:
env.init_state()

with torch.inference_mode():
env.run_playback(
num_steps=num_steps,
initialize=lambda: np.asarray(
env.reset(np.arange(play_env_num, dtype=np.int32))[0]["obs"],
dtype=np.float32,
initialize=lambda: _motrix_playback_obs(
env.reset(torch.arange(play_env_num, dtype=torch.int64, device=env.device))
),
step=lambda obs_np: np.asarray(
env.step(
actor(
TensorDict(
{"policy": torch.from_numpy(obs_np).to(device)}, batch_size=play_env_num
)
)
.cpu()
.numpy()
.astype(np.float32)
).obs["obs"],
dtype=np.float32,
step=lambda obs: _motrix_playback_step(
env=env,
actor=actor,
device=device,
play_env_num=play_env_num,
obs=obs,
),
)


def _motrix_playback_obs(state: Any) -> torch.Tensor:
"""Extract the actor observation tensor from a TorchEnv reset result."""
obs = state[0].obs if hasattr(state[0], "obs") else state[0]["obs"]
if not isinstance(obs, torch.Tensor):
raise TypeError(
f"APPO playback reset observation must be a tensor, got {type(obs).__name__}"
)
return obs


def _motrix_playback_step(
*,
env: Any,
actor: Callable[[Any], Any],
device: str,
play_env_num: int,
obs: Any,
) -> torch.Tensor:
"""Run one APPO policy step and return the next actor observation tensor."""
from tensordict import TensorDict

if not isinstance(obs, torch.Tensor):
raise TypeError(
f"APPO playback step observation must be a tensor, got {type(obs).__name__}"
)
action = actor(
TensorDict(
{"policy": obs.to(device=device, dtype=torch.float32)},
batch_size=play_env_num,
)
).to(device=env.device, dtype=torch.float32)
return _motrix_playback_obs((env.step(action.contiguous()).obs, None))


def _get_log_root(cfg: DictConfig) -> str:
return str(get_log_root(Path.cwd(), cfg))

Expand Down
67 changes: 44 additions & 23 deletions src/unilab/visualization/interactive_playback.py
Original file line number Diff line number Diff line change
Expand Up @@ -287,35 +287,41 @@ def __init__(
actor_algo_type: str,
normalizer: Any | None,
num_envs: int,
obs_extractor: Callable[[dict[str, np.ndarray]], np.ndarray],
obs_extractor: Callable[[dict[str, torch.Tensor]], torch.Tensor],
) -> None:
self.env = env
self.device = device
# A backend may remap a generic CUDA request to a concrete device.
# ``env.device`` is authoritative once the environment exists.
self.policy_device = (
torch.device(device)
if getattr(env, "device", None) is None
else torch.device(env.device)
)
self.action_mode = action_mode
self.actor = actor
self.actor_algo_type = str(actor_algo_type)
self.normalizer = normalizer
self.num_envs = int(num_envs)
self.obs_extractor = obs_extractor
self.obs: np.ndarray | None = None
self.obs: torch.Tensor | None = None
self.step_count = 0

def reset(self) -> np.ndarray:
def reset(self) -> torch.Tensor:
if self.env.state is None:
self.env.init_state()
env_indices = np.arange(self.num_envs, dtype=np.int32)
env_indices = torch.arange(self.num_envs, dtype=torch.int64, device=self.env.device)
reset_result = self.env.reset(env_indices)
if not isinstance(reset_result, tuple) or len(reset_result) != 2:
raise ValueError(f"Unexpected env.reset return format: {type(reset_result)!r}")
obs_out, _ = reset_result
self.obs = np.asarray(self.obs_extractor(obs_out), dtype=np.float32)
self.obs = _to_play_tensor(self.obs_extractor(obs_out), device=self.policy_device)
self.step_count = 0
return self.obs

def step_once(self) -> np.ndarray:
def step_once(self) -> torch.Tensor:
actions = self._build_actions()
state = self.env.step(actions)
self.obs = np.asarray(self.obs_extractor(state.obs), dtype=np.float32)
self.obs = _to_play_tensor(self.obs_extractor(state.obs), device=self.policy_device)
self.step_count += 1
return self.obs

Expand All @@ -334,26 +340,33 @@ def info(self) -> dict[str, Any]:
info = getattr(state, "info", None)
return info if isinstance(info, dict) else {}

def _build_actions(self) -> np.ndarray:
def _build_actions(self) -> torch.Tensor:
if self.obs is None:
raise RuntimeError("Playback session must be reset before stepping.")
action_space = self.env.action_space
action_dim = int(action_space.shape[0])
if self.action_mode == "policy" and self.actor is not None:
obs_torch = torch.from_numpy(self.obs).to(self.device)
if obs_torch.dtype != torch.float32:
obs_torch = obs_torch.float()
obs_torch = self.obs.to(device=self.policy_device, dtype=torch.float32)
if self.normalizer is not None:
obs_torch = self.normalizer(obs_torch, update=False)
actions = self.actor.explore(obs_torch, deterministic=True)
return actions.detach().cpu().numpy().astype(np.float32)
return actions.detach().to(device=self.env.device, dtype=torch.float32).contiguous()
if self.action_mode == "random":
return np.random.uniform(
action_space.low,
action_space.high,
size=(self.num_envs, action_dim),
).astype(np.float32)
return np.zeros((self.num_envs, action_dim), dtype=np.float32)
low = torch.as_tensor(action_space.low, device=self.env.device).to(torch.float32)
high = torch.as_tensor(action_space.high, device=self.env.device).to(torch.float32)
return low + (high - low) * torch.rand(
(self.num_envs, action_dim), device=self.env.device, dtype=torch.float32
)
return torch.zeros((self.num_envs, action_dim), dtype=torch.float32, device=self.env.device)


def _to_play_tensor(value: Any, *, device: torch.device) -> torch.Tensor:
if not isinstance(value, torch.Tensor):
raise TypeError(
"Off-policy playback observations must be torch tensors; "
f"received {type(value).__name__}"
)
return value.to(device=device, dtype=torch.float32).contiguous()


def make_sim2sim_preflight(
Expand Down Expand Up @@ -766,13 +779,20 @@ def resolve_play_obs_dims(obs_groups_spec: dict[str, int]) -> tuple[int, int]:
return int(obs_dim), int(critic_obs_dim)


def extract_play_obs(obs_dict):
def extract_play_obs(obs_dict: dict[str, torch.Tensor]) -> torch.Tensor:
from uni_rl.utils.observations import split_obs_dict

obs_out, _ = split_obs_dict(obs_dict)
return obs_out


def resolve_playback_device(env_device: Any, requested_device: str) -> torch.device:
"""Resolve actor placement from the authoritative TorchEnv device."""
if env_device is None:
return torch.device(requested_device)
return torch.device(env_device)


def resolve_play_actor_spec(
algo_name: str,
cfg: DictConfig,
Expand Down Expand Up @@ -934,6 +954,7 @@ def create_sac_playback_session(
if action_shape is None:
raise ValueError("env.action_space.shape must be defined")
action_dim = int(action_shape[0])
session_device = resolve_playback_device(getattr(env, "device", None), device_name)
actor_algo_type, actor_kwargs = resolve_play_actor_spec(
algo_name,
cfg,
Expand All @@ -955,15 +976,15 @@ def create_sac_playback_session(
if bool(getattr(cfg.algo, "obs_normalization", False)):
from uni_rl.algos.common.normalization import EmpiricalNormalization

normalizer = EmpiricalNormalization(shape=obs_dim, device=device_name)
normalizer = EmpiricalNormalization(shape=obs_dim, device=session_device)
if playback_cfg.action_mode == "policy":
actor = build_actor(
actor_algo_type,
obs_dim,
action_dim,
cfg.algo.actor_hidden_dim,
cfg.algo.use_layer_norm,
device_name,
str(session_device),
**actor_kwargs,
)
actor.eval()
Expand All @@ -986,7 +1007,7 @@ def create_sac_playback_session(
algo_name=algo_name,
strict=bool(getattr(cfg.training, "sim2sim_strict", True)),
)
checkpoint = torch.load(checkpoint_path, map_location=device_name, weights_only=True)
checkpoint = torch.load(checkpoint_path, map_location=session_device, weights_only=True)
with policy_load_dim_guard(
env_obs_dim=obs_dim,
env_action_dim=action_dim,
Expand Down
22 changes: 12 additions & 10 deletions tests/scripts/test_train_scripts.py
Original file line number Diff line number Diff line change
Expand Up @@ -18,6 +18,7 @@

import numpy as np
import pytest
import torch
from hydra import compose, initialize_config_dir
from hydra.core.global_hydra import GlobalHydra
from omegaconf import OmegaConf
Expand Down Expand Up @@ -1220,14 +1221,12 @@ def test_build_appo_runner_kwargs_forwards_sim_backend():


def test_run_motrix_play_loop_runs_without_physics_state():
import numpy as np
import torch

mod = _train_appo()

class FakeActor:
def __call__(self, td):
batch = td.batch_size[0]
assert td["policy"].device.type == "cpu"
return torch.zeros((batch, 3), dtype=torch.float32)

class FakeBackend:
Expand All @@ -1244,10 +1243,11 @@ def render(self):

class FakeState:
def __init__(self):
self.obs = {"obs": np.ones((2, 5), dtype=np.float32)}
self.obs = {"obs": torch.ones((2, 5), dtype=torch.float32)}

class FakeEnv:
def __init__(self):
self.device = torch.device("cpu")
self.state = None
self._renderer = FakeBackend()
self.init_state_calls = 0
Expand All @@ -1261,11 +1261,15 @@ def init_state(self):
def reset(self, env_indices):
self.reset_calls += 1
assert env_indices.shape == (2,)
return {"obs": np.ones((2, 5), dtype=np.float32)}, {}
assert env_indices.dtype == torch.int64
assert env_indices.device == self.device
return {"obs": torch.ones((2, 5), dtype=torch.float32, device=self.device)}, {}

def step(self, actions):
self.step_calls += 1
assert actions.shape == (2, 3)
assert actions.dtype == torch.float32
assert actions.device == self.device
return FakeState()

def init_play_renderer(self, render_spacing=None, render_offset_mode=None):
Expand Down Expand Up @@ -1640,19 +1644,17 @@ def test_offpolicy_resolve_play_obs_dim_ignores_critic():


def test_offpolicy_extract_play_obs_uses_obs_group_only():
import numpy as np

from unilab.visualization.interactive_playback import extract_play_obs

obs = {
"obs": np.ones((2, 98), dtype=np.float32),
"critic": np.full((2, 101), 2.0, dtype=np.float32),
"obs": torch.ones((2, 98), dtype=torch.float32),
"critic": torch.full((2, 101), 2.0, dtype=torch.float32),
}

play_obs = extract_play_obs(obs)

assert play_obs.shape == (2, 98)
assert np.allclose(play_obs, 1.0)
assert torch.allclose(play_obs, torch.ones_like(play_obs))


def test_offpolicy_play_actor_spec_keeps_standard_sac_and_flashsac():
Expand Down
Loading