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
3 changes: 2 additions & 1 deletion conf/offpolicy/config.yaml
Original file line number Diff line number Diff line change
Expand Up @@ -21,7 +21,8 @@ training:
play_only: false
no_play: false
play_env_num: 16
play_steps: 200
play_steps: 800
cam_distance: 6.0
cam_elevation: -20.0
cam_azimuth: 90.0
log_dir: null
Expand Down
37 changes: 37 additions & 0 deletions conf/offpolicy/task/sac/g1_sac_wbt/mujoco.yaml
Original file line number Diff line number Diff line change
@@ -0,0 +1,37 @@
# @package _global_
# G1 Whole-Body Tracking (WBT) with FastSAC on MuJoCo.
# Hyperparameters aligned with holosoma g1-29dof-wbt-fast-sac.
training:
task_name: G1MotionTrackingSAC
sim_backend: mujoco
algo:
num_envs: 2048
max_iterations: 100000
save_interval: 1000
# --- holosoma WBT-specific overrides (vs sac.yaml defaults) ---
gamma: 0.99
tau: 0.05
num_atoms: 501
updates_per_step: 4
policy_frequency: 2
use_symmetry: false
algo_params:
alpha_init: 0.02
target_entropy_ratio: 0.5
max_grad_norm: 10.0
env:
control_config:
action_scale: 2.0
reward:
scales:
motion_global_root_pos: 1.0
motion_global_root_ori: 0.5
motion_body_pos: 2.0
motion_body_ori: 1.0
motion_body_lin_vel: 1.0
motion_body_ang_vel: 1.0
motion_joint_pos: 0.0
motion_joint_vel: 0.0
action_rate_l2: -1.0
joint_limit: -10.0
undesired_contacts: -0.1
1 change: 1 addition & 0 deletions docs/users/zh_CN/02-simulation-backends.md
Original file line number Diff line number Diff line change
Expand Up @@ -67,6 +67,7 @@ uv run scripts/generate_support_matrix.py --write
| APPO (torch) | `allegro_inhand` (Allegro in-hand) | Tested | Registered |
| SAC (torch) | `g1_walk_flat` (G1 walk flat) | Tested | Tested |
| SAC (torch) | `g1_walk_rough` (G1 walk rough) | Tested | Registered |
| SAC (torch) | `g1_sac_wbt` (g1 sac wbt) | Tested | - |
| TD3 (torch) | `g1_walk_flat` (G1 walk flat) | Tested | Registered |

### Source Index
Expand Down
39 changes: 36 additions & 3 deletions docs/users/zh_CN/05-motion-tracking.md
Original file line number Diff line number Diff line change
Expand Up @@ -2,17 +2,20 @@

语言: 简体中文

UniLab 当前提供两个 G1 whole-body motion tracking task family:
UniLab 当前提供两个 G1 whole-body motion tracking task family 和一个 FastSAC off-policy 入口:

- task family:`g1_motion_tracking`
- 注册环境名:`G1MotionTracking`
- task family:`g1_flip_tracking`(flip 专用 profile)
- 注册环境名:`G1FlipTracking`
- 后端注册:`mujoco` 和 `motrix`
- task family:`g1_sac_wbt`(FastSAC off-policy,从 holosoma 迁移)
- 注册环境名:`G1MotionTrackingSAC`
- 后端注册:`mujoco` 和 `motrix`(`g1_sac_wbt` 目前仅 `mujoco`)
- 已提交的 Motrix 特化配置:PPO 和 APPO 的 motion-tracking reward
- `g1_motion_tracking` 默认 motion:`src/unilab/assets/motions/g1/dance1_subject2_part.npz`
- `g1_flip_tracking` 默认 motion:`src/unilab/assets/motions/g1/flip_360_001__A304.npz`
- 实际训练入口统一写成 `task=<family>/<backend>`
- `g1_sac_wbt` 默认 motion:与 `g1_motion_tracking` 相同
- 实际训练入口统一写成 `task=<family>/<backend>`(off-policy 需加 `algo` 前缀)

## Environment Entrypoints

Expand All @@ -35,6 +38,12 @@ uv run python scripts/train_appo.py task=g1_motion_tracking/mujoco
# APPO (Motrix)
uv run python scripts/train_appo.py task=g1_motion_tracking/motrix

# FastSAC (MuJoCo, holosoma-aligned WBT)
uv run scripts/train_offpolicy.py algo=sac task=sac/g1_sac_wbt/mujoco

# FastSAC with AMP (recommended for CUDA)
uv run scripts/train_offpolicy.py algo=sac task=sac/g1_sac_wbt/mujoco training.use_amp=true

# 回放最新 checkpoint
uv run python scripts/train_rsl_rl.py task=g1_motion_tracking/mujoco training.play_only=true

Expand Down Expand Up @@ -118,6 +127,30 @@ uv run python scripts/motion/replay_npz.py \
--speed 0.5
```

## FastSAC WBT (holosoma Migration)

`task=sac/g1_sac_wbt/mujoco` 提供从 holosoma 迁移的 G1 whole-body tracking FastSAC 训练。超参数对齐 holosoma `exp:g1-29dof-wbt-fast-sac`(`gamma=0.99, tau=0.05, num_atoms=501, target_entropy_ratio=0.5`)。环境 `G1MotionTrackingSAC` 在 PPO 版 `G1MotionTracking` 基础上为 critic 增加了 `base_lin_vel`(asymmetric actor-critic)。

```bash
# 默认训练
uv run scripts/train_offpolicy.py algo=sac task=sac/g1_sac_wbt/mujoco

# 推荐:开启 AMP 加速
uv run scripts/train_offpolicy.py algo=sac task=sac/g1_sac_wbt/mujoco training.use_amp=true

# 自定义并行环境数和迭代次数
uv run scripts/train_offpolicy.py algo=sac task=sac/g1_sac_wbt/mujoco \
algo.num_envs=4096 algo.max_iterations=10000 training.use_amp=true

# 使用 wandb 记录
uv run scripts/train_offpolicy.py algo=sac task=sac/g1_sac_wbt/mujoco \
training.use_amp=true training.logger=wandb

# 指定 motion 文件
uv run scripts/train_offpolicy.py algo=sac task=sac/g1_sac_wbt/mujoco \
env.motion_file=src/unilab/assets/motions/g1/dance1_subject2_part.npz
```

## Configuration Note

`task=g1_motion_tracking/mujoco` 默认读取环境配置里的单个 `motion_file`(历史默认是 `dance1_subject2_part.npz`)。`task=g1_flip_tracking/mujoco` 提供 flip 专用默认 profile(更保守的 reset 随机化与 termination)。
Expand Down
5 changes: 5 additions & 0 deletions scripts/train_offpolicy.py
Original file line number Diff line number Diff line change
Expand Up @@ -438,6 +438,11 @@ def _policy_step(obs_np: np.ndarray) -> np.ndarray:
dtype=np.float32,
),
step=_policy_step,
camera_kwargs={
"cam_distance": cfg.training.cam_distance,
"cam_elevation": cfg.training.cam_elevation,
"cam_azimuth": cfg.training.cam_azimuth,
},
)
print(f"Saving video to {output_video} ...")
print("Done.")
Expand Down
Binary file not shown.
4 changes: 4 additions & 0 deletions src/unilab/envs/motion_tracking/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -9,12 +9,16 @@
G1MotionTrackingCfg,
G1MotionTrackingEnv,
G1MotionTrackingEnvCfg,
G1MotionTrackingSACCfg,
G1MotionTrackingSACEnv,
)

__all__ = [
"G1MotionTrackingCfg",
"G1MotionTrackingEnv",
"G1MotionTrackingEnvCfg",
"G1MotionTrackingSACCfg",
"G1MotionTrackingSACEnv",
"G1FlipTrackingCfg",
"G1FlipTrackingEnv",
"G1FlipTrackingEnvCfg",
Expand Down
3 changes: 3 additions & 0 deletions src/unilab/envs/motion_tracking/g1/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -2,11 +2,14 @@

from .flip_tracking import G1FlipTrackingCfg, G1FlipTrackingEnv, G1FlipTrackingEnvCfg
from .tracking import G1MotionTrackingCfg, G1MotionTrackingEnv, G1MotionTrackingEnvCfg
from .tracking_sac import G1MotionTrackingSACCfg, G1MotionTrackingSACEnv

__all__ = [
"G1MotionTrackingCfg",
"G1MotionTrackingEnv",
"G1MotionTrackingEnvCfg",
"G1MotionTrackingSACCfg",
"G1MotionTrackingSACEnv",
"G1FlipTrackingCfg",
"G1FlipTrackingEnv",
"G1FlipTrackingEnvCfg",
Expand Down
66 changes: 66 additions & 0 deletions src/unilab/envs/motion_tracking/g1/tracking_sac.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,66 @@
"""G1 Motion Tracking SAC Environment — thin SAC wrapper over G1MotionTrackingEnv.

Differences from the PPO base:
- Critic observations additionally include ``base_lin_vel`` (3 dims),
matching holosoma's asymmetric actor-critic design for WBT.
- Registered under a separate name so it can be paired with FastSAC
configs without affecting the PPO motion-tracking pipeline.
"""

from __future__ import annotations

from dataclasses import dataclass

import numpy as np

from unilab.base import registry
from unilab.base.dtype_config import get_global_dtype

from .tracking import G1MotionTrackingCfg, G1MotionTrackingEnv


@registry.envcfg("G1MotionTrackingSAC")
@dataclass
class G1MotionTrackingSACCfg(G1MotionTrackingCfg):
"""Config for SAC-based motion tracking (identical fields, separate registry entry)."""


@registry.env("G1MotionTrackingSAC", sim_backend="mujoco")
class G1MotionTrackingSACEnv(G1MotionTrackingEnv):
"""G1 Motion Tracking environment for FastSAC training.

Extends the PPO motion-tracking environment with ``base_lin_vel``
appended to the critic observation, matching holosoma's asymmetric
actor-critic WBT design.
"""

@property
def obs_groups_spec(self) -> dict[str, int]:
spec = super().obs_groups_spec
# Append base_lin_vel (3) to critic observations.
return {**spec, "critic": spec["critic"] + 3}

def _compute_obs(
self,
info: dict,
motion_data,
linvel: np.ndarray,
gyro: np.ndarray,
dof_pos: np.ndarray,
dof_vel: np.ndarray,
robot_body_pos_w: np.ndarray,
robot_body_quat_w: np.ndarray,
) -> dict[str, np.ndarray]:
obs = super()._compute_obs( # pyright: ignore[reportAttributeAccessIssue]
info,
motion_data,
linvel,
gyro,
dof_pos,
dof_vel,
robot_body_pos_w,
robot_body_quat_w,
)
# Append base_lin_vel to critic observations.
obs["critic"] = np.concatenate([obs["critic"], linvel], axis=1, dtype=get_global_dtype()) # type: ignore[call-overload]
return obs
2 changes: 1 addition & 1 deletion src/unilab/utils/render_many.py
Original file line number Diff line number Diff line change
Expand Up @@ -270,7 +270,7 @@ def set_state(model, d, s, offset=None):
center_x = np.mean(offsets[:, 0])
center_y = np.mean(offsets[:, 1])
if cam_lookat is None:
cam.lookat = [center_x, center_y, 0.0]
cam.lookat = [center_x, center_y, 0.75]
else:
cam.lookat = [float(cam_lookat[0]), float(cam_lookat[1]), float(cam_lookat[2])]
cam.distance = cam_distance
Expand Down