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
8 changes: 4 additions & 4 deletions conf/offpolicy/algo/flashsac.yaml
Original file line number Diff line number Diff line change
Expand Up @@ -3,16 +3,16 @@ algo: flashsac
algo_log_name: flash_sac
load_run: "-1"
seed: 1
num_envs: 2048
num_envs: 1024
batch_size: 2048
replay_buffer_n: 512
updates_per_step: 1
warmup_steps: 10000
updates_per_step: 2
warmup_steps: 100000
policy_frequency: 2
env_steps_per_sync: 1
max_iterations: 5000
save_interval: 1000
gamma: 0.99
gamma: 0.97
tau: 0.01
actor_lr: 3.0e-4
critic_lr: 3.0e-4
Expand Down
39 changes: 39 additions & 0 deletions conf/offpolicy/task/flashsac/g1_joystick/mujoco.yaml
Original file line number Diff line number Diff line change
@@ -0,0 +1,39 @@
# @package _global_
training:
task_name: G1JoystickFlatTerrain
sim_backend: mujoco
use_amp: true
algo:
num_envs: 1024
replay_buffer_n: 9766
warmup_steps: 100000
updates_per_step: 2
max_iterations: 48829
save_interval: 5000
gamma: 0.97
env:
control_config:
action_scale: 0.5
commands:
vel_limit:
- [-1.0, -0.5, -1.0]
- [1.0, 0.5, 1.0]
reward:
scales:
tracking_lin_vel: 1.0
tracking_ang_vel: 0.75
feet_phase: 1.0
lin_vel_z: 0.0
ang_vel_xy: -0.15
base_height: 0.0
orientation: -2.0
action_rate: 0.0
pose: -0.1
tracking_sigma: 0.25
gait_frequency: 1.5
feet_phase_swing_height: 0.15
feet_phase_tracking_sigma: 0.01
base_height_target: 0.754
min_base_height: 0.55
max_tilt_deg: 25.0
pose_weights: [0.01, 1.0, 5.0, 0.01, 5.0, 5.0, 0.01, 1.0, 5.0, 0.01, 5.0, 5.0, 50.0, 50.0, 50.0, 50.0, 50.0, 50.0, 50.0, 50.0, 50.0, 50.0, 50.0, 50.0, 50.0, 50.0, 50.0, 50.0, 50.0]
10 changes: 6 additions & 4 deletions docs/zh_CN/02-simulation-backends.md
Original file line number Diff line number Diff line change
Expand Up @@ -45,23 +45,25 @@ uv run python scripts/generate_support_matrix.py --write
| PPO (torch) | `g1_joystick` (G1 joystick) | Tested | Tested |
| PPO (torch) | `g1_motion_tracking` (G1 motion tracking) | Tested | Tested |
| PPO (torch) | `g1_flip_tracking` (G1 flip tracking) | Tested | Tested |
| PPO (torch) | `allegro_inhand` (Allegro in-hand) | Tested | - |
| PPO (torch) | `allegro_inhand` (Allegro in-hand) | Tested | Tested |
| PPO (torch) | `allegro_inhand_grasp` (allegro inhand grasp) | Tested | Tested |
| PPO (mlx) | `go1_joystick` (Go1 joystick) | Tested | Tested |
| PPO (mlx) | `go2_joystick` (Go2 joystick) | Tested | Tested |
| PPO (mlx) | `g1_joystick` (G1 joystick) | Tested | Tested |
| PPO (mlx) | `g1_motion_tracking` (G1 motion tracking) | Configured | Configured |
| PPO (mlx) | `g1_flip_tracking` (G1 flip tracking) | Configured | Configured |
| PPO (mlx) | `allegro_inhand` (Allegro in-hand) | Configured | - |
| PPO (mlx) | `allegro_inhand` (Allegro in-hand) | Configured | Configured |
| PPO (mlx) | `allegro_inhand_grasp` (allegro inhand grasp) | Configured | Configured |
| APPO (torch) | `go1_joystick` (Go1 joystick) | Tested | Registered |
| APPO (torch) | `go2_joystick` (Go2 joystick) | Tested | Registered |
| APPO (torch) | `g1_joystick` (G1 joystick) | Tested | Registered |
| APPO (torch) | `g1_motion_tracking` (G1 motion tracking) | Tested | Tested |
| APPO (torch) | `g1_flip_tracking` (G1 flip tracking) | Tested | Tested |
| APPO (torch) | `allegro_inhand` (Allegro in-hand) | Tested | - |
| APPO (torch) | `allegro_inhand` (Allegro in-hand) | Tested | Registered |
| SAC (torch) | `go1_joystick` (Go1 joystick) | Tested | Tested |
| SAC (torch) | `go2_joystick` (Go2 joystick) | Tested | Tested |
| SAC (torch) | `g1_sac` (G1 SAC locomotion) | Tested | Tested |
| SAC (torch) | `allegro_sac` (Allegro SAC in-hand) | Tested | - |
| SAC (torch) | `allegro_sac` (Allegro SAC in-hand) | - | - |
| TD3 (torch) | `go1_joystick` (Go1 joystick) | Tested | Tested |
| TD3 (torch) | `go2_joystick` (Go2 joystick) | Tested | Tested |

Expand Down
4 changes: 2 additions & 2 deletions docs/zh_CN/07-domain-randomization.md
Original file line number Diff line number Diff line change
Expand Up @@ -25,7 +25,7 @@
| `G1WalkTaskMjSAC` | 是 | 是:复用 [`G1JoystickDomainRandomizationProvider`](../../src/unilab/envs/locomotion/g1/joystick.py) | 任务状态采样 + common payload | push | [`g1/joystick_sac.py`](../../src/unilab/envs/locomotion/g1/joystick_sac.py) |
| `G1MotionTracking` | 是 | 是:`Domain_Rand + Provider + ResetPlan` | 大量任务特有 reset 采样 + common payload | push | [`motion_tracking/g1/tracking.py`](../../src/unilab/envs/motion_tracking/g1/tracking.py) |
| `AllegroInhandRotation` | 是 | 是:`DomainRandConfig + Provider + ResetPlan` | 纯任务特有 reset 采样,`randomization=None` | 无 | [`inhand_rot_allegro/rotation.py`](../../src/unilab/envs/manipulation/inhand_rot_allegro/rotation.py) |
| `AllegroInhandRotationSac` | 是 | 是:`DomainRandConfigSac + Provider + ResetPlan` | 纯任务特有 reset 采样,`randomization=None` | 无 | [`inhand_rot_allegro/rotation_sac.py`](../../src/unilab/envs/manipulation/inhand_rot_allegro/rotation_sac.py) |
| `AllegroInhandRotationSac` | 是 | 是:`DomainRandConfigSac + Provider + ResetPlan` | 纯任务特有 reset 采样,`randomization=None` | 无 | [`inhand_rot_allegro/README.md`](../../src/unilab/envs/manipulation/inhand_rot_allegro/README.md) |

## 任务域随机化清单

Expand Down Expand Up @@ -100,7 +100,7 @@ backend capability 当前是:
- [`go1/joystick.py`](../../src/unilab/envs/locomotion/go1/joystick.py)
- [`go2/joystick.py`](../../src/unilab/envs/locomotion/go2/joystick.py)
- [`inhand_rot_allegro/rotation.py`](../../src/unilab/envs/manipulation/inhand_rot_allegro/rotation.py)
- [`inhand_rot_allegro/rotation_sac.py`](../../src/unilab/envs/manipulation/inhand_rot_allegro/rotation_sac.py)
- [`inhand_rot_allegro/README.md`](../../src/unilab/envs/manipulation/inhand_rot_allegro/README.md)

因此从“当前任务 DR 状态”角度,`kp/kd` 还不能算各任务已经统一接入的域随机项。

Expand Down
3 changes: 3 additions & 0 deletions scripts/train_offpolicy.py
Original file line number Diff line number Diff line change
Expand Up @@ -16,6 +16,7 @@

from unilab.training import (
BackendAdapter,
assert_offpolicy_task_choice_matches_algo,
create_env,
ensure_registries,
get_log_root,
Expand Down Expand Up @@ -76,6 +77,7 @@ def extract_play_obs(obs_dict):


def build_offpolicy_env_cfg_override(algo_name: str, cfg: DictConfig) -> dict[str, Any] | None:
assert_offpolicy_task_choice_matches_algo(cfg, algo_name=algo_name)
return cast(
dict[str, Any] | None,
BackendAdapter(cfg, root_dir=ROOT_DIR, algo_name=algo_name).build_task_env_cfg_override(),
Expand Down Expand Up @@ -444,6 +446,7 @@ def main(cfg: DictConfig) -> None:

algo_name = cfg.algo.algo
task_name = cfg.training.task_name
assert_offpolicy_task_choice_matches_algo(cfg, algo_name=algo_name)

if cfg.training.log_dir is None:
timestamp = datetime.datetime.now().strftime("%Y-%m-%d_%H-%M-%S")
Expand Down
67 changes: 55 additions & 12 deletions src/unilab/algos/torch/flash_sac/learner.py
Original file line number Diff line number Diff line change
Expand Up @@ -68,38 +68,72 @@ def load_state_dict(self, state_dict: dict[str, torch.Tensor]) -> None:


class RewardNormalizer:
"""Adaptive reward scaling with running reward statistics."""
"""Adaptive reward scaling with running discounted-return statistics."""

def __init__(self, g_max: float, device: torch.device, eps: float = 1e-8):
def __init__(
self,
gamma: float,
g_max: float,
device: torch.device,
eps: float = 1e-8,
):
self.gamma = gamma
self.g_max = g_max
self.eps = eps
self.device = device
self.rms = RunningMeanStd.create(device)
self.reward_abs_max = torch.tensor(0.0, device=device, dtype=torch.float32)
self.g_r = torch.zeros(0, device=device, dtype=torch.float32)
self.g_r_max = torch.tensor(0.0, device=device, dtype=torch.float32)

def _ensure_g_r_shape(self, num_envs: int) -> None:
if self.g_r.shape == (num_envs,):
return
self.g_r = torch.zeros(num_envs, device=self.device, dtype=torch.float32)

def update(self, rewards: torch.Tensor) -> None:
rewards = rewards.reshape(-1).to(device=self.device, dtype=torch.float32)
def update_from_transitions(
self,
rewards: torch.Tensor,
terminated: torch.Tensor,
truncated: torch.Tensor,
) -> None:
rewards = rewards.to(device=self.device, dtype=torch.float32)
terminated = terminated.to(device=self.device, dtype=torch.float32)
truncated = truncated.to(device=self.device, dtype=torch.float32)

if rewards.ndim == 1:
rewards = rewards.unsqueeze(0)
terminated = terminated.unsqueeze(0)
truncated = truncated.unsqueeze(0)
if rewards.numel() == 0:
return
self.rms.update(rewards)
self.reward_abs_max = torch.maximum(self.reward_abs_max, rewards.abs().max())

num_envs = int(rewards.shape[-1])
self._ensure_g_r_shape(num_envs)
done = torch.clamp(terminated + truncated, min=0.0, max=1.0)

for step in range(rewards.shape[0]):
self.g_r = self.gamma * (1.0 - done[step]) * self.g_r + rewards[step]
self.g_r_max = torch.maximum(self.g_r_max, self.g_r.abs().max())
self.rms.update(self.g_r)

def normalize(self, rewards: torch.Tensor) -> torch.Tensor:
denominator = torch.maximum(
torch.sqrt(self.rms.var + self.eps),
self.reward_abs_max / max(self.g_max, self.eps),
self.g_r_max / max(self.g_max, self.eps),
)
return rewards / denominator

def state_dict(self) -> dict[str, Any]:
return {
"rms": self.rms.state_dict(),
"reward_abs_max": self.reward_abs_max,
"g_r": self.g_r,
"g_r_max": self.g_r_max,
}

def load_state_dict(self, state_dict: dict[str, Any]) -> None:
self.rms.load_state_dict(state_dict["rms"])
self.reward_abs_max = state_dict["reward_abs_max"]
self.g_r = state_dict["g_r"]
self.g_r_max = state_dict["g_r_max"]


class FlashSACLearner:
Expand Down Expand Up @@ -188,7 +222,7 @@ def __init__(
self.obs_normalizer = nn.Identity()

self.reward_normalizer = (
RewardNormalizer(g_max=normalized_g_max, device=self.device)
RewardNormalizer(gamma=self.gamma, g_max=normalized_g_max, device=self.device)
if normalize_reward
else None
)
Expand Down Expand Up @@ -237,6 +271,16 @@ def _build_critic_obs(
def _autocast(self):
return torch.autocast(device_type="cuda", dtype=torch.float16, enabled=self.use_amp)

def update_reward_stats(
self,
rewards: torch.Tensor,
terminated: torch.Tensor,
truncated: torch.Tensor,
) -> None:
if self.reward_normalizer is None:
return
self.reward_normalizer.update_from_transitions(rewards, terminated, truncated)

@staticmethod
def _set_requires_grad(module: nn.Module, requires_grad: bool) -> None:
for param in module.parameters():
Expand All @@ -260,7 +304,6 @@ def update_critic(self, batch: dict[str, torch.Tensor]) -> dict[str, float]:
critic_next_obs = self._build_critic_obs(next_obs, next_privileged)

if self.reward_normalizer is not None:
self.reward_normalizer.update(rewards)
rewards = self.reward_normalizer.normalize(rewards)

gamma = self.gamma**self.n_step
Expand Down
58 changes: 58 additions & 0 deletions src/unilab/algos/torch/offpolicy/runner.py
Original file line number Diff line number Diff line change
Expand Up @@ -5,6 +5,7 @@
import sys
import time
from collections import deque
from typing import cast

import torch

Expand Down Expand Up @@ -78,6 +79,57 @@ def _build_learner(self):
def _collector_fn(self, stop_event, **kwargs):
off_policy_collector_fn(stop_event=stop_event, **kwargs)

@staticmethod
def _read_recent_replay_field(
replay_buffer, field_name: str, start_ptr: int, count: int
) -> torch.Tensor:
idx = start_ptr % replay_buffer.capacity

if hasattr(replay_buffer, field_name):
source = getattr(replay_buffer, field_name)
else:
packed_key = {
"rewards": "_rew_col",
"dones": "_done_col",
"truncated": "_trunc_col",
}[field_name]
source = replay_buffer._storage[:, getattr(replay_buffer, packed_key)]

if idx + count <= replay_buffer.capacity:
return cast(torch.Tensor, source[idx : idx + count].clone())

split = replay_buffer.capacity - idx
return cast(torch.Tensor, torch.cat([source[idx:], source[: count - split]], dim=0).clone())

def _update_reward_stats_from_replay(self, replay_buffer, start_ptr: int, end_ptr: int) -> int:
if not hasattr(self.learner, "update_reward_stats"):
return end_ptr
if getattr(self.learner, "reward_normalizer", None) is None:
return end_ptr

count = end_ptr - start_ptr
if count <= 0:
return end_ptr
if count > replay_buffer.capacity:
count = replay_buffer.capacity
start_ptr = end_ptr - count
if count % self.num_envs != 0:
count -= count % self.num_envs
start_ptr = end_ptr - count
if count <= 0:
return end_ptr

rewards = self._read_recent_replay_field(replay_buffer, "rewards", start_ptr, count)
dones = self._read_recent_replay_field(replay_buffer, "dones", start_ptr, count)
truncated = self._read_recent_replay_field(replay_buffer, "truncated", start_ptr, count)
num_steps = count // self.num_envs
self.learner.update_reward_stats(
rewards.view(num_steps, self.num_envs),
dones.view(num_steps, self.num_envs),
truncated.view(num_steps, self.num_envs),
)
return end_ptr

def learn(
self,
max_iterations: int = 1500,
Expand Down Expand Up @@ -180,6 +232,7 @@ def learn(
latest_reward_components: dict[str, float] = {}
last_buf_log = 0
write_read_ema = 0.0
reward_stats_ptr = 0

# Training loop
for iteration in range(1, max_iterations + 1):
Expand Down Expand Up @@ -245,6 +298,11 @@ def learn(
wait_time = time.time() - wait_start
self._drain_metrics(metrics_queue, reward_history, latest_reward_components, logger)
collect_time = time.time() - iter_start
reward_stats_ptr = self._update_reward_stats_from_replay(
replay_buffer,
reward_stats_ptr,
int(replay_buffer.ptr[0]),
)

train_start = time.time()
from collections import defaultdict
Expand Down
28 changes: 23 additions & 5 deletions src/unilab/algos/torch/offpolicy/worker.py
Original file line number Diff line number Diff line change
Expand Up @@ -6,6 +6,7 @@

import queue
import sys
from typing import cast

import numpy as np
import torch
Expand Down Expand Up @@ -39,6 +40,21 @@ def resolve_collector_actor_dims(
return obs_dim, action_dim


def sample_offpolicy_actions(
actor,
algo_type: str,
obs_torch: torch.Tensor,
prev_dones_torch: torch.Tensor,
) -> torch.Tensor:
"""Sample collector actions using the algorithm's exploration policy."""
if algo_type in ("sac", "td3", "flashsac"):
return cast(
torch.Tensor,
actor.explore(obs_torch, dones=prev_dones_torch, deterministic=False),
)
raise ValueError(f"Unsupported off-policy algo_type for collector action sampling: {algo_type}")


def off_policy_collector_fn(
stop_event,
env_name: str,
Expand Down Expand Up @@ -245,11 +261,13 @@ def _run_collector(
else:
_t_infer = _time.perf_counter()
obs_torch = torch.from_numpy(obs_np_input)
if algo_type in ("sac", "td3"):
dones_torch = torch.from_numpy(prev_dones_np)
actions_torch = actor.explore(obs_torch, dones=dones_torch, deterministic=False)
else:
actions_torch = torch.zeros((num_envs, action_dim))
dones_torch = torch.from_numpy(prev_dones_np)
actions_torch = sample_offpolicy_actions(
actor=actor,
algo_type=algo_type,
obs_torch=obs_torch,
prev_dones_torch=dones_torch,
)
actions_np = actions_torch.numpy()
timing_accum_ms["mlp_infer_ms"] += (_time.perf_counter() - _t_infer) * 1000

Expand Down
Loading