add scripts

This commit is contained in:
Can Chen 2026-05-04 10:26:20 +08:00
parent bccbb759d8
commit e1e58a6574
336 changed files with 32001 additions and 11 deletions

159
LICENSE-CC-BY-NC-4.0.md Normal file
View File

@ -0,0 +1,159 @@
# Attribution-NonCommercial 4.0 International
> *Creative Commons Corporation (“Creative Commons”) is not a law firm and does not provide legal services or legal advice. Distribution of Creative Commons public licenses does not create a lawyer-client or other relationship. Creative Commons makes its licenses and related information available on an “as-is” basis. Creative Commons gives no warranties regarding its licenses, any material licensed under their terms and conditions, or any related information. Creative Commons disclaims all liability for damages resulting from their use to the fullest extent possible.*
>
> ### Using Creative Commons Public Licenses
>
> Creative Commons public licenses provide a standard set of terms and conditions that creators and other rights holders may use to share original works of authorship and other material subject to copyright and certain other rights specified in the public license below. The following considerations are for informational purposes only, are not exhaustive, and do not form part of our licenses.
>
> * __Considerations for licensors:__ Our public licenses are intended for use by those authorized to give the public permission to use material in ways otherwise restricted by copyright and certain other rights. Our licenses are irrevocable. Licensors should read and understand the terms and conditions of the license they choose before applying it. Licensors should also secure all rights necessary before applying our licenses so that the public can reuse the material as expected. Licensors should clearly mark any material not subject to the license. This includes other CC-licensed material, or material used under an exception or limitation to copyright. [More considerations for licensors](http://wiki.creativecommons.org/Considerations_for_licensors_and_licensees#Considerations_for_licensors).
>
> * __Considerations for the public:__ By using one of our public licenses, a licensor grants the public permission to use the licensed material under specified terms and conditions. If the licensors permission is not necessary for any reasonfor example, because of any applicable exception or limitation to copyrightthen that use is not regulated by the license. Our licenses grant only permissions under copyright and certain other rights that a licensor has authority to grant. Use of the licensed material may still be restricted for other reasons, including because others have copyright or other rights in the material. A licensor may make special requests, such as asking that all changes be marked or described. Although not required by our licenses, you are encouraged to respect those requests where reasonable. [More considerations for the public](http://wiki.creativecommons.org/Considerations_for_licensors_and_licensees#Considerations_for_licensees).
## Creative Commons Attribution-NonCommercial 4.0 International Public License
By exercising the Licensed Rights (defined below), You accept and agree to be bound by the terms and conditions of this Creative Commons Attribution-NonCommercial 4.0 International Public License ("Public License"). To the extent this Public License may be interpreted as a contract, You are granted the Licensed Rights in consideration of Your acceptance of these terms and conditions, and the Licensor grants You such rights in consideration of benefits the Licensor receives from making the Licensed Material available under these terms and conditions.
### Section 1 Definitions.
a. __Adapted Material__ means material subject to Copyright and Similar Rights that is derived from or based upon the Licensed Material and in which the Licensed Material is translated, altered, arranged, transformed, or otherwise modified in a manner requiring permission under the Copyright and Similar Rights held by the Licensor. For purposes of this Public License, where the Licensed Material is a musical work, performance, or sound recording, Adapted Material is always produced where the Licensed Material is synched in timed relation with a moving image.
b. __Adapter's License__ means the license You apply to Your Copyright and Similar Rights in Your contributions to Adapted Material in accordance with the terms and conditions of this Public License.
c. __Copyright and Similar Rights__ means copyright and/or similar rights closely related to copyright including, without limitation, performance, broadcast, sound recording, and Sui Generis Database Rights, without regard to how the rights are labeled or categorized. For purposes of this Public License, the rights specified in Section 2(b)(1)-(2) are not Copyright and Similar Rights.
d. __Effective Technological Measures__ means those measures that, in the absence of proper authority, may not be circumvented under laws fulfilling obligations under Article 11 of the WIPO Copyright Treaty adopted on December 20, 1996, and/or similar international agreements.
e. __Exceptions and Limitations__ means fair use, fair dealing, and/or any other exception or limitation to Copyright and Similar Rights that applies to Your use of the Licensed Material.
f. __Licensed Material__ means the artistic or literary work, database, or other material to which the Licensor applied this Public License.
g. __Licensed Rights__ means the rights granted to You subject to the terms and conditions of this Public License, which are limited to all Copyright and Similar Rights that apply to Your use of the Licensed Material and that the Licensor has authority to license.
h. __Licensor__ means the individual(s) or entity(ies) granting rights under this Public License.
i. __NonCommercial__ means not primarily intended for or directed towards commercial advantage or monetary compensation. For purposes of this Public License, the exchange of the Licensed Material for other material subject to Copyright and Similar Rights by digital file-sharing or similar means is NonCommercial provided there is no payment of monetary compensation in connection with the exchange.
j. __Share__ means to provide material to the public by any means or process that requires permission under the Licensed Rights, such as reproduction, public display, public performance, distribution, dissemination, communication, or importation, and to make material available to the public including in ways that members of the public may access the material from a place and at a time individually chosen by them.
k. __Sui Generis Database Rights__ means rights other than copyright resulting from Directive 96/9/EC of the European Parliament and of the Council of 11 March 1996 on the legal protection of databases, as amended and/or succeeded, as well as other essentially equivalent rights anywhere in the world.
l. __You__ means the individual or entity exercising the Licensed Rights under this Public License. Your has a corresponding meaning.
### Section 2 Scope.
a. ___License grant.___
1. Subject to the terms and conditions of this Public License, the Licensor hereby grants You a worldwide, royalty-free, non-sublicensable, non-exclusive, irrevocable license to exercise the Licensed Rights in the Licensed Material to:
A. reproduce and Share the Licensed Material, in whole or in part, for NonCommercial purposes only; and
B. produce, reproduce, and Share Adapted Material for NonCommercial purposes only.
2. __Exceptions and Limitations.__ For the avoidance of doubt, where Exceptions and Limitations apply to Your use, this Public License does not apply, and You do not need to comply with its terms and conditions.
3. __Term.__ The term of this Public License is specified in Section 6(a).
4. __Media and formats; technical modifications allowed.__ The Licensor authorizes You to exercise the Licensed Rights in all media and formats whether now known or hereafter created, and to make technical modifications necessary to do so. The Licensor waives and/or agrees not to assert any right or authority to forbid You from making technical modifications necessary to exercise the Licensed Rights, including technical modifications necessary to circumvent Effective Technological Measures. For purposes of this Public License, simply making modifications authorized by this Section 2(a)(4) never produces Adapted Material.
5. __Downstream recipients.__
A. __Offer from the Licensor Licensed Material.__ Every recipient of the Licensed Material automatically receives an offer from the Licensor to exercise the Licensed Rights under the terms and conditions of this Public License.
B. __No downstream restrictions.__ You may not offer or impose any additional or different terms or conditions on, or apply any Effective Technological Measures to, the Licensed Material if doing so restricts exercise of the Licensed Rights by any recipient of the Licensed Material.
6. __No endorsement.__ Nothing in this Public License constitutes or may be construed as permission to assert or imply that You are, or that Your use of the Licensed Material is, connected with, or sponsored, endorsed, or granted official status by, the Licensor or others designated to receive attribution as provided in Section 3(a)(1)(A)(i).
b. ___Other rights.___
1. Moral rights, such as the right of integrity, are not licensed under this Public License, nor are publicity, privacy, and/or other similar personality rights; however, to the extent possible, the Licensor waives and/or agrees not to assert any such rights held by the Licensor to the limited extent necessary to allow You to exercise the Licensed Rights, but not otherwise.
2. Patent and trademark rights are not licensed under this Public License.
3. To the extent possible, the Licensor waives any right to collect royalties from You for the exercise of the Licensed Rights, whether directly or through a collecting society under any voluntary or waivable statutory or compulsory licensing scheme. In all other cases the Licensor expressly reserves any right to collect such royalties, including when the Licensed Material is used other than for NonCommercial purposes.
### Section 3 License Conditions.
Your exercise of the Licensed Rights is expressly made subject to the following conditions.
a. ___Attribution.___
1. If You Share the Licensed Material (including in modified form), You must:
A. retain the following if it is supplied by the Licensor with the Licensed Material:
i. identification of the creator(s) of the Licensed Material and any others designated to receive attribution, in any reasonable manner requested by the Licensor (including by pseudonym if designated);
ii. a copyright notice;
iii. a notice that refers to this Public License;
iv. a notice that refers to the disclaimer of warranties;
v. a URI or hyperlink to the Licensed Material to the extent reasonably practicable;
B. indicate if You modified the Licensed Material and retain an indication of any previous modifications; and
C. indicate the Licensed Material is licensed under this Public License, and include the text of, or the URI or hyperlink to, this Public License.
2. You may satisfy the conditions in Section 3(a)(1) in any reasonable manner based on the medium, means, and context in which You Share the Licensed Material. For example, it may be reasonable to satisfy the conditions by providing a URI or hyperlink to a resource that includes the required information.
3. If requested by the Licensor, You must remove any of the information required by Section 3(a)(1)(A) to the extent reasonably practicable.
4. If You Share Adapted Material You produce, the Adapter's License You apply must not prevent recipients of the Adapted Material from complying with this Public License.
### Section 4 Sui Generis Database Rights.
Where the Licensed Rights include Sui Generis Database Rights that apply to Your use of the Licensed Material:
a. for the avoidance of doubt, Section 2(a)(1) grants You the right to extract, reuse, reproduce, and Share all or a substantial portion of the contents of the database for NonCommercial purposes only;
b. if You include all or a substantial portion of the database contents in a database in which You have Sui Generis Database Rights, then the database in which You have Sui Generis Database Rights (but not its individual contents) is Adapted Material; and
c. You must comply with the conditions in Section 3(a) if You Share all or a substantial portion of the contents of the database.
For the avoidance of doubt, this Section 4 supplements and does not replace Your obligations under this Public License where the Licensed Rights include other Copyright and Similar Rights.
### Section 5 Disclaimer of Warranties and Limitation of Liability.
a. __Unless otherwise separately undertaken by the Licensor, to the extent possible, the Licensor offers the Licensed Material as-is and as-available, and makes no representations or warranties of any kind concerning the Licensed Material, whether express, implied, statutory, or other. This includes, without limitation, warranties of title, merchantability, fitness for a particular purpose, non-infringement, absence of latent or other defects, accuracy, or the presence or absence of errors, whether or not known or discoverable. Where disclaimers of warranties are not allowed in full or in part, this disclaimer may not apply to You.__
b. __To the extent possible, in no event will the Licensor be liable to You on any legal theory (including, without limitation, negligence) or otherwise for any direct, special, indirect, incidental, consequential, punitive, exemplary, or other losses, costs, expenses, or damages arising out of this Public License or use of the Licensed Material, even if the Licensor has been advised of the possibility of such losses, costs, expenses, or damages. Where a limitation of liability is not allowed in full or in part, this limitation may not apply to You.__
c. The disclaimer of warranties and limitation of liability provided above shall be interpreted in a manner that, to the extent possible, most closely approximates an absolute disclaimer and waiver of all liability.
### Section 6 Term and Termination.
a. This Public License applies for the term of the Copyright and Similar Rights licensed here. However, if You fail to comply with this Public License, then Your rights under this Public License terminate automatically.
b. Where Your right to use the Licensed Material has terminated under Section 6(a), it reinstates:
1. automatically as of the date the violation is cured, provided it is cured within 30 days of Your discovery of the violation; or
2. upon express reinstatement by the Licensor.
For the avoidance of doubt, this Section 6(b) does not affect any right the Licensor may have to seek remedies for Your violations of this Public License.
c. For the avoidance of doubt, the Licensor may also offer the Licensed Material under separate terms or conditions or stop distributing the Licensed Material at any time; however, doing so will not terminate this Public License.
d. Sections 1, 5, 6, 7, and 8 survive termination of this Public License.
### Section 7 Other Terms and Conditions.
a. The Licensor shall not be bound by any additional or different terms or conditions communicated by You unless expressly agreed.
b. Any arrangements, understandings, or agreements regarding the Licensed Material not stated herein are separate from and independent of the terms and conditions of this Public License.
### Section 8 Interpretation.
a. For the avoidance of doubt, this Public License does not, and shall not be interpreted to, reduce, limit, restrict, or impose conditions on any use of the Licensed Material that could lawfully be made without permission under this Public License.
b. To the extent possible, if any provision of this Public License is deemed unenforceable, it shall be automatically reformed to the minimum extent necessary to make it enforceable. If the provision cannot be reformed, it shall be severed from this Public License without affecting the enforceability of the remaining terms and conditions.
c. No term or condition of this Public License will be waived and no failure to comply consented to unless expressly agreed to by the Licensor.
d. Nothing in this Public License constitutes or may be interpreted as a limitation upon, or waiver of, any privileges and immunities that apply to the Licensor or You, including from the legal processes of any jurisdiction or authority.
> Creative Commons is not a party to its public licenses. Notwithstanding, Creative Commons may elect to apply one of its public licenses to material it publishes and in those instances will be considered the “Licensor.” Except for the limited purpose of indicating that material is shared under a Creative Commons public license or as otherwise permitted by the Creative Commons policies published at [creativecommons.org/policies](http://creativecommons.org/policies), Creative Commons does not authorize the use of the trademark “Creative Commons” or any other trademark or logo of Creative Commons without its prior written consent including, without limitation, in connection with any unauthorized modifications to any of its public licenses or any other arrangements, understandings, or agreements concerning use of licensed material. For the avoidance of doubt, this paragraph does not form part of the public licenses.
>
> Creative Commons may be contacted at creativecommons.org

298
README.md
View File

@ -1,18 +1,294 @@
# 开始拉取代码
# HUSKY: Humanoid Skateboarding System via Physics-Aware Whole-Body Control
基于 [HUSKY](https://arxiv.org/abs/2602.03205) 思路的人形滑板全身控制实验代码mjlab 训练、`rsl_rl` 与 MuJoCo 评测脚本。本仓库包含个人开发与 **Docker** 封装。
**目录:** [`src/mjlab_husky`](src/mjlab_husky) · [`rsl_rl/`](rsl_rl/) · [`dataset/`](dataset/) · [`test_scene/`](test_scene/) · [`ckpts/`](ckpts/)
---
```
git clone http://git-test.databall.tech:3000/autocc008/robot_skating_03.git
cd robot_skating_03
## 本地安装Ubuntu 22.04,推荐 `uv`
# add your files to repo
git push
```bash
curl -LsSf https://astral.sh/uv/install.sh | sh
git clone https://github.com/<你的用户名>/humanoid_skateboarding.git
cd humanoid_skateboarding
uv sync && uv pip install -e .
```
**可选LeRobot v3 导出 / 边播边录** 需要额外安装 `lerobot`(不在默认 `pyproject` 依赖里):
```bash
uv pip install lerobot
```
若安装后出现 `import torch` 报 NCCL 符号错误,可尝试:
```bash
uv pip uninstall nvidia-nccl-cu12
uv pip install --force-reinstall "nvidia-nccl-cu13>=2.29"
```
---
## 训练
```bash
cd humanoid_skateboarding
uv run train Mjlab-Skater-Flat-Unitree-G1 --env.scene.num-envs 4096
```
查看全部参数:
```bash
uv run train Mjlab-Skater-Flat-Unitree-G1 --help
```
---
## 回放 `play`
任务名固定为 **`Mjlab-Skater-Flat-Unitree-G1`**(注册在 `mjlab_husky.tasks`)。
### 通用
```bash
uv run play Mjlab-Skater-Flat-Unitree-G1 --checkpoint_file ckpts/test.pt
```
- **`--viewer auto`**(默认):有 `DISPLAY` / `WAYLAND_DISPLAY` 时用 **native**,否则 **rerun**
- **`--viewer native`**:本机有图形界面时使用 MuJoCo 原生窗口。
- **`--viewer rerun`**Rerun Web Viewer无头服务器常用
- **`--viewer viser`**Viser。
完整参数:
```bash
uv run play Mjlab-Skater-Flat-Unitree-G1 --help
```
### 无头 OpenGLMuJoCo 离屏相机)
在无 `DISPLAY` 的 Linux 上,`play` 会在导入 MuJoCo 前尽量设置 **`MUJOCO_GL=egl`**(见 `mjlab_husky/mujoco_gl.py`)。若仍失败可手动指定:
```bash
export MUJOCO_GL=egl # GPU 无头(常见)
# 或
export MUJOCO_GL=osmesa # 纯 CPU 软件光栅(更慢)
```
### Rerun端口与远程浏览器
Rerun 需要 **两个端口****Web**(默认 `8080`+ **gRPC**(默认多为 `9876`,以终端打印为准)。
**浏览器与 `play` 在同一台机器**:直接打开终端里 **`http://127.0.0.1:<web_port>/?url=...`** 完整链接(不要只打开无 `?url=` 的首页)。
**浏览器在自己电脑、`play` 在云主机**:必须在本机做 **SSH 双端口转发**(把 `user@host` 换成你的登录方式,端口与 `play` 一致):
```bash
ssh -N -L 8080:127.0.0.1:8080 -L 9876:127.0.0.1:9876 user@云主机IP
```
若使用 `~/.ssh/config` 里的 `Host` 别名(例如 `Seoul`
```bash
ssh -N -L 8080:127.0.0.1:8080 -L 9876:127.0.0.1:9876 Seoul
```
指定密钥时:
```bash
ssh -i ~/.ssh/你的_key -N -L 8080:127.0.0.1:8080 -L 9876:127.0.0.1:9876 ubuntu@云主机IP
```
**本机 8080/9876 已被占用**时,改用空闲本地端口,并同时改 `?url=` 里 gRPC 端口,例如:
```bash
ssh -N -L 18080:127.0.0.1:18080 -L 19876:127.0.0.1:19876 user@云主机IP
```
云主机上 `play` 需一致:
```bash
uv run play Mjlab-Skater-Flat-Unitree-G1 --checkpoint_file ckpts/test.pt \
--viewer rerun \
--rerun-web-port 18080 \
--rerun-grpc-port 19876
```
**不经 SSH、浏览器直连公网**:安全组放行 Web + gRPC 端口,并指定(示例):
```bash
uv run play ... --viewer rerun --rerun-connect-host <云主机公网IP>
```
### Rerun 常用性能参数(可选)
```bash
uv run play Mjlab-Skater-Flat-Unitree-G1 --checkpoint_file ckpts/test.pt \
--viewer rerun \
--rerun-viewer-width 640 --rerun-viewer-height 360 \
--rerun-camera-log-stride 4 --rerun-qpos-log-stride 8 \
--rerun-camera-max-side 480 \
--no-rerun-open-browser
```
说明:本项目 CLI 使用 **tyro**,布尔开关一般为 **`--xxx` / `--no-xxx`**(例如 `--lerobot-record`、`--no-rerun-open-browser`),不要写成 `--lerobot-record True`
---
## LeRobot v3 数据(`lerobot_data/`
LeRobot **v3****Parquet + `meta/`**(不是 HDF5。本仓库提供两种方式写入 **`observation.state`qposfloat32 向量)**。
### 1离线批量导出不跑 Rerun
需已 `uv pip install lerobot`
```bash
uv run python -m mjlab_husky.scripts.export_lerobot_qpos \
--task-id Mjlab-Skater-Flat-Unitree-G1 \
--checkpoint-file ckpts/test.pt \
--out-dir lerobot_data \
--dataset-name mjlab_husky_skater_qpos \
--episodes 1 \
--steps-per-episode 1000 \
--overwrite
```
### 2`play` + Rerun 同时边播边录
```bash
uv run play Mjlab-Skater-Flat-Unitree-G1 \
--checkpoint_file ckpts/test.pt \
--viewer rerun \
--rerun-web-port 18080 \
--rerun-grpc-port 19876 \
--lerobot-record \
--lerobot-out-dir lerobot_data \
--lerobot-dataset-name mjlab_husky_live \
--lerobot-overwrite
```
要点:
- **`--lerobot-overwrite`**:每次启动会 **删除** 同名数据集目录;要 **累积** 多次运行,请 **去掉** 该参数,或换 `--lerobot-dataset-name`
- 默认每录满 **`--lerobot-steps-per-episode`**(默认 1000帧会 `save_episode()` 一次;仿真里多次 `reset` **不会**自动切分,除非打开 **`--lerobot-save-on-env-reset`**。
- 退出 `play`(如 Ctrl+C时会 `finalize()`,避免 Parquet 不完整。
按仿真每次 `done -> reset` 存成一个 LeRobot episode
```bash
uv run play Mjlab-Skater-Flat-Unitree-G1 ... --lerobot-record --lerobot-save-on-env-reset
```
### 检查数据集是否可读(行数 / episode
```bash
uv run python -c "
from pathlib import Path
import json
info = json.loads(Path('lerobot_data/mjlab_husky_live/meta/info.json').read_text())
print('total_episodes', info.get('total_episodes'), 'total_frames', info.get('total_frames'))
"
```
---
## Docker推荐
基础环境Ubuntu 22.04、CUDA 13、`uv` 与项目依赖。镜像 **`MUJOCO_GL=egl`**,默认 **`CMD`** 为 **Rerun** 回放(`--no-rerun-open-browser`)。
**构建**
```bash
docker build -t husky-skate:latest .
```
**GPU 运行**(需 [NVIDIA Container Toolkit](https://docs.nvidia.com/datacenter/cloud-native/container-toolkit/install-guide.html)
```bash
docker run --gpus all -it --rm -p 8080:8080 -p 9876:9876 husky-skate:latest
```
在**宿主机浏览器**打开容器日志里打印的 **`http://127.0.0.1:8080/?url=...`**(若浏览器不在宿主机,需自行把对应端口转发到本机)。
**宿主机 8080/9876 已被占用**(例如已有其他容器映射):换主机端口 + 覆盖容器内 `play` 端口,例如:
```bash
docker run --gpus all -it --rm -p 18080:18080 -p 19876:19876 husky-skate:latest \
uv run play Mjlab-Skater-Flat-Unitree-G1 --checkpoint_file ckpts/test.pt \
--viewer rerun \
--rerun-web-port 18080 \
--rerun-grpc-port 19876 \
--no-rerun-open-browser
```
**仅 CPU**(较慢)
```bash
docker run -it --rm -p 8080:8080 -p 9876:9876 husky-skate:latest
```
**进入容器 Shell**
```bash
docker run --gpus all -it --rm --entrypoint /bin/bash husky-skate:latest
```
**容器内训练**
```bash
docker run --gpus all -it --rm husky-skate:latest \
uv run train Mjlab-Skater-Flat-Unitree-G1 --env.scene.num-envs 4096
```
**(可选)容器内录 LeRobot**:需先安装 `lerobot`,并把目录挂载出来,例如:
```bash
docker run --gpus all -it --rm \
-p 18080:18080 -p 19876:19876 \
-v "$(pwd)/lerobot_data:/app/lerobot_data" \
husky-skate:latest \
bash -lc 'uv pip install lerobot && uv run play Mjlab-Skater-Flat-Unitree-G1 \
--checkpoint_file ckpts/test.pt --viewer rerun \
--rerun-web-port 18080 --rerun-grpc-port 19876 --no-rerun-open-browser \
--lerobot-record --lerobot-out-dir lerobot_data --lerobot-dataset-name mjlab_docker_live \
--lerobot-overwrite'
```
---
## PyTorch / CUDA 提示
若日志出现 **driver too old如 12080**`torch.cuda.is_available()``False`,多为 **PyTorch cu13x 与当前驱动 API 不匹配**。可选:
- 安装与驱动匹配的 **cu12x** 轮子,例如:
`uv pip install torch torchvision --index-url https://download.pytorch.org/whl/cu124`
- 或升级宿主机 NVIDIA 驱动以匹配当前 PyTorch 所要求的 CUDA。
仿真侧 **Warp/MuJoCo** 仍可能显示 `cpu`,与 **`torch.cuda.is_available()``play` 选的 device** 一致。
---
## 轻量 MuJoCo 评测
```bash
bash test_scene/sim.sh your-onnx-path
```
| Viser | MuJoCo |
|-------|--------|
| ![](media/viser.gif) | ![](media/mjc.gif) |
---
## 论文引用(原论文)
```bibtex
@article{han2026husky,
title={HUSKY: Humanoid Skateboarding System via Physics-Aware Whole-Body Control},
author={Jinrui Han and Dewei Wang and Chenyun Zhang and Xinzhe Liu and Ping Luo and Chenjia Bai and Xuelong Li},
journal={arXiv preprint arXiv:2602.03205},
year={2026}
}
```
cd existing_repo
git remote add origin http://git-test.databall.tech:3000/autocc008/robot_skating_03.git
git branch -M main
git push -uf origin main
```

9
autorun.sh Executable file
View File

@ -0,0 +1,9 @@
#!/bin/bash
export PATH="$PATH:/opt/miniconda/bin"
# 加载 conda 的 shell 函数(不需要 conda init 和 source ~/.bashrc
eval "$(conda shell.bash hook)"
conda activate pms-py312
export PYTHONPATH="/opt/vendor/python:$CONDA_PREFIX/lib/python3.12/site-packages:/workspace/src:/workspace/rsl_rl"
python -m mjlab_husky.scripts.play Mjlab-Skater-Flat-Unitree-G1 --checkpoint_file ckpts/test.pt --viewer rerun_viser --rerun-web-port 8081 --rerun-grpc-port 9876 --viser-port 8082

BIN
ckpts/test.onnx Normal file

Binary file not shown.

BIN
ckpts/test.pt Normal file

Binary file not shown.

BIN
dataset/.DS_Store vendored Normal file

Binary file not shown.

Binary file not shown.

Binary file not shown.

Binary file not shown.

Binary file not shown.

BIN
lerobot_data/.DS_Store vendored Normal file

Binary file not shown.

BIN
lerobot_data/mjlab_husky_live/.DS_Store vendored Normal file

Binary file not shown.

View File

@ -0,0 +1,60 @@
{
"codebase_version": "v3.0",
"robot_type": "mjlab_husky",
"total_episodes": 85,
"total_frames": 85000,
"total_tasks": 1,
"chunks_size": 1000,
"data_files_size_in_mb": 100,
"video_files_size_in_mb": 200,
"fps": 50,
"splits": {
"train": "0:85"
},
"data_path": "data/chunk-{chunk_index:03d}/file-{file_index:03d}.parquet",
"video_path": null,
"features": {
"observation.state": {
"dtype": "float32",
"shape": [
44
],
"names": null
},
"timestamp": {
"dtype": "float32",
"shape": [
1
],
"names": null
},
"frame_index": {
"dtype": "int64",
"shape": [
1
],
"names": null
},
"episode_index": {
"dtype": "int64",
"shape": [
1
],
"names": null
},
"index": {
"dtype": "int64",
"shape": [
1
],
"names": null
},
"task_index": {
"dtype": "int64",
"shape": [
1
],
"names": null
}
}
}

View File

@ -0,0 +1,581 @@
{
"frame_index": {
"min": [
0
],
"max": [
999
],
"mean": [
499.5
],
"std": [
288.6749902572095
],
"count": [
85000
],
"q01": [
9.190799999901826
],
"q10": [
99.1007999999198
],
"q50": [
499.1003999999992
],
"q90": [
899.1000000000806
],
"q99": [
989.0100000000987
]
},
"episode_index": {
"min": [
0
],
"max": [
84
],
"mean": [
42.0
],
"std": [
24.53568829277059
],
"count": [
85000
],
"q01": [
42.0
],
"q10": [
42.00000000000001
],
"q50": [
42.000000000000014
],
"q90": [
42.00000000000003
],
"q99": [
42.00000000000003
]
},
"observation.state": {
"min": [
-0.07113656401634216,
-70.33460998535156,
0.0829688161611557,
0.10346301645040512,
-0.5429396033287048,
-0.7814562320709229,
-0.9942379593849182,
-1.3800686597824097,
-0.36666902899742126,
-0.5998758673667908,
-0.12973685562610626,
-1.0291401147842407,
-0.30399635434150696,
-1.6473090648651123,
-0.854110062122345,
-0.3959100544452667,
0.12965033948421478,
-0.9802296161651611,
-0.2205560803413391,
-0.6682753562927246,
-0.6041809916496277,
-0.5346789956092834,
-0.9392914175987244,
0.12718690931797028,
-0.630613386631012,
-0.13077163696289062,
-1.645430564880371,
-1.463309407234192,
-0.6851853728294373,
-0.019867030903697014,
0.0,
-70.32968139648438,
0.0039673116989433765,
-0.2486332356929779,
-0.7575073838233948,
-0.09620650112628937,
-0.9997463822364807,
-0.2166292518377304,
-0.12168516218662262,
0.0,
0.0,
-0.12534703314304352,
0.0,
0.0
],
"max": [
75.91880798339844,
6.302745819091797,
0.8795564770698547,
1.0,
0.27965977787971497,
0.7949097156524658,
0.7174645066261292,
0.4639802575111389,
0.6382898688316345,
0.6904199123382568,
2.7510406970977783,
0.44586628675460815,
0.22848469018936157,
0.20644104480743408,
0.3050974905490875,
0.8214103579521179,
2.1551849842071533,
0.5406686067581177,
0.30794599652290344,
0.9843693375587463,
0.5602899789810181,
0.6002522706985474,
1.3705785274505615,
2.2902050018310547,
0.58452969789505,
2.0403666496276855,
0.9998741149902344,
-0.04809024930000305,
0.9831812977790833,
1.8275797367095947,
75.97547149658203,
6.415305137634277,
0.17043574154376984,
1.0,
0.00961267203092575,
0.7573930025100708,
0.24439620971679688,
0.13883520662784576,
0.14782053232192993,
3994.871337890625,
3002.75146484375,
0.15178228914737701,
4508.38818359375,
5436.13671875
],
"mean": [
23.25296789302546,
-14.87063912240898,
0.7819512668777915,
0.865570327464272,
0.01820748194586485,
0.06374942751341114,
-0.1310197852771072,
-0.3828951469677337,
0.1711113832452718,
0.03854638235369588,
0.6976036687107648,
-0.3548905068679767,
-0.020054900660248034,
-0.779085868334069,
-0.06996407535146264,
0.02570628469727714,
0.9095189350492814,
-0.3067187977297342,
0.01838621137242573,
-0.1969201212100041,
-0.051224787414813525,
0.1772239558035782,
-0.17888976724489644,
0.6842897481778086,
-0.23786781745798447,
0.5855691568925977,
0.04866229006472756,
-0.5830464264925787,
0.07088629282490515,
0.7353050638647644,
23.28549128586755,
-14.991129989364966,
0.09786677574191023,
0.8207214686085268,
-0.014800452884902131,
0.012796162317574593,
-0.3960311810705153,
-0.004298211561913111,
0.004420292135738932,
1144.7475955318005,
1128.1047784496752,
0.005849475387538619,
1147.7585460326252,
1155.0732623577117
],
"std": [
15.179133948547573,
15.20011979901754,
0.1454885092628437,
0.16651873814941476,
0.10197849704538488,
0.1578681255507523,
0.40764685053480554,
0.2708406675734598,
0.1379569269283228,
0.13227254048456757,
0.2890878111557964,
0.19052926227978695,
0.08294386306895925,
0.3608905403828366,
0.14723826734479117,
0.13379134266282225,
0.38752975069890283,
0.22396932802488995,
0.0868930687538448,
0.14074262142096175,
0.1886161113878726,
0.11362784145745387,
0.2621380919964973,
0.2556253485331334,
0.1126427040214631,
0.3526049895335999,
0.3032407931542005,
0.11991412831623412,
0.3064957559934792,
0.3248620681656251,
15.189723680756549,
15.223858376740733,
0.012404234495425687,
0.22559996637236646,
0.1049420524852779,
0.09063896603938991,
0.31475311295102665,
0.01805742472878103,
0.018927681317975016,
763.1105549460016,
734.9102254844961,
0.025587495248672498,
777.1573189417905,
805.3248523580429
],
"count": [
85000
],
"q01": [
15.518811176371708,
-21.800787278659204,
0.6495954289599476,
0.6718457682167782,
-0.08911537257755264,
-0.09254132599076804,
-0.5442103528143726,
-0.8142312556505203,
-0.2746849088984377,
-0.37409252572026064,
0.30180432632857673,
-0.761493222503101,
-0.18507926103823322,
-1.3567008463775416,
-0.5036703249110895,
-0.28734303055440685,
0.28153787482213344,
-0.6747250879512114,
-0.1396642809843316,
-0.4284897106535294,
-0.42121099120915334,
-0.051688335814020214,
-0.7924406105542886,
0.31562799811363224,
-0.37630962519084704,
0.07698490332592935,
-0.3561562091112136,
-0.8444881160119001,
-0.3497024485512692,
0.22996248002437986,
15.576270143400837,
-21.96829468889173,
0.09735096952713591,
0.7123482675212154,
-0.018511691378479748,
0.006816049990297945,
-0.5658663816106864,
-0.04238463985989345,
-0.03602184447898145,
762.0721309375697,
753.0855775318535,
-0.0461389991139774,
763.0219338896693,
766.1330671246137
],
"q10": [
17.13613685887827,
-20.258861770384623,
0.7103266863599536,
0.7094701728107882,
-0.060989131895872205,
-0.029106133983091073,
-0.504221053586579,
-0.6624676102662788,
0.027281049858121312,
-0.030302501289641702,
0.4506148300529305,
-0.6142739851392953,
-0.11152195437164868,
-1.1864144497536786,
-0.22565803413881974,
-0.13463749580289525,
0.42177760002733145,
-0.537168428028581,
-0.0704819829193603,
-0.3211724682765849,
-0.31891577460017867,
0.05491297600119163,
-0.5311642064492241,
0.5432212542085089,
-0.3212178046207084,
0.30091702385342545,
-0.13884038798072756,
-0.7119056475773254,
-0.21982713513540358,
0.3988324834181499,
17.1882684046792,
-20.450890284314816,
0.09739910307028207,
0.7412608709828881,
-0.01803065373541732,
0.007440648461717085,
-0.5365755949991055,
-0.025924483710719686,
-0.013262635344442894,
836.1614272927652,
826.3365670006482,
-0.018764696247214368,
837.2647427414483,
841.0888403746054
],
"q50": [
23.762409283764843,
-14.563460296239693,
0.7827337545593831,
0.8979291666198453,
0.011350150855273758,
0.06532513037080345,
-0.12206655524496907,
-0.4096627717294724,
0.18403609829247833,
0.033636638799904965,
0.6160347013881324,
-0.33975082198066753,
-0.018045310162279112,
-0.8163800627140383,
-0.09180862832401136,
0.018329558631927374,
0.9480252705071734,
-0.3165479713502877,
-0.004242551547214943,
-0.1876591242178994,
-0.002054170142044198,
0.1841658829086471,
-0.1383427245692955,
0.7028915360612112,
-0.2458919337199281,
0.5460589266844458,
0.0677177166342812,
-0.5704874180267574,
0.17712665092709504,
0.6865986548212377,
23.765883293045736,
-14.70968608737983,
0.0974536156380619,
0.8232050776237005,
-0.0177864022127682,
0.015337733944697386,
-0.3919809371150244,
-0.0017066563739786982,
0.0019141628737798318,
1146.1196838017734,
1129.962561927736,
0.0012937257823444764,
1148.892501481066,
1155.902512883781
],
"q90": [
28.635734952869846,
-10.041157697359191,
0.8461450801716657,
0.9661417386890236,
0.10726293688968701,
0.15083795809839803,
0.2146376604382313,
-0.06311925098884336,
0.3076535279334336,
0.1550071704183138,
1.0939077803548265,
-0.15099184687477152,
0.06143217126236244,
-0.373247745245278,
0.09912829460511562,
0.17111820184132637,
1.333676097759791,
-0.060704099186905884,
0.12421973123064842,
-0.086594612443564,
0.11381211995410652,
0.2819234085433624,
0.0842839852017181,
0.8042130172922879,
-0.1492722455033644,
0.907436110787069,
0.19406254114557048,
-0.473296269342011,
0.31749770715945097,
1.1360847358593222,
28.65005411403944,
-10.114658482182755,
0.09869820489690093,
0.8902162918360124,
-0.008530837538709927,
0.015510477181937963,
-0.2611029862148012,
0.012261228386334136,
0.026974081897148575,
1453.8357671251485,
1429.161401413135,
0.03886226614772836,
1459.450008259809,
1470.675159581278
],
"q99": [
29.426830072364876,
-8.898120419735253,
0.8527026614373789,
0.9813330734009837,
0.14279720059212514,
0.19780941931640406,
0.31998690711516964,
0.10990977427538703,
0.34699264243537326,
0.22542898462099187,
1.2790248403157376,
-0.0044630865840350874,
0.1601653475971783,
-0.23684657999113493,
0.14029283857395664,
0.2544758809367996,
1.6298833114916786,
0.0403672861702302,
0.16848280955763428,
-0.007583600633284624,
0.20617429927775735,
0.3594991666429182,
0.20428104321746265,
0.9497321564037832,
0.044834615903742187,
1.179834731440334,
0.3116965556845947,
-0.38883973184753867,
0.40519437930163216,
1.2721450699430548,
29.490910510154002,
-8.960652529571895,
0.09934294331161415,
0.9097403214357518,
-0.008158426082638257,
0.016114267153009183,
-0.21360506679982902,
0.03291656909620061,
0.043081196781028716,
1524.6454015946413,
1498.215975788589,
0.059730523218585053,
1530.8834143921215,
1543.3855692444952
]
},
"task_index": {
"min": [
0
],
"max": [
0
],
"mean": [
0.0
],
"std": [
0.0
],
"count": [
85000
],
"q01": [
3.9999999999994176e-16
],
"q10": [
3.999999999999417e-15
],
"q50": [
1.999999999999706e-14
],
"q90": [
3.599999999999482e-14
],
"q99": [
3.959999999999426e-14
]
},
"index": {
"min": [
0
],
"max": [
84999
],
"mean": [
42499.5
],
"std": [
24537.38643886101
],
"count": [
85000
],
"q01": [
42009.19079999987
],
"q10": [
42099.100799999884
],
"q50": [
42499.100400000025
],
"q90": [
42899.10000000006
],
"q99": [
42989.01000000005
]
},
"timestamp": {
"min": [
0.0
],
"max": [
19.98
],
"mean": [
9.99
],
"std": [
5.773499805144186
],
"count": [
85000
],
"q01": [
0.1838159999018404
],
"q10": [
1.9820159999198408
],
"q50": [
9.982007999999915
],
"q90": [
17.982000000080028
],
"q99": [
19.780200000097988
]
}
}

Binary file not shown.

BIN
media/method.jpg Normal file

Binary file not shown.

After

Width:  |  Height:  |  Size: 3.4 MiB

BIN
media/mjc.gif Normal file

Binary file not shown.

After

Width:  |  Height:  |  Size: 10 MiB

BIN
media/teaser.jpg Normal file

Binary file not shown.

After

Width:  |  Height:  |  Size: 156 KiB

BIN
media/viser.gif Normal file

Binary file not shown.

After

Width:  |  Height:  |  Size: 5.3 MiB

42
pyproject.toml Normal file
View File

@ -0,0 +1,42 @@
[build-system]
requires = ["uv_build>=0.8.19,<0.9.0"]
build-backend = "uv_build"
[project]
name = "mjlab_husky"
version = "0.1.0"
license = "Apache-2.0"
description = "Humanoid Skateboarding System via Physics-Aware Whole-Body Control"
readme = "README.md"
requires-python = ">=3.12,<3.14"
dependencies = [
"mjlab",
"onnxruntime",
"pybullet",
"rsl-rl-lib",
"pynput",
"rerun-sdk"
]
[project.entry-points."mjlab.tasks"]
mjlab_husky = "mjlab_husky.tasks"
[project.scripts]
train = "mjlab_husky.scripts.train:main"
play = "mjlab_husky.scripts.play:main"
rerun-lerobot = "mjlab_husky.scripts.rerun_lerobot_dataset:main"
play-lerobot-rerun = "mjlab_husky.scripts.play_lerobot_rerun:main"
[tool.ruff]
src = ["src"] # Helpful for recognizing first-party imports.
indent-width = 4
[tool.uv]
# py.mujoco.org 上的 3.7.0.dev* 预发布包会被撤下wheel 404强制使用 PyPI 稳定版。
override-dependencies = ["mujoco>=3.8.0,<3.9"]
# `mjlab` 声明的 mujoco-warp git rev1dc288c依赖 mjENBL_MULTICCDPyPI 的 mujoco 3.8.x 尚无该枚举。
# `uv.lock` 将 mujoco-warp 固定为上游标签 v3.8.0(与 mujoco 3.8.0 对齐);升级 mjlab / 跑 `uv lock` 后若冲突请复查该包。
[tool.uv.sources]
mjlab = { git = "https://github.com/mujocolab/mjlab.git", rev = "13212ad" }
rsl-rl-lib = { path = "rsl_rl" }

BIN
rsl_rl/.DS_Store vendored Normal file

Binary file not shown.

6
rsl_rl/__init__.py Normal file
View File

@ -0,0 +1,6 @@
# Copyright (c) 2021-2025, ETH Zurich and NVIDIA CORPORATION
# All rights reserved.
#
# SPDX-License-Identifier: BSD-3-Clause
"""Main module for the rsl_rl package."""

Binary file not shown.

View File

@ -0,0 +1,11 @@
# Copyright (c) 2021-2025, ETH Zurich and NVIDIA CORPORATION
# All rights reserved.
#
# SPDX-License-Identifier: BSD-3-Clause
"""Implementation of different RL agents."""
from .distillation import Distillation
from .ppo import PPO
from .amp_ppo import AMP_PPO
__all__ = ["PPO", "Distillation", "AMP_PPO"]

Binary file not shown.

Binary file not shown.

Binary file not shown.

View File

@ -0,0 +1,571 @@
# Copyright (c) 2021-2025, ETH Zurich and NVIDIA CORPORATION
# All rights reserved.
#
# SPDX-License-Identifier: BSD-3-Clause
from __future__ import annotations
from torch._tensor import Tensor
from torch._tensor import Tensor
from typing import Any
from copy import deepcopy
import torch
import torch.nn as nn
import torch.optim as optim
from itertools import chain
from rsl_rl.modules import ActorCritic
from rsl_rl.modules.rnd import RandomNetworkDistillation
from rsl_rl.storage import RolloutStorage, ReplayBufferMulti
from rsl_rl.utils import string_to_callable
class AMP_PPO:
"""Proximal Policy Optimization algorithm (https://arxiv.org/abs/1707.06347)."""
policy: ActorCritic
"""The actor critic module."""
def __init__(
self,
policy,
discriminator,
amp_data,
amp_normalizer,
amp_num_frames=1,
amp_replay_buffer_size=100000,
num_learning_epochs=5,
num_mini_batches=4,
clip_param=0.2,
gamma=0.99,
lam=0.95,
value_loss_coef=1.0,
entropy_coef=0.01,
learning_rate=0.001,
max_grad_norm=1.0,
use_clipped_value_loss=True,
schedule="adaptive",
desired_kl=0.01,
device="cpu",
normalize_advantage_per_mini_batch=False,
# RND parameters
rnd_cfg: dict | None = None,
# Symmetry parameters
symmetry_cfg: dict | None = None,
# Distributed training parameters
multi_gpu_cfg: dict | None = None,
):
# device-related parameters
self.device = device
self.is_multi_gpu = multi_gpu_cfg is not None
# Multi-GPU parameters
if multi_gpu_cfg is not None:
self.gpu_global_rank = multi_gpu_cfg["global_rank"]
self.gpu_world_size = multi_gpu_cfg["world_size"]
else:
self.gpu_global_rank = 0
self.gpu_world_size = 1
# RND components
if rnd_cfg is not None:
# Extract parameters used in ppo
rnd_lr = rnd_cfg.pop("learning_rate", 1e-3)
# Create RND module
self.rnd = RandomNetworkDistillation(device=self.device, **rnd_cfg)
# Create RND optimizer
params = self.rnd.predictor.parameters()
self.rnd_optimizer = optim.Adam(params, lr=rnd_lr)
else:
self.rnd = None
self.rnd_optimizer = None
# Symmetry components
if symmetry_cfg is not None:
# Check if symmetry is enabled
use_symmetry = symmetry_cfg["use_data_augmentation"] or symmetry_cfg["use_mirror_loss"]
# Print that we are not using symmetry
if not use_symmetry:
print("Symmetry not used for learning. We will use it for logging instead.")
# If function is a string then resolve it to a function
if isinstance(symmetry_cfg["data_augmentation_func"], str):
symmetry_cfg["data_augmentation_func"] = string_to_callable(symmetry_cfg["data_augmentation_func"])
# Check valid configuration
if symmetry_cfg["use_data_augmentation"] and not callable(symmetry_cfg["data_augmentation_func"]):
raise ValueError(
"Data augmentation enabled but the function is not callable:"
f" {symmetry_cfg['data_augmentation_func']}"
)
# Store symmetry configuration
self.symmetry = symmetry_cfg
else:
self.symmetry = None
## AMP components
self.discriminator = discriminator
self.discriminator.to(self.device)
self.amp_storage = ReplayBufferMulti(discriminator.state_dim, amp_replay_buffer_size, amp_num_frames, device)
self.amp_data = amp_data
self.amp_normalizer = amp_normalizer
# PPO components
self.policy = policy
self.policy.to(self.device)
# Create rollout storage
self.storage: RolloutStorage = None # type: ignore
self.transition = RolloutStorage.Transition()
self.amp_transition = RolloutStorage.Transition()
params = [
{'params': self.policy.parameters(), 'name': 'policy'},
]
params.append({
'params': self.discriminator.trunk.parameters(),
'weight_decay': 10e-4,
'name': f'amp_trunk'
})
params.append({
'params': self.discriminator.amp_linear.parameters(),
'weight_decay': 10e-2,
'name': f'amp_head'
})
# Create optimizer
self.optimizer = optim.Adam(params, lr=learning_rate)
# PPO parameters
self.clip_param = clip_param
self.num_learning_epochs = num_learning_epochs
self.num_mini_batches = num_mini_batches
self.value_loss_coef = value_loss_coef
self.entropy_coef = entropy_coef
self.gamma = gamma
self.lam = lam
self.max_grad_norm = max_grad_norm
self.use_clipped_value_loss = use_clipped_value_loss
self.desired_kl = desired_kl
self.schedule = schedule
self.learning_rate = learning_rate
self.normalize_advantage_per_mini_batch = normalize_advantage_per_mini_batch
def init_storage(self, training_type, num_envs, num_transitions_per_env, obs, actions_shape):
# create rollout storage
self.storage = RolloutStorage(
training_type,
num_envs,
num_transitions_per_env,
obs,
actions_shape,
self.device,
)
def act(self, obs, amp_obs):
if self.policy.is_recurrent:
self.transition.hidden_states = self.policy.get_hidden_states()
# compute the actions and values
self.transition.actions = self.policy.act(obs).detach()
self.transition.values = self.policy.evaluate(obs).detach()
self.transition.actions_log_prob = self.policy.get_actions_log_prob(self.transition.actions).detach()
self.transition.action_mean = self.policy.action_mean.detach()
self.transition.action_sigma = self.policy.action_std.detach()
# need to record obs before env.step()
self.transition.observations = obs
self.amp_transition.observations = amp_obs
return self.transition.actions
def process_env_step(self, obs, rewards, dones, extras,amp_obs, amp_obs_frames=None):
# update the normalizers
self.policy.update_normalization(obs)
if self.rnd:
self.rnd.update_normalization(obs)
# Record the rewards and dones
# Note: we clone here because later on we bootstrap the rewards based on timeouts
self.transition.rewards = rewards.clone()
self.transition.dones = dones
# Compute the intrinsic rewards and add to extrinsic rewards
if self.rnd:
# Compute the intrinsic rewards
self.intrinsic_rewards = self.rnd.get_intrinsic_reward(obs)
# Add intrinsic rewards to extrinsic rewards
self.transition.rewards += self.intrinsic_rewards
# Bootstrapping on time outs
if "time_outs" in extras:
self.transition.rewards += self.gamma * torch.squeeze(
self.transition.values * extras["time_outs"].unsqueeze(1).to(self.device), 1
)
if amp_obs_frames is not None:
self.amp_storage.insert(amp_obs_frames)
else:
self.amp_storage.insert(self.amp_transition.observations, amp_obs)
# record the transition
self.storage.add_transitions(self.transition)
self.transition.clear()
self.amp_transition.clear()
self.policy.reset(dones)
def compute_returns(self, obs):
# compute value for the last step
last_values = self.policy.evaluate(obs).detach()
self.storage.compute_returns(
last_values, self.gamma, self.lam, normalize_advantage=not self.normalize_advantage_per_mini_batch
)
def update(self): # noqa: C901
mean_value_loss = 0
mean_surrogate_loss = 0
mean_entropy = 0
mean_amp_loss = 0
mean_grad_pen_loss = 0
mean_policy_pred = 0
mean_expert_pred = 0
# -- RND loss
if self.rnd:
mean_rnd_loss = 0
else:
mean_rnd_loss = None
# -- Symmetry loss
if self.symmetry:
mean_symmetry_loss = 0
else:
mean_symmetry_loss = None
# generator for mini batches
if self.policy.is_recurrent:
generator = self.storage.recurrent_mini_batch_generator(self.num_mini_batches, self.num_learning_epochs)
else:
generator = self.storage.mini_batch_generator(self.num_mini_batches, self.num_learning_epochs)
amp_policy_generator = self.amp_storage.feed_forward_generator(
self.num_learning_epochs * self.num_mini_batches,
self.storage.num_envs * self.storage.num_transitions_per_env // self.num_mini_batches,
)
amp_expert_generator = self.amp_data.feed_forward_generator_23dof_multi(
self.num_learning_epochs * self.num_mini_batches,
self.storage.num_envs * self.storage.num_transitions_per_env // self.num_mini_batches,
)
# iterate over batches
for sample, sample_amp_policy, sample_amp_expert in zip(generator, amp_policy_generator, amp_expert_generator):
(
obs_batch,
actions_batch,
target_values_batch,
advantages_batch,
returns_batch,
old_actions_log_prob_batch,
old_mu_batch,
old_sigma_batch,
hid_states_batch,
masks_batch,
) = sample
# number of augmentations per sample
# we start with 1 and increase it if we use symmetry augmentation
num_aug = 1
# original batch size
# we assume policy group is always there and needs augmentation
original_batch_size = obs_batch.batch_size[0]
# check if we should normalize advantages per mini batch
if self.normalize_advantage_per_mini_batch:
with torch.no_grad():
advantages_batch = (advantages_batch - advantages_batch.mean()) / (advantages_batch.std() + 1e-8)
# Perform symmetric augmentation
if self.symmetry and self.symmetry["use_data_augmentation"]:
# augmentation using symmetry
data_augmentation_func = self.symmetry["data_augmentation_func"]
# returned shape: [batch_size * num_aug, ...]
obs_batch, actions_batch = data_augmentation_func(
obs=obs_batch,
actions=actions_batch,
env=self.symmetry["_env"],
)
# compute number of augmentations per sample
# we assume policy group is always there and needs augmentation
num_aug = int(obs_batch.batch_size[0] / original_batch_size)
# repeat the rest of the batch
# -- actor
old_actions_log_prob_batch = old_actions_log_prob_batch.repeat(num_aug, 1)
# -- critic
target_values_batch = target_values_batch.repeat(num_aug, 1)
advantages_batch = advantages_batch.repeat(num_aug, 1)
returns_batch = returns_batch.repeat(num_aug, 1)
# Recompute actions log prob and entropy for current batch of transitions
# Note: we need to do this because we updated the policy with the new parameters
# -- actor
self.policy.act(obs_batch, masks=masks_batch, hidden_states=hid_states_batch[0])
actions_log_prob_batch = self.policy.get_actions_log_prob(actions_batch)
# -- critic
value_batch = self.policy.evaluate(obs_batch, masks=masks_batch, hidden_states=hid_states_batch[1])
# -- entropy
# we only keep the entropy of the first augmentation (the original one)
mu_batch = self.policy.action_mean[:original_batch_size]
sigma_batch = self.policy.action_std[:original_batch_size]
entropy_batch = self.policy.entropy[:original_batch_size]
# KL
if self.desired_kl is not None and self.schedule == "adaptive":
with torch.inference_mode():
kl = torch.sum(
torch.log(sigma_batch / old_sigma_batch + 1.0e-5)
+ (torch.square(old_sigma_batch) + torch.square(old_mu_batch - mu_batch))
/ (2.0 * torch.square(sigma_batch))
- 0.5,
axis=-1,
)
kl_mean = torch.mean(kl)
# Reduce the KL divergence across all GPUs
if self.is_multi_gpu:
torch.distributed.all_reduce(kl_mean, op=torch.distributed.ReduceOp.SUM)
kl_mean /= self.gpu_world_size
# Update the learning rate
# Perform this adaptation only on the main process
# TODO: Is this needed? If KL-divergence is the "same" across all GPUs,
# then the learning rate should be the same across all GPUs.
if self.gpu_global_rank == 0:
if kl_mean > self.desired_kl * 2.0:
self.learning_rate = max(1e-5, self.learning_rate / 1.5)
elif kl_mean < self.desired_kl / 2.0 and kl_mean > 0.0:
self.learning_rate = min(1e-2, self.learning_rate * 1.5)
# Update the learning rate for all GPUs
if self.is_multi_gpu:
lr_tensor = torch.tensor(self.learning_rate, device=self.device)
torch.distributed.broadcast(lr_tensor, src=0)
self.learning_rate = lr_tensor.item()
# Update the learning rate for all parameter groups
for param_group in self.optimizer.param_groups:
param_group["lr"] = self.learning_rate
# Surrogate loss
ratio = torch.exp(actions_log_prob_batch - torch.squeeze(old_actions_log_prob_batch))
surrogate = -torch.squeeze(advantages_batch) * ratio
surrogate_clipped = -torch.squeeze(advantages_batch) * torch.clamp(
ratio, 1.0 - self.clip_param, 1.0 + self.clip_param
)
surrogate_loss = torch.max(surrogate, surrogate_clipped).mean()
# Value function loss
if self.use_clipped_value_loss:
value_clipped = target_values_batch + (value_batch - target_values_batch).clamp(
-self.clip_param, self.clip_param
)
value_losses = (value_batch - returns_batch).pow(2)
value_losses_clipped = (value_clipped - returns_batch).pow(2)
value_loss = torch.max(value_losses, value_losses_clipped).mean()
else:
value_loss = (returns_batch - value_batch).pow(2).mean()
loss = surrogate_loss + self.value_loss_coef * value_loss - self.entropy_coef * entropy_batch.mean()
# Symmetry loss
if self.symmetry:
# obtain the symmetric actions
# if we did augmentation before then we don't need to augment again
if not self.symmetry["use_data_augmentation"]:
data_augmentation_func = self.symmetry["data_augmentation_func"]
obs_batch, _ = data_augmentation_func(obs=obs_batch, actions=None, env=self.symmetry["_env"])
# compute number of augmentations per sample
num_aug = int(obs_batch.shape[0] / original_batch_size)
# actions predicted by the actor for symmetrically-augmented observations
mean_actions_batch = self.policy.act_inference(obs_batch.detach().clone())
# compute the symmetrically augmented actions
# note: we are assuming the first augmentation is the original one.
# We do not use the action_batch from earlier since that action was sampled from the distribution.
# However, the symmetry loss is computed using the mean of the distribution.
action_mean_orig = mean_actions_batch[:original_batch_size]
_, actions_mean_symm_batch = data_augmentation_func(
obs=None, actions=action_mean_orig, env=self.symmetry["_env"]
)
# compute the loss (we skip the first augmentation as it is the original one)
mse_loss = torch.nn.MSELoss()
symmetry_loss = mse_loss(
mean_actions_batch[original_batch_size:], actions_mean_symm_batch.detach()[original_batch_size:]
)
# add the loss to the total loss
if self.symmetry["use_mirror_loss"]:
loss += self.symmetry["mirror_loss_coeff"] * symmetry_loss
else:
symmetry_loss = symmetry_loss.detach()
# Random Network Distillation loss
# TODO: Move this processing to inside RND module.
if self.rnd:
# extract the rnd_state
# TODO: Check if we still need torch no grad. It is just an affine transformation.
with torch.no_grad():
rnd_state_batch = self.rnd.get_rnd_state(obs_batch[:original_batch_size])
rnd_state_batch = self.rnd.state_normalizer(rnd_state_batch)
# predict the embedding and the target
predicted_embedding = self.rnd.predictor(rnd_state_batch)
target_embedding = self.rnd.target(rnd_state_batch).detach()
# compute the loss as the mean squared error
mseloss = torch.nn.MSELoss()
rnd_loss = mseloss(predicted_embedding, target_embedding)
expert_states = sample_amp_expert
policy_states = sample_amp_policy
with torch.no_grad():
expert_states = self.amp_normalizer.normalize_torch(expert_states.to(self.device), self.device)
policy_states = self.amp_normalizer.normalize_torch(policy_states, self.device)
contact_phase_push = obs_batch['critic'][:, -4]
mask_push = contact_phase_push == 1.
if mask_push.any():
policy_d = self.discriminator(policy_states.flatten(1))
expert_states = expert_states.to(self.device)
expert_d = self.discriminator(expert_states.flatten(1))
expert_loss = torch.nn.MSELoss()(expert_d, torch.ones(expert_d.size(), device=self.device))
policy_loss = torch.nn.MSELoss()(policy_d, -1 * torch.ones(policy_d.size(), device=self.device))
amp_loss = 0.5 * (expert_loss + policy_loss)
# grad penalty
grad_pen_loss = self.discriminator.compute_grad_pen(expert_states, lambda_=5)
else:
amp_loss = torch.tensor(0.0, device=self.device)
grad_pen_loss = torch.tensor(0.0, device=self.device)
expert_loss = torch.tensor(0.0, device=self.device)
policy_loss = torch.tensor(0.0, device=self.device)
loss += (amp_loss + grad_pen_loss)
self.amp_normalizer.update(policy_states.cpu().numpy())
self.amp_normalizer.update(expert_states.cpu().numpy())
# Compute the gradients
# -- For PPO
self.optimizer.zero_grad()
loss.backward()
# -- For RND
if self.rnd:
self.rnd_optimizer.zero_grad() # type: ignore
rnd_loss.backward()
# Collect gradients from all GPUs
if self.is_multi_gpu:
self.reduce_parameters()
# Apply the gradients
# -- For PPO
nn.utils.clip_grad_norm_(self.policy.parameters(), self.max_grad_norm)
self.optimizer.step()
# -- For RND
if self.rnd_optimizer:
self.rnd_optimizer.step()
# Store the losses
mean_value_loss += value_loss.item()
mean_surrogate_loss += surrogate_loss.item()
mean_entropy += entropy_batch.mean().item()
mean_amp_loss += amp_loss.item()
mean_grad_pen_loss += grad_pen_loss.item()
mean_policy_pred += policy_loss.mean().item()
mean_expert_pred += expert_loss.mean().item()
# -- RND loss
if mean_rnd_loss is not None:
mean_rnd_loss += rnd_loss.item()
# -- Symmetry loss
if mean_symmetry_loss is not None:
mean_symmetry_loss += symmetry_loss.item()
# -- For PPO
num_updates = self.num_learning_epochs * self.num_mini_batches
mean_value_loss /= num_updates
mean_surrogate_loss /= num_updates
mean_entropy /= num_updates
mean_amp_loss /= num_updates
mean_grad_pen_loss /= num_updates
mean_policy_pred /= num_updates
mean_expert_pred /= num_updates
# -- For RND
if mean_rnd_loss is not None:
mean_rnd_loss /= num_updates
# -- For Symmetry
if mean_symmetry_loss is not None:
mean_symmetry_loss /= num_updates
# -- Clear the storage
self.storage.clear()
# construct the loss dictionary
loss_dict = {
"value_function": mean_value_loss,
"surrogate": mean_surrogate_loss,
"entropy": mean_entropy,
"amp": mean_amp_loss,
"amp_grad_pen": mean_grad_pen_loss,
"amp_policy_pred": mean_policy_pred,
"amp_expert_pred": mean_expert_pred,
}
if self.rnd:
loss_dict["rnd"] = mean_rnd_loss
if self.symmetry:
loss_dict["symmetry"] = mean_symmetry_loss
return loss_dict
"""
Helper functions
"""
def broadcast_parameters(self):
"""Broadcast model parameters to all GPUs."""
# obtain the model parameters on current GPU
model_params = [self.policy.state_dict()]
if self.rnd:
model_params.append(self.rnd.predictor.state_dict())
# broadcast the model parameters
torch.distributed.broadcast_object_list(model_params, src=0)
# load the model parameters on all GPUs from source GPU
self.policy.load_state_dict(model_params[0])
if self.rnd:
self.rnd.predictor.load_state_dict(model_params[1])
def reduce_parameters(self):
"""Collect gradients from all GPUs and average them.
This function is called after the backward pass to synchronize the gradients across all GPUs.
"""
# Create a tensor to store the gradients
grads = [param.grad.view(-1) for param in self.policy.parameters() if param.grad is not None]
if self.rnd:
grads += [param.grad.view(-1) for param in self.rnd.parameters() if param.grad is not None]
all_grads = torch.cat(grads)
# Average the gradients across all GPUs
torch.distributed.all_reduce(all_grads, op=torch.distributed.ReduceOp.SUM)
all_grads /= self.gpu_world_size
# Get all parameters
all_params = self.policy.parameters()
if self.rnd:
all_params = chain(all_params, self.rnd.parameters())
# Update the gradients for all parameters with the reduced gradients
offset = 0
for param in all_params:
if param.grad is not None:
numel = param.numel()
# copy data back from shared buffer
param.grad.data.copy_(all_grads[offset : offset + numel].view_as(param.grad.data))
# update the offset for the next parameter
offset += numel

View File

@ -0,0 +1,185 @@
# Copyright (c) 2021-2025, ETH Zurich and NVIDIA CORPORATION
# All rights reserved.
#
# SPDX-License-Identifier: BSD-3-Clause
import torch
import torch.nn as nn
from rsl_rl.modules import StudentTeacher, StudentTeacherRecurrent
from rsl_rl.storage import RolloutStorage
from rsl_rl.utils import resolve_optimizer
class Distillation:
"""Distillation algorithm for training a student model to mimic a teacher model."""
policy: StudentTeacher | StudentTeacherRecurrent
"""The student teacher model."""
def __init__(
self,
policy,
num_learning_epochs=1,
gradient_length=15,
learning_rate=1e-3,
max_grad_norm=None,
loss_type="mse",
optimizer="adam",
device="cpu",
# Distributed training parameters
multi_gpu_cfg: dict | None = None,
):
# device-related parameters
self.device = device
self.is_multi_gpu = multi_gpu_cfg is not None
# Multi-GPU parameters
if multi_gpu_cfg is not None:
self.gpu_global_rank = multi_gpu_cfg["global_rank"]
self.gpu_world_size = multi_gpu_cfg["world_size"]
else:
self.gpu_global_rank = 0
self.gpu_world_size = 1
# distillation components
self.policy = policy
self.policy.to(self.device)
self.storage = None # initialized later
# initialize the optimizer
self.optimizer = resolve_optimizer(optimizer)(self.policy.parameters(), lr=learning_rate)
# initialize the transition
self.transition = RolloutStorage.Transition()
self.last_hidden_states = None
# distillation parameters
self.num_learning_epochs = num_learning_epochs
self.gradient_length = gradient_length
self.learning_rate = learning_rate
self.max_grad_norm = max_grad_norm
# initialize the loss function
loss_fn_dict = {
"mse": nn.functional.mse_loss,
"huber": nn.functional.huber_loss,
}
if loss_type in loss_fn_dict:
self.loss_fn = loss_fn_dict[loss_type]
else:
raise ValueError(f"Unknown loss type: {loss_type}. Supported types are: {list(loss_fn_dict.keys())}")
self.num_updates = 0
def init_storage(self, training_type, num_envs, num_transitions_per_env, obs, actions_shape):
# create rollout storage
self.storage = RolloutStorage(
training_type,
num_envs,
num_transitions_per_env,
obs,
actions_shape,
self.device,
)
def act(self, obs):
# compute the actions
self.transition.actions = self.policy.act(obs).detach()
self.transition.privileged_actions = self.policy.evaluate(obs).detach()
# record the observations
self.transition.observations = obs
return self.transition.actions
def process_env_step(self, obs, rewards, dones, extras):
# update the normalizers
self.policy.update_normalization(obs)
# record the rewards and dones
self.transition.rewards = rewards
self.transition.dones = dones
# record the transition
self.storage.add_transitions(self.transition)
self.transition.clear()
self.policy.reset(dones)
def update(self):
self.num_updates += 1
mean_behavior_loss = 0
loss = 0
cnt = 0
for epoch in range(self.num_learning_epochs):
self.policy.reset(hidden_states=self.last_hidden_states)
self.policy.detach_hidden_states()
for obs, _, privileged_actions, dones in self.storage.generator():
# inference the student for gradient computation
actions = self.policy.act_inference(obs)
# behavior cloning loss
behavior_loss = self.loss_fn(actions, privileged_actions)
# total loss
loss = loss + behavior_loss
mean_behavior_loss += behavior_loss.item()
cnt += 1
# gradient step
if cnt % self.gradient_length == 0:
self.optimizer.zero_grad()
loss.backward()
if self.is_multi_gpu:
self.reduce_parameters()
if self.max_grad_norm:
nn.utils.clip_grad_norm_(self.policy.student.parameters(), self.max_grad_norm)
self.optimizer.step()
self.policy.detach_hidden_states()
loss = 0
# reset dones
self.policy.reset(dones.view(-1))
self.policy.detach_hidden_states(dones.view(-1))
mean_behavior_loss /= cnt
self.storage.clear()
self.last_hidden_states = self.policy.get_hidden_states()
self.policy.detach_hidden_states()
# construct the loss dictionary
loss_dict = {"behavior": mean_behavior_loss}
return loss_dict
"""
Helper functions
"""
def broadcast_parameters(self):
"""Broadcast model parameters to all GPUs."""
# obtain the model parameters on current GPU
model_params = [self.policy.state_dict()]
# broadcast the model parameters
torch.distributed.broadcast_object_list(model_params, src=0)
# load the model parameters on all GPUs from source GPU
self.policy.load_state_dict(model_params[0])
def reduce_parameters(self):
"""Collect gradients from all GPUs and average them.
This function is called after the backward pass to synchronize the gradients across all GPUs.
"""
# Create a tensor to store the gradients
grads = [param.grad.view(-1) for param in self.policy.parameters() if param.grad is not None]
all_grads = torch.cat(grads)
# Average the gradients across all GPUs
torch.distributed.all_reduce(all_grads, op=torch.distributed.ReduceOp.SUM)
all_grads /= self.gpu_world_size
# Update the gradients for all parameters with the reduced gradients
offset = 0
for param in self.policy.parameters():
if param.grad is not None:
numel = param.numel()
# copy data back from shared buffer
param.grad.data.copy_(all_grads[offset : offset + numel].view_as(param.grad.data))
# update the offset for the next parameter
offset += numel

469
rsl_rl/algorithms/ppo.py Normal file
View File

@ -0,0 +1,469 @@
# Copyright (c) 2021-2025, ETH Zurich and NVIDIA CORPORATION
# All rights reserved.
#
# SPDX-License-Identifier: BSD-3-Clause
from __future__ import annotations
import torch
import torch.nn as nn
import torch.optim as optim
from itertools import chain
from rsl_rl.modules import ActorCritic
from rsl_rl.modules.rnd import RandomNetworkDistillation
from rsl_rl.storage import RolloutStorage
from rsl_rl.utils import string_to_callable
class PPO:
"""Proximal Policy Optimization algorithm (https://arxiv.org/abs/1707.06347)."""
policy: ActorCritic
"""The actor critic module."""
def __init__(
self,
policy,
num_learning_epochs=5,
num_mini_batches=4,
clip_param=0.2,
gamma=0.99,
lam=0.95,
value_loss_coef=1.0,
entropy_coef=0.01,
learning_rate=0.001,
max_grad_norm=1.0,
use_clipped_value_loss=True,
schedule="adaptive",
desired_kl=0.01,
device="cpu",
normalize_advantage_per_mini_batch=False,
# RND parameters
rnd_cfg: dict | None = None,
# Symmetry parameters
symmetry_cfg: dict | None = None,
# Distributed training parameters
multi_gpu_cfg: dict | None = None,
):
# device-related parameters
self.device = device
self.is_multi_gpu = multi_gpu_cfg is not None
# Multi-GPU parameters
if multi_gpu_cfg is not None:
self.gpu_global_rank = multi_gpu_cfg["global_rank"]
self.gpu_world_size = multi_gpu_cfg["world_size"]
else:
self.gpu_global_rank = 0
self.gpu_world_size = 1
# RND components
if rnd_cfg is not None:
# Extract parameters used in ppo
rnd_lr = rnd_cfg.pop("learning_rate", 1e-3)
# Create RND module
self.rnd = RandomNetworkDistillation(device=self.device, **rnd_cfg)
# Create RND optimizer
params = self.rnd.predictor.parameters()
self.rnd_optimizer = optim.Adam(params, lr=rnd_lr)
else:
self.rnd = None
self.rnd_optimizer = None
# Symmetry components
if symmetry_cfg is not None:
# Check if symmetry is enabled
use_symmetry = symmetry_cfg["use_data_augmentation"] or symmetry_cfg["use_mirror_loss"]
# Print that we are not using symmetry
if not use_symmetry:
print("Symmetry not used for learning. We will use it for logging instead.")
# If function is a string then resolve it to a function
if isinstance(symmetry_cfg["data_augmentation_func"], str):
symmetry_cfg["data_augmentation_func"] = string_to_callable(symmetry_cfg["data_augmentation_func"])
# Check valid configuration
if symmetry_cfg["use_data_augmentation"] and not callable(symmetry_cfg["data_augmentation_func"]):
raise ValueError(
"Data augmentation enabled but the function is not callable:"
f" {symmetry_cfg['data_augmentation_func']}"
)
# Store symmetry configuration
self.symmetry = symmetry_cfg
else:
self.symmetry = None
# PPO components
self.policy = policy
self.policy.to(self.device)
# Create optimizer
self.optimizer = optim.Adam(self.policy.parameters(), lr=learning_rate)
# Create rollout storage
self.storage: RolloutStorage = None # type: ignore
self.transition = RolloutStorage.Transition()
# PPO parameters
self.clip_param = clip_param
self.num_learning_epochs = num_learning_epochs
self.num_mini_batches = num_mini_batches
self.value_loss_coef = value_loss_coef
self.entropy_coef = entropy_coef
self.gamma = gamma
self.lam = lam
self.max_grad_norm = max_grad_norm
self.use_clipped_value_loss = use_clipped_value_loss
self.desired_kl = desired_kl
self.schedule = schedule
self.learning_rate = learning_rate
self.normalize_advantage_per_mini_batch = normalize_advantage_per_mini_batch
def init_storage(self, training_type, num_envs, num_transitions_per_env, obs, actions_shape):
# create rollout storage
self.storage = RolloutStorage(
training_type,
num_envs,
num_transitions_per_env,
obs,
actions_shape,
self.device,
)
def act(self, obs):
if self.policy.is_recurrent:
self.transition.hidden_states = self.policy.get_hidden_states()
# compute the actions and values
self.transition.actions = self.policy.act(obs).detach()
self.transition.values = self.policy.evaluate(obs).detach()
self.transition.actions_log_prob = self.policy.get_actions_log_prob(self.transition.actions).detach()
self.transition.action_mean = self.policy.action_mean.detach()
self.transition.action_sigma = self.policy.action_std.detach()
# need to record obs before env.step()
self.transition.observations = obs
return self.transition.actions
def process_env_step(self, obs, rewards, dones, extras):
# update the normalizers
self.policy.update_normalization(obs)
if self.rnd:
self.rnd.update_normalization(obs)
# Record the rewards and dones
# Note: we clone here because later on we bootstrap the rewards based on timeouts
self.transition.rewards = rewards.clone()
self.transition.dones = dones
# Compute the intrinsic rewards and add to extrinsic rewards
if self.rnd:
# Compute the intrinsic rewards
self.intrinsic_rewards = self.rnd.get_intrinsic_reward(obs)
# Add intrinsic rewards to extrinsic rewards
self.transition.rewards += self.intrinsic_rewards
# Bootstrapping on time outs
if "time_outs" in extras:
self.transition.rewards += self.gamma * torch.squeeze(
self.transition.values * extras["time_outs"].unsqueeze(1).to(self.device), 1
)
# record the transition
self.storage.add_transitions(self.transition)
self.transition.clear()
self.policy.reset(dones)
def compute_returns(self, obs):
# compute value for the last step
last_values = self.policy.evaluate(obs).detach()
self.storage.compute_returns(
last_values, self.gamma, self.lam, normalize_advantage=not self.normalize_advantage_per_mini_batch
)
def update(self): # noqa: C901
mean_value_loss = 0
mean_surrogate_loss = 0
mean_entropy = 0
# -- RND loss
if self.rnd:
mean_rnd_loss = 0
else:
mean_rnd_loss = None
# -- Symmetry loss
if self.symmetry:
mean_symmetry_loss = 0
else:
mean_symmetry_loss = None
# generator for mini batches
if self.policy.is_recurrent:
generator = self.storage.recurrent_mini_batch_generator(self.num_mini_batches, self.num_learning_epochs)
else:
generator = self.storage.mini_batch_generator(self.num_mini_batches, self.num_learning_epochs)
# iterate over batches
for (
obs_batch,
actions_batch,
target_values_batch,
advantages_batch,
returns_batch,
old_actions_log_prob_batch,
old_mu_batch,
old_sigma_batch,
hid_states_batch,
masks_batch,
) in generator:
# number of augmentations per sample
# we start with 1 and increase it if we use symmetry augmentation
num_aug = 1
# original batch size
# we assume policy group is always there and needs augmentation
original_batch_size = obs_batch.batch_size[0]
# check if we should normalize advantages per mini batch
if self.normalize_advantage_per_mini_batch:
with torch.no_grad():
advantages_batch = (advantages_batch - advantages_batch.mean()) / (advantages_batch.std() + 1e-8)
# Perform symmetric augmentation
if self.symmetry and self.symmetry["use_data_augmentation"]:
# augmentation using symmetry
data_augmentation_func = self.symmetry["data_augmentation_func"]
# returned shape: [batch_size * num_aug, ...]
obs_batch, actions_batch = data_augmentation_func(
obs=obs_batch,
actions=actions_batch,
env=self.symmetry["_env"],
)
# compute number of augmentations per sample
# we assume policy group is always there and needs augmentation
num_aug = int(obs_batch.batch_size[0] / original_batch_size)
# repeat the rest of the batch
# -- actor
old_actions_log_prob_batch = old_actions_log_prob_batch.repeat(num_aug, 1)
# -- critic
target_values_batch = target_values_batch.repeat(num_aug, 1)
advantages_batch = advantages_batch.repeat(num_aug, 1)
returns_batch = returns_batch.repeat(num_aug, 1)
# Recompute actions log prob and entropy for current batch of transitions
# Note: we need to do this because we updated the policy with the new parameters
# -- actor
self.policy.act(obs_batch, masks=masks_batch, hidden_states=hid_states_batch[0])
actions_log_prob_batch = self.policy.get_actions_log_prob(actions_batch)
# -- critic
value_batch = self.policy.evaluate(obs_batch, masks=masks_batch, hidden_states=hid_states_batch[1])
# -- entropy
# we only keep the entropy of the first augmentation (the original one)
mu_batch = self.policy.action_mean[:original_batch_size]
sigma_batch = self.policy.action_std[:original_batch_size]
entropy_batch = self.policy.entropy[:original_batch_size]
# KL
if self.desired_kl is not None and self.schedule == "adaptive":
with torch.inference_mode():
kl = torch.sum(
torch.log(sigma_batch / old_sigma_batch + 1.0e-5)
+ (torch.square(old_sigma_batch) + torch.square(old_mu_batch - mu_batch))
/ (2.0 * torch.square(sigma_batch))
- 0.5,
axis=-1,
)
kl_mean = torch.mean(kl)
# Reduce the KL divergence across all GPUs
if self.is_multi_gpu:
torch.distributed.all_reduce(kl_mean, op=torch.distributed.ReduceOp.SUM)
kl_mean /= self.gpu_world_size
# Update the learning rate
# Perform this adaptation only on the main process
# TODO: Is this needed? If KL-divergence is the "same" across all GPUs,
# then the learning rate should be the same across all GPUs.
if self.gpu_global_rank == 0:
if kl_mean > self.desired_kl * 2.0:
self.learning_rate = max(1e-5, self.learning_rate / 1.5)
elif kl_mean < self.desired_kl / 2.0 and kl_mean > 0.0:
self.learning_rate = min(1e-2, self.learning_rate * 1.5)
# Update the learning rate for all GPUs
if self.is_multi_gpu:
lr_tensor = torch.tensor(self.learning_rate, device=self.device)
torch.distributed.broadcast(lr_tensor, src=0)
self.learning_rate = lr_tensor.item()
# Update the learning rate for all parameter groups
for param_group in self.optimizer.param_groups:
param_group["lr"] = self.learning_rate
# Surrogate loss
ratio = torch.exp(actions_log_prob_batch - torch.squeeze(old_actions_log_prob_batch))
surrogate = -torch.squeeze(advantages_batch) * ratio
surrogate_clipped = -torch.squeeze(advantages_batch) * torch.clamp(
ratio, 1.0 - self.clip_param, 1.0 + self.clip_param
)
surrogate_loss = torch.max(surrogate, surrogate_clipped).mean()
# Value function loss
if self.use_clipped_value_loss:
value_clipped = target_values_batch + (value_batch - target_values_batch).clamp(
-self.clip_param, self.clip_param
)
value_losses = (value_batch - returns_batch).pow(2)
value_losses_clipped = (value_clipped - returns_batch).pow(2)
value_loss = torch.max(value_losses, value_losses_clipped).mean()
else:
value_loss = (returns_batch - value_batch).pow(2).mean()
loss = surrogate_loss + self.value_loss_coef * value_loss - self.entropy_coef * entropy_batch.mean()
# Symmetry loss
if self.symmetry:
# obtain the symmetric actions
# if we did augmentation before then we don't need to augment again
if not self.symmetry["use_data_augmentation"]:
data_augmentation_func = self.symmetry["data_augmentation_func"]
obs_batch, _ = data_augmentation_func(obs=obs_batch, actions=None, env=self.symmetry["_env"])
# compute number of augmentations per sample
num_aug = int(obs_batch.shape[0] / original_batch_size)
# actions predicted by the actor for symmetrically-augmented observations
mean_actions_batch = self.policy.act_inference(obs_batch.detach().clone())
# compute the symmetrically augmented actions
# note: we are assuming the first augmentation is the original one.
# We do not use the action_batch from earlier since that action was sampled from the distribution.
# However, the symmetry loss is computed using the mean of the distribution.
action_mean_orig = mean_actions_batch[:original_batch_size]
_, actions_mean_symm_batch = data_augmentation_func(
obs=None, actions=action_mean_orig, env=self.symmetry["_env"]
)
# compute the loss (we skip the first augmentation as it is the original one)
mse_loss = torch.nn.MSELoss()
symmetry_loss = mse_loss(
mean_actions_batch[original_batch_size:], actions_mean_symm_batch.detach()[original_batch_size:]
)
# add the loss to the total loss
if self.symmetry["use_mirror_loss"]:
loss += self.symmetry["mirror_loss_coeff"] * symmetry_loss
else:
symmetry_loss = symmetry_loss.detach()
# Random Network Distillation loss
# TODO: Move this processing to inside RND module.
if self.rnd:
# extract the rnd_state
# TODO: Check if we still need torch no grad. It is just an affine transformation.
with torch.no_grad():
rnd_state_batch = self.rnd.get_rnd_state(obs_batch[:original_batch_size])
rnd_state_batch = self.rnd.state_normalizer(rnd_state_batch)
# predict the embedding and the target
predicted_embedding = self.rnd.predictor(rnd_state_batch)
target_embedding = self.rnd.target(rnd_state_batch).detach()
# compute the loss as the mean squared error
mseloss = torch.nn.MSELoss()
rnd_loss = mseloss(predicted_embedding, target_embedding)
# Compute the gradients
# -- For PPO
self.optimizer.zero_grad()
loss.backward()
# -- For RND
if self.rnd:
self.rnd_optimizer.zero_grad() # type: ignore
rnd_loss.backward()
# Collect gradients from all GPUs
if self.is_multi_gpu:
self.reduce_parameters()
# Apply the gradients
# -- For PPO
nn.utils.clip_grad_norm_(self.policy.parameters(), self.max_grad_norm)
self.optimizer.step()
# -- For RND
if self.rnd_optimizer:
self.rnd_optimizer.step()
# Store the losses
mean_value_loss += value_loss.item()
mean_surrogate_loss += surrogate_loss.item()
mean_entropy += entropy_batch.mean().item()
# -- RND loss
if mean_rnd_loss is not None:
mean_rnd_loss += rnd_loss.item()
# -- Symmetry loss
if mean_symmetry_loss is not None:
mean_symmetry_loss += symmetry_loss.item()
# -- For PPO
num_updates = self.num_learning_epochs * self.num_mini_batches
mean_value_loss /= num_updates
mean_surrogate_loss /= num_updates
mean_entropy /= num_updates
# -- For RND
if mean_rnd_loss is not None:
mean_rnd_loss /= num_updates
# -- For Symmetry
if mean_symmetry_loss is not None:
mean_symmetry_loss /= num_updates
# -- Clear the storage
self.storage.clear()
# construct the loss dictionary
loss_dict = {
"value_function": mean_value_loss,
"surrogate": mean_surrogate_loss,
"entropy": mean_entropy,
}
if self.rnd:
loss_dict["rnd"] = mean_rnd_loss
if self.symmetry:
loss_dict["symmetry"] = mean_symmetry_loss
return loss_dict
"""
Helper functions
"""
def broadcast_parameters(self):
"""Broadcast model parameters to all GPUs."""
# obtain the model parameters on current GPU
model_params = [self.policy.state_dict()]
if self.rnd:
model_params.append(self.rnd.predictor.state_dict())
# broadcast the model parameters
torch.distributed.broadcast_object_list(model_params, src=0)
# load the model parameters on all GPUs from source GPU
self.policy.load_state_dict(model_params[0])
if self.rnd:
self.rnd.predictor.load_state_dict(model_params[1])
def reduce_parameters(self):
"""Collect gradients from all GPUs and average them.
This function is called after the backward pass to synchronize the gradients across all GPUs.
"""
# Create a tensor to store the gradients
grads = [param.grad.view(-1) for param in self.policy.parameters() if param.grad is not None]
if self.rnd:
grads += [param.grad.view(-1) for param in self.rnd.parameters() if param.grad is not None]
all_grads = torch.cat(grads)
# Average the gradients across all GPUs
torch.distributed.all_reduce(all_grads, op=torch.distributed.ReduceOp.SUM)
all_grads /= self.gpu_world_size
# Get all parameters
all_params = self.policy.parameters()
if self.rnd:
all_params = chain(all_params, self.rnd.parameters())
# Update the gradients for all parameters with the reduced gradients
offset = 0
for param in all_params:
if param.grad is not None:
numel = param.numel()
# copy data back from shared buffer
param.grad.data.copy_(all_grads[offset : offset + numel].view_as(param.grad.data))
# update the offset for the next parameter
offset += numel

View File

@ -0,0 +1,6 @@
# Copyright (c) 2021-2025, ETH Zurich and NVIDIA CORPORATION
# All rights reserved.
#
# SPDX-License-Identifier: BSD-3-Clause
"""Main module for the rsl_rl package."""

View File

@ -0,0 +1,11 @@
# Copyright (c) 2021-2025, ETH Zurich and NVIDIA CORPORATION
# All rights reserved.
#
# SPDX-License-Identifier: BSD-3-Clause
"""Implementation of different RL agents."""
from .distillation import Distillation
from .ppo import PPO
from .amp_ppo import AMP_PPO
__all__ = ["PPO", "Distillation", "AMP_PPO"]

View File

@ -0,0 +1,571 @@
# Copyright (c) 2021-2025, ETH Zurich and NVIDIA CORPORATION
# All rights reserved.
#
# SPDX-License-Identifier: BSD-3-Clause
from __future__ import annotations
from torch._tensor import Tensor
from torch._tensor import Tensor
from typing import Any
from copy import deepcopy
import torch
import torch.nn as nn
import torch.optim as optim
from itertools import chain
from rsl_rl.modules import ActorCritic
from rsl_rl.modules.rnd import RandomNetworkDistillation
from rsl_rl.storage import RolloutStorage, ReplayBufferMulti
from rsl_rl.utils import string_to_callable
class AMP_PPO:
"""Proximal Policy Optimization algorithm (https://arxiv.org/abs/1707.06347)."""
policy: ActorCritic
"""The actor critic module."""
def __init__(
self,
policy,
discriminator,
amp_data,
amp_normalizer,
amp_num_frames=1,
amp_replay_buffer_size=100000,
num_learning_epochs=5,
num_mini_batches=4,
clip_param=0.2,
gamma=0.99,
lam=0.95,
value_loss_coef=1.0,
entropy_coef=0.01,
learning_rate=0.001,
max_grad_norm=1.0,
use_clipped_value_loss=True,
schedule="adaptive",
desired_kl=0.01,
device="cpu",
normalize_advantage_per_mini_batch=False,
# RND parameters
rnd_cfg: dict | None = None,
# Symmetry parameters
symmetry_cfg: dict | None = None,
# Distributed training parameters
multi_gpu_cfg: dict | None = None,
):
# device-related parameters
self.device = device
self.is_multi_gpu = multi_gpu_cfg is not None
# Multi-GPU parameters
if multi_gpu_cfg is not None:
self.gpu_global_rank = multi_gpu_cfg["global_rank"]
self.gpu_world_size = multi_gpu_cfg["world_size"]
else:
self.gpu_global_rank = 0
self.gpu_world_size = 1
# RND components
if rnd_cfg is not None:
# Extract parameters used in ppo
rnd_lr = rnd_cfg.pop("learning_rate", 1e-3)
# Create RND module
self.rnd = RandomNetworkDistillation(device=self.device, **rnd_cfg)
# Create RND optimizer
params = self.rnd.predictor.parameters()
self.rnd_optimizer = optim.Adam(params, lr=rnd_lr)
else:
self.rnd = None
self.rnd_optimizer = None
# Symmetry components
if symmetry_cfg is not None:
# Check if symmetry is enabled
use_symmetry = symmetry_cfg["use_data_augmentation"] or symmetry_cfg["use_mirror_loss"]
# Print that we are not using symmetry
if not use_symmetry:
print("Symmetry not used for learning. We will use it for logging instead.")
# If function is a string then resolve it to a function
if isinstance(symmetry_cfg["data_augmentation_func"], str):
symmetry_cfg["data_augmentation_func"] = string_to_callable(symmetry_cfg["data_augmentation_func"])
# Check valid configuration
if symmetry_cfg["use_data_augmentation"] and not callable(symmetry_cfg["data_augmentation_func"]):
raise ValueError(
"Data augmentation enabled but the function is not callable:"
f" {symmetry_cfg['data_augmentation_func']}"
)
# Store symmetry configuration
self.symmetry = symmetry_cfg
else:
self.symmetry = None
## AMP components
self.discriminator = discriminator
self.discriminator.to(self.device)
self.amp_storage = ReplayBufferMulti(discriminator.state_dim, amp_replay_buffer_size, amp_num_frames, device)
self.amp_data = amp_data
self.amp_normalizer = amp_normalizer
# PPO components
self.policy = policy
self.policy.to(self.device)
# Create rollout storage
self.storage: RolloutStorage = None # type: ignore
self.transition = RolloutStorage.Transition()
self.amp_transition = RolloutStorage.Transition()
params = [
{'params': self.policy.parameters(), 'name': 'policy'},
]
params.append({
'params': self.discriminator.trunk.parameters(),
'weight_decay': 10e-4,
'name': f'amp_trunk'
})
params.append({
'params': self.discriminator.amp_linear.parameters(),
'weight_decay': 10e-2,
'name': f'amp_head'
})
# Create optimizer
self.optimizer = optim.Adam(params, lr=learning_rate)
# PPO parameters
self.clip_param = clip_param
self.num_learning_epochs = num_learning_epochs
self.num_mini_batches = num_mini_batches
self.value_loss_coef = value_loss_coef
self.entropy_coef = entropy_coef
self.gamma = gamma
self.lam = lam
self.max_grad_norm = max_grad_norm
self.use_clipped_value_loss = use_clipped_value_loss
self.desired_kl = desired_kl
self.schedule = schedule
self.learning_rate = learning_rate
self.normalize_advantage_per_mini_batch = normalize_advantage_per_mini_batch
def init_storage(self, training_type, num_envs, num_transitions_per_env, obs, actions_shape):
# create rollout storage
self.storage = RolloutStorage(
training_type,
num_envs,
num_transitions_per_env,
obs,
actions_shape,
self.device,
)
def act(self, obs, amp_obs):
if self.policy.is_recurrent:
self.transition.hidden_states = self.policy.get_hidden_states()
# compute the actions and values
self.transition.actions = self.policy.act(obs).detach()
self.transition.values = self.policy.evaluate(obs).detach()
self.transition.actions_log_prob = self.policy.get_actions_log_prob(self.transition.actions).detach()
self.transition.action_mean = self.policy.action_mean.detach()
self.transition.action_sigma = self.policy.action_std.detach()
# need to record obs before env.step()
self.transition.observations = obs
self.amp_transition.observations = amp_obs
return self.transition.actions
def process_env_step(self, obs, rewards, dones, extras,amp_obs, amp_obs_frames=None):
# update the normalizers
self.policy.update_normalization(obs)
if self.rnd:
self.rnd.update_normalization(obs)
# Record the rewards and dones
# Note: we clone here because later on we bootstrap the rewards based on timeouts
self.transition.rewards = rewards.clone()
self.transition.dones = dones
# Compute the intrinsic rewards and add to extrinsic rewards
if self.rnd:
# Compute the intrinsic rewards
self.intrinsic_rewards = self.rnd.get_intrinsic_reward(obs)
# Add intrinsic rewards to extrinsic rewards
self.transition.rewards += self.intrinsic_rewards
# Bootstrapping on time outs
if "time_outs" in extras:
self.transition.rewards += self.gamma * torch.squeeze(
self.transition.values * extras["time_outs"].unsqueeze(1).to(self.device), 1
)
if amp_obs_frames is not None:
self.amp_storage.insert(amp_obs_frames)
else:
self.amp_storage.insert(self.amp_transition.observations, amp_obs)
# record the transition
self.storage.add_transitions(self.transition)
self.transition.clear()
self.amp_transition.clear()
self.policy.reset(dones)
def compute_returns(self, obs):
# compute value for the last step
last_values = self.policy.evaluate(obs).detach()
self.storage.compute_returns(
last_values, self.gamma, self.lam, normalize_advantage=not self.normalize_advantage_per_mini_batch
)
def update(self): # noqa: C901
mean_value_loss = 0
mean_surrogate_loss = 0
mean_entropy = 0
mean_amp_loss = 0
mean_grad_pen_loss = 0
mean_policy_pred = 0
mean_expert_pred = 0
# -- RND loss
if self.rnd:
mean_rnd_loss = 0
else:
mean_rnd_loss = None
# -- Symmetry loss
if self.symmetry:
mean_symmetry_loss = 0
else:
mean_symmetry_loss = None
# generator for mini batches
if self.policy.is_recurrent:
generator = self.storage.recurrent_mini_batch_generator(self.num_mini_batches, self.num_learning_epochs)
else:
generator = self.storage.mini_batch_generator(self.num_mini_batches, self.num_learning_epochs)
amp_policy_generator = self.amp_storage.feed_forward_generator(
self.num_learning_epochs * self.num_mini_batches,
self.storage.num_envs * self.storage.num_transitions_per_env // self.num_mini_batches,
)
amp_expert_generator = self.amp_data.feed_forward_generator_23dof_multi(
self.num_learning_epochs * self.num_mini_batches,
self.storage.num_envs * self.storage.num_transitions_per_env // self.num_mini_batches,
)
# iterate over batches
for sample, sample_amp_policy, sample_amp_expert in zip(generator, amp_policy_generator, amp_expert_generator):
(
obs_batch,
actions_batch,
target_values_batch,
advantages_batch,
returns_batch,
old_actions_log_prob_batch,
old_mu_batch,
old_sigma_batch,
hid_states_batch,
masks_batch,
) = sample
# number of augmentations per sample
# we start with 1 and increase it if we use symmetry augmentation
num_aug = 1
# original batch size
# we assume policy group is always there and needs augmentation
original_batch_size = obs_batch.batch_size[0]
# check if we should normalize advantages per mini batch
if self.normalize_advantage_per_mini_batch:
with torch.no_grad():
advantages_batch = (advantages_batch - advantages_batch.mean()) / (advantages_batch.std() + 1e-8)
# Perform symmetric augmentation
if self.symmetry and self.symmetry["use_data_augmentation"]:
# augmentation using symmetry
data_augmentation_func = self.symmetry["data_augmentation_func"]
# returned shape: [batch_size * num_aug, ...]
obs_batch, actions_batch = data_augmentation_func(
obs=obs_batch,
actions=actions_batch,
env=self.symmetry["_env"],
)
# compute number of augmentations per sample
# we assume policy group is always there and needs augmentation
num_aug = int(obs_batch.batch_size[0] / original_batch_size)
# repeat the rest of the batch
# -- actor
old_actions_log_prob_batch = old_actions_log_prob_batch.repeat(num_aug, 1)
# -- critic
target_values_batch = target_values_batch.repeat(num_aug, 1)
advantages_batch = advantages_batch.repeat(num_aug, 1)
returns_batch = returns_batch.repeat(num_aug, 1)
# Recompute actions log prob and entropy for current batch of transitions
# Note: we need to do this because we updated the policy with the new parameters
# -- actor
self.policy.act(obs_batch, masks=masks_batch, hidden_states=hid_states_batch[0])
actions_log_prob_batch = self.policy.get_actions_log_prob(actions_batch)
# -- critic
value_batch = self.policy.evaluate(obs_batch, masks=masks_batch, hidden_states=hid_states_batch[1])
# -- entropy
# we only keep the entropy of the first augmentation (the original one)
mu_batch = self.policy.action_mean[:original_batch_size]
sigma_batch = self.policy.action_std[:original_batch_size]
entropy_batch = self.policy.entropy[:original_batch_size]
# KL
if self.desired_kl is not None and self.schedule == "adaptive":
with torch.inference_mode():
kl = torch.sum(
torch.log(sigma_batch / old_sigma_batch + 1.0e-5)
+ (torch.square(old_sigma_batch) + torch.square(old_mu_batch - mu_batch))
/ (2.0 * torch.square(sigma_batch))
- 0.5,
axis=-1,
)
kl_mean = torch.mean(kl)
# Reduce the KL divergence across all GPUs
if self.is_multi_gpu:
torch.distributed.all_reduce(kl_mean, op=torch.distributed.ReduceOp.SUM)
kl_mean /= self.gpu_world_size
# Update the learning rate
# Perform this adaptation only on the main process
# TODO: Is this needed? If KL-divergence is the "same" across all GPUs,
# then the learning rate should be the same across all GPUs.
if self.gpu_global_rank == 0:
if kl_mean > self.desired_kl * 2.0:
self.learning_rate = max(1e-5, self.learning_rate / 1.5)
elif kl_mean < self.desired_kl / 2.0 and kl_mean > 0.0:
self.learning_rate = min(1e-2, self.learning_rate * 1.5)
# Update the learning rate for all GPUs
if self.is_multi_gpu:
lr_tensor = torch.tensor(self.learning_rate, device=self.device)
torch.distributed.broadcast(lr_tensor, src=0)
self.learning_rate = lr_tensor.item()
# Update the learning rate for all parameter groups
for param_group in self.optimizer.param_groups:
param_group["lr"] = self.learning_rate
# Surrogate loss
ratio = torch.exp(actions_log_prob_batch - torch.squeeze(old_actions_log_prob_batch))
surrogate = -torch.squeeze(advantages_batch) * ratio
surrogate_clipped = -torch.squeeze(advantages_batch) * torch.clamp(
ratio, 1.0 - self.clip_param, 1.0 + self.clip_param
)
surrogate_loss = torch.max(surrogate, surrogate_clipped).mean()
# Value function loss
if self.use_clipped_value_loss:
value_clipped = target_values_batch + (value_batch - target_values_batch).clamp(
-self.clip_param, self.clip_param
)
value_losses = (value_batch - returns_batch).pow(2)
value_losses_clipped = (value_clipped - returns_batch).pow(2)
value_loss = torch.max(value_losses, value_losses_clipped).mean()
else:
value_loss = (returns_batch - value_batch).pow(2).mean()
loss = surrogate_loss + self.value_loss_coef * value_loss - self.entropy_coef * entropy_batch.mean()
# Symmetry loss
if self.symmetry:
# obtain the symmetric actions
# if we did augmentation before then we don't need to augment again
if not self.symmetry["use_data_augmentation"]:
data_augmentation_func = self.symmetry["data_augmentation_func"]
obs_batch, _ = data_augmentation_func(obs=obs_batch, actions=None, env=self.symmetry["_env"])
# compute number of augmentations per sample
num_aug = int(obs_batch.shape[0] / original_batch_size)
# actions predicted by the actor for symmetrically-augmented observations
mean_actions_batch = self.policy.act_inference(obs_batch.detach().clone())
# compute the symmetrically augmented actions
# note: we are assuming the first augmentation is the original one.
# We do not use the action_batch from earlier since that action was sampled from the distribution.
# However, the symmetry loss is computed using the mean of the distribution.
action_mean_orig = mean_actions_batch[:original_batch_size]
_, actions_mean_symm_batch = data_augmentation_func(
obs=None, actions=action_mean_orig, env=self.symmetry["_env"]
)
# compute the loss (we skip the first augmentation as it is the original one)
mse_loss = torch.nn.MSELoss()
symmetry_loss = mse_loss(
mean_actions_batch[original_batch_size:], actions_mean_symm_batch.detach()[original_batch_size:]
)
# add the loss to the total loss
if self.symmetry["use_mirror_loss"]:
loss += self.symmetry["mirror_loss_coeff"] * symmetry_loss
else:
symmetry_loss = symmetry_loss.detach()
# Random Network Distillation loss
# TODO: Move this processing to inside RND module.
if self.rnd:
# extract the rnd_state
# TODO: Check if we still need torch no grad. It is just an affine transformation.
with torch.no_grad():
rnd_state_batch = self.rnd.get_rnd_state(obs_batch[:original_batch_size])
rnd_state_batch = self.rnd.state_normalizer(rnd_state_batch)
# predict the embedding and the target
predicted_embedding = self.rnd.predictor(rnd_state_batch)
target_embedding = self.rnd.target(rnd_state_batch).detach()
# compute the loss as the mean squared error
mseloss = torch.nn.MSELoss()
rnd_loss = mseloss(predicted_embedding, target_embedding)
expert_states = sample_amp_expert
policy_states = sample_amp_policy
with torch.no_grad():
expert_states = self.amp_normalizer.normalize_torch(expert_states.to(self.device), self.device)
policy_states = self.amp_normalizer.normalize_torch(policy_states, self.device)
contact_phase_push = obs_batch['critic'][:, -4]
mask_push = contact_phase_push == 1.
if mask_push.any():
policy_d = self.discriminator(policy_states.flatten(1))
expert_states = expert_states.to(self.device)
expert_d = self.discriminator(expert_states.flatten(1))
expert_loss = torch.nn.MSELoss()(expert_d, torch.ones(expert_d.size(), device=self.device))
policy_loss = torch.nn.MSELoss()(policy_d, -1 * torch.ones(policy_d.size(), device=self.device))
amp_loss = 0.5 * (expert_loss + policy_loss)
# grad penalty
grad_pen_loss = self.discriminator.compute_grad_pen(expert_states, lambda_=5)
else:
amp_loss = torch.tensor(0.0, device=self.device)
grad_pen_loss = torch.tensor(0.0, device=self.device)
expert_loss = torch.tensor(0.0, device=self.device)
policy_loss = torch.tensor(0.0, device=self.device)
loss += (amp_loss + grad_pen_loss)
self.amp_normalizer.update(policy_states.cpu().numpy())
self.amp_normalizer.update(expert_states.cpu().numpy())
# Compute the gradients
# -- For PPO
self.optimizer.zero_grad()
loss.backward()
# -- For RND
if self.rnd:
self.rnd_optimizer.zero_grad() # type: ignore
rnd_loss.backward()
# Collect gradients from all GPUs
if self.is_multi_gpu:
self.reduce_parameters()
# Apply the gradients
# -- For PPO
nn.utils.clip_grad_norm_(self.policy.parameters(), self.max_grad_norm)
self.optimizer.step()
# -- For RND
if self.rnd_optimizer:
self.rnd_optimizer.step()
# Store the losses
mean_value_loss += value_loss.item()
mean_surrogate_loss += surrogate_loss.item()
mean_entropy += entropy_batch.mean().item()
mean_amp_loss += amp_loss.item()
mean_grad_pen_loss += grad_pen_loss.item()
mean_policy_pred += policy_loss.mean().item()
mean_expert_pred += expert_loss.mean().item()
# -- RND loss
if mean_rnd_loss is not None:
mean_rnd_loss += rnd_loss.item()
# -- Symmetry loss
if mean_symmetry_loss is not None:
mean_symmetry_loss += symmetry_loss.item()
# -- For PPO
num_updates = self.num_learning_epochs * self.num_mini_batches
mean_value_loss /= num_updates
mean_surrogate_loss /= num_updates
mean_entropy /= num_updates
mean_amp_loss /= num_updates
mean_grad_pen_loss /= num_updates
mean_policy_pred /= num_updates
mean_expert_pred /= num_updates
# -- For RND
if mean_rnd_loss is not None:
mean_rnd_loss /= num_updates
# -- For Symmetry
if mean_symmetry_loss is not None:
mean_symmetry_loss /= num_updates
# -- Clear the storage
self.storage.clear()
# construct the loss dictionary
loss_dict = {
"value_function": mean_value_loss,
"surrogate": mean_surrogate_loss,
"entropy": mean_entropy,
"amp": mean_amp_loss,
"amp_grad_pen": mean_grad_pen_loss,
"amp_policy_pred": mean_policy_pred,
"amp_expert_pred": mean_expert_pred,
}
if self.rnd:
loss_dict["rnd"] = mean_rnd_loss
if self.symmetry:
loss_dict["symmetry"] = mean_symmetry_loss
return loss_dict
"""
Helper functions
"""
def broadcast_parameters(self):
"""Broadcast model parameters to all GPUs."""
# obtain the model parameters on current GPU
model_params = [self.policy.state_dict()]
if self.rnd:
model_params.append(self.rnd.predictor.state_dict())
# broadcast the model parameters
torch.distributed.broadcast_object_list(model_params, src=0)
# load the model parameters on all GPUs from source GPU
self.policy.load_state_dict(model_params[0])
if self.rnd:
self.rnd.predictor.load_state_dict(model_params[1])
def reduce_parameters(self):
"""Collect gradients from all GPUs and average them.
This function is called after the backward pass to synchronize the gradients across all GPUs.
"""
# Create a tensor to store the gradients
grads = [param.grad.view(-1) for param in self.policy.parameters() if param.grad is not None]
if self.rnd:
grads += [param.grad.view(-1) for param in self.rnd.parameters() if param.grad is not None]
all_grads = torch.cat(grads)
# Average the gradients across all GPUs
torch.distributed.all_reduce(all_grads, op=torch.distributed.ReduceOp.SUM)
all_grads /= self.gpu_world_size
# Get all parameters
all_params = self.policy.parameters()
if self.rnd:
all_params = chain(all_params, self.rnd.parameters())
# Update the gradients for all parameters with the reduced gradients
offset = 0
for param in all_params:
if param.grad is not None:
numel = param.numel()
# copy data back from shared buffer
param.grad.data.copy_(all_grads[offset : offset + numel].view_as(param.grad.data))
# update the offset for the next parameter
offset += numel

View File

@ -0,0 +1,185 @@
# Copyright (c) 2021-2025, ETH Zurich and NVIDIA CORPORATION
# All rights reserved.
#
# SPDX-License-Identifier: BSD-3-Clause
import torch
import torch.nn as nn
from rsl_rl.modules import StudentTeacher, StudentTeacherRecurrent
from rsl_rl.storage import RolloutStorage
from rsl_rl.utils import resolve_optimizer
class Distillation:
"""Distillation algorithm for training a student model to mimic a teacher model."""
policy: StudentTeacher | StudentTeacherRecurrent
"""The student teacher model."""
def __init__(
self,
policy,
num_learning_epochs=1,
gradient_length=15,
learning_rate=1e-3,
max_grad_norm=None,
loss_type="mse",
optimizer="adam",
device="cpu",
# Distributed training parameters
multi_gpu_cfg: dict | None = None,
):
# device-related parameters
self.device = device
self.is_multi_gpu = multi_gpu_cfg is not None
# Multi-GPU parameters
if multi_gpu_cfg is not None:
self.gpu_global_rank = multi_gpu_cfg["global_rank"]
self.gpu_world_size = multi_gpu_cfg["world_size"]
else:
self.gpu_global_rank = 0
self.gpu_world_size = 1
# distillation components
self.policy = policy
self.policy.to(self.device)
self.storage = None # initialized later
# initialize the optimizer
self.optimizer = resolve_optimizer(optimizer)(self.policy.parameters(), lr=learning_rate)
# initialize the transition
self.transition = RolloutStorage.Transition()
self.last_hidden_states = None
# distillation parameters
self.num_learning_epochs = num_learning_epochs
self.gradient_length = gradient_length
self.learning_rate = learning_rate
self.max_grad_norm = max_grad_norm
# initialize the loss function
loss_fn_dict = {
"mse": nn.functional.mse_loss,
"huber": nn.functional.huber_loss,
}
if loss_type in loss_fn_dict:
self.loss_fn = loss_fn_dict[loss_type]
else:
raise ValueError(f"Unknown loss type: {loss_type}. Supported types are: {list(loss_fn_dict.keys())}")
self.num_updates = 0
def init_storage(self, training_type, num_envs, num_transitions_per_env, obs, actions_shape):
# create rollout storage
self.storage = RolloutStorage(
training_type,
num_envs,
num_transitions_per_env,
obs,
actions_shape,
self.device,
)
def act(self, obs):
# compute the actions
self.transition.actions = self.policy.act(obs).detach()
self.transition.privileged_actions = self.policy.evaluate(obs).detach()
# record the observations
self.transition.observations = obs
return self.transition.actions
def process_env_step(self, obs, rewards, dones, extras):
# update the normalizers
self.policy.update_normalization(obs)
# record the rewards and dones
self.transition.rewards = rewards
self.transition.dones = dones
# record the transition
self.storage.add_transitions(self.transition)
self.transition.clear()
self.policy.reset(dones)
def update(self):
self.num_updates += 1
mean_behavior_loss = 0
loss = 0
cnt = 0
for epoch in range(self.num_learning_epochs):
self.policy.reset(hidden_states=self.last_hidden_states)
self.policy.detach_hidden_states()
for obs, _, privileged_actions, dones in self.storage.generator():
# inference the student for gradient computation
actions = self.policy.act_inference(obs)
# behavior cloning loss
behavior_loss = self.loss_fn(actions, privileged_actions)
# total loss
loss = loss + behavior_loss
mean_behavior_loss += behavior_loss.item()
cnt += 1
# gradient step
if cnt % self.gradient_length == 0:
self.optimizer.zero_grad()
loss.backward()
if self.is_multi_gpu:
self.reduce_parameters()
if self.max_grad_norm:
nn.utils.clip_grad_norm_(self.policy.student.parameters(), self.max_grad_norm)
self.optimizer.step()
self.policy.detach_hidden_states()
loss = 0
# reset dones
self.policy.reset(dones.view(-1))
self.policy.detach_hidden_states(dones.view(-1))
mean_behavior_loss /= cnt
self.storage.clear()
self.last_hidden_states = self.policy.get_hidden_states()
self.policy.detach_hidden_states()
# construct the loss dictionary
loss_dict = {"behavior": mean_behavior_loss}
return loss_dict
"""
Helper functions
"""
def broadcast_parameters(self):
"""Broadcast model parameters to all GPUs."""
# obtain the model parameters on current GPU
model_params = [self.policy.state_dict()]
# broadcast the model parameters
torch.distributed.broadcast_object_list(model_params, src=0)
# load the model parameters on all GPUs from source GPU
self.policy.load_state_dict(model_params[0])
def reduce_parameters(self):
"""Collect gradients from all GPUs and average them.
This function is called after the backward pass to synchronize the gradients across all GPUs.
"""
# Create a tensor to store the gradients
grads = [param.grad.view(-1) for param in self.policy.parameters() if param.grad is not None]
all_grads = torch.cat(grads)
# Average the gradients across all GPUs
torch.distributed.all_reduce(all_grads, op=torch.distributed.ReduceOp.SUM)
all_grads /= self.gpu_world_size
# Update the gradients for all parameters with the reduced gradients
offset = 0
for param in self.policy.parameters():
if param.grad is not None:
numel = param.numel()
# copy data back from shared buffer
param.grad.data.copy_(all_grads[offset : offset + numel].view_as(param.grad.data))
# update the offset for the next parameter
offset += numel

View File

@ -0,0 +1,469 @@
# Copyright (c) 2021-2025, ETH Zurich and NVIDIA CORPORATION
# All rights reserved.
#
# SPDX-License-Identifier: BSD-3-Clause
from __future__ import annotations
import torch
import torch.nn as nn
import torch.optim as optim
from itertools import chain
from rsl_rl.modules import ActorCritic
from rsl_rl.modules.rnd import RandomNetworkDistillation
from rsl_rl.storage import RolloutStorage
from rsl_rl.utils import string_to_callable
class PPO:
"""Proximal Policy Optimization algorithm (https://arxiv.org/abs/1707.06347)."""
policy: ActorCritic
"""The actor critic module."""
def __init__(
self,
policy,
num_learning_epochs=5,
num_mini_batches=4,
clip_param=0.2,
gamma=0.99,
lam=0.95,
value_loss_coef=1.0,
entropy_coef=0.01,
learning_rate=0.001,
max_grad_norm=1.0,
use_clipped_value_loss=True,
schedule="adaptive",
desired_kl=0.01,
device="cpu",
normalize_advantage_per_mini_batch=False,
# RND parameters
rnd_cfg: dict | None = None,
# Symmetry parameters
symmetry_cfg: dict | None = None,
# Distributed training parameters
multi_gpu_cfg: dict | None = None,
):
# device-related parameters
self.device = device
self.is_multi_gpu = multi_gpu_cfg is not None
# Multi-GPU parameters
if multi_gpu_cfg is not None:
self.gpu_global_rank = multi_gpu_cfg["global_rank"]
self.gpu_world_size = multi_gpu_cfg["world_size"]
else:
self.gpu_global_rank = 0
self.gpu_world_size = 1
# RND components
if rnd_cfg is not None:
# Extract parameters used in ppo
rnd_lr = rnd_cfg.pop("learning_rate", 1e-3)
# Create RND module
self.rnd = RandomNetworkDistillation(device=self.device, **rnd_cfg)
# Create RND optimizer
params = self.rnd.predictor.parameters()
self.rnd_optimizer = optim.Adam(params, lr=rnd_lr)
else:
self.rnd = None
self.rnd_optimizer = None
# Symmetry components
if symmetry_cfg is not None:
# Check if symmetry is enabled
use_symmetry = symmetry_cfg["use_data_augmentation"] or symmetry_cfg["use_mirror_loss"]
# Print that we are not using symmetry
if not use_symmetry:
print("Symmetry not used for learning. We will use it for logging instead.")
# If function is a string then resolve it to a function
if isinstance(symmetry_cfg["data_augmentation_func"], str):
symmetry_cfg["data_augmentation_func"] = string_to_callable(symmetry_cfg["data_augmentation_func"])
# Check valid configuration
if symmetry_cfg["use_data_augmentation"] and not callable(symmetry_cfg["data_augmentation_func"]):
raise ValueError(
"Data augmentation enabled but the function is not callable:"
f" {symmetry_cfg['data_augmentation_func']}"
)
# Store symmetry configuration
self.symmetry = symmetry_cfg
else:
self.symmetry = None
# PPO components
self.policy = policy
self.policy.to(self.device)
# Create optimizer
self.optimizer = optim.Adam(self.policy.parameters(), lr=learning_rate)
# Create rollout storage
self.storage: RolloutStorage = None # type: ignore
self.transition = RolloutStorage.Transition()
# PPO parameters
self.clip_param = clip_param
self.num_learning_epochs = num_learning_epochs
self.num_mini_batches = num_mini_batches
self.value_loss_coef = value_loss_coef
self.entropy_coef = entropy_coef
self.gamma = gamma
self.lam = lam
self.max_grad_norm = max_grad_norm
self.use_clipped_value_loss = use_clipped_value_loss
self.desired_kl = desired_kl
self.schedule = schedule
self.learning_rate = learning_rate
self.normalize_advantage_per_mini_batch = normalize_advantage_per_mini_batch
def init_storage(self, training_type, num_envs, num_transitions_per_env, obs, actions_shape):
# create rollout storage
self.storage = RolloutStorage(
training_type,
num_envs,
num_transitions_per_env,
obs,
actions_shape,
self.device,
)
def act(self, obs):
if self.policy.is_recurrent:
self.transition.hidden_states = self.policy.get_hidden_states()
# compute the actions and values
self.transition.actions = self.policy.act(obs).detach()
self.transition.values = self.policy.evaluate(obs).detach()
self.transition.actions_log_prob = self.policy.get_actions_log_prob(self.transition.actions).detach()
self.transition.action_mean = self.policy.action_mean.detach()
self.transition.action_sigma = self.policy.action_std.detach()
# need to record obs before env.step()
self.transition.observations = obs
return self.transition.actions
def process_env_step(self, obs, rewards, dones, extras):
# update the normalizers
self.policy.update_normalization(obs)
if self.rnd:
self.rnd.update_normalization(obs)
# Record the rewards and dones
# Note: we clone here because later on we bootstrap the rewards based on timeouts
self.transition.rewards = rewards.clone()
self.transition.dones = dones
# Compute the intrinsic rewards and add to extrinsic rewards
if self.rnd:
# Compute the intrinsic rewards
self.intrinsic_rewards = self.rnd.get_intrinsic_reward(obs)
# Add intrinsic rewards to extrinsic rewards
self.transition.rewards += self.intrinsic_rewards
# Bootstrapping on time outs
if "time_outs" in extras:
self.transition.rewards += self.gamma * torch.squeeze(
self.transition.values * extras["time_outs"].unsqueeze(1).to(self.device), 1
)
# record the transition
self.storage.add_transitions(self.transition)
self.transition.clear()
self.policy.reset(dones)
def compute_returns(self, obs):
# compute value for the last step
last_values = self.policy.evaluate(obs).detach()
self.storage.compute_returns(
last_values, self.gamma, self.lam, normalize_advantage=not self.normalize_advantage_per_mini_batch
)
def update(self): # noqa: C901
mean_value_loss = 0
mean_surrogate_loss = 0
mean_entropy = 0
# -- RND loss
if self.rnd:
mean_rnd_loss = 0
else:
mean_rnd_loss = None
# -- Symmetry loss
if self.symmetry:
mean_symmetry_loss = 0
else:
mean_symmetry_loss = None
# generator for mini batches
if self.policy.is_recurrent:
generator = self.storage.recurrent_mini_batch_generator(self.num_mini_batches, self.num_learning_epochs)
else:
generator = self.storage.mini_batch_generator(self.num_mini_batches, self.num_learning_epochs)
# iterate over batches
for (
obs_batch,
actions_batch,
target_values_batch,
advantages_batch,
returns_batch,
old_actions_log_prob_batch,
old_mu_batch,
old_sigma_batch,
hid_states_batch,
masks_batch,
) in generator:
# number of augmentations per sample
# we start with 1 and increase it if we use symmetry augmentation
num_aug = 1
# original batch size
# we assume policy group is always there and needs augmentation
original_batch_size = obs_batch.batch_size[0]
# check if we should normalize advantages per mini batch
if self.normalize_advantage_per_mini_batch:
with torch.no_grad():
advantages_batch = (advantages_batch - advantages_batch.mean()) / (advantages_batch.std() + 1e-8)
# Perform symmetric augmentation
if self.symmetry and self.symmetry["use_data_augmentation"]:
# augmentation using symmetry
data_augmentation_func = self.symmetry["data_augmentation_func"]
# returned shape: [batch_size * num_aug, ...]
obs_batch, actions_batch = data_augmentation_func(
obs=obs_batch,
actions=actions_batch,
env=self.symmetry["_env"],
)
# compute number of augmentations per sample
# we assume policy group is always there and needs augmentation
num_aug = int(obs_batch.batch_size[0] / original_batch_size)
# repeat the rest of the batch
# -- actor
old_actions_log_prob_batch = old_actions_log_prob_batch.repeat(num_aug, 1)
# -- critic
target_values_batch = target_values_batch.repeat(num_aug, 1)
advantages_batch = advantages_batch.repeat(num_aug, 1)
returns_batch = returns_batch.repeat(num_aug, 1)
# Recompute actions log prob and entropy for current batch of transitions
# Note: we need to do this because we updated the policy with the new parameters
# -- actor
self.policy.act(obs_batch, masks=masks_batch, hidden_states=hid_states_batch[0])
actions_log_prob_batch = self.policy.get_actions_log_prob(actions_batch)
# -- critic
value_batch = self.policy.evaluate(obs_batch, masks=masks_batch, hidden_states=hid_states_batch[1])
# -- entropy
# we only keep the entropy of the first augmentation (the original one)
mu_batch = self.policy.action_mean[:original_batch_size]
sigma_batch = self.policy.action_std[:original_batch_size]
entropy_batch = self.policy.entropy[:original_batch_size]
# KL
if self.desired_kl is not None and self.schedule == "adaptive":
with torch.inference_mode():
kl = torch.sum(
torch.log(sigma_batch / old_sigma_batch + 1.0e-5)
+ (torch.square(old_sigma_batch) + torch.square(old_mu_batch - mu_batch))
/ (2.0 * torch.square(sigma_batch))
- 0.5,
axis=-1,
)
kl_mean = torch.mean(kl)
# Reduce the KL divergence across all GPUs
if self.is_multi_gpu:
torch.distributed.all_reduce(kl_mean, op=torch.distributed.ReduceOp.SUM)
kl_mean /= self.gpu_world_size
# Update the learning rate
# Perform this adaptation only on the main process
# TODO: Is this needed? If KL-divergence is the "same" across all GPUs,
# then the learning rate should be the same across all GPUs.
if self.gpu_global_rank == 0:
if kl_mean > self.desired_kl * 2.0:
self.learning_rate = max(1e-5, self.learning_rate / 1.5)
elif kl_mean < self.desired_kl / 2.0 and kl_mean > 0.0:
self.learning_rate = min(1e-2, self.learning_rate * 1.5)
# Update the learning rate for all GPUs
if self.is_multi_gpu:
lr_tensor = torch.tensor(self.learning_rate, device=self.device)
torch.distributed.broadcast(lr_tensor, src=0)
self.learning_rate = lr_tensor.item()
# Update the learning rate for all parameter groups
for param_group in self.optimizer.param_groups:
param_group["lr"] = self.learning_rate
# Surrogate loss
ratio = torch.exp(actions_log_prob_batch - torch.squeeze(old_actions_log_prob_batch))
surrogate = -torch.squeeze(advantages_batch) * ratio
surrogate_clipped = -torch.squeeze(advantages_batch) * torch.clamp(
ratio, 1.0 - self.clip_param, 1.0 + self.clip_param
)
surrogate_loss = torch.max(surrogate, surrogate_clipped).mean()
# Value function loss
if self.use_clipped_value_loss:
value_clipped = target_values_batch + (value_batch - target_values_batch).clamp(
-self.clip_param, self.clip_param
)
value_losses = (value_batch - returns_batch).pow(2)
value_losses_clipped = (value_clipped - returns_batch).pow(2)
value_loss = torch.max(value_losses, value_losses_clipped).mean()
else:
value_loss = (returns_batch - value_batch).pow(2).mean()
loss = surrogate_loss + self.value_loss_coef * value_loss - self.entropy_coef * entropy_batch.mean()
# Symmetry loss
if self.symmetry:
# obtain the symmetric actions
# if we did augmentation before then we don't need to augment again
if not self.symmetry["use_data_augmentation"]:
data_augmentation_func = self.symmetry["data_augmentation_func"]
obs_batch, _ = data_augmentation_func(obs=obs_batch, actions=None, env=self.symmetry["_env"])
# compute number of augmentations per sample
num_aug = int(obs_batch.shape[0] / original_batch_size)
# actions predicted by the actor for symmetrically-augmented observations
mean_actions_batch = self.policy.act_inference(obs_batch.detach().clone())
# compute the symmetrically augmented actions
# note: we are assuming the first augmentation is the original one.
# We do not use the action_batch from earlier since that action was sampled from the distribution.
# However, the symmetry loss is computed using the mean of the distribution.
action_mean_orig = mean_actions_batch[:original_batch_size]
_, actions_mean_symm_batch = data_augmentation_func(
obs=None, actions=action_mean_orig, env=self.symmetry["_env"]
)
# compute the loss (we skip the first augmentation as it is the original one)
mse_loss = torch.nn.MSELoss()
symmetry_loss = mse_loss(
mean_actions_batch[original_batch_size:], actions_mean_symm_batch.detach()[original_batch_size:]
)
# add the loss to the total loss
if self.symmetry["use_mirror_loss"]:
loss += self.symmetry["mirror_loss_coeff"] * symmetry_loss
else:
symmetry_loss = symmetry_loss.detach()
# Random Network Distillation loss
# TODO: Move this processing to inside RND module.
if self.rnd:
# extract the rnd_state
# TODO: Check if we still need torch no grad. It is just an affine transformation.
with torch.no_grad():
rnd_state_batch = self.rnd.get_rnd_state(obs_batch[:original_batch_size])
rnd_state_batch = self.rnd.state_normalizer(rnd_state_batch)
# predict the embedding and the target
predicted_embedding = self.rnd.predictor(rnd_state_batch)
target_embedding = self.rnd.target(rnd_state_batch).detach()
# compute the loss as the mean squared error
mseloss = torch.nn.MSELoss()
rnd_loss = mseloss(predicted_embedding, target_embedding)
# Compute the gradients
# -- For PPO
self.optimizer.zero_grad()
loss.backward()
# -- For RND
if self.rnd:
self.rnd_optimizer.zero_grad() # type: ignore
rnd_loss.backward()
# Collect gradients from all GPUs
if self.is_multi_gpu:
self.reduce_parameters()
# Apply the gradients
# -- For PPO
nn.utils.clip_grad_norm_(self.policy.parameters(), self.max_grad_norm)
self.optimizer.step()
# -- For RND
if self.rnd_optimizer:
self.rnd_optimizer.step()
# Store the losses
mean_value_loss += value_loss.item()
mean_surrogate_loss += surrogate_loss.item()
mean_entropy += entropy_batch.mean().item()
# -- RND loss
if mean_rnd_loss is not None:
mean_rnd_loss += rnd_loss.item()
# -- Symmetry loss
if mean_symmetry_loss is not None:
mean_symmetry_loss += symmetry_loss.item()
# -- For PPO
num_updates = self.num_learning_epochs * self.num_mini_batches
mean_value_loss /= num_updates
mean_surrogate_loss /= num_updates
mean_entropy /= num_updates
# -- For RND
if mean_rnd_loss is not None:
mean_rnd_loss /= num_updates
# -- For Symmetry
if mean_symmetry_loss is not None:
mean_symmetry_loss /= num_updates
# -- Clear the storage
self.storage.clear()
# construct the loss dictionary
loss_dict = {
"value_function": mean_value_loss,
"surrogate": mean_surrogate_loss,
"entropy": mean_entropy,
}
if self.rnd:
loss_dict["rnd"] = mean_rnd_loss
if self.symmetry:
loss_dict["symmetry"] = mean_symmetry_loss
return loss_dict
"""
Helper functions
"""
def broadcast_parameters(self):
"""Broadcast model parameters to all GPUs."""
# obtain the model parameters on current GPU
model_params = [self.policy.state_dict()]
if self.rnd:
model_params.append(self.rnd.predictor.state_dict())
# broadcast the model parameters
torch.distributed.broadcast_object_list(model_params, src=0)
# load the model parameters on all GPUs from source GPU
self.policy.load_state_dict(model_params[0])
if self.rnd:
self.rnd.predictor.load_state_dict(model_params[1])
def reduce_parameters(self):
"""Collect gradients from all GPUs and average them.
This function is called after the backward pass to synchronize the gradients across all GPUs.
"""
# Create a tensor to store the gradients
grads = [param.grad.view(-1) for param in self.policy.parameters() if param.grad is not None]
if self.rnd:
grads += [param.grad.view(-1) for param in self.rnd.parameters() if param.grad is not None]
all_grads = torch.cat(grads)
# Average the gradients across all GPUs
torch.distributed.all_reduce(all_grads, op=torch.distributed.ReduceOp.SUM)
all_grads /= self.gpu_world_size
# Get all parameters
all_params = self.policy.parameters()
if self.rnd:
all_params = chain(all_params, self.rnd.parameters())
# Update the gradients for all parameters with the reduced gradients
offset = 0
for param in all_params:
if param.grad is not None:
numel = param.numel()
# copy data back from shared buffer
param.grad.data.copy_(all_grads[offset : offset + numel].view_as(param.grad.data))
# update the offset for the next parameter
offset += numel

View File

@ -0,0 +1,6 @@
# Copyright (c) 2021-2025, ETH Zurich and NVIDIA CORPORATION
# All rights reserved.
#
# SPDX-License-Identifier: BSD-3-Clause
"""Main module for the rsl_rl package."""

View File

@ -0,0 +1,11 @@
# Copyright (c) 2021-2025, ETH Zurich and NVIDIA CORPORATION
# All rights reserved.
#
# SPDX-License-Identifier: BSD-3-Clause
"""Implementation of different RL agents."""
from .distillation import Distillation
from .ppo import PPO
from .amp_ppo import AMP_PPO
__all__ = ["PPO", "Distillation", "AMP_PPO"]

View File

@ -0,0 +1,571 @@
# Copyright (c) 2021-2025, ETH Zurich and NVIDIA CORPORATION
# All rights reserved.
#
# SPDX-License-Identifier: BSD-3-Clause
from __future__ import annotations
from torch._tensor import Tensor
from torch._tensor import Tensor
from typing import Any
from copy import deepcopy
import torch
import torch.nn as nn
import torch.optim as optim
from itertools import chain
from rsl_rl.modules import ActorCritic
from rsl_rl.modules.rnd import RandomNetworkDistillation
from rsl_rl.storage import RolloutStorage, ReplayBufferMulti
from rsl_rl.utils import string_to_callable
class AMP_PPO:
"""Proximal Policy Optimization algorithm (https://arxiv.org/abs/1707.06347)."""
policy: ActorCritic
"""The actor critic module."""
def __init__(
self,
policy,
discriminator,
amp_data,
amp_normalizer,
amp_num_frames=1,
amp_replay_buffer_size=100000,
num_learning_epochs=5,
num_mini_batches=4,
clip_param=0.2,
gamma=0.99,
lam=0.95,
value_loss_coef=1.0,
entropy_coef=0.01,
learning_rate=0.001,
max_grad_norm=1.0,
use_clipped_value_loss=True,
schedule="adaptive",
desired_kl=0.01,
device="cpu",
normalize_advantage_per_mini_batch=False,
# RND parameters
rnd_cfg: dict | None = None,
# Symmetry parameters
symmetry_cfg: dict | None = None,
# Distributed training parameters
multi_gpu_cfg: dict | None = None,
):
# device-related parameters
self.device = device
self.is_multi_gpu = multi_gpu_cfg is not None
# Multi-GPU parameters
if multi_gpu_cfg is not None:
self.gpu_global_rank = multi_gpu_cfg["global_rank"]
self.gpu_world_size = multi_gpu_cfg["world_size"]
else:
self.gpu_global_rank = 0
self.gpu_world_size = 1
# RND components
if rnd_cfg is not None:
# Extract parameters used in ppo
rnd_lr = rnd_cfg.pop("learning_rate", 1e-3)
# Create RND module
self.rnd = RandomNetworkDistillation(device=self.device, **rnd_cfg)
# Create RND optimizer
params = self.rnd.predictor.parameters()
self.rnd_optimizer = optim.Adam(params, lr=rnd_lr)
else:
self.rnd = None
self.rnd_optimizer = None
# Symmetry components
if symmetry_cfg is not None:
# Check if symmetry is enabled
use_symmetry = symmetry_cfg["use_data_augmentation"] or symmetry_cfg["use_mirror_loss"]
# Print that we are not using symmetry
if not use_symmetry:
print("Symmetry not used for learning. We will use it for logging instead.")
# If function is a string then resolve it to a function
if isinstance(symmetry_cfg["data_augmentation_func"], str):
symmetry_cfg["data_augmentation_func"] = string_to_callable(symmetry_cfg["data_augmentation_func"])
# Check valid configuration
if symmetry_cfg["use_data_augmentation"] and not callable(symmetry_cfg["data_augmentation_func"]):
raise ValueError(
"Data augmentation enabled but the function is not callable:"
f" {symmetry_cfg['data_augmentation_func']}"
)
# Store symmetry configuration
self.symmetry = symmetry_cfg
else:
self.symmetry = None
## AMP components
self.discriminator = discriminator
self.discriminator.to(self.device)
self.amp_storage = ReplayBufferMulti(discriminator.state_dim, amp_replay_buffer_size, amp_num_frames, device)
self.amp_data = amp_data
self.amp_normalizer = amp_normalizer
# PPO components
self.policy = policy
self.policy.to(self.device)
# Create rollout storage
self.storage: RolloutStorage = None # type: ignore
self.transition = RolloutStorage.Transition()
self.amp_transition = RolloutStorage.Transition()
params = [
{'params': self.policy.parameters(), 'name': 'policy'},
]
params.append({
'params': self.discriminator.trunk.parameters(),
'weight_decay': 10e-4,
'name': f'amp_trunk'
})
params.append({
'params': self.discriminator.amp_linear.parameters(),
'weight_decay': 10e-2,
'name': f'amp_head'
})
# Create optimizer
self.optimizer = optim.Adam(params, lr=learning_rate)
# PPO parameters
self.clip_param = clip_param
self.num_learning_epochs = num_learning_epochs
self.num_mini_batches = num_mini_batches
self.value_loss_coef = value_loss_coef
self.entropy_coef = entropy_coef
self.gamma = gamma
self.lam = lam
self.max_grad_norm = max_grad_norm
self.use_clipped_value_loss = use_clipped_value_loss
self.desired_kl = desired_kl
self.schedule = schedule
self.learning_rate = learning_rate
self.normalize_advantage_per_mini_batch = normalize_advantage_per_mini_batch
def init_storage(self, training_type, num_envs, num_transitions_per_env, obs, actions_shape):
# create rollout storage
self.storage = RolloutStorage(
training_type,
num_envs,
num_transitions_per_env,
obs,
actions_shape,
self.device,
)
def act(self, obs, amp_obs):
if self.policy.is_recurrent:
self.transition.hidden_states = self.policy.get_hidden_states()
# compute the actions and values
self.transition.actions = self.policy.act(obs).detach()
self.transition.values = self.policy.evaluate(obs).detach()
self.transition.actions_log_prob = self.policy.get_actions_log_prob(self.transition.actions).detach()
self.transition.action_mean = self.policy.action_mean.detach()
self.transition.action_sigma = self.policy.action_std.detach()
# need to record obs before env.step()
self.transition.observations = obs
self.amp_transition.observations = amp_obs
return self.transition.actions
def process_env_step(self, obs, rewards, dones, extras,amp_obs, amp_obs_frames=None):
# update the normalizers
self.policy.update_normalization(obs)
if self.rnd:
self.rnd.update_normalization(obs)
# Record the rewards and dones
# Note: we clone here because later on we bootstrap the rewards based on timeouts
self.transition.rewards = rewards.clone()
self.transition.dones = dones
# Compute the intrinsic rewards and add to extrinsic rewards
if self.rnd:
# Compute the intrinsic rewards
self.intrinsic_rewards = self.rnd.get_intrinsic_reward(obs)
# Add intrinsic rewards to extrinsic rewards
self.transition.rewards += self.intrinsic_rewards
# Bootstrapping on time outs
if "time_outs" in extras:
self.transition.rewards += self.gamma * torch.squeeze(
self.transition.values * extras["time_outs"].unsqueeze(1).to(self.device), 1
)
if amp_obs_frames is not None:
self.amp_storage.insert(amp_obs_frames)
else:
self.amp_storage.insert(self.amp_transition.observations, amp_obs)
# record the transition
self.storage.add_transitions(self.transition)
self.transition.clear()
self.amp_transition.clear()
self.policy.reset(dones)
def compute_returns(self, obs):
# compute value for the last step
last_values = self.policy.evaluate(obs).detach()
self.storage.compute_returns(
last_values, self.gamma, self.lam, normalize_advantage=not self.normalize_advantage_per_mini_batch
)
def update(self): # noqa: C901
mean_value_loss = 0
mean_surrogate_loss = 0
mean_entropy = 0
mean_amp_loss = 0
mean_grad_pen_loss = 0
mean_policy_pred = 0
mean_expert_pred = 0
# -- RND loss
if self.rnd:
mean_rnd_loss = 0
else:
mean_rnd_loss = None
# -- Symmetry loss
if self.symmetry:
mean_symmetry_loss = 0
else:
mean_symmetry_loss = None
# generator for mini batches
if self.policy.is_recurrent:
generator = self.storage.recurrent_mini_batch_generator(self.num_mini_batches, self.num_learning_epochs)
else:
generator = self.storage.mini_batch_generator(self.num_mini_batches, self.num_learning_epochs)
amp_policy_generator = self.amp_storage.feed_forward_generator(
self.num_learning_epochs * self.num_mini_batches,
self.storage.num_envs * self.storage.num_transitions_per_env // self.num_mini_batches,
)
amp_expert_generator = self.amp_data.feed_forward_generator_23dof_multi(
self.num_learning_epochs * self.num_mini_batches,
self.storage.num_envs * self.storage.num_transitions_per_env // self.num_mini_batches,
)
# iterate over batches
for sample, sample_amp_policy, sample_amp_expert in zip(generator, amp_policy_generator, amp_expert_generator):
(
obs_batch,
actions_batch,
target_values_batch,
advantages_batch,
returns_batch,
old_actions_log_prob_batch,
old_mu_batch,
old_sigma_batch,
hid_states_batch,
masks_batch,
) = sample
# number of augmentations per sample
# we start with 1 and increase it if we use symmetry augmentation
num_aug = 1
# original batch size
# we assume policy group is always there and needs augmentation
original_batch_size = obs_batch.batch_size[0]
# check if we should normalize advantages per mini batch
if self.normalize_advantage_per_mini_batch:
with torch.no_grad():
advantages_batch = (advantages_batch - advantages_batch.mean()) / (advantages_batch.std() + 1e-8)
# Perform symmetric augmentation
if self.symmetry and self.symmetry["use_data_augmentation"]:
# augmentation using symmetry
data_augmentation_func = self.symmetry["data_augmentation_func"]
# returned shape: [batch_size * num_aug, ...]
obs_batch, actions_batch = data_augmentation_func(
obs=obs_batch,
actions=actions_batch,
env=self.symmetry["_env"],
)
# compute number of augmentations per sample
# we assume policy group is always there and needs augmentation
num_aug = int(obs_batch.batch_size[0] / original_batch_size)
# repeat the rest of the batch
# -- actor
old_actions_log_prob_batch = old_actions_log_prob_batch.repeat(num_aug, 1)
# -- critic
target_values_batch = target_values_batch.repeat(num_aug, 1)
advantages_batch = advantages_batch.repeat(num_aug, 1)
returns_batch = returns_batch.repeat(num_aug, 1)
# Recompute actions log prob and entropy for current batch of transitions
# Note: we need to do this because we updated the policy with the new parameters
# -- actor
self.policy.act(obs_batch, masks=masks_batch, hidden_states=hid_states_batch[0])
actions_log_prob_batch = self.policy.get_actions_log_prob(actions_batch)
# -- critic
value_batch = self.policy.evaluate(obs_batch, masks=masks_batch, hidden_states=hid_states_batch[1])
# -- entropy
# we only keep the entropy of the first augmentation (the original one)
mu_batch = self.policy.action_mean[:original_batch_size]
sigma_batch = self.policy.action_std[:original_batch_size]
entropy_batch = self.policy.entropy[:original_batch_size]
# KL
if self.desired_kl is not None and self.schedule == "adaptive":
with torch.inference_mode():
kl = torch.sum(
torch.log(sigma_batch / old_sigma_batch + 1.0e-5)
+ (torch.square(old_sigma_batch) + torch.square(old_mu_batch - mu_batch))
/ (2.0 * torch.square(sigma_batch))
- 0.5,
axis=-1,
)
kl_mean = torch.mean(kl)
# Reduce the KL divergence across all GPUs
if self.is_multi_gpu:
torch.distributed.all_reduce(kl_mean, op=torch.distributed.ReduceOp.SUM)
kl_mean /= self.gpu_world_size
# Update the learning rate
# Perform this adaptation only on the main process
# TODO: Is this needed? If KL-divergence is the "same" across all GPUs,
# then the learning rate should be the same across all GPUs.
if self.gpu_global_rank == 0:
if kl_mean > self.desired_kl * 2.0:
self.learning_rate = max(1e-5, self.learning_rate / 1.5)
elif kl_mean < self.desired_kl / 2.0 and kl_mean > 0.0:
self.learning_rate = min(1e-2, self.learning_rate * 1.5)
# Update the learning rate for all GPUs
if self.is_multi_gpu:
lr_tensor = torch.tensor(self.learning_rate, device=self.device)
torch.distributed.broadcast(lr_tensor, src=0)
self.learning_rate = lr_tensor.item()
# Update the learning rate for all parameter groups
for param_group in self.optimizer.param_groups:
param_group["lr"] = self.learning_rate
# Surrogate loss
ratio = torch.exp(actions_log_prob_batch - torch.squeeze(old_actions_log_prob_batch))
surrogate = -torch.squeeze(advantages_batch) * ratio
surrogate_clipped = -torch.squeeze(advantages_batch) * torch.clamp(
ratio, 1.0 - self.clip_param, 1.0 + self.clip_param
)
surrogate_loss = torch.max(surrogate, surrogate_clipped).mean()
# Value function loss
if self.use_clipped_value_loss:
value_clipped = target_values_batch + (value_batch - target_values_batch).clamp(
-self.clip_param, self.clip_param
)
value_losses = (value_batch - returns_batch).pow(2)
value_losses_clipped = (value_clipped - returns_batch).pow(2)
value_loss = torch.max(value_losses, value_losses_clipped).mean()
else:
value_loss = (returns_batch - value_batch).pow(2).mean()
loss = surrogate_loss + self.value_loss_coef * value_loss - self.entropy_coef * entropy_batch.mean()
# Symmetry loss
if self.symmetry:
# obtain the symmetric actions
# if we did augmentation before then we don't need to augment again
if not self.symmetry["use_data_augmentation"]:
data_augmentation_func = self.symmetry["data_augmentation_func"]
obs_batch, _ = data_augmentation_func(obs=obs_batch, actions=None, env=self.symmetry["_env"])
# compute number of augmentations per sample
num_aug = int(obs_batch.shape[0] / original_batch_size)
# actions predicted by the actor for symmetrically-augmented observations
mean_actions_batch = self.policy.act_inference(obs_batch.detach().clone())
# compute the symmetrically augmented actions
# note: we are assuming the first augmentation is the original one.
# We do not use the action_batch from earlier since that action was sampled from the distribution.
# However, the symmetry loss is computed using the mean of the distribution.
action_mean_orig = mean_actions_batch[:original_batch_size]
_, actions_mean_symm_batch = data_augmentation_func(
obs=None, actions=action_mean_orig, env=self.symmetry["_env"]
)
# compute the loss (we skip the first augmentation as it is the original one)
mse_loss = torch.nn.MSELoss()
symmetry_loss = mse_loss(
mean_actions_batch[original_batch_size:], actions_mean_symm_batch.detach()[original_batch_size:]
)
# add the loss to the total loss
if self.symmetry["use_mirror_loss"]:
loss += self.symmetry["mirror_loss_coeff"] * symmetry_loss
else:
symmetry_loss = symmetry_loss.detach()
# Random Network Distillation loss
# TODO: Move this processing to inside RND module.
if self.rnd:
# extract the rnd_state
# TODO: Check if we still need torch no grad. It is just an affine transformation.
with torch.no_grad():
rnd_state_batch = self.rnd.get_rnd_state(obs_batch[:original_batch_size])
rnd_state_batch = self.rnd.state_normalizer(rnd_state_batch)
# predict the embedding and the target
predicted_embedding = self.rnd.predictor(rnd_state_batch)
target_embedding = self.rnd.target(rnd_state_batch).detach()
# compute the loss as the mean squared error
mseloss = torch.nn.MSELoss()
rnd_loss = mseloss(predicted_embedding, target_embedding)
expert_states = sample_amp_expert
policy_states = sample_amp_policy
with torch.no_grad():
expert_states = self.amp_normalizer.normalize_torch(expert_states.to(self.device), self.device)
policy_states = self.amp_normalizer.normalize_torch(policy_states, self.device)
contact_phase_push = obs_batch['critic'][:, -4]
mask_push = contact_phase_push == 1.
if mask_push.any():
policy_d = self.discriminator(policy_states.flatten(1))
expert_states = expert_states.to(self.device)
expert_d = self.discriminator(expert_states.flatten(1))
expert_loss = torch.nn.MSELoss()(expert_d, torch.ones(expert_d.size(), device=self.device))
policy_loss = torch.nn.MSELoss()(policy_d, -1 * torch.ones(policy_d.size(), device=self.device))
amp_loss = 0.5 * (expert_loss + policy_loss)
# grad penalty
grad_pen_loss = self.discriminator.compute_grad_pen(expert_states, lambda_=5)
else:
amp_loss = torch.tensor(0.0, device=self.device)
grad_pen_loss = torch.tensor(0.0, device=self.device)
expert_loss = torch.tensor(0.0, device=self.device)
policy_loss = torch.tensor(0.0, device=self.device)
loss += (amp_loss + grad_pen_loss)
self.amp_normalizer.update(policy_states.cpu().numpy())
self.amp_normalizer.update(expert_states.cpu().numpy())
# Compute the gradients
# -- For PPO
self.optimizer.zero_grad()
loss.backward()
# -- For RND
if self.rnd:
self.rnd_optimizer.zero_grad() # type: ignore
rnd_loss.backward()
# Collect gradients from all GPUs
if self.is_multi_gpu:
self.reduce_parameters()
# Apply the gradients
# -- For PPO
nn.utils.clip_grad_norm_(self.policy.parameters(), self.max_grad_norm)
self.optimizer.step()
# -- For RND
if self.rnd_optimizer:
self.rnd_optimizer.step()
# Store the losses
mean_value_loss += value_loss.item()
mean_surrogate_loss += surrogate_loss.item()
mean_entropy += entropy_batch.mean().item()
mean_amp_loss += amp_loss.item()
mean_grad_pen_loss += grad_pen_loss.item()
mean_policy_pred += policy_loss.mean().item()
mean_expert_pred += expert_loss.mean().item()
# -- RND loss
if mean_rnd_loss is not None:
mean_rnd_loss += rnd_loss.item()
# -- Symmetry loss
if mean_symmetry_loss is not None:
mean_symmetry_loss += symmetry_loss.item()
# -- For PPO
num_updates = self.num_learning_epochs * self.num_mini_batches
mean_value_loss /= num_updates
mean_surrogate_loss /= num_updates
mean_entropy /= num_updates
mean_amp_loss /= num_updates
mean_grad_pen_loss /= num_updates
mean_policy_pred /= num_updates
mean_expert_pred /= num_updates
# -- For RND
if mean_rnd_loss is not None:
mean_rnd_loss /= num_updates
# -- For Symmetry
if mean_symmetry_loss is not None:
mean_symmetry_loss /= num_updates
# -- Clear the storage
self.storage.clear()
# construct the loss dictionary
loss_dict = {
"value_function": mean_value_loss,
"surrogate": mean_surrogate_loss,
"entropy": mean_entropy,
"amp": mean_amp_loss,
"amp_grad_pen": mean_grad_pen_loss,
"amp_policy_pred": mean_policy_pred,
"amp_expert_pred": mean_expert_pred,
}
if self.rnd:
loss_dict["rnd"] = mean_rnd_loss
if self.symmetry:
loss_dict["symmetry"] = mean_symmetry_loss
return loss_dict
"""
Helper functions
"""
def broadcast_parameters(self):
"""Broadcast model parameters to all GPUs."""
# obtain the model parameters on current GPU
model_params = [self.policy.state_dict()]
if self.rnd:
model_params.append(self.rnd.predictor.state_dict())
# broadcast the model parameters
torch.distributed.broadcast_object_list(model_params, src=0)
# load the model parameters on all GPUs from source GPU
self.policy.load_state_dict(model_params[0])
if self.rnd:
self.rnd.predictor.load_state_dict(model_params[1])
def reduce_parameters(self):
"""Collect gradients from all GPUs and average them.
This function is called after the backward pass to synchronize the gradients across all GPUs.
"""
# Create a tensor to store the gradients
grads = [param.grad.view(-1) for param in self.policy.parameters() if param.grad is not None]
if self.rnd:
grads += [param.grad.view(-1) for param in self.rnd.parameters() if param.grad is not None]
all_grads = torch.cat(grads)
# Average the gradients across all GPUs
torch.distributed.all_reduce(all_grads, op=torch.distributed.ReduceOp.SUM)
all_grads /= self.gpu_world_size
# Get all parameters
all_params = self.policy.parameters()
if self.rnd:
all_params = chain(all_params, self.rnd.parameters())
# Update the gradients for all parameters with the reduced gradients
offset = 0
for param in all_params:
if param.grad is not None:
numel = param.numel()
# copy data back from shared buffer
param.grad.data.copy_(all_grads[offset : offset + numel].view_as(param.grad.data))
# update the offset for the next parameter
offset += numel

View File

@ -0,0 +1,185 @@
# Copyright (c) 2021-2025, ETH Zurich and NVIDIA CORPORATION
# All rights reserved.
#
# SPDX-License-Identifier: BSD-3-Clause
import torch
import torch.nn as nn
from rsl_rl.modules import StudentTeacher, StudentTeacherRecurrent
from rsl_rl.storage import RolloutStorage
from rsl_rl.utils import resolve_optimizer
class Distillation:
"""Distillation algorithm for training a student model to mimic a teacher model."""
policy: StudentTeacher | StudentTeacherRecurrent
"""The student teacher model."""
def __init__(
self,
policy,
num_learning_epochs=1,
gradient_length=15,
learning_rate=1e-3,
max_grad_norm=None,
loss_type="mse",
optimizer="adam",
device="cpu",
# Distributed training parameters
multi_gpu_cfg: dict | None = None,
):
# device-related parameters
self.device = device
self.is_multi_gpu = multi_gpu_cfg is not None
# Multi-GPU parameters
if multi_gpu_cfg is not None:
self.gpu_global_rank = multi_gpu_cfg["global_rank"]
self.gpu_world_size = multi_gpu_cfg["world_size"]
else:
self.gpu_global_rank = 0
self.gpu_world_size = 1
# distillation components
self.policy = policy
self.policy.to(self.device)
self.storage = None # initialized later
# initialize the optimizer
self.optimizer = resolve_optimizer(optimizer)(self.policy.parameters(), lr=learning_rate)
# initialize the transition
self.transition = RolloutStorage.Transition()
self.last_hidden_states = None
# distillation parameters
self.num_learning_epochs = num_learning_epochs
self.gradient_length = gradient_length
self.learning_rate = learning_rate
self.max_grad_norm = max_grad_norm
# initialize the loss function
loss_fn_dict = {
"mse": nn.functional.mse_loss,
"huber": nn.functional.huber_loss,
}
if loss_type in loss_fn_dict:
self.loss_fn = loss_fn_dict[loss_type]
else:
raise ValueError(f"Unknown loss type: {loss_type}. Supported types are: {list(loss_fn_dict.keys())}")
self.num_updates = 0
def init_storage(self, training_type, num_envs, num_transitions_per_env, obs, actions_shape):
# create rollout storage
self.storage = RolloutStorage(
training_type,
num_envs,
num_transitions_per_env,
obs,
actions_shape,
self.device,
)
def act(self, obs):
# compute the actions
self.transition.actions = self.policy.act(obs).detach()
self.transition.privileged_actions = self.policy.evaluate(obs).detach()
# record the observations
self.transition.observations = obs
return self.transition.actions
def process_env_step(self, obs, rewards, dones, extras):
# update the normalizers
self.policy.update_normalization(obs)
# record the rewards and dones
self.transition.rewards = rewards
self.transition.dones = dones
# record the transition
self.storage.add_transitions(self.transition)
self.transition.clear()
self.policy.reset(dones)
def update(self):
self.num_updates += 1
mean_behavior_loss = 0
loss = 0
cnt = 0
for epoch in range(self.num_learning_epochs):
self.policy.reset(hidden_states=self.last_hidden_states)
self.policy.detach_hidden_states()
for obs, _, privileged_actions, dones in self.storage.generator():
# inference the student for gradient computation
actions = self.policy.act_inference(obs)
# behavior cloning loss
behavior_loss = self.loss_fn(actions, privileged_actions)
# total loss
loss = loss + behavior_loss
mean_behavior_loss += behavior_loss.item()
cnt += 1
# gradient step
if cnt % self.gradient_length == 0:
self.optimizer.zero_grad()
loss.backward()
if self.is_multi_gpu:
self.reduce_parameters()
if self.max_grad_norm:
nn.utils.clip_grad_norm_(self.policy.student.parameters(), self.max_grad_norm)
self.optimizer.step()
self.policy.detach_hidden_states()
loss = 0
# reset dones
self.policy.reset(dones.view(-1))
self.policy.detach_hidden_states(dones.view(-1))
mean_behavior_loss /= cnt
self.storage.clear()
self.last_hidden_states = self.policy.get_hidden_states()
self.policy.detach_hidden_states()
# construct the loss dictionary
loss_dict = {"behavior": mean_behavior_loss}
return loss_dict
"""
Helper functions
"""
def broadcast_parameters(self):
"""Broadcast model parameters to all GPUs."""
# obtain the model parameters on current GPU
model_params = [self.policy.state_dict()]
# broadcast the model parameters
torch.distributed.broadcast_object_list(model_params, src=0)
# load the model parameters on all GPUs from source GPU
self.policy.load_state_dict(model_params[0])
def reduce_parameters(self):
"""Collect gradients from all GPUs and average them.
This function is called after the backward pass to synchronize the gradients across all GPUs.
"""
# Create a tensor to store the gradients
grads = [param.grad.view(-1) for param in self.policy.parameters() if param.grad is not None]
all_grads = torch.cat(grads)
# Average the gradients across all GPUs
torch.distributed.all_reduce(all_grads, op=torch.distributed.ReduceOp.SUM)
all_grads /= self.gpu_world_size
# Update the gradients for all parameters with the reduced gradients
offset = 0
for param in self.policy.parameters():
if param.grad is not None:
numel = param.numel()
# copy data back from shared buffer
param.grad.data.copy_(all_grads[offset : offset + numel].view_as(param.grad.data))
# update the offset for the next parameter
offset += numel

View File

@ -0,0 +1,469 @@
# Copyright (c) 2021-2025, ETH Zurich and NVIDIA CORPORATION
# All rights reserved.
#
# SPDX-License-Identifier: BSD-3-Clause
from __future__ import annotations
import torch
import torch.nn as nn
import torch.optim as optim
from itertools import chain
from rsl_rl.modules import ActorCritic
from rsl_rl.modules.rnd import RandomNetworkDistillation
from rsl_rl.storage import RolloutStorage
from rsl_rl.utils import string_to_callable
class PPO:
"""Proximal Policy Optimization algorithm (https://arxiv.org/abs/1707.06347)."""
policy: ActorCritic
"""The actor critic module."""
def __init__(
self,
policy,
num_learning_epochs=5,
num_mini_batches=4,
clip_param=0.2,
gamma=0.99,
lam=0.95,
value_loss_coef=1.0,
entropy_coef=0.01,
learning_rate=0.001,
max_grad_norm=1.0,
use_clipped_value_loss=True,
schedule="adaptive",
desired_kl=0.01,
device="cpu",
normalize_advantage_per_mini_batch=False,
# RND parameters
rnd_cfg: dict | None = None,
# Symmetry parameters
symmetry_cfg: dict | None = None,
# Distributed training parameters
multi_gpu_cfg: dict | None = None,
):
# device-related parameters
self.device = device
self.is_multi_gpu = multi_gpu_cfg is not None
# Multi-GPU parameters
if multi_gpu_cfg is not None:
self.gpu_global_rank = multi_gpu_cfg["global_rank"]
self.gpu_world_size = multi_gpu_cfg["world_size"]
else:
self.gpu_global_rank = 0
self.gpu_world_size = 1
# RND components
if rnd_cfg is not None:
# Extract parameters used in ppo
rnd_lr = rnd_cfg.pop("learning_rate", 1e-3)
# Create RND module
self.rnd = RandomNetworkDistillation(device=self.device, **rnd_cfg)
# Create RND optimizer
params = self.rnd.predictor.parameters()
self.rnd_optimizer = optim.Adam(params, lr=rnd_lr)
else:
self.rnd = None
self.rnd_optimizer = None
# Symmetry components
if symmetry_cfg is not None:
# Check if symmetry is enabled
use_symmetry = symmetry_cfg["use_data_augmentation"] or symmetry_cfg["use_mirror_loss"]
# Print that we are not using symmetry
if not use_symmetry:
print("Symmetry not used for learning. We will use it for logging instead.")
# If function is a string then resolve it to a function
if isinstance(symmetry_cfg["data_augmentation_func"], str):
symmetry_cfg["data_augmentation_func"] = string_to_callable(symmetry_cfg["data_augmentation_func"])
# Check valid configuration
if symmetry_cfg["use_data_augmentation"] and not callable(symmetry_cfg["data_augmentation_func"]):
raise ValueError(
"Data augmentation enabled but the function is not callable:"
f" {symmetry_cfg['data_augmentation_func']}"
)
# Store symmetry configuration
self.symmetry = symmetry_cfg
else:
self.symmetry = None
# PPO components
self.policy = policy
self.policy.to(self.device)
# Create optimizer
self.optimizer = optim.Adam(self.policy.parameters(), lr=learning_rate)
# Create rollout storage
self.storage: RolloutStorage = None # type: ignore
self.transition = RolloutStorage.Transition()
# PPO parameters
self.clip_param = clip_param
self.num_learning_epochs = num_learning_epochs
self.num_mini_batches = num_mini_batches
self.value_loss_coef = value_loss_coef
self.entropy_coef = entropy_coef
self.gamma = gamma
self.lam = lam
self.max_grad_norm = max_grad_norm
self.use_clipped_value_loss = use_clipped_value_loss
self.desired_kl = desired_kl
self.schedule = schedule
self.learning_rate = learning_rate
self.normalize_advantage_per_mini_batch = normalize_advantage_per_mini_batch
def init_storage(self, training_type, num_envs, num_transitions_per_env, obs, actions_shape):
# create rollout storage
self.storage = RolloutStorage(
training_type,
num_envs,
num_transitions_per_env,
obs,
actions_shape,
self.device,
)
def act(self, obs):
if self.policy.is_recurrent:
self.transition.hidden_states = self.policy.get_hidden_states()
# compute the actions and values
self.transition.actions = self.policy.act(obs).detach()
self.transition.values = self.policy.evaluate(obs).detach()
self.transition.actions_log_prob = self.policy.get_actions_log_prob(self.transition.actions).detach()
self.transition.action_mean = self.policy.action_mean.detach()
self.transition.action_sigma = self.policy.action_std.detach()
# need to record obs before env.step()
self.transition.observations = obs
return self.transition.actions
def process_env_step(self, obs, rewards, dones, extras):
# update the normalizers
self.policy.update_normalization(obs)
if self.rnd:
self.rnd.update_normalization(obs)
# Record the rewards and dones
# Note: we clone here because later on we bootstrap the rewards based on timeouts
self.transition.rewards = rewards.clone()
self.transition.dones = dones
# Compute the intrinsic rewards and add to extrinsic rewards
if self.rnd:
# Compute the intrinsic rewards
self.intrinsic_rewards = self.rnd.get_intrinsic_reward(obs)
# Add intrinsic rewards to extrinsic rewards
self.transition.rewards += self.intrinsic_rewards
# Bootstrapping on time outs
if "time_outs" in extras:
self.transition.rewards += self.gamma * torch.squeeze(
self.transition.values * extras["time_outs"].unsqueeze(1).to(self.device), 1
)
# record the transition
self.storage.add_transitions(self.transition)
self.transition.clear()
self.policy.reset(dones)
def compute_returns(self, obs):
# compute value for the last step
last_values = self.policy.evaluate(obs).detach()
self.storage.compute_returns(
last_values, self.gamma, self.lam, normalize_advantage=not self.normalize_advantage_per_mini_batch
)
def update(self): # noqa: C901
mean_value_loss = 0
mean_surrogate_loss = 0
mean_entropy = 0
# -- RND loss
if self.rnd:
mean_rnd_loss = 0
else:
mean_rnd_loss = None
# -- Symmetry loss
if self.symmetry:
mean_symmetry_loss = 0
else:
mean_symmetry_loss = None
# generator for mini batches
if self.policy.is_recurrent:
generator = self.storage.recurrent_mini_batch_generator(self.num_mini_batches, self.num_learning_epochs)
else:
generator = self.storage.mini_batch_generator(self.num_mini_batches, self.num_learning_epochs)
# iterate over batches
for (
obs_batch,
actions_batch,
target_values_batch,
advantages_batch,
returns_batch,
old_actions_log_prob_batch,
old_mu_batch,
old_sigma_batch,
hid_states_batch,
masks_batch,
) in generator:
# number of augmentations per sample
# we start with 1 and increase it if we use symmetry augmentation
num_aug = 1
# original batch size
# we assume policy group is always there and needs augmentation
original_batch_size = obs_batch.batch_size[0]
# check if we should normalize advantages per mini batch
if self.normalize_advantage_per_mini_batch:
with torch.no_grad():
advantages_batch = (advantages_batch - advantages_batch.mean()) / (advantages_batch.std() + 1e-8)
# Perform symmetric augmentation
if self.symmetry and self.symmetry["use_data_augmentation"]:
# augmentation using symmetry
data_augmentation_func = self.symmetry["data_augmentation_func"]
# returned shape: [batch_size * num_aug, ...]
obs_batch, actions_batch = data_augmentation_func(
obs=obs_batch,
actions=actions_batch,
env=self.symmetry["_env"],
)
# compute number of augmentations per sample
# we assume policy group is always there and needs augmentation
num_aug = int(obs_batch.batch_size[0] / original_batch_size)
# repeat the rest of the batch
# -- actor
old_actions_log_prob_batch = old_actions_log_prob_batch.repeat(num_aug, 1)
# -- critic
target_values_batch = target_values_batch.repeat(num_aug, 1)
advantages_batch = advantages_batch.repeat(num_aug, 1)
returns_batch = returns_batch.repeat(num_aug, 1)
# Recompute actions log prob and entropy for current batch of transitions
# Note: we need to do this because we updated the policy with the new parameters
# -- actor
self.policy.act(obs_batch, masks=masks_batch, hidden_states=hid_states_batch[0])
actions_log_prob_batch = self.policy.get_actions_log_prob(actions_batch)
# -- critic
value_batch = self.policy.evaluate(obs_batch, masks=masks_batch, hidden_states=hid_states_batch[1])
# -- entropy
# we only keep the entropy of the first augmentation (the original one)
mu_batch = self.policy.action_mean[:original_batch_size]
sigma_batch = self.policy.action_std[:original_batch_size]
entropy_batch = self.policy.entropy[:original_batch_size]
# KL
if self.desired_kl is not None and self.schedule == "adaptive":
with torch.inference_mode():
kl = torch.sum(
torch.log(sigma_batch / old_sigma_batch + 1.0e-5)
+ (torch.square(old_sigma_batch) + torch.square(old_mu_batch - mu_batch))
/ (2.0 * torch.square(sigma_batch))
- 0.5,
axis=-1,
)
kl_mean = torch.mean(kl)
# Reduce the KL divergence across all GPUs
if self.is_multi_gpu:
torch.distributed.all_reduce(kl_mean, op=torch.distributed.ReduceOp.SUM)
kl_mean /= self.gpu_world_size
# Update the learning rate
# Perform this adaptation only on the main process
# TODO: Is this needed? If KL-divergence is the "same" across all GPUs,
# then the learning rate should be the same across all GPUs.
if self.gpu_global_rank == 0:
if kl_mean > self.desired_kl * 2.0:
self.learning_rate = max(1e-5, self.learning_rate / 1.5)
elif kl_mean < self.desired_kl / 2.0 and kl_mean > 0.0:
self.learning_rate = min(1e-2, self.learning_rate * 1.5)
# Update the learning rate for all GPUs
if self.is_multi_gpu:
lr_tensor = torch.tensor(self.learning_rate, device=self.device)
torch.distributed.broadcast(lr_tensor, src=0)
self.learning_rate = lr_tensor.item()
# Update the learning rate for all parameter groups
for param_group in self.optimizer.param_groups:
param_group["lr"] = self.learning_rate
# Surrogate loss
ratio = torch.exp(actions_log_prob_batch - torch.squeeze(old_actions_log_prob_batch))
surrogate = -torch.squeeze(advantages_batch) * ratio
surrogate_clipped = -torch.squeeze(advantages_batch) * torch.clamp(
ratio, 1.0 - self.clip_param, 1.0 + self.clip_param
)
surrogate_loss = torch.max(surrogate, surrogate_clipped).mean()
# Value function loss
if self.use_clipped_value_loss:
value_clipped = target_values_batch + (value_batch - target_values_batch).clamp(
-self.clip_param, self.clip_param
)
value_losses = (value_batch - returns_batch).pow(2)
value_losses_clipped = (value_clipped - returns_batch).pow(2)
value_loss = torch.max(value_losses, value_losses_clipped).mean()
else:
value_loss = (returns_batch - value_batch).pow(2).mean()
loss = surrogate_loss + self.value_loss_coef * value_loss - self.entropy_coef * entropy_batch.mean()
# Symmetry loss
if self.symmetry:
# obtain the symmetric actions
# if we did augmentation before then we don't need to augment again
if not self.symmetry["use_data_augmentation"]:
data_augmentation_func = self.symmetry["data_augmentation_func"]
obs_batch, _ = data_augmentation_func(obs=obs_batch, actions=None, env=self.symmetry["_env"])
# compute number of augmentations per sample
num_aug = int(obs_batch.shape[0] / original_batch_size)
# actions predicted by the actor for symmetrically-augmented observations
mean_actions_batch = self.policy.act_inference(obs_batch.detach().clone())
# compute the symmetrically augmented actions
# note: we are assuming the first augmentation is the original one.
# We do not use the action_batch from earlier since that action was sampled from the distribution.
# However, the symmetry loss is computed using the mean of the distribution.
action_mean_orig = mean_actions_batch[:original_batch_size]
_, actions_mean_symm_batch = data_augmentation_func(
obs=None, actions=action_mean_orig, env=self.symmetry["_env"]
)
# compute the loss (we skip the first augmentation as it is the original one)
mse_loss = torch.nn.MSELoss()
symmetry_loss = mse_loss(
mean_actions_batch[original_batch_size:], actions_mean_symm_batch.detach()[original_batch_size:]
)
# add the loss to the total loss
if self.symmetry["use_mirror_loss"]:
loss += self.symmetry["mirror_loss_coeff"] * symmetry_loss
else:
symmetry_loss = symmetry_loss.detach()
# Random Network Distillation loss
# TODO: Move this processing to inside RND module.
if self.rnd:
# extract the rnd_state
# TODO: Check if we still need torch no grad. It is just an affine transformation.
with torch.no_grad():
rnd_state_batch = self.rnd.get_rnd_state(obs_batch[:original_batch_size])
rnd_state_batch = self.rnd.state_normalizer(rnd_state_batch)
# predict the embedding and the target
predicted_embedding = self.rnd.predictor(rnd_state_batch)
target_embedding = self.rnd.target(rnd_state_batch).detach()
# compute the loss as the mean squared error
mseloss = torch.nn.MSELoss()
rnd_loss = mseloss(predicted_embedding, target_embedding)
# Compute the gradients
# -- For PPO
self.optimizer.zero_grad()
loss.backward()
# -- For RND
if self.rnd:
self.rnd_optimizer.zero_grad() # type: ignore
rnd_loss.backward()
# Collect gradients from all GPUs
if self.is_multi_gpu:
self.reduce_parameters()
# Apply the gradients
# -- For PPO
nn.utils.clip_grad_norm_(self.policy.parameters(), self.max_grad_norm)
self.optimizer.step()
# -- For RND
if self.rnd_optimizer:
self.rnd_optimizer.step()
# Store the losses
mean_value_loss += value_loss.item()
mean_surrogate_loss += surrogate_loss.item()
mean_entropy += entropy_batch.mean().item()
# -- RND loss
if mean_rnd_loss is not None:
mean_rnd_loss += rnd_loss.item()
# -- Symmetry loss
if mean_symmetry_loss is not None:
mean_symmetry_loss += symmetry_loss.item()
# -- For PPO
num_updates = self.num_learning_epochs * self.num_mini_batches
mean_value_loss /= num_updates
mean_surrogate_loss /= num_updates
mean_entropy /= num_updates
# -- For RND
if mean_rnd_loss is not None:
mean_rnd_loss /= num_updates
# -- For Symmetry
if mean_symmetry_loss is not None:
mean_symmetry_loss /= num_updates
# -- Clear the storage
self.storage.clear()
# construct the loss dictionary
loss_dict = {
"value_function": mean_value_loss,
"surrogate": mean_surrogate_loss,
"entropy": mean_entropy,
}
if self.rnd:
loss_dict["rnd"] = mean_rnd_loss
if self.symmetry:
loss_dict["symmetry"] = mean_symmetry_loss
return loss_dict
"""
Helper functions
"""
def broadcast_parameters(self):
"""Broadcast model parameters to all GPUs."""
# obtain the model parameters on current GPU
model_params = [self.policy.state_dict()]
if self.rnd:
model_params.append(self.rnd.predictor.state_dict())
# broadcast the model parameters
torch.distributed.broadcast_object_list(model_params, src=0)
# load the model parameters on all GPUs from source GPU
self.policy.load_state_dict(model_params[0])
if self.rnd:
self.rnd.predictor.load_state_dict(model_params[1])
def reduce_parameters(self):
"""Collect gradients from all GPUs and average them.
This function is called after the backward pass to synchronize the gradients across all GPUs.
"""
# Create a tensor to store the gradients
grads = [param.grad.view(-1) for param in self.policy.parameters() if param.grad is not None]
if self.rnd:
grads += [param.grad.view(-1) for param in self.rnd.parameters() if param.grad is not None]
all_grads = torch.cat(grads)
# Average the gradients across all GPUs
torch.distributed.all_reduce(all_grads, op=torch.distributed.ReduceOp.SUM)
all_grads /= self.gpu_world_size
# Get all parameters
all_params = self.policy.parameters()
if self.rnd:
all_params = chain(all_params, self.rnd.parameters())
# Update the gradients for all parameters with the reduced gradients
offset = 0
for param in all_params:
if param.grad is not None:
numel = param.numel()
# copy data back from shared buffer
param.grad.data.copy_(all_grads[offset : offset + numel].view_as(param.grad.data))
# update the offset for the next parameter
offset += numel

View File

@ -0,0 +1,6 @@
# Copyright (c) 2021-2025, ETH Zurich and NVIDIA CORPORATION
# All rights reserved.
#
# SPDX-License-Identifier: BSD-3-Clause
"""Main module for the rsl_rl package."""

View File

@ -0,0 +1,11 @@
# Copyright (c) 2021-2025, ETH Zurich and NVIDIA CORPORATION
# All rights reserved.
#
# SPDX-License-Identifier: BSD-3-Clause
"""Implementation of different RL agents."""
from .distillation import Distillation
from .ppo import PPO
from .amp_ppo import AMP_PPO
__all__ = ["PPO", "Distillation", "AMP_PPO"]

View File

@ -0,0 +1,571 @@
# Copyright (c) 2021-2025, ETH Zurich and NVIDIA CORPORATION
# All rights reserved.
#
# SPDX-License-Identifier: BSD-3-Clause
from __future__ import annotations
from torch._tensor import Tensor
from torch._tensor import Tensor
from typing import Any
from copy import deepcopy
import torch
import torch.nn as nn
import torch.optim as optim
from itertools import chain
from rsl_rl.modules import ActorCritic
from rsl_rl.modules.rnd import RandomNetworkDistillation
from rsl_rl.storage import RolloutStorage, ReplayBufferMulti
from rsl_rl.utils import string_to_callable
class AMP_PPO:
"""Proximal Policy Optimization algorithm (https://arxiv.org/abs/1707.06347)."""
policy: ActorCritic
"""The actor critic module."""
def __init__(
self,
policy,
discriminator,
amp_data,
amp_normalizer,
amp_num_frames=1,
amp_replay_buffer_size=100000,
num_learning_epochs=5,
num_mini_batches=4,
clip_param=0.2,
gamma=0.99,
lam=0.95,
value_loss_coef=1.0,
entropy_coef=0.01,
learning_rate=0.001,
max_grad_norm=1.0,
use_clipped_value_loss=True,
schedule="adaptive",
desired_kl=0.01,
device="cpu",
normalize_advantage_per_mini_batch=False,
# RND parameters
rnd_cfg: dict | None = None,
# Symmetry parameters
symmetry_cfg: dict | None = None,
# Distributed training parameters
multi_gpu_cfg: dict | None = None,
):
# device-related parameters
self.device = device
self.is_multi_gpu = multi_gpu_cfg is not None
# Multi-GPU parameters
if multi_gpu_cfg is not None:
self.gpu_global_rank = multi_gpu_cfg["global_rank"]
self.gpu_world_size = multi_gpu_cfg["world_size"]
else:
self.gpu_global_rank = 0
self.gpu_world_size = 1
# RND components
if rnd_cfg is not None:
# Extract parameters used in ppo
rnd_lr = rnd_cfg.pop("learning_rate", 1e-3)
# Create RND module
self.rnd = RandomNetworkDistillation(device=self.device, **rnd_cfg)
# Create RND optimizer
params = self.rnd.predictor.parameters()
self.rnd_optimizer = optim.Adam(params, lr=rnd_lr)
else:
self.rnd = None
self.rnd_optimizer = None
# Symmetry components
if symmetry_cfg is not None:
# Check if symmetry is enabled
use_symmetry = symmetry_cfg["use_data_augmentation"] or symmetry_cfg["use_mirror_loss"]
# Print that we are not using symmetry
if not use_symmetry:
print("Symmetry not used for learning. We will use it for logging instead.")
# If function is a string then resolve it to a function
if isinstance(symmetry_cfg["data_augmentation_func"], str):
symmetry_cfg["data_augmentation_func"] = string_to_callable(symmetry_cfg["data_augmentation_func"])
# Check valid configuration
if symmetry_cfg["use_data_augmentation"] and not callable(symmetry_cfg["data_augmentation_func"]):
raise ValueError(
"Data augmentation enabled but the function is not callable:"
f" {symmetry_cfg['data_augmentation_func']}"
)
# Store symmetry configuration
self.symmetry = symmetry_cfg
else:
self.symmetry = None
## AMP components
self.discriminator = discriminator
self.discriminator.to(self.device)
self.amp_storage = ReplayBufferMulti(discriminator.state_dim, amp_replay_buffer_size, amp_num_frames, device)
self.amp_data = amp_data
self.amp_normalizer = amp_normalizer
# PPO components
self.policy = policy
self.policy.to(self.device)
# Create rollout storage
self.storage: RolloutStorage = None # type: ignore
self.transition = RolloutStorage.Transition()
self.amp_transition = RolloutStorage.Transition()
params = [
{'params': self.policy.parameters(), 'name': 'policy'},
]
params.append({
'params': self.discriminator.trunk.parameters(),
'weight_decay': 10e-4,
'name': f'amp_trunk'
})
params.append({
'params': self.discriminator.amp_linear.parameters(),
'weight_decay': 10e-2,
'name': f'amp_head'
})
# Create optimizer
self.optimizer = optim.Adam(params, lr=learning_rate)
# PPO parameters
self.clip_param = clip_param
self.num_learning_epochs = num_learning_epochs
self.num_mini_batches = num_mini_batches
self.value_loss_coef = value_loss_coef
self.entropy_coef = entropy_coef
self.gamma = gamma
self.lam = lam
self.max_grad_norm = max_grad_norm
self.use_clipped_value_loss = use_clipped_value_loss
self.desired_kl = desired_kl
self.schedule = schedule
self.learning_rate = learning_rate
self.normalize_advantage_per_mini_batch = normalize_advantage_per_mini_batch
def init_storage(self, training_type, num_envs, num_transitions_per_env, obs, actions_shape):
# create rollout storage
self.storage = RolloutStorage(
training_type,
num_envs,
num_transitions_per_env,
obs,
actions_shape,
self.device,
)
def act(self, obs, amp_obs):
if self.policy.is_recurrent:
self.transition.hidden_states = self.policy.get_hidden_states()
# compute the actions and values
self.transition.actions = self.policy.act(obs).detach()
self.transition.values = self.policy.evaluate(obs).detach()
self.transition.actions_log_prob = self.policy.get_actions_log_prob(self.transition.actions).detach()
self.transition.action_mean = self.policy.action_mean.detach()
self.transition.action_sigma = self.policy.action_std.detach()
# need to record obs before env.step()
self.transition.observations = obs
self.amp_transition.observations = amp_obs
return self.transition.actions
def process_env_step(self, obs, rewards, dones, extras,amp_obs, amp_obs_frames=None):
# update the normalizers
self.policy.update_normalization(obs)
if self.rnd:
self.rnd.update_normalization(obs)
# Record the rewards and dones
# Note: we clone here because later on we bootstrap the rewards based on timeouts
self.transition.rewards = rewards.clone()
self.transition.dones = dones
# Compute the intrinsic rewards and add to extrinsic rewards
if self.rnd:
# Compute the intrinsic rewards
self.intrinsic_rewards = self.rnd.get_intrinsic_reward(obs)
# Add intrinsic rewards to extrinsic rewards
self.transition.rewards += self.intrinsic_rewards
# Bootstrapping on time outs
if "time_outs" in extras:
self.transition.rewards += self.gamma * torch.squeeze(
self.transition.values * extras["time_outs"].unsqueeze(1).to(self.device), 1
)
if amp_obs_frames is not None:
self.amp_storage.insert(amp_obs_frames)
else:
self.amp_storage.insert(self.amp_transition.observations, amp_obs)
# record the transition
self.storage.add_transitions(self.transition)
self.transition.clear()
self.amp_transition.clear()
self.policy.reset(dones)
def compute_returns(self, obs):
# compute value for the last step
last_values = self.policy.evaluate(obs).detach()
self.storage.compute_returns(
last_values, self.gamma, self.lam, normalize_advantage=not self.normalize_advantage_per_mini_batch
)
def update(self): # noqa: C901
mean_value_loss = 0
mean_surrogate_loss = 0
mean_entropy = 0
mean_amp_loss = 0
mean_grad_pen_loss = 0
mean_policy_pred = 0
mean_expert_pred = 0
# -- RND loss
if self.rnd:
mean_rnd_loss = 0
else:
mean_rnd_loss = None
# -- Symmetry loss
if self.symmetry:
mean_symmetry_loss = 0
else:
mean_symmetry_loss = None
# generator for mini batches
if self.policy.is_recurrent:
generator = self.storage.recurrent_mini_batch_generator(self.num_mini_batches, self.num_learning_epochs)
else:
generator = self.storage.mini_batch_generator(self.num_mini_batches, self.num_learning_epochs)
amp_policy_generator = self.amp_storage.feed_forward_generator(
self.num_learning_epochs * self.num_mini_batches,
self.storage.num_envs * self.storage.num_transitions_per_env // self.num_mini_batches,
)
amp_expert_generator = self.amp_data.feed_forward_generator_23dof_multi(
self.num_learning_epochs * self.num_mini_batches,
self.storage.num_envs * self.storage.num_transitions_per_env // self.num_mini_batches,
)
# iterate over batches
for sample, sample_amp_policy, sample_amp_expert in zip(generator, amp_policy_generator, amp_expert_generator):
(
obs_batch,
actions_batch,
target_values_batch,
advantages_batch,
returns_batch,
old_actions_log_prob_batch,
old_mu_batch,
old_sigma_batch,
hid_states_batch,
masks_batch,
) = sample
# number of augmentations per sample
# we start with 1 and increase it if we use symmetry augmentation
num_aug = 1
# original batch size
# we assume policy group is always there and needs augmentation
original_batch_size = obs_batch.batch_size[0]
# check if we should normalize advantages per mini batch
if self.normalize_advantage_per_mini_batch:
with torch.no_grad():
advantages_batch = (advantages_batch - advantages_batch.mean()) / (advantages_batch.std() + 1e-8)
# Perform symmetric augmentation
if self.symmetry and self.symmetry["use_data_augmentation"]:
# augmentation using symmetry
data_augmentation_func = self.symmetry["data_augmentation_func"]
# returned shape: [batch_size * num_aug, ...]
obs_batch, actions_batch = data_augmentation_func(
obs=obs_batch,
actions=actions_batch,
env=self.symmetry["_env"],
)
# compute number of augmentations per sample
# we assume policy group is always there and needs augmentation
num_aug = int(obs_batch.batch_size[0] / original_batch_size)
# repeat the rest of the batch
# -- actor
old_actions_log_prob_batch = old_actions_log_prob_batch.repeat(num_aug, 1)
# -- critic
target_values_batch = target_values_batch.repeat(num_aug, 1)
advantages_batch = advantages_batch.repeat(num_aug, 1)
returns_batch = returns_batch.repeat(num_aug, 1)
# Recompute actions log prob and entropy for current batch of transitions
# Note: we need to do this because we updated the policy with the new parameters
# -- actor
self.policy.act(obs_batch, masks=masks_batch, hidden_states=hid_states_batch[0])
actions_log_prob_batch = self.policy.get_actions_log_prob(actions_batch)
# -- critic
value_batch = self.policy.evaluate(obs_batch, masks=masks_batch, hidden_states=hid_states_batch[1])
# -- entropy
# we only keep the entropy of the first augmentation (the original one)
mu_batch = self.policy.action_mean[:original_batch_size]
sigma_batch = self.policy.action_std[:original_batch_size]
entropy_batch = self.policy.entropy[:original_batch_size]
# KL
if self.desired_kl is not None and self.schedule == "adaptive":
with torch.inference_mode():
kl = torch.sum(
torch.log(sigma_batch / old_sigma_batch + 1.0e-5)
+ (torch.square(old_sigma_batch) + torch.square(old_mu_batch - mu_batch))
/ (2.0 * torch.square(sigma_batch))
- 0.5,
axis=-1,
)
kl_mean = torch.mean(kl)
# Reduce the KL divergence across all GPUs
if self.is_multi_gpu:
torch.distributed.all_reduce(kl_mean, op=torch.distributed.ReduceOp.SUM)
kl_mean /= self.gpu_world_size
# Update the learning rate
# Perform this adaptation only on the main process
# TODO: Is this needed? If KL-divergence is the "same" across all GPUs,
# then the learning rate should be the same across all GPUs.
if self.gpu_global_rank == 0:
if kl_mean > self.desired_kl * 2.0:
self.learning_rate = max(1e-5, self.learning_rate / 1.5)
elif kl_mean < self.desired_kl / 2.0 and kl_mean > 0.0:
self.learning_rate = min(1e-2, self.learning_rate * 1.5)
# Update the learning rate for all GPUs
if self.is_multi_gpu:
lr_tensor = torch.tensor(self.learning_rate, device=self.device)
torch.distributed.broadcast(lr_tensor, src=0)
self.learning_rate = lr_tensor.item()
# Update the learning rate for all parameter groups
for param_group in self.optimizer.param_groups:
param_group["lr"] = self.learning_rate
# Surrogate loss
ratio = torch.exp(actions_log_prob_batch - torch.squeeze(old_actions_log_prob_batch))
surrogate = -torch.squeeze(advantages_batch) * ratio
surrogate_clipped = -torch.squeeze(advantages_batch) * torch.clamp(
ratio, 1.0 - self.clip_param, 1.0 + self.clip_param
)
surrogate_loss = torch.max(surrogate, surrogate_clipped).mean()
# Value function loss
if self.use_clipped_value_loss:
value_clipped = target_values_batch + (value_batch - target_values_batch).clamp(
-self.clip_param, self.clip_param
)
value_losses = (value_batch - returns_batch).pow(2)
value_losses_clipped = (value_clipped - returns_batch).pow(2)
value_loss = torch.max(value_losses, value_losses_clipped).mean()
else:
value_loss = (returns_batch - value_batch).pow(2).mean()
loss = surrogate_loss + self.value_loss_coef * value_loss - self.entropy_coef * entropy_batch.mean()
# Symmetry loss
if self.symmetry:
# obtain the symmetric actions
# if we did augmentation before then we don't need to augment again
if not self.symmetry["use_data_augmentation"]:
data_augmentation_func = self.symmetry["data_augmentation_func"]
obs_batch, _ = data_augmentation_func(obs=obs_batch, actions=None, env=self.symmetry["_env"])
# compute number of augmentations per sample
num_aug = int(obs_batch.shape[0] / original_batch_size)
# actions predicted by the actor for symmetrically-augmented observations
mean_actions_batch = self.policy.act_inference(obs_batch.detach().clone())
# compute the symmetrically augmented actions
# note: we are assuming the first augmentation is the original one.
# We do not use the action_batch from earlier since that action was sampled from the distribution.
# However, the symmetry loss is computed using the mean of the distribution.
action_mean_orig = mean_actions_batch[:original_batch_size]
_, actions_mean_symm_batch = data_augmentation_func(
obs=None, actions=action_mean_orig, env=self.symmetry["_env"]
)
# compute the loss (we skip the first augmentation as it is the original one)
mse_loss = torch.nn.MSELoss()
symmetry_loss = mse_loss(
mean_actions_batch[original_batch_size:], actions_mean_symm_batch.detach()[original_batch_size:]
)
# add the loss to the total loss
if self.symmetry["use_mirror_loss"]:
loss += self.symmetry["mirror_loss_coeff"] * symmetry_loss
else:
symmetry_loss = symmetry_loss.detach()
# Random Network Distillation loss
# TODO: Move this processing to inside RND module.
if self.rnd:
# extract the rnd_state
# TODO: Check if we still need torch no grad. It is just an affine transformation.
with torch.no_grad():
rnd_state_batch = self.rnd.get_rnd_state(obs_batch[:original_batch_size])
rnd_state_batch = self.rnd.state_normalizer(rnd_state_batch)
# predict the embedding and the target
predicted_embedding = self.rnd.predictor(rnd_state_batch)
target_embedding = self.rnd.target(rnd_state_batch).detach()
# compute the loss as the mean squared error
mseloss = torch.nn.MSELoss()
rnd_loss = mseloss(predicted_embedding, target_embedding)
expert_states = sample_amp_expert
policy_states = sample_amp_policy
with torch.no_grad():
expert_states = self.amp_normalizer.normalize_torch(expert_states.to(self.device), self.device)
policy_states = self.amp_normalizer.normalize_torch(policy_states, self.device)
contact_phase_push = obs_batch['critic'][:, -4]
mask_push = contact_phase_push == 1.
if mask_push.any():
policy_d = self.discriminator(policy_states.flatten(1))
expert_states = expert_states.to(self.device)
expert_d = self.discriminator(expert_states.flatten(1))
expert_loss = torch.nn.MSELoss()(expert_d, torch.ones(expert_d.size(), device=self.device))
policy_loss = torch.nn.MSELoss()(policy_d, -1 * torch.ones(policy_d.size(), device=self.device))
amp_loss = 0.5 * (expert_loss + policy_loss)
# grad penalty
grad_pen_loss = self.discriminator.compute_grad_pen(expert_states, lambda_=5)
else:
amp_loss = torch.tensor(0.0, device=self.device)
grad_pen_loss = torch.tensor(0.0, device=self.device)
expert_loss = torch.tensor(0.0, device=self.device)
policy_loss = torch.tensor(0.0, device=self.device)
loss += (amp_loss + grad_pen_loss)
self.amp_normalizer.update(policy_states.cpu().numpy())
self.amp_normalizer.update(expert_states.cpu().numpy())
# Compute the gradients
# -- For PPO
self.optimizer.zero_grad()
loss.backward()
# -- For RND
if self.rnd:
self.rnd_optimizer.zero_grad() # type: ignore
rnd_loss.backward()
# Collect gradients from all GPUs
if self.is_multi_gpu:
self.reduce_parameters()
# Apply the gradients
# -- For PPO
nn.utils.clip_grad_norm_(self.policy.parameters(), self.max_grad_norm)
self.optimizer.step()
# -- For RND
if self.rnd_optimizer:
self.rnd_optimizer.step()
# Store the losses
mean_value_loss += value_loss.item()
mean_surrogate_loss += surrogate_loss.item()
mean_entropy += entropy_batch.mean().item()
mean_amp_loss += amp_loss.item()
mean_grad_pen_loss += grad_pen_loss.item()
mean_policy_pred += policy_loss.mean().item()
mean_expert_pred += expert_loss.mean().item()
# -- RND loss
if mean_rnd_loss is not None:
mean_rnd_loss += rnd_loss.item()
# -- Symmetry loss
if mean_symmetry_loss is not None:
mean_symmetry_loss += symmetry_loss.item()
# -- For PPO
num_updates = self.num_learning_epochs * self.num_mini_batches
mean_value_loss /= num_updates
mean_surrogate_loss /= num_updates
mean_entropy /= num_updates
mean_amp_loss /= num_updates
mean_grad_pen_loss /= num_updates
mean_policy_pred /= num_updates
mean_expert_pred /= num_updates
# -- For RND
if mean_rnd_loss is not None:
mean_rnd_loss /= num_updates
# -- For Symmetry
if mean_symmetry_loss is not None:
mean_symmetry_loss /= num_updates
# -- Clear the storage
self.storage.clear()
# construct the loss dictionary
loss_dict = {
"value_function": mean_value_loss,
"surrogate": mean_surrogate_loss,
"entropy": mean_entropy,
"amp": mean_amp_loss,
"amp_grad_pen": mean_grad_pen_loss,
"amp_policy_pred": mean_policy_pred,
"amp_expert_pred": mean_expert_pred,
}
if self.rnd:
loss_dict["rnd"] = mean_rnd_loss
if self.symmetry:
loss_dict["symmetry"] = mean_symmetry_loss
return loss_dict
"""
Helper functions
"""
def broadcast_parameters(self):
"""Broadcast model parameters to all GPUs."""
# obtain the model parameters on current GPU
model_params = [self.policy.state_dict()]
if self.rnd:
model_params.append(self.rnd.predictor.state_dict())
# broadcast the model parameters
torch.distributed.broadcast_object_list(model_params, src=0)
# load the model parameters on all GPUs from source GPU
self.policy.load_state_dict(model_params[0])
if self.rnd:
self.rnd.predictor.load_state_dict(model_params[1])
def reduce_parameters(self):
"""Collect gradients from all GPUs and average them.
This function is called after the backward pass to synchronize the gradients across all GPUs.
"""
# Create a tensor to store the gradients
grads = [param.grad.view(-1) for param in self.policy.parameters() if param.grad is not None]
if self.rnd:
grads += [param.grad.view(-1) for param in self.rnd.parameters() if param.grad is not None]
all_grads = torch.cat(grads)
# Average the gradients across all GPUs
torch.distributed.all_reduce(all_grads, op=torch.distributed.ReduceOp.SUM)
all_grads /= self.gpu_world_size
# Get all parameters
all_params = self.policy.parameters()
if self.rnd:
all_params = chain(all_params, self.rnd.parameters())
# Update the gradients for all parameters with the reduced gradients
offset = 0
for param in all_params:
if param.grad is not None:
numel = param.numel()
# copy data back from shared buffer
param.grad.data.copy_(all_grads[offset : offset + numel].view_as(param.grad.data))
# update the offset for the next parameter
offset += numel

View File

@ -0,0 +1,185 @@
# Copyright (c) 2021-2025, ETH Zurich and NVIDIA CORPORATION
# All rights reserved.
#
# SPDX-License-Identifier: BSD-3-Clause
import torch
import torch.nn as nn
from rsl_rl.modules import StudentTeacher, StudentTeacherRecurrent
from rsl_rl.storage import RolloutStorage
from rsl_rl.utils import resolve_optimizer
class Distillation:
"""Distillation algorithm for training a student model to mimic a teacher model."""
policy: StudentTeacher | StudentTeacherRecurrent
"""The student teacher model."""
def __init__(
self,
policy,
num_learning_epochs=1,
gradient_length=15,
learning_rate=1e-3,
max_grad_norm=None,
loss_type="mse",
optimizer="adam",
device="cpu",
# Distributed training parameters
multi_gpu_cfg: dict | None = None,
):
# device-related parameters
self.device = device
self.is_multi_gpu = multi_gpu_cfg is not None
# Multi-GPU parameters
if multi_gpu_cfg is not None:
self.gpu_global_rank = multi_gpu_cfg["global_rank"]
self.gpu_world_size = multi_gpu_cfg["world_size"]
else:
self.gpu_global_rank = 0
self.gpu_world_size = 1
# distillation components
self.policy = policy
self.policy.to(self.device)
self.storage = None # initialized later
# initialize the optimizer
self.optimizer = resolve_optimizer(optimizer)(self.policy.parameters(), lr=learning_rate)
# initialize the transition
self.transition = RolloutStorage.Transition()
self.last_hidden_states = None
# distillation parameters
self.num_learning_epochs = num_learning_epochs
self.gradient_length = gradient_length
self.learning_rate = learning_rate
self.max_grad_norm = max_grad_norm
# initialize the loss function
loss_fn_dict = {
"mse": nn.functional.mse_loss,
"huber": nn.functional.huber_loss,
}
if loss_type in loss_fn_dict:
self.loss_fn = loss_fn_dict[loss_type]
else:
raise ValueError(f"Unknown loss type: {loss_type}. Supported types are: {list(loss_fn_dict.keys())}")
self.num_updates = 0
def init_storage(self, training_type, num_envs, num_transitions_per_env, obs, actions_shape):
# create rollout storage
self.storage = RolloutStorage(
training_type,
num_envs,
num_transitions_per_env,
obs,
actions_shape,
self.device,
)
def act(self, obs):
# compute the actions
self.transition.actions = self.policy.act(obs).detach()
self.transition.privileged_actions = self.policy.evaluate(obs).detach()
# record the observations
self.transition.observations = obs
return self.transition.actions
def process_env_step(self, obs, rewards, dones, extras):
# update the normalizers
self.policy.update_normalization(obs)
# record the rewards and dones
self.transition.rewards = rewards
self.transition.dones = dones
# record the transition
self.storage.add_transitions(self.transition)
self.transition.clear()
self.policy.reset(dones)
def update(self):
self.num_updates += 1
mean_behavior_loss = 0
loss = 0
cnt = 0
for epoch in range(self.num_learning_epochs):
self.policy.reset(hidden_states=self.last_hidden_states)
self.policy.detach_hidden_states()
for obs, _, privileged_actions, dones in self.storage.generator():
# inference the student for gradient computation
actions = self.policy.act_inference(obs)
# behavior cloning loss
behavior_loss = self.loss_fn(actions, privileged_actions)
# total loss
loss = loss + behavior_loss
mean_behavior_loss += behavior_loss.item()
cnt += 1
# gradient step
if cnt % self.gradient_length == 0:
self.optimizer.zero_grad()
loss.backward()
if self.is_multi_gpu:
self.reduce_parameters()
if self.max_grad_norm:
nn.utils.clip_grad_norm_(self.policy.student.parameters(), self.max_grad_norm)
self.optimizer.step()
self.policy.detach_hidden_states()
loss = 0
# reset dones
self.policy.reset(dones.view(-1))
self.policy.detach_hidden_states(dones.view(-1))
mean_behavior_loss /= cnt
self.storage.clear()
self.last_hidden_states = self.policy.get_hidden_states()
self.policy.detach_hidden_states()
# construct the loss dictionary
loss_dict = {"behavior": mean_behavior_loss}
return loss_dict
"""
Helper functions
"""
def broadcast_parameters(self):
"""Broadcast model parameters to all GPUs."""
# obtain the model parameters on current GPU
model_params = [self.policy.state_dict()]
# broadcast the model parameters
torch.distributed.broadcast_object_list(model_params, src=0)
# load the model parameters on all GPUs from source GPU
self.policy.load_state_dict(model_params[0])
def reduce_parameters(self):
"""Collect gradients from all GPUs and average them.
This function is called after the backward pass to synchronize the gradients across all GPUs.
"""
# Create a tensor to store the gradients
grads = [param.grad.view(-1) for param in self.policy.parameters() if param.grad is not None]
all_grads = torch.cat(grads)
# Average the gradients across all GPUs
torch.distributed.all_reduce(all_grads, op=torch.distributed.ReduceOp.SUM)
all_grads /= self.gpu_world_size
# Update the gradients for all parameters with the reduced gradients
offset = 0
for param in self.policy.parameters():
if param.grad is not None:
numel = param.numel()
# copy data back from shared buffer
param.grad.data.copy_(all_grads[offset : offset + numel].view_as(param.grad.data))
# update the offset for the next parameter
offset += numel

View File

@ -0,0 +1,469 @@
# Copyright (c) 2021-2025, ETH Zurich and NVIDIA CORPORATION
# All rights reserved.
#
# SPDX-License-Identifier: BSD-3-Clause
from __future__ import annotations
import torch
import torch.nn as nn
import torch.optim as optim
from itertools import chain
from rsl_rl.modules import ActorCritic
from rsl_rl.modules.rnd import RandomNetworkDistillation
from rsl_rl.storage import RolloutStorage
from rsl_rl.utils import string_to_callable
class PPO:
"""Proximal Policy Optimization algorithm (https://arxiv.org/abs/1707.06347)."""
policy: ActorCritic
"""The actor critic module."""
def __init__(
self,
policy,
num_learning_epochs=5,
num_mini_batches=4,
clip_param=0.2,
gamma=0.99,
lam=0.95,
value_loss_coef=1.0,
entropy_coef=0.01,
learning_rate=0.001,
max_grad_norm=1.0,
use_clipped_value_loss=True,
schedule="adaptive",
desired_kl=0.01,
device="cpu",
normalize_advantage_per_mini_batch=False,
# RND parameters
rnd_cfg: dict | None = None,
# Symmetry parameters
symmetry_cfg: dict | None = None,
# Distributed training parameters
multi_gpu_cfg: dict | None = None,
):
# device-related parameters
self.device = device
self.is_multi_gpu = multi_gpu_cfg is not None
# Multi-GPU parameters
if multi_gpu_cfg is not None:
self.gpu_global_rank = multi_gpu_cfg["global_rank"]
self.gpu_world_size = multi_gpu_cfg["world_size"]
else:
self.gpu_global_rank = 0
self.gpu_world_size = 1
# RND components
if rnd_cfg is not None:
# Extract parameters used in ppo
rnd_lr = rnd_cfg.pop("learning_rate", 1e-3)
# Create RND module
self.rnd = RandomNetworkDistillation(device=self.device, **rnd_cfg)
# Create RND optimizer
params = self.rnd.predictor.parameters()
self.rnd_optimizer = optim.Adam(params, lr=rnd_lr)
else:
self.rnd = None
self.rnd_optimizer = None
# Symmetry components
if symmetry_cfg is not None:
# Check if symmetry is enabled
use_symmetry = symmetry_cfg["use_data_augmentation"] or symmetry_cfg["use_mirror_loss"]
# Print that we are not using symmetry
if not use_symmetry:
print("Symmetry not used for learning. We will use it for logging instead.")
# If function is a string then resolve it to a function
if isinstance(symmetry_cfg["data_augmentation_func"], str):
symmetry_cfg["data_augmentation_func"] = string_to_callable(symmetry_cfg["data_augmentation_func"])
# Check valid configuration
if symmetry_cfg["use_data_augmentation"] and not callable(symmetry_cfg["data_augmentation_func"]):
raise ValueError(
"Data augmentation enabled but the function is not callable:"
f" {symmetry_cfg['data_augmentation_func']}"
)
# Store symmetry configuration
self.symmetry = symmetry_cfg
else:
self.symmetry = None
# PPO components
self.policy = policy
self.policy.to(self.device)
# Create optimizer
self.optimizer = optim.Adam(self.policy.parameters(), lr=learning_rate)
# Create rollout storage
self.storage: RolloutStorage = None # type: ignore
self.transition = RolloutStorage.Transition()
# PPO parameters
self.clip_param = clip_param
self.num_learning_epochs = num_learning_epochs
self.num_mini_batches = num_mini_batches
self.value_loss_coef = value_loss_coef
self.entropy_coef = entropy_coef
self.gamma = gamma
self.lam = lam
self.max_grad_norm = max_grad_norm
self.use_clipped_value_loss = use_clipped_value_loss
self.desired_kl = desired_kl
self.schedule = schedule
self.learning_rate = learning_rate
self.normalize_advantage_per_mini_batch = normalize_advantage_per_mini_batch
def init_storage(self, training_type, num_envs, num_transitions_per_env, obs, actions_shape):
# create rollout storage
self.storage = RolloutStorage(
training_type,
num_envs,
num_transitions_per_env,
obs,
actions_shape,
self.device,
)
def act(self, obs):
if self.policy.is_recurrent:
self.transition.hidden_states = self.policy.get_hidden_states()
# compute the actions and values
self.transition.actions = self.policy.act(obs).detach()
self.transition.values = self.policy.evaluate(obs).detach()
self.transition.actions_log_prob = self.policy.get_actions_log_prob(self.transition.actions).detach()
self.transition.action_mean = self.policy.action_mean.detach()
self.transition.action_sigma = self.policy.action_std.detach()
# need to record obs before env.step()
self.transition.observations = obs
return self.transition.actions
def process_env_step(self, obs, rewards, dones, extras):
# update the normalizers
self.policy.update_normalization(obs)
if self.rnd:
self.rnd.update_normalization(obs)
# Record the rewards and dones
# Note: we clone here because later on we bootstrap the rewards based on timeouts
self.transition.rewards = rewards.clone()
self.transition.dones = dones
# Compute the intrinsic rewards and add to extrinsic rewards
if self.rnd:
# Compute the intrinsic rewards
self.intrinsic_rewards = self.rnd.get_intrinsic_reward(obs)
# Add intrinsic rewards to extrinsic rewards
self.transition.rewards += self.intrinsic_rewards
# Bootstrapping on time outs
if "time_outs" in extras:
self.transition.rewards += self.gamma * torch.squeeze(
self.transition.values * extras["time_outs"].unsqueeze(1).to(self.device), 1
)
# record the transition
self.storage.add_transitions(self.transition)
self.transition.clear()
self.policy.reset(dones)
def compute_returns(self, obs):
# compute value for the last step
last_values = self.policy.evaluate(obs).detach()
self.storage.compute_returns(
last_values, self.gamma, self.lam, normalize_advantage=not self.normalize_advantage_per_mini_batch
)
def update(self): # noqa: C901
mean_value_loss = 0
mean_surrogate_loss = 0
mean_entropy = 0
# -- RND loss
if self.rnd:
mean_rnd_loss = 0
else:
mean_rnd_loss = None
# -- Symmetry loss
if self.symmetry:
mean_symmetry_loss = 0
else:
mean_symmetry_loss = None
# generator for mini batches
if self.policy.is_recurrent:
generator = self.storage.recurrent_mini_batch_generator(self.num_mini_batches, self.num_learning_epochs)
else:
generator = self.storage.mini_batch_generator(self.num_mini_batches, self.num_learning_epochs)
# iterate over batches
for (
obs_batch,
actions_batch,
target_values_batch,
advantages_batch,
returns_batch,
old_actions_log_prob_batch,
old_mu_batch,
old_sigma_batch,
hid_states_batch,
masks_batch,
) in generator:
# number of augmentations per sample
# we start with 1 and increase it if we use symmetry augmentation
num_aug = 1
# original batch size
# we assume policy group is always there and needs augmentation
original_batch_size = obs_batch.batch_size[0]
# check if we should normalize advantages per mini batch
if self.normalize_advantage_per_mini_batch:
with torch.no_grad():
advantages_batch = (advantages_batch - advantages_batch.mean()) / (advantages_batch.std() + 1e-8)
# Perform symmetric augmentation
if self.symmetry and self.symmetry["use_data_augmentation"]:
# augmentation using symmetry
data_augmentation_func = self.symmetry["data_augmentation_func"]
# returned shape: [batch_size * num_aug, ...]
obs_batch, actions_batch = data_augmentation_func(
obs=obs_batch,
actions=actions_batch,
env=self.symmetry["_env"],
)
# compute number of augmentations per sample
# we assume policy group is always there and needs augmentation
num_aug = int(obs_batch.batch_size[0] / original_batch_size)
# repeat the rest of the batch
# -- actor
old_actions_log_prob_batch = old_actions_log_prob_batch.repeat(num_aug, 1)
# -- critic
target_values_batch = target_values_batch.repeat(num_aug, 1)
advantages_batch = advantages_batch.repeat(num_aug, 1)
returns_batch = returns_batch.repeat(num_aug, 1)
# Recompute actions log prob and entropy for current batch of transitions
# Note: we need to do this because we updated the policy with the new parameters
# -- actor
self.policy.act(obs_batch, masks=masks_batch, hidden_states=hid_states_batch[0])
actions_log_prob_batch = self.policy.get_actions_log_prob(actions_batch)
# -- critic
value_batch = self.policy.evaluate(obs_batch, masks=masks_batch, hidden_states=hid_states_batch[1])
# -- entropy
# we only keep the entropy of the first augmentation (the original one)
mu_batch = self.policy.action_mean[:original_batch_size]
sigma_batch = self.policy.action_std[:original_batch_size]
entropy_batch = self.policy.entropy[:original_batch_size]
# KL
if self.desired_kl is not None and self.schedule == "adaptive":
with torch.inference_mode():
kl = torch.sum(
torch.log(sigma_batch / old_sigma_batch + 1.0e-5)
+ (torch.square(old_sigma_batch) + torch.square(old_mu_batch - mu_batch))
/ (2.0 * torch.square(sigma_batch))
- 0.5,
axis=-1,
)
kl_mean = torch.mean(kl)
# Reduce the KL divergence across all GPUs
if self.is_multi_gpu:
torch.distributed.all_reduce(kl_mean, op=torch.distributed.ReduceOp.SUM)
kl_mean /= self.gpu_world_size
# Update the learning rate
# Perform this adaptation only on the main process
# TODO: Is this needed? If KL-divergence is the "same" across all GPUs,
# then the learning rate should be the same across all GPUs.
if self.gpu_global_rank == 0:
if kl_mean > self.desired_kl * 2.0:
self.learning_rate = max(1e-5, self.learning_rate / 1.5)
elif kl_mean < self.desired_kl / 2.0 and kl_mean > 0.0:
self.learning_rate = min(1e-2, self.learning_rate * 1.5)
# Update the learning rate for all GPUs
if self.is_multi_gpu:
lr_tensor = torch.tensor(self.learning_rate, device=self.device)
torch.distributed.broadcast(lr_tensor, src=0)
self.learning_rate = lr_tensor.item()
# Update the learning rate for all parameter groups
for param_group in self.optimizer.param_groups:
param_group["lr"] = self.learning_rate
# Surrogate loss
ratio = torch.exp(actions_log_prob_batch - torch.squeeze(old_actions_log_prob_batch))
surrogate = -torch.squeeze(advantages_batch) * ratio
surrogate_clipped = -torch.squeeze(advantages_batch) * torch.clamp(
ratio, 1.0 - self.clip_param, 1.0 + self.clip_param
)
surrogate_loss = torch.max(surrogate, surrogate_clipped).mean()
# Value function loss
if self.use_clipped_value_loss:
value_clipped = target_values_batch + (value_batch - target_values_batch).clamp(
-self.clip_param, self.clip_param
)
value_losses = (value_batch - returns_batch).pow(2)
value_losses_clipped = (value_clipped - returns_batch).pow(2)
value_loss = torch.max(value_losses, value_losses_clipped).mean()
else:
value_loss = (returns_batch - value_batch).pow(2).mean()
loss = surrogate_loss + self.value_loss_coef * value_loss - self.entropy_coef * entropy_batch.mean()
# Symmetry loss
if self.symmetry:
# obtain the symmetric actions
# if we did augmentation before then we don't need to augment again
if not self.symmetry["use_data_augmentation"]:
data_augmentation_func = self.symmetry["data_augmentation_func"]
obs_batch, _ = data_augmentation_func(obs=obs_batch, actions=None, env=self.symmetry["_env"])
# compute number of augmentations per sample
num_aug = int(obs_batch.shape[0] / original_batch_size)
# actions predicted by the actor for symmetrically-augmented observations
mean_actions_batch = self.policy.act_inference(obs_batch.detach().clone())
# compute the symmetrically augmented actions
# note: we are assuming the first augmentation is the original one.
# We do not use the action_batch from earlier since that action was sampled from the distribution.
# However, the symmetry loss is computed using the mean of the distribution.
action_mean_orig = mean_actions_batch[:original_batch_size]
_, actions_mean_symm_batch = data_augmentation_func(
obs=None, actions=action_mean_orig, env=self.symmetry["_env"]
)
# compute the loss (we skip the first augmentation as it is the original one)
mse_loss = torch.nn.MSELoss()
symmetry_loss = mse_loss(
mean_actions_batch[original_batch_size:], actions_mean_symm_batch.detach()[original_batch_size:]
)
# add the loss to the total loss
if self.symmetry["use_mirror_loss"]:
loss += self.symmetry["mirror_loss_coeff"] * symmetry_loss
else:
symmetry_loss = symmetry_loss.detach()
# Random Network Distillation loss
# TODO: Move this processing to inside RND module.
if self.rnd:
# extract the rnd_state
# TODO: Check if we still need torch no grad. It is just an affine transformation.
with torch.no_grad():
rnd_state_batch = self.rnd.get_rnd_state(obs_batch[:original_batch_size])
rnd_state_batch = self.rnd.state_normalizer(rnd_state_batch)
# predict the embedding and the target
predicted_embedding = self.rnd.predictor(rnd_state_batch)
target_embedding = self.rnd.target(rnd_state_batch).detach()
# compute the loss as the mean squared error
mseloss = torch.nn.MSELoss()
rnd_loss = mseloss(predicted_embedding, target_embedding)
# Compute the gradients
# -- For PPO
self.optimizer.zero_grad()
loss.backward()
# -- For RND
if self.rnd:
self.rnd_optimizer.zero_grad() # type: ignore
rnd_loss.backward()
# Collect gradients from all GPUs
if self.is_multi_gpu:
self.reduce_parameters()
# Apply the gradients
# -- For PPO
nn.utils.clip_grad_norm_(self.policy.parameters(), self.max_grad_norm)
self.optimizer.step()
# -- For RND
if self.rnd_optimizer:
self.rnd_optimizer.step()
# Store the losses
mean_value_loss += value_loss.item()
mean_surrogate_loss += surrogate_loss.item()
mean_entropy += entropy_batch.mean().item()
# -- RND loss
if mean_rnd_loss is not None:
mean_rnd_loss += rnd_loss.item()
# -- Symmetry loss
if mean_symmetry_loss is not None:
mean_symmetry_loss += symmetry_loss.item()
# -- For PPO
num_updates = self.num_learning_epochs * self.num_mini_batches
mean_value_loss /= num_updates
mean_surrogate_loss /= num_updates
mean_entropy /= num_updates
# -- For RND
if mean_rnd_loss is not None:
mean_rnd_loss /= num_updates
# -- For Symmetry
if mean_symmetry_loss is not None:
mean_symmetry_loss /= num_updates
# -- Clear the storage
self.storage.clear()
# construct the loss dictionary
loss_dict = {
"value_function": mean_value_loss,
"surrogate": mean_surrogate_loss,
"entropy": mean_entropy,
}
if self.rnd:
loss_dict["rnd"] = mean_rnd_loss
if self.symmetry:
loss_dict["symmetry"] = mean_symmetry_loss
return loss_dict
"""
Helper functions
"""
def broadcast_parameters(self):
"""Broadcast model parameters to all GPUs."""
# obtain the model parameters on current GPU
model_params = [self.policy.state_dict()]
if self.rnd:
model_params.append(self.rnd.predictor.state_dict())
# broadcast the model parameters
torch.distributed.broadcast_object_list(model_params, src=0)
# load the model parameters on all GPUs from source GPU
self.policy.load_state_dict(model_params[0])
if self.rnd:
self.rnd.predictor.load_state_dict(model_params[1])
def reduce_parameters(self):
"""Collect gradients from all GPUs and average them.
This function is called after the backward pass to synchronize the gradients across all GPUs.
"""
# Create a tensor to store the gradients
grads = [param.grad.view(-1) for param in self.policy.parameters() if param.grad is not None]
if self.rnd:
grads += [param.grad.view(-1) for param in self.rnd.parameters() if param.grad is not None]
all_grads = torch.cat(grads)
# Average the gradients across all GPUs
torch.distributed.all_reduce(all_grads, op=torch.distributed.ReduceOp.SUM)
all_grads /= self.gpu_world_size
# Get all parameters
all_params = self.policy.parameters()
if self.rnd:
all_params = chain(all_params, self.rnd.parameters())
# Update the gradients for all parameters with the reduced gradients
offset = 0
for param in all_params:
if param.grad is not None:
numel = param.numel()
# copy data back from shared buffer
param.grad.data.copy_(all_grads[offset : offset + numel].view_as(param.grad.data))
# update the offset for the next parameter
offset += numel

View File

@ -0,0 +1,10 @@
# Copyright (c) 2021-2025, ETH Zurich and NVIDIA CORPORATION
# All rights reserved.
#
# SPDX-License-Identifier: BSD-3-Clause
"""Submodule defining the environment definitions."""
from .vec_env import VecEnv
__all__ = ["VecEnv"]

View File

@ -0,0 +1,113 @@
# Copyright (c) 2021-2025, ETH Zurich and NVIDIA CORPORATION
# All rights reserved.
#
# SPDX-License-Identifier: BSD-3-Clause
from __future__ import annotations
import torch
from abc import ABC, abstractmethod
from tensordict import TensorDict
class VecEnv(ABC):
"""Abstract class for a vectorized environment.
The vectorized environment is a collection of environments that are synchronized. This means that
the same type of action is applied to all environments and the same type of observation is returned from all
environments.
"""
num_envs: int
"""Number of environments."""
num_actions: int
"""Number of actions."""
max_episode_length: int | torch.Tensor
max_episode_length_s: float
"""Maximum episode length.
The maximum episode length can be a scalar or a tensor. If it is a scalar, it is the same for all environments.
If it is a tensor, it is the maximum episode length for each environment. This is useful for dynamic episode
lengths.
"""
episode_length_buf: torch.Tensor
"""Buffer for current episode lengths."""
device: torch.device | str
"""Device to use."""
cfg: dict | object
"""Configuration object."""
reset_env_ids: torch.Tensor | None = None
contact_phase: torch.Tensor | None = None
"""
Operations.
"""
@abstractmethod
def get_observations(self) -> TensorDict:
"""Return the current observations.
Returns:
observations (TensorDict): Observations from the environment.
"""
raise NotImplementedError
@abstractmethod
def get_amp_observations(self) -> TensorDict:
"""Return the current AMP observations.
Returns:
observations (TensorDict): Observations from the environment.
"""
raise NotImplementedError
@abstractmethod
def step(self, actions: torch.Tensor) -> tuple[TensorDict, torch.Tensor, torch.Tensor, dict]:
"""Apply input action to the environment.
Args:
actions (torch.Tensor): Input actions to apply. Shape: (num_envs, num_actions)
Returns:
observations (TensorDict): Observations from the environment.
rewards (torch.Tensor): Rewards from the environment. Shape: (num_envs,)
dones (torch.Tensor): Done flags from the environment. Shape: (num_envs,)
extras (dict): Extra information from the environment.
Observations:
The observations TensorDict usually contains multiple observation groups. The `obs_groups`
dictionary of the runner configuration specifies which observation groups are used for which
purpose, i.e., it maps the available observation groups to observation sets. The observation sets
(keys of the `obs_groups` dictionary) currently used by rsl_rl are:
- "policy": Specified observation groups are used as input to the actor/student network.
- "critic": Specified observation groups are used as input to the critic network.
- "teacher": Specified observation groups are used as input to the teacher network.
- "rnd_state": Specified observation groups are used as input to the RND network.
Incomplete or incorrect configurations are handled in the `resolve_obs_groups()` function in
`rsl_rl/utils/utils.py`.
Extras:
The extras dictionary includes metrics such as the episode reward, episode length, etc. The following
dictionary keys are used by rsl_rl:
- "time_outs" (torch.Tensor): Timeouts for the environments. These correspond to terminations that
happen due to time limits and not due to the environment reaching a terminal state. This is useful
for environments that have a fixed episode length.
- "log" (dict[str, float | torch.Tensor]): Additional information for logging and debugging purposes.
The key should be a string and start with "/" for namespacing. The value can be a scalar or a
tensor. If it is a tensor, the mean of the tensor is used for logging.
"""
raise NotImplementedError

View File

@ -0,0 +1,21 @@
# Copyright (c) 2021-2025, ETH Zurich and NVIDIA CORPORATION
# All rights reserved.
#
# SPDX-License-Identifier: BSD-3-Clause
"""Definitions for neural-network components for RL-agents."""
from .actor_critic import ActorCritic
from .actor_critic_recurrent import ActorCriticRecurrent
from .rnd import *
from .student_teacher import StudentTeacher
from .student_teacher_recurrent import StudentTeacherRecurrent
from .symmetry import *
from .discriminator_multi import DiscriminatorMulti
__all__ = [
"ActorCritic",
"ActorCriticRecurrent",
"StudentTeacher",
"StudentTeacherRecurrent",
"DiscriminatorMulti",
]

View File

@ -0,0 +1,195 @@
# Copyright (c) 2021-2025, ETH Zurich and NVIDIA CORPORATION
# All rights reserved.
#
# SPDX-License-Identifier: BSD-3-Clause
from __future__ import annotations
import torch
import torch.nn as nn
from torch.distributions import Normal
from rsl_rl.networks import MLP, EmpiricalNormalization
class ActorCritic(nn.Module):
is_recurrent = False
def __init__(
self,
obs,
obs_groups,
num_actions,
actor_obs_normalization=False,
critic_obs_normalization=False,
actor_hidden_dims=[256, 256, 256],
critic_hidden_dims=[256, 256, 256],
activation="elu",
init_noise_std=1.0,
noise_std_type: str = "scalar",
state_dependent_std=False,
**kwargs,
):
if kwargs:
print(
"ActorCritic.__init__ got unexpected arguments, which will be ignored: "
+ str([key for key in kwargs.keys()])
)
super().__init__()
# get the observation dimensions
self.obs_groups = obs_groups
num_actor_obs = 0
for obs_group in obs_groups["policy"]:
assert len(obs[obs_group].shape) == 2, "The ActorCritic module only supports 1D observations."
num_actor_obs += obs[obs_group].shape[-1]
num_critic_obs = 0
for obs_group in obs_groups["critic"]:
assert len(obs[obs_group].shape) == 2, "The ActorCritic module only supports 1D observations."
num_critic_obs += obs[obs_group].shape[-1]
self.state_dependent_std = state_dependent_std
# actor
if self.state_dependent_std:
self.actor = MLP(num_actor_obs, [2, num_actions], actor_hidden_dims, activation)
else:
self.actor = MLP(num_actor_obs, num_actions, actor_hidden_dims, activation)
# actor observation normalization
self.actor_obs_normalization = actor_obs_normalization
if actor_obs_normalization:
self.actor_obs_normalizer = EmpiricalNormalization(num_actor_obs)
else:
self.actor_obs_normalizer = torch.nn.Identity()
print(f"Actor MLP: {self.actor}")
# critic
self.critic = MLP(num_critic_obs, 1, critic_hidden_dims, activation)
# critic observation normalization
self.critic_obs_normalization = critic_obs_normalization
if critic_obs_normalization:
self.critic_obs_normalizer = EmpiricalNormalization(num_critic_obs)
else:
self.critic_obs_normalizer = torch.nn.Identity()
print(f"Critic MLP: {self.critic}")
# Action noise
self.noise_std_type = noise_std_type
if self.state_dependent_std:
torch.nn.init.zeros_(self.actor[-2].weight[num_actions:])
if self.noise_std_type == "scalar":
torch.nn.init.constant_(self.actor[-2].bias[num_actions:], init_noise_std)
elif self.noise_std_type == "log":
torch.nn.init.constant_(
self.actor[-2].bias[num_actions:], torch.log(torch.tensor(init_noise_std + 1e-7))
)
else:
raise ValueError(f"Unknown standard deviation type: {self.noise_std_type}. Should be 'scalar' or 'log'")
else:
if self.noise_std_type == "scalar":
self.std = nn.Parameter(init_noise_std * torch.ones(num_actions))
elif self.noise_std_type == "log":
self.log_std = nn.Parameter(torch.log(init_noise_std * torch.ones(num_actions)))
else:
raise ValueError(f"Unknown standard deviation type: {self.noise_std_type}. Should be 'scalar' or 'log'")
# Action distribution (populated in update_distribution)
self.distribution = None
# disable args validation for speedup
Normal.set_default_validate_args(False)
def reset(self, dones=None):
pass
def forward(self):
raise NotImplementedError
@property
def action_mean(self):
return self.distribution.mean
@property
def action_std(self):
return self.distribution.stddev
@property
def entropy(self):
return self.distribution.entropy().sum(dim=-1)
def update_distribution(self, obs):
if self.state_dependent_std:
# compute mean and standard deviation
mean_and_std = self.actor(obs)
if self.noise_std_type == "scalar":
mean, std = torch.unbind(mean_and_std, dim=-2)
elif self.noise_std_type == "log":
mean, log_std = torch.unbind(mean_and_std, dim=-2)
std = torch.exp(log_std)
else:
raise ValueError(f"Unknown standard deviation type: {self.noise_std_type}. Should be 'scalar' or 'log'")
else:
# compute mean
mean = self.actor(obs)
# compute standard deviation
if self.noise_std_type == "scalar":
std = self.std.expand_as(mean)
elif self.noise_std_type == "log":
std = torch.exp(self.log_std).expand_as(mean)
else:
raise ValueError(f"Unknown standard deviation type: {self.noise_std_type}. Should be 'scalar' or 'log'")
# create distribution
self.distribution = Normal(mean, std)
def act(self, obs, **kwargs):
obs = self.get_actor_obs(obs)
obs = self.actor_obs_normalizer(obs)
self.update_distribution(obs)
return self.distribution.sample()
def act_inference(self, obs):
obs = self.get_actor_obs(obs)
obs = self.actor_obs_normalizer(obs)
return self.actor(obs)
def evaluate(self, obs, **kwargs):
obs = self.get_critic_obs(obs)
obs = self.critic_obs_normalizer(obs)
return self.critic(obs)
def get_actor_obs(self, obs):
obs_list = []
for obs_group in self.obs_groups["policy"]:
obs_list.append(obs[obs_group])
return torch.cat(obs_list, dim=-1)
def get_critic_obs(self, obs):
obs_list = []
for obs_group in self.obs_groups["critic"]:
obs_list.append(obs[obs_group])
return torch.cat(obs_list, dim=-1)
def get_actions_log_prob(self, actions):
return self.distribution.log_prob(actions).sum(dim=-1)
def update_normalization(self, obs):
if self.actor_obs_normalization:
actor_obs = self.get_actor_obs(obs)
self.actor_obs_normalizer.update(actor_obs)
if self.critic_obs_normalization:
critic_obs = self.get_critic_obs(obs)
self.critic_obs_normalizer.update(critic_obs)
def load_state_dict(self, state_dict, strict=True):
"""Load the parameters of the actor-critic model.
Args:
state_dict (dict): State dictionary of the model.
strict (bool): Whether to strictly enforce that the keys in state_dict match the keys returned by this
module's state_dict() function.
Returns:
bool: Whether this training resumes a previous training. This flag is used by the `load()` function of
`OnPolicyRunner` to determine how to load further parameters (relevant for, e.g., distillation).
"""
super().load_state_dict(state_dict, strict=strict)
return True # training resumes

View File

@ -0,0 +1,218 @@
# Copyright (c) 2021-2025, ETH Zurich and NVIDIA CORPORATION
# All rights reserved.
#
# SPDX-License-Identifier: BSD-3-Clause
from __future__ import annotations
import torch
import torch.nn as nn
import warnings
from torch.distributions import Normal
from rsl_rl.networks import MLP, EmpiricalNormalization, Memory
class ActorCriticRecurrent(nn.Module):
is_recurrent = True
def __init__(
self,
obs,
obs_groups,
num_actions,
actor_obs_normalization=False,
critic_obs_normalization=False,
actor_hidden_dims=[256, 256, 256],
critic_hidden_dims=[256, 256, 256],
activation="elu",
init_noise_std=1.0,
noise_std_type: str = "scalar",
state_dependent_std=False,
rnn_type="lstm",
rnn_hidden_dim=256,
rnn_num_layers=1,
**kwargs,
):
if "rnn_hidden_size" in kwargs:
warnings.warn(
"The argument `rnn_hidden_size` is deprecated and will be removed in a future version. "
"Please use `rnn_hidden_dim` instead.",
DeprecationWarning,
)
if rnn_hidden_dim == 256: # Only override if the new argument is at its default
rnn_hidden_dim = kwargs.pop("rnn_hidden_size")
if kwargs:
print(
"ActorCriticRecurrent.__init__ got unexpected arguments, which will be ignored: " + str(kwargs.keys()),
)
super().__init__()
# get the observation dimensions
self.obs_groups = obs_groups
num_actor_obs = 0
for obs_group in obs_groups["policy"]:
assert len(obs[obs_group].shape) == 2, "The ActorCriticRecurrent module only supports 1D observations."
num_actor_obs += obs[obs_group].shape[-1]
num_critic_obs = 0
for obs_group in obs_groups["critic"]:
assert len(obs[obs_group].shape) == 2, "The ActorCriticRecurrent module only supports 1D observations."
num_critic_obs += obs[obs_group].shape[-1]
self.state_dependent_std = state_dependent_std
# actor
self.memory_a = Memory(num_actor_obs, type=rnn_type, num_layers=rnn_num_layers, hidden_size=rnn_hidden_dim)
if self.state_dependent_std:
self.actor = MLP(rnn_hidden_dim, [2, num_actions], actor_hidden_dims, activation)
else:
self.actor = MLP(rnn_hidden_dim, num_actions, actor_hidden_dims, activation)
# actor observation normalization
self.actor_obs_normalization = actor_obs_normalization
if actor_obs_normalization:
self.actor_obs_normalizer = EmpiricalNormalization(num_actor_obs)
else:
self.actor_obs_normalizer = torch.nn.Identity()
print(f"Actor RNN: {self.memory_a}")
print(f"Actor MLP: {self.actor}")
# critic
self.memory_c = Memory(num_critic_obs, type=rnn_type, num_layers=rnn_num_layers, hidden_size=rnn_hidden_dim)
self.critic = MLP(rnn_hidden_dim, 1, critic_hidden_dims, activation)
# critic observation normalization
self.critic_obs_normalization = critic_obs_normalization
if critic_obs_normalization:
self.critic_obs_normalizer = EmpiricalNormalization(num_critic_obs)
else:
self.critic_obs_normalizer = torch.nn.Identity()
print(f"Critic RNN: {self.memory_c}")
print(f"Critic MLP: {self.critic}")
# Action noise
self.noise_std_type = noise_std_type
if self.state_dependent_std:
torch.nn.init.zeros_(self.actor[-2].weight[num_actions:])
if self.noise_std_type == "scalar":
torch.nn.init.constant_(self.actor[-2].bias[num_actions:], init_noise_std)
elif self.noise_std_type == "log":
torch.nn.init.constant_(
self.actor[-2].bias[num_actions:], torch.log(torch.tensor(init_noise_std + 1e-7))
)
else:
raise ValueError(f"Unknown standard deviation type: {self.noise_std_type}. Should be 'scalar' or 'log'")
else:
if self.noise_std_type == "scalar":
self.std = nn.Parameter(init_noise_std * torch.ones(num_actions))
elif self.noise_std_type == "log":
self.log_std = nn.Parameter(torch.log(init_noise_std * torch.ones(num_actions)))
else:
raise ValueError(f"Unknown standard deviation type: {self.noise_std_type}. Should be 'scalar' or 'log'")
# Action distribution (populated in update_distribution)
self.distribution = None
# disable args validation for speedup
Normal.set_default_validate_args(False)
@property
def action_mean(self):
return self.distribution.mean
@property
def action_std(self):
return self.distribution.stddev
@property
def entropy(self):
return self.distribution.entropy().sum(dim=-1)
def reset(self, dones=None):
self.memory_a.reset(dones)
self.memory_c.reset(dones)
def forward(self):
raise NotImplementedError
def update_distribution(self, obs):
if self.state_dependent_std:
# compute mean and standard deviation
mean_and_std = self.actor(obs)
if self.noise_std_type == "scalar":
mean, std = torch.unbind(mean_and_std, dim=-2)
elif self.noise_std_type == "log":
mean, log_std = torch.unbind(mean_and_std, dim=-2)
std = torch.exp(log_std)
else:
raise ValueError(f"Unknown standard deviation type: {self.noise_std_type}. Should be 'scalar' or 'log'")
else:
# compute mean
mean = self.actor(obs)
# compute standard deviation
if self.noise_std_type == "scalar":
std = self.std.expand_as(mean)
elif self.noise_std_type == "log":
std = torch.exp(self.log_std).expand_as(mean)
else:
raise ValueError(f"Unknown standard deviation type: {self.noise_std_type}. Should be 'scalar' or 'log'")
# create distribution
self.distribution = Normal(mean, std)
def act(self, obs, masks=None, hidden_states=None):
obs = self.get_actor_obs(obs)
obs = self.actor_obs_normalizer(obs)
out_mem = self.memory_a(obs, masks, hidden_states).squeeze(0)
self.update_distribution(out_mem)
return self.distribution.sample()
def act_inference(self, obs):
obs = self.get_actor_obs(obs)
obs = self.actor_obs_normalizer(obs)
out_mem = self.memory_a(obs).squeeze(0)
return self.actor(out_mem)
def evaluate(self, obs, masks=None, hidden_states=None):
obs = self.get_critic_obs(obs)
obs = self.critic_obs_normalizer(obs)
out_mem = self.memory_c(obs, masks, hidden_states).squeeze(0)
return self.critic(out_mem)
def get_actor_obs(self, obs):
obs_list = []
for obs_group in self.obs_groups["policy"]:
obs_list.append(obs[obs_group])
return torch.cat(obs_list, dim=-1)
def get_critic_obs(self, obs):
obs_list = []
for obs_group in self.obs_groups["critic"]:
obs_list.append(obs[obs_group])
return torch.cat(obs_list, dim=-1)
def get_actions_log_prob(self, actions):
return self.distribution.log_prob(actions).sum(dim=-1)
def get_hidden_states(self):
return self.memory_a.hidden_states, self.memory_c.hidden_states
def update_normalization(self, obs):
if self.actor_obs_normalization:
actor_obs = self.get_actor_obs(obs)
self.actor_obs_normalizer.update(actor_obs)
if self.critic_obs_normalization:
critic_obs = self.get_critic_obs(obs)
self.critic_obs_normalizer.update(critic_obs)
def load_state_dict(self, state_dict, strict=True):
"""Load the parameters of the actor-critic model.
Args:
state_dict (dict): State dictionary of the model.
strict (bool): Whether to strictly enforce that the keys in state_dict match the keys returned by this
module's state_dict() function.
Returns:
bool: Whether this training resumes a previous training. This flag is used by the `load()` function of
`OnPolicyRunner` to determine how to load further parameters (relevant for, e.g., distillation).
"""
super().load_state_dict(state_dict, strict=strict)
return True

View File

@ -0,0 +1,102 @@
import torch
import torch.nn as nn
from torch import autograd
import torch.nn.utils.spectral_norm as spectral_norm
class DiscriminatorMulti(nn.Module):
def __init__(
self, state_dim, amp_reward_coef, hidden_layer_sizes, device,
num_frames=2, task_reward_lerp=0.0, use_lerp=True):
super(DiscriminatorMulti, self).__init__()
self.device = device
self.state_dim = state_dim
self.use_lerp = use_lerp
self.num_frames = num_frames # 存储帧数参数
self.amp_reward_coef = amp_reward_coef
amp_layers = []
curr_in_dim = state_dim * num_frames
for hidden_dim in hidden_layer_sizes:
amp_layers.append(spectral_norm(nn.Linear(curr_in_dim, hidden_dim)))
amp_layers.append(nn.ReLU())
curr_in_dim = hidden_dim
self.trunk = nn.Sequential(*amp_layers).to(device)
self.amp_linear = spectral_norm(nn.Linear(hidden_layer_sizes[-1], 1)).to(device)
self.trunk.train()
self.amp_linear.train()
self.task_reward_lerp = task_reward_lerp
def forward(self, x):
h = self.trunk(x)
d = self.amp_linear(h)
return d
def compute_grad_pen(self,
expert_states, # 改为接收多帧状态列表
lambda_=10):
# 将多帧状态沿最后一个维度拼接
expert_data = expert_states.flatten(1)
expert_data.requires_grad = True
disc = self.amp_linear(self.trunk(expert_data))
ones = torch.ones(disc.size(), device=disc.device)
grad = autograd.grad(
outputs=disc, inputs=expert_data,
grad_outputs=ones, create_graph=True,
retain_graph=True, only_inputs=True)[0]
# Enforce that the grad norm approaches 0.
grad_pen = lambda_ * (grad.norm(2, dim=1) - 0).pow(2).mean()
return grad_pen
def get_disc_weights(self):
weights = []
for m in self.trunk.modules():
if isinstance(m, nn.Linear):
weights.append(torch.flatten(m.weight))
weights.append(torch.flatten(self.amp_linear.weight))
return weights
def get_disc_logit_weights(self):
return torch.flatten(self.amp_linear.weight)
def predict_amp_reward(
self, states, # 改为接收多帧状态列表
task_reward, normalizer=None):
"""
states: torch.Tensor, shape=(num_envs, num_frames, state_dim)
task_reward: torch.Tensor, shape=(num_envs, 1)
"""
# import ipdb; ipdb.set_trace()
with torch.no_grad():
self.eval()
if normalizer is not None:
# 对每一帧状态进行归一化
states = normalizer.normalize_torch(states, self.device)
# 拼接多帧状态
state_cat = states.flatten(1)
d = self.amp_linear(self.trunk(state_cat))
disc_reward = self.amp_reward_coef * torch.clamp(1 - (1/4) * torch.square(d - 1), min=0)
if self.use_lerp:
if self.task_reward_lerp > 0:
reward = self._lerp_reward(disc_reward, task_reward.unsqueeze(-1))
self.train()
return reward.squeeze(), d, disc_reward.squeeze() * (1.0 - self.task_reward_lerp)
else:
disc_reward *= 0.02
reward = task_reward.unsqueeze(-1) + disc_reward
self.train()
return reward.squeeze(), d, disc_reward.squeeze()
def _lerp_reward(self, disc_r, task_r):
r = (1.0 - self.task_reward_lerp) * disc_r + self.task_reward_lerp * task_r
return r

View File

@ -0,0 +1,209 @@
# Copyright (c) 2021-2025, ETH Zurich and NVIDIA CORPORATION
# All rights reserved.
#
# SPDX-License-Identifier: BSD-3-Clause
from __future__ import annotations
import torch
import torch.nn as nn
from rsl_rl.networks import MLP, EmpiricalDiscountedVariationNormalization, EmpiricalNormalization
class RandomNetworkDistillation(nn.Module):
"""Implementation of Random Network Distillation (RND) [1]
References:
.. [1] Burda, Yuri, et al. "Exploration by random network distillation." arXiv preprint arXiv:1810.12894 (2018).
"""
def __init__(
self,
num_states: int,
obs_groups: dict,
num_outputs: int,
predictor_hidden_dims: list[int],
target_hidden_dims: list[int],
activation: str = "elu",
weight: float = 0.0,
state_normalization: bool = False,
reward_normalization: bool = False,
device: str = "cpu",
weight_schedule: dict | None = None,
):
"""Initialize the RND module.
- If :attr:`state_normalization` is True, then the input state is normalized using an Empirical Normalization layer.
- If :attr:`reward_normalization` is True, then the intrinsic reward is normalized using an Empirical Discounted
Variation Normalization layer.
.. note::
If the hidden dimensions are -1 in the predictor and target networks configuration, then the number of states
is used as the hidden dimension.
Args:
num_states: Number of states/inputs to the predictor and target networks.
num_outputs: Number of outputs (embedding size) of the predictor and target networks.
predictor_hidden_dims: List of hidden dimensions of the predictor network.
target_hidden_dims: List of hidden dimensions of the target network.
activation: Activation function. Defaults to "elu".
weight: Scaling factor of the intrinsic reward. Defaults to 0.0.
state_normalization: Whether to normalize the input state. Defaults to False.
reward_normalization: Whether to normalize the intrinsic reward. Defaults to False.
device: Device to use. Defaults to "cpu".
weight_schedule: The type of schedule to use for the RND weight parameter.
Defaults to None, in which case the weight parameter is constant.
It is a dictionary with the following keys:
- "mode": The type of schedule to use for the RND weight parameter.
- "constant": Constant weight schedule.
- "step": Step weight schedule.
- "linear": Linear weight schedule.
For the "step" weight schedule, the following parameters are required:
- "final_step": The step at which the weight parameter is set to the final value.
- "final_value": The final value of the weight parameter.
For the "linear" weight schedule, the following parameters are required:
- "initial_step": The step at which the weight parameter is set to the initial value.
- "final_step": The step at which the weight parameter is set to the final value.
- "final_value": The final value of the weight parameter.
"""
# initialize parent class
super().__init__()
# Store parameters
self.num_states = num_states
self.obs_groups = obs_groups
self.num_outputs = num_outputs
self.initial_weight = weight
self.device = device
self.state_normalization = state_normalization
self.reward_normalization = reward_normalization
# Normalization of input gates
if state_normalization:
self.state_normalizer = EmpiricalNormalization(shape=[self.num_states], until=1.0e8).to(self.device)
else:
self.state_normalizer = torch.nn.Identity()
# Normalization of intrinsic reward
if reward_normalization:
self.reward_normalizer = EmpiricalDiscountedVariationNormalization(shape=[], until=1.0e8).to(self.device)
else:
self.reward_normalizer = torch.nn.Identity()
# counter for the number of updates
self.update_counter = 0
# resolve weight schedule
if weight_schedule is not None:
self.weight_scheduler_params = weight_schedule
self.weight_scheduler = getattr(self, f"_{weight_schedule['mode']}_weight_schedule")
else:
self.weight_scheduler = None
# Create network architecture
self.predictor = MLP(num_states, num_outputs, predictor_hidden_dims, activation).to(self.device)
self.target = MLP(num_states, num_outputs, target_hidden_dims, activation).to(self.device)
# make target network not trainable
self.target.eval()
def get_intrinsic_reward(self, obs) -> torch.Tensor:
# Note: the counter is updated number of env steps per learning iteration
self.update_counter += 1
# Extract the rnd state from the observation
rnd_state = self.get_rnd_state(obs)
rnd_state = self.state_normalizer(rnd_state)
# Obtain the embedding of the rnd state from the target and predictor networks
target_embedding = self.target(rnd_state).detach()
predictor_embedding = self.predictor(rnd_state).detach()
# Compute the intrinsic reward as the distance between the embeddings
intrinsic_reward = torch.linalg.norm(target_embedding - predictor_embedding, dim=1)
# Normalize intrinsic reward
intrinsic_reward = self.reward_normalizer(intrinsic_reward)
# Check the weight schedule
if self.weight_scheduler is not None:
self.weight = self.weight_scheduler(step=self.update_counter, **self.weight_scheduler_params)
else:
self.weight = self.initial_weight
# Scale intrinsic reward
intrinsic_reward *= self.weight
return intrinsic_reward
def forward(self, *args, **kwargs):
raise RuntimeError("Forward method is not implemented. Use get_intrinsic_reward instead.")
def train(self, mode: bool = True):
# sets module into training mode
self.predictor.train(mode)
if self.state_normalization:
self.state_normalizer.train(mode)
if self.reward_normalization:
self.reward_normalizer.train(mode)
return self
def eval(self):
return self.train(False)
def get_rnd_state(self, obs):
obs_list = []
for obs_group in self.obs_groups["rnd_state"]:
obs_list.append(obs[obs_group])
return torch.cat(obs_list, dim=-1)
def update_normalization(self, obs):
# Normalize the state
if self.state_normalization:
rnd_state = self.get_rnd_state(obs)
self.state_normalizer.update(rnd_state)
"""
Different weight schedules.
"""
def _constant_weight_schedule(self, step: int, **kwargs):
return self.initial_weight
def _step_weight_schedule(self, step: int, final_step: int, final_value: float, **kwargs):
return self.initial_weight if step < final_step else final_value
def _linear_weight_schedule(self, step: int, initial_step: int, final_step: int, final_value: float, **kwargs):
if step < initial_step:
return self.initial_weight
elif step > final_step:
return final_value
else:
return self.initial_weight + (final_value - self.initial_weight) * (step - initial_step) / (
final_step - initial_step
)
def resolve_rnd_config(alg_cfg, obs, obs_groups, env):
"""Resolve the RND configuration.
Args:
alg_cfg: The algorithm configuration dictionary.
obs: The observation dictionary.
obs_groups: The observation groups dictionary.
env: The environment.
Returns:
The resolved algorithm configuration dictionary.
"""
# resolve dimension of rnd gated state
if "rnd_cfg" in alg_cfg and alg_cfg["rnd_cfg"] is not None:
# get dimension of rnd gated state
num_rnd_state = 0
for obs_group in obs_groups["rnd_state"]:
assert len(obs[obs_group].shape) == 2, "The RND module only supports 1D observations."
num_rnd_state += obs[obs_group].shape[-1]
# add rnd gated state to config
alg_cfg["rnd_cfg"]["num_states"] = num_rnd_state
alg_cfg["rnd_cfg"]["obs_groups"] = obs_groups
# scale down the rnd weight with timestep
alg_cfg["rnd_cfg"]["weight"] *= env.unwrapped.step_dt
return alg_cfg

View File

@ -0,0 +1,206 @@
# Copyright (c) 2021-2025, ETH Zurich and NVIDIA CORPORATION
# All rights reserved.
#
# SPDX-License-Identifier: BSD-3-Clause
from __future__ import annotations
import torch
import torch.nn as nn
from torch.distributions import Normal
from rsl_rl.networks import MLP, EmpiricalNormalization
class StudentTeacher(nn.Module):
is_recurrent = False
def __init__(
self,
obs,
obs_groups,
num_actions,
student_obs_normalization=False,
teacher_obs_normalization=False,
student_hidden_dims=[256, 256, 256],
teacher_hidden_dims=[256, 256, 256],
activation="elu",
init_noise_std=0.1,
noise_std_type: str = "scalar",
**kwargs,
):
if kwargs:
print(
"StudentTeacher.__init__ got unexpected arguments, which will be ignored: "
+ str([key for key in kwargs.keys()])
)
super().__init__()
self.loaded_teacher = False # indicates if teacher has been loaded
# get the observation dimensions
self.obs_groups = obs_groups
num_student_obs = 0
for obs_group in obs_groups["policy"]:
assert len(obs[obs_group].shape) == 2, "The StudentTeacher module only supports 1D observations."
num_student_obs += obs[obs_group].shape[-1]
num_teacher_obs = 0
for obs_group in obs_groups["teacher"]:
assert len(obs[obs_group].shape) == 2, "The StudentTeacher module only supports 1D observations."
num_teacher_obs += obs[obs_group].shape[-1]
# student
self.student = MLP(num_student_obs, num_actions, student_hidden_dims, activation)
# student observation normalization
self.student_obs_normalization = student_obs_normalization
if student_obs_normalization:
self.student_obs_normalizer = EmpiricalNormalization(num_student_obs)
else:
self.student_obs_normalizer = torch.nn.Identity()
print(f"Student MLP: {self.student}")
# teacher
self.teacher = MLP(num_teacher_obs, num_actions, teacher_hidden_dims, activation)
self.teacher.eval()
# teacher observation normalization
self.teacher_obs_normalization = teacher_obs_normalization
if teacher_obs_normalization:
self.teacher_obs_normalizer = EmpiricalNormalization(num_teacher_obs)
else:
self.teacher_obs_normalizer = torch.nn.Identity()
print(f"Teacher MLP: {self.teacher}")
# action noise
self.noise_std_type = noise_std_type
if self.noise_std_type == "scalar":
self.std = nn.Parameter(init_noise_std * torch.ones(num_actions))
elif self.noise_std_type == "log":
self.log_std = nn.Parameter(torch.log(init_noise_std * torch.ones(num_actions)))
else:
raise ValueError(f"Unknown standard deviation type: {self.noise_std_type}. Should be 'scalar' or 'log'")
# action distribution (populated in update_distribution)
self.distribution = None
# disable args validation for speedup
Normal.set_default_validate_args(False)
def reset(self, dones=None, hidden_states=None):
pass
def forward(self):
raise NotImplementedError
@property
def action_mean(self):
return self.distribution.mean
@property
def action_std(self):
return self.distribution.stddev
@property
def entropy(self):
return self.distribution.entropy().sum(dim=-1)
def update_distribution(self, obs):
# compute mean
mean = self.student(obs)
# compute standard deviation
if self.noise_std_type == "scalar":
std = self.std.expand_as(mean)
elif self.noise_std_type == "log":
std = torch.exp(self.log_std).expand_as(mean)
else:
raise ValueError(f"Unknown standard deviation type: {self.noise_std_type}. Should be 'scalar' or 'log'")
# create distribution
self.distribution = Normal(mean, std)
def act(self, obs):
obs = self.get_student_obs(obs)
obs = self.student_obs_normalizer(obs)
self.update_distribution(obs)
return self.distribution.sample()
def act_inference(self, obs):
obs = self.get_student_obs(obs)
obs = self.student_obs_normalizer(obs)
return self.student(obs)
def evaluate(self, obs):
obs = self.get_teacher_obs(obs)
obs = self.teacher_obs_normalizer(obs)
with torch.no_grad():
return self.teacher(obs)
def get_student_obs(self, obs):
obs_list = []
for obs_group in self.obs_groups["policy"]:
obs_list.append(obs[obs_group])
return torch.cat(obs_list, dim=-1)
def get_teacher_obs(self, obs):
obs_list = []
for obs_group in self.obs_groups["teacher"]:
obs_list.append(obs[obs_group])
return torch.cat(obs_list, dim=-1)
def get_hidden_states(self):
return None
def detach_hidden_states(self, dones=None):
pass
def train(self, mode=True):
super().train(mode)
# make sure teacher is in eval mode
self.teacher.eval()
self.teacher_obs_normalizer.eval()
def update_normalization(self, obs):
if self.student_obs_normalization:
student_obs = self.get_student_obs(obs)
self.student_obs_normalizer.update(student_obs)
def load_state_dict(self, state_dict, strict=True):
"""Load the parameters of the student and teacher networks.
Args:
state_dict (dict): State dictionary of the model.
strict (bool): Whether to strictly enforce that the keys in state_dict match the keys returned by this
module's state_dict() function.
Returns:
bool: Whether this training resumes a previous training. This flag is used by the `load()` function of
`OnPolicyRunner` to determine how to load further parameters.
"""
# check if state_dict contains teacher and student or just teacher parameters
if any("actor" in key for key in state_dict.keys()): # loading parameters from rl training
# rename keys to match teacher and remove critic parameters
teacher_state_dict = {}
teacher_obs_normalizer_state_dict = {}
for key, value in state_dict.items():
if "actor." in key:
teacher_state_dict[key.replace("actor.", "")] = value
if "actor_obs_normalizer." in key:
teacher_obs_normalizer_state_dict[key.replace("actor_obs_normalizer.", "")] = value
self.teacher.load_state_dict(teacher_state_dict, strict=strict)
self.teacher_obs_normalizer.load_state_dict(teacher_obs_normalizer_state_dict, strict=strict)
# set flag for successfully loading the parameters
self.loaded_teacher = True
self.teacher.eval()
self.teacher_obs_normalizer.eval()
return False # training does not resume
elif any("student" in key for key in state_dict.keys()): # loading parameters from distillation training
super().load_state_dict(state_dict, strict=strict)
# set flag for successfully loading the parameters
self.loaded_teacher = True
self.teacher.eval()
self.teacher_obs_normalizer.eval()
return True # training resumes
else:
raise ValueError("state_dict does not contain student or teacher parameters")

View File

@ -0,0 +1,249 @@
# Copyright (c) 2021-2025, ETH Zurich and NVIDIA CORPORATION
# All rights reserved.
#
# SPDX-License-Identifier: BSD-3-Clause
from __future__ import annotations
import torch
import torch.nn as nn
import warnings
from torch.distributions import Normal
from rsl_rl.networks import MLP, EmpiricalNormalization, Memory
class StudentTeacherRecurrent(nn.Module):
is_recurrent = True
def __init__(
self,
obs,
obs_groups,
num_actions,
student_obs_normalization=False,
teacher_obs_normalization=False,
student_hidden_dims=[256, 256, 256],
teacher_hidden_dims=[256, 256, 256],
activation="elu",
init_noise_std=0.1,
noise_std_type: str = "scalar",
rnn_type="lstm",
rnn_hidden_dim=256,
rnn_num_layers=1,
teacher_recurrent=False,
**kwargs,
):
if "rnn_hidden_size" in kwargs:
warnings.warn(
"The argument `rnn_hidden_size` is deprecated and will be removed in a future version. "
"Please use `rnn_hidden_dim` instead.",
DeprecationWarning,
)
if rnn_hidden_dim == 256: # Only override if the new argument is at its default
rnn_hidden_dim = kwargs.pop("rnn_hidden_size")
if kwargs:
print(
"StudentTeacherRecurrent.__init__ got unexpected arguments, which will be ignored: "
+ str(kwargs.keys()),
)
super().__init__()
self.loaded_teacher = False # indicates if teacher has been loaded
self.teacher_recurrent = teacher_recurrent # indicates if teacher is recurrent too
# get the observation dimensions
self.obs_groups = obs_groups
num_student_obs = 0
for obs_group in obs_groups["policy"]:
assert len(obs[obs_group].shape) == 2, "The StudentTeacher module only supports 1D observations."
num_student_obs += obs[obs_group].shape[-1]
num_teacher_obs = 0
for obs_group in obs_groups["teacher"]:
assert len(obs[obs_group].shape) == 2, "The StudentTeacher module only supports 1D observations."
num_teacher_obs += obs[obs_group].shape[-1]
# student
self.memory_s = Memory(num_student_obs, type=rnn_type, num_layers=rnn_num_layers, hidden_size=rnn_hidden_dim)
self.student = MLP(rnn_hidden_dim, num_actions, student_hidden_dims, activation)
# student observation normalization
self.student_obs_normalization = student_obs_normalization
if student_obs_normalization:
self.student_obs_normalizer = EmpiricalNormalization(num_student_obs)
else:
self.student_obs_normalizer = torch.nn.Identity()
print(f"Student RNN: {self.memory_s}")
print(f"Student MLP: {self.student}")
# teacher
if self.teacher_recurrent:
self.memory_t = Memory(
num_teacher_obs, type=rnn_type, num_layers=rnn_num_layers, hidden_size=rnn_hidden_dim
)
num_teacher_obs = rnn_hidden_dim
self.teacher = MLP(num_teacher_obs, num_actions, teacher_hidden_dims, activation)
# teacher observation normalization
self.teacher_obs_normalization = teacher_obs_normalization
if teacher_obs_normalization:
self.teacher_obs_normalizer = EmpiricalNormalization(num_teacher_obs)
else:
self.teacher_obs_normalizer = torch.nn.Identity()
if self.teacher_recurrent:
print(f"Teacher RNN: {self.memory_t}")
print(f"Teacher MLP: {self.teacher}")
# action noise
self.noise_std_type = noise_std_type
if self.noise_std_type == "scalar":
self.std = nn.Parameter(init_noise_std * torch.ones(num_actions))
elif self.noise_std_type == "log":
self.log_std = nn.Parameter(torch.log(init_noise_std * torch.ones(num_actions)))
else:
raise ValueError(f"Unknown standard deviation type: {self.noise_std_type}. Should be 'scalar' or 'log'")
# action distribution (populated in update_distribution)
self.distribution = None
# disable args validation for speedup
Normal.set_default_validate_args(False)
def reset(self, dones=None, hidden_states=None):
if hidden_states is None:
hidden_states = (None, None)
self.memory_s.reset(dones, hidden_states[0])
if self.teacher_recurrent:
self.memory_t.reset(dones, hidden_states[1])
def forward(self):
raise NotImplementedError
@property
def action_mean(self):
return self.distribution.mean
@property
def action_std(self):
return self.distribution.stddev
@property
def entropy(self):
return self.distribution.entropy().sum(dim=-1)
def update_distribution(self, obs):
# compute mean
mean = self.student(obs)
# compute standard deviation
if self.noise_std_type == "scalar":
std = self.std.expand_as(mean)
elif self.noise_std_type == "log":
std = torch.exp(self.log_std).expand_as(mean)
else:
raise ValueError(f"Unknown standard deviation type: {self.noise_std_type}. Should be 'scalar' or 'log'")
# create distribution
self.distribution = Normal(mean, std)
def act(self, obs):
obs = self.get_student_obs(obs)
obs = self.student_obs_normalizer(obs)
out_mem = self.memory_s(obs).squeeze(0)
self.update_distribution(out_mem)
return self.distribution.sample()
def act_inference(self, obs):
obs = self.get_student_obs(obs)
obs = self.student_obs_normalizer(obs)
out_mem = self.memory_s(obs).squeeze(0)
return self.student(out_mem)
def evaluate(self, obs):
obs = self.get_teacher_obs(obs)
obs = self.teacher_obs_normalizer(obs)
with torch.no_grad():
if self.teacher_recurrent:
self.memory_t.eval()
obs = self.memory_t(obs).squeeze(0)
return self.teacher(obs)
def get_student_obs(self, obs):
obs_list = []
for obs_group in self.obs_groups["policy"]:
obs_list.append(obs[obs_group])
return torch.cat(obs_list, dim=-1)
def get_teacher_obs(self, obs):
obs_list = []
for obs_group in self.obs_groups["teacher"]:
obs_list.append(obs[obs_group])
return torch.cat(obs_list, dim=-1)
def get_hidden_states(self):
if self.teacher_recurrent:
return self.memory_s.hidden_states, self.memory_t.hidden_states
else:
return self.memory_s.hidden_states, None
def detach_hidden_states(self, dones=None):
self.memory_s.detach_hidden_states(dones)
if self.teacher_recurrent:
self.memory_t.detach_hidden_states(dones)
def train(self, mode=True):
super().train(mode)
# make sure teacher is in eval mode
self.teacher.eval()
self.teacher_obs_normalizer.eval()
def update_normalization(self, obs):
if self.student_obs_normalization:
student_obs = self.get_student_obs(obs)
self.student_obs_normalizer.update(student_obs)
def load_state_dict(self, state_dict, strict=True):
"""Load the parameters of the student and teacher networks.
Args:
state_dict (dict): State dictionary of the model.
strict (bool): Whether to strictly enforce that the keys in state_dict match the keys returned by this
module's state_dict() function.
Returns:
bool: Whether this training resumes a previous training. This flag is used by the `load()` function of
`OnPolicyRunner` to determine how to load further parameters.
"""
# check if state_dict contains teacher and student or just teacher parameters
if any("actor" in key for key in state_dict.keys()): # loading parameters from rl training
# rename keys to match teacher and remove critic parameters
teacher_state_dict = {}
teacher_obs_normalizer_state_dict = {}
for key, value in state_dict.items():
if "actor." in key:
teacher_state_dict[key.replace("actor.", "")] = value
if "actor_obs_normalizer." in key:
teacher_obs_normalizer_state_dict[key.replace("actor_obs_normalizer.", "")] = value
self.teacher.load_state_dict(teacher_state_dict, strict=strict)
self.teacher_obs_normalizer.load_state_dict(teacher_obs_normalizer_state_dict, strict=strict)
# also load recurrent memory if teacher is recurrent
if self.teacher_recurrent:
memory_t_state_dict = {}
for key, value in state_dict.items():
if "memory_a." in key:
memory_t_state_dict[key.replace("memory_a.", "")] = value
self.memory_t.load_state_dict(memory_t_state_dict, strict=strict)
# set flag for successfully loading the parameters
self.loaded_teacher = True
self.teacher.eval()
self.teacher_obs_normalizer.eval()
return False # training does not resume
elif any("student" in key for key in state_dict.keys()): # loading parameters from distillation training
super().load_state_dict(state_dict, strict=strict)
# set flag for successfully loading the parameters
self.loaded_teacher = True
self.teacher.eval()
self.teacher_obs_normalizer.eval()
return True # training resumes
else:
raise ValueError("state_dict does not contain student or teacher parameters")

View File

@ -0,0 +1,24 @@
# Copyright (c) 2021-2025, ETH Zurich and NVIDIA CORPORATION
# All rights reserved.
#
# SPDX-License-Identifier: BSD-3-Clause
from __future__ import annotations
def resolve_symmetry_config(alg_cfg, env):
"""Resolve the symmetry configuration.
Args:
alg_cfg: The algorithm configuration dictionary.
env: The environment.
Returns:
The resolved algorithm configuration dictionary.
"""
# if using symmetry then pass the environment config object
if "symmetry_cfg" in alg_cfg and alg_cfg["symmetry_cfg"] is not None:
# this is used by the symmetry function for handling different observation terms
alg_cfg["symmetry_cfg"]["_env"] = env
return alg_cfg

View File

@ -0,0 +1,10 @@
# Copyright (c) 2021-2025, ETH Zurich and NVIDIA CORPORATION
# All rights reserved.
#
# SPDX-License-Identifier: BSD-3-Clause
"""Definitions for components of modules."""
from .memory import Memory
from .mlp import MLP
from .normalization import EmpiricalDiscountedVariationNormalization, EmpiricalNormalization

View File

@ -0,0 +1,70 @@
# Copyright (c) 2021-2025, ETH Zurich and NVIDIA CORPORATION
# All rights reserved.
#
# SPDX-License-Identifier: BSD-3-Clause
from __future__ import annotations
import torch.nn as nn
from rsl_rl.utils import unpad_trajectories
class Memory(nn.Module):
"""Memory module for recurrent networks.
This module is used to store the hidden states of the policy.
Currently only supports GRU and LSTM.
"""
def __init__(self, input_size, type="lstm", num_layers=1, hidden_size=256):
super().__init__()
# RNN
rnn_cls = nn.GRU if type.lower() == "gru" else nn.LSTM
self.rnn = rnn_cls(input_size=input_size, hidden_size=hidden_size, num_layers=num_layers)
self.hidden_states = None
def forward(self, input, masks=None, hidden_states=None):
batch_mode = masks is not None
if batch_mode:
# batch mode: needs saved hidden states
if hidden_states is None:
raise ValueError("Hidden states not passed to memory module during policy update")
out, _ = self.rnn(input, hidden_states)
out = unpad_trajectories(out, masks)
else:
# inference/distillation mode: uses hidden states of last step
out, self.hidden_states = self.rnn(input.unsqueeze(0), self.hidden_states)
return out
def reset(self, dones=None, hidden_states=None):
if dones is None: # reset all hidden states
if hidden_states is None:
self.hidden_states = None
else:
self.hidden_states = hidden_states
elif self.hidden_states is not None: # reset hidden states of done environments
if hidden_states is None:
if isinstance(self.hidden_states, tuple): # tuple in case of LSTM
for hidden_state in self.hidden_states:
hidden_state[..., dones == 1, :] = 0.0
else:
self.hidden_states[..., dones == 1, :] = 0.0
else:
NotImplementedError(
"Resetting hidden states of done environments with custom hidden states is not implemented"
)
def detach_hidden_states(self, dones=None):
if self.hidden_states is not None:
if dones is None: # detach all hidden states
if isinstance(self.hidden_states, tuple): # tuple in case of LSTM
self.hidden_states = tuple(hidden_state.detach() for hidden_state in self.hidden_states)
else:
self.hidden_states = self.hidden_states.detach()
else: # detach hidden states of done environments
if isinstance(self.hidden_states, tuple): # tuple in case of LSTM
for hidden_state in self.hidden_states:
hidden_state[..., dones == 1, :] = hidden_state[..., dones == 1, :].detach()
else:
self.hidden_states[..., dones == 1, :] = self.hidden_states[..., dones == 1, :].detach()

View File

@ -0,0 +1,120 @@
# Copyright (c) 2021-2025, ETH Zurich and NVIDIA CORPORATION
# All rights reserved.
#
# SPDX-License-Identifier: BSD-3-Clause
from __future__ import annotations
import torch
import torch.nn as nn
from functools import reduce
from rsl_rl.utils import resolve_nn_activation
class MLP(nn.Sequential):
"""Multi-layer perceptron.
The MLP network is a sequence of linear layers and activation functions. The
last layer is a linear layer that outputs the desired dimension unless the
last activation function is specified.
It provides additional conveniences:
- If the hidden dimensions have a value of ``-1``, the dimension is inferred
from the input dimension.
- If the output dimension is a tuple, the output is reshaped to the desired
shape.
"""
def __init__(
self,
input_dim: int,
output_dim: int | tuple[int] | list[int],
hidden_dims: tuple[int] | list[int],
activation: str = "elu",
last_activation: str | None = None,
):
"""Initialize the MLP.
Args:
input_dim: Dimension of the input.
output_dim: Dimension of the output.
hidden_dims: Dimensions of the hidden layers. A value of ``-1`` indicates
that the dimension should be inferred from the input dimension.
activation: Activation function. Defaults to "elu".
last_activation: Activation function of the last layer. Defaults to None,
in which case the last layer is linear.
"""
super().__init__()
# resolve activation functions
activation_mod = resolve_nn_activation(activation)
last_activation_mod = resolve_nn_activation(last_activation) if last_activation is not None else None
# resolve number of hidden dims if they are -1
hidden_dims_processed = [input_dim if dim == -1 else dim for dim in hidden_dims]
# create layers sequentially
layers = []
layers.append(nn.Linear(input_dim, hidden_dims_processed[0]))
layers.append(activation_mod)
for layer_index in range(len(hidden_dims_processed) - 1):
layers.append(nn.Linear(hidden_dims_processed[layer_index], hidden_dims_processed[layer_index + 1]))
layers.append(activation_mod)
# add last layer
if isinstance(output_dim, int):
layers.append(nn.Linear(hidden_dims_processed[-1], output_dim))
else:
# compute the total output dimension
total_out_dim = reduce(lambda x, y: x * y, output_dim)
# add a layer to reshape the output to the desired shape
layers.append(nn.Linear(hidden_dims_processed[-1], total_out_dim))
layers.append(nn.Unflatten(dim=-1, unflattened_size=output_dim))
# add last activation function if specified
if last_activation_mod is not None:
layers.append(last_activation_mod)
# register the layers
for idx, layer in enumerate(layers):
self.add_module(f"{idx}", layer)
def init_weights(self, scales: float | tuple[float]):
"""Initialize the weights of the MLP.
Args:
scales: Scale factor for the weights.
"""
def get_scale(idx) -> float:
"""Get the scale factor for the weights of the MLP.
Args:
idx: Index of the layer.
"""
return scales[idx] if isinstance(scales, (list, tuple)) else scales
# initialize the weights
for idx, module in enumerate(self):
if isinstance(module, nn.Linear):
nn.init.orthogonal_(module.weight, gain=get_scale(idx))
nn.init.zeros_(module.bias)
def forward(self, x: torch.Tensor) -> torch.Tensor:
"""Forward pass of the MLP.
Args:
x: Input tensor.
"""
for layer in self:
x = layer(x)
return x
def reset(self, dones=None, hidden_states=None):
pass
def detach_hidden_states(self, dones=None):
pass

View File

@ -0,0 +1,130 @@
# Copyright (c) 2021-2025, ETH Zurich and NVIDIA CORPORATION
# All rights reserved.
#
# SPDX-License-Identifier: BSD-3-Clause
# Copyright (c) 2020 Preferred Networks, Inc.
from __future__ import annotations
import torch
from torch import nn
class EmpiricalNormalization(nn.Module):
"""Normalize mean and variance of values based on empirical values."""
def __init__(self, shape, eps=1e-2, until=None):
"""Initialize EmpiricalNormalization module.
Args:
shape (int or tuple of int): Shape of input values except batch axis.
eps (float): Small value for stability.
until (int or None): If this arg is specified, the module learns input values until the sum of batch sizes
exceeds it.
Note: The normalization parameters are computed over the whole batch, not for each environment separately.
"""
super().__init__()
self.eps = eps
self.until = until
self.register_buffer("_mean", torch.zeros(shape).unsqueeze(0))
self.register_buffer("_var", torch.ones(shape).unsqueeze(0))
self.register_buffer("_std", torch.ones(shape).unsqueeze(0))
self.register_buffer("count", torch.tensor(0, dtype=torch.long))
@property
def mean(self):
return self._mean.squeeze(0).clone()
@property
def std(self):
return self._std.squeeze(0).clone()
def forward(self, x):
"""Normalize mean and variance of values based on empirical values."""
return (x - self._mean) / (self._std + self.eps)
@torch.jit.unused
def update(self, x):
"""Learn input values without computing the output values of them"""
if not self.training:
return
if self.until is not None and self.count >= self.until:
return
count_x = x.shape[0]
self.count += count_x
rate = count_x / self.count
var_x = torch.var(x, dim=0, unbiased=False, keepdim=True)
mean_x = torch.mean(x, dim=0, keepdim=True)
delta_mean = mean_x - self._mean
self._mean += rate * delta_mean
self._var += rate * (var_x - self._var + delta_mean * (mean_x - self._mean))
self._std = torch.sqrt(self._var)
@torch.jit.unused
def inverse(self, y):
"""De-normalize values based on empirical values."""
return y * (self._std + self.eps) + self._mean
class EmpiricalDiscountedVariationNormalization(nn.Module):
"""Reward normalization from Pathak's large scale study on PPO.
Reward normalization. Since the reward function is non-stationary, it is useful to normalize
the scale of the rewards so that the value function can learn quickly. We did this by dividing
the rewards by a running estimate of the standard deviation of the sum of discounted rewards.
"""
def __init__(self, shape, eps=1e-2, gamma=0.99, until=None):
super().__init__()
self.emp_norm = EmpiricalNormalization(shape, eps, until)
self.disc_avg = _DiscountedAverage(gamma)
def forward(self, rew):
if self.training:
# update discounted rewards
avg = self.disc_avg.update(rew)
# update moments from discounted rewards
self.emp_norm.update(avg)
# normalize rewards with the empirical std
if self.emp_norm._std > 0:
return rew / self.emp_norm._std
else:
return rew
"""
Helper class.
"""
class _DiscountedAverage:
r"""Discounted average of rewards.
The discounted average is defined as:
.. math::
\bar{R}_t = \gamma \bar{R}_{t-1} + r_t
Args:
gamma (float): Discount factor.
"""
def __init__(self, gamma):
self.avg = None
self.gamma = gamma
def update(self, rew: torch.Tensor) -> torch.Tensor:
if self.avg is None:
self.avg = rew
else:
self.avg = self.avg * self.gamma + rew
return self.avg

View File

@ -0,0 +1,12 @@
# Copyright (c) 2021-2025, ETH Zurich and NVIDIA CORPORATION
# All rights reserved.
#
# SPDX-License-Identifier: BSD-3-Clause
"""Implementation of runners for environment-agent interaction."""
from .on_policy_runner import OnPolicyRunner # isort:skip
from .distillation_runner import DistillationRunner
from .amp_on_policy_runner import AMPOnPolicyRunner
__all__ = ["OnPolicyRunner", "DistillationRunner", "AMPOnPolicyRunner"]

View File

@ -0,0 +1,521 @@
# Copyright (c) 2021-2025, ETH Zurich and NVIDIA CORPORATION
# All rights reserved.
#
# SPDX-License-Identifier: BSD-3-Clause
from __future__ import annotations
import os
import statistics
import time
import torch
import warnings
from collections import deque
import rsl_rl
from rsl_rl.algorithms import AMP_PPO
from rsl_rl.env import VecEnv
from rsl_rl.modules import ActorCritic, ActorCriticRecurrent,DiscriminatorMulti, resolve_rnd_config, resolve_symmetry_config
from rsl_rl.utils import resolve_obs_groups, store_code_state, Normalizer, G1_AMPLoader
class AMPOnPolicyRunner:
"""On-policy runner for training and evaluation of actor-critic methods."""
def __init__(self, env: VecEnv, train_cfg: dict, log_dir: str | None = None, device="cpu"):
self.cfg = train_cfg
self.alg_cfg = train_cfg["algorithm"]
self.policy_cfg = train_cfg["policy"]
self.device = device
self.env = env
# check if multi-gpu is enabled
self._configure_multi_gpu()
# store training configuration
self.num_steps_per_env = self.cfg["num_steps_per_env"]
self.save_interval = self.cfg["save_interval"]
# query observations from environment for algorithm construction
obs = self.env.get_observations()
default_sets = ["critic"]
if "rnd_cfg" in self.alg_cfg and self.alg_cfg["rnd_cfg"] is not None:
default_sets.append("rnd_state")
self.cfg["obs_groups"] = resolve_obs_groups(obs, self.cfg["obs_groups"], default_sets)
self.amp_data = G1_AMPLoader(
device,
time_between_frames=1/50.0,
preload_transitions=True,
num_preload_transitions=train_cfg["amp_num_preload_transitions"],
motion_files=train_cfg["amp_motion_files"],
num_frames=train_cfg['amp_num_frames']
)
self.amp_observation_dim = self.amp_data.observation_dim if self.cfg["amp_num_obs"] == 0 else self.cfg["amp_num_obs"]
self.amp_num_frames = 0 if self.cfg["amp_num_frames"] == 0 else self.cfg["amp_num_frames"]
self.amp_normalizer = Normalizer(self.amp_observation_dim)
self.discriminator = DiscriminatorMulti(
self.amp_observation_dim,
train_cfg["amp_reward_coef"],
train_cfg["amp_discr_hidden_dims"],
device,
train_cfg["amp_num_frames"],
train_cfg["amp_task_reward_lerp"],
train_cfg['use_lerp'],
).to(self.device)
# create the algorithm
self.alg = self._construct_algorithm(obs)
# Decide whether to disable logging
# We only log from the process with rank 0 (main process)
self.disable_logs = self.is_distributed and self.gpu_global_rank != 0
# Logging
self.log_dir = log_dir
self.writer = None
self.tot_timesteps = 0
self.tot_time = 0
self.current_learning_iteration = 0
self.git_status_repos = [rsl_rl.__file__]
def learn(self, num_learning_iterations: int, init_at_random_ep_len: bool = False): # noqa: C901
# initialize writer
self._prepare_logging_writer()
# randomize initial episode lengths (for exploration)
if init_at_random_ep_len:
self.env.episode_length_buf = torch.randint_like(
self.env.episode_length_buf, high=int(self.env.max_episode_length)
)
# start learning
obs = self.env.get_observations().to(self.device)
self.train_mode() # switch to train mode (for dropout for example)
amp_obs = self.env.get_amp_observations()
amp_obs = amp_obs.to(self.device)
if self.amp_num_frames != 0:
self.amp_obs_frames = torch.zeros(size=(self.env.num_envs, self.amp_num_frames, self.amp_observation_dim), device=self.device)
self.amp_obs_frames = torch.concat((self.amp_obs_frames[:, 1:], amp_obs.unsqueeze(1)), dim=1)
# Book keeping
ep_infos = []
rewbuffer = deque(maxlen=100)
lenbuffer = deque(maxlen=100)
step_discrewbuffer = deque(maxlen=100)
cur_reward_sum = torch.zeros(self.env.num_envs, dtype=torch.float, device=self.device)
cur_episode_length = torch.zeros(self.env.num_envs, dtype=torch.float, device=self.device)
cur_single_step_disc_rew = torch.zeros(self.env.num_envs, dtype=torch.float, device=self.device)
# create buffers for logging extrinsic and intrinsic rewards
if self.alg.rnd:
erewbuffer = deque(maxlen=100)
irewbuffer = deque(maxlen=100)
cur_ereward_sum = torch.zeros(self.env.num_envs, dtype=torch.float, device=self.device)
cur_ireward_sum = torch.zeros(self.env.num_envs, dtype=torch.float, device=self.device)
# Ensure all parameters are in-synced
if self.is_distributed:
print(f"Synchronizing parameters for rank {self.gpu_global_rank}...")
self.alg.broadcast_parameters()
# Start training
start_iter = self.current_learning_iteration
tot_iter = start_iter + num_learning_iterations
for it in range(start_iter, tot_iter):
start = time.time()
# Rollout
with torch.inference_mode():
for _ in range(self.num_steps_per_env):
# Sample actions
actions = self.alg.act(obs,amp_obs)
# Step the environment
obs, rewards, dones, extras = self.env.step(actions.to(self.env.device))
# Move to device
obs, rewards, dones = (obs.to(self.device), rewards.to(self.device), dones.to(self.device))
next_amp_obs = self.env.get_amp_observations()
next_amp_obs = next_amp_obs.to(self.device)
next_amp_obs_with_term = torch.clone(next_amp_obs)
reset_env_ids = self.env.reset_env_ids
terminal_amp_states = self.env.get_amp_observations()[reset_env_ids]
next_amp_obs_with_term[reset_env_ids] = terminal_amp_states
self.amp_obs_frames = torch.concat((self.amp_obs_frames[:, 1:], next_amp_obs_with_term.unsqueeze(1)), dim=1)
amp_reward = torch.zeros(self.env.num_envs, device=obs.device)
mask = self.env.contact_phase[:, 0] == 1.0
if mask.any():
rewards[mask], logit, disc_reward = self.alg.discriminator.predict_amp_reward(
self.amp_obs_frames[mask], rewards[mask], normalizer=self.alg.amp_normalizer
)
amp_reward[mask] += disc_reward
# process the step
self.alg.process_env_step(obs, rewards, dones, extras, next_amp_obs_with_term, self.amp_obs_frames)
self.amp_obs_frames[reset_env_ids] = 0
amp_obs = torch.clone(next_amp_obs)
# Extract intrinsic rewards (only for logging)
intrinsic_rewards = self.alg.intrinsic_rewards if self.alg.rnd else None
# book keeping
if self.log_dir is not None:
if "episode" in extras:
ep_infos.append(extras["episode"])
elif "log" in extras:
ep_infos.append(extras["log"])
# Update rewards
if self.alg.rnd:
cur_ereward_sum += rewards
cur_ireward_sum += intrinsic_rewards # type: ignore
cur_reward_sum += rewards + intrinsic_rewards
else:
cur_reward_sum += rewards
# Update episode length
cur_episode_length += 1
cur_single_step_disc_rew += amp_reward
# Clear data for completed episodes
# -- common
new_ids = (dones > 0).nonzero(as_tuple=False)
rewbuffer.extend(cur_reward_sum[new_ids][:, 0].cpu().numpy().tolist())
lenbuffer.extend(cur_episode_length[new_ids][:, 0].cpu().numpy().tolist())
cur_reward_sum[new_ids] = 0
cur_episode_length[new_ids] = 0
to_extend_disc = (cur_single_step_disc_rew[new_ids] / self.env.max_episode_length_s)[:, 0].cpu().numpy()
step_discrewbuffer.extend(to_extend_disc.tolist())
cur_single_step_disc_rew[new_ids] = 0
# -- intrinsic and extrinsic rewards
if self.alg.rnd:
erewbuffer.extend(cur_ereward_sum[new_ids][:, 0].cpu().numpy().tolist())
irewbuffer.extend(cur_ireward_sum[new_ids][:, 0].cpu().numpy().tolist())
cur_ereward_sum[new_ids] = 0
cur_ireward_sum[new_ids] = 0
stop = time.time()
collection_time = stop - start
start = stop
# compute returns
self.alg.compute_returns(obs)
# update policy
loss_dict = self.alg.update()
stop = time.time()
learn_time = stop - start
self.current_learning_iteration = it
# log info
if self.log_dir is not None and not self.disable_logs:
# Log information
self.log(locals())
# Save model
if it % self.save_interval == 0:
self.save(os.path.join(self.log_dir, f"model_{it}.pt"))
# Clear episode infos
ep_infos.clear()
# Save code state
if it == start_iter and not self.disable_logs:
# obtain all the diff files
git_file_paths = store_code_state(self.log_dir, self.git_status_repos)
# if possible store them to wandb
if self.logger_type in ["wandb", "neptune"] and git_file_paths:
for path in git_file_paths:
self.writer.save_file(path)
# Save the final model after training
if self.log_dir is not None and not self.disable_logs:
self.save(os.path.join(self.log_dir, f"model_{self.current_learning_iteration}.pt"))
def log(self, locs: dict, width: int = 80, pad: int = 35):
# Compute the collection size
collection_size = self.num_steps_per_env * self.env.num_envs * self.gpu_world_size
# Update total time-steps and time
self.tot_timesteps += collection_size
self.tot_time += locs["collection_time"] + locs["learn_time"]
iteration_time = locs["collection_time"] + locs["learn_time"]
# -- Episode info
ep_string = ""
if locs["ep_infos"]:
for key in locs["ep_infos"][0]:
infotensor = torch.tensor([], device=self.device)
for ep_info in locs["ep_infos"]:
# handle scalar and zero dimensional tensor infos
if key not in ep_info:
continue
if not isinstance(ep_info[key], torch.Tensor):
ep_info[key] = torch.Tensor([ep_info[key]])
if len(ep_info[key].shape) == 0:
ep_info[key] = ep_info[key].unsqueeze(0)
infotensor = torch.cat((infotensor, ep_info[key].to(self.device)))
value = torch.mean(infotensor)
# log to logger and terminal
if "/" in key:
self.writer.add_scalar(key, value, locs["it"])
ep_string += f"""{f'{key}:':>{pad}} {value:.4f}\n"""
else:
self.writer.add_scalar("Episode/" + key, value, locs["it"])
ep_string += f"""{f'Mean episode {key}:':>{pad}} {value:.4f}\n"""
mean_std = self.alg.policy.action_std.mean()
fps = int(collection_size / (locs["collection_time"] + locs["learn_time"]))
# -- Losses
for key, value in locs["loss_dict"].items():
self.writer.add_scalar(f"Loss/{key}", value, locs["it"])
self.writer.add_scalar("Loss/learning_rate", self.alg.learning_rate, locs["it"])
# -- Policy
self.writer.add_scalar("Policy/mean_noise_std", mean_std.item(), locs["it"])
# -- Performance
self.writer.add_scalar("Perf/total_fps", fps, locs["it"])
self.writer.add_scalar("Perf/collection time", locs["collection_time"], locs["it"])
self.writer.add_scalar("Perf/learning_time", locs["learn_time"], locs["it"])
# -- Training
if len(locs["rewbuffer"]) > 0:
# separate logging for intrinsic and extrinsic rewards
if hasattr(self.alg, "rnd") and self.alg.rnd:
self.writer.add_scalar("Rnd/mean_extrinsic_reward", statistics.mean(locs["erewbuffer"]), locs["it"])
self.writer.add_scalar("Rnd/mean_intrinsic_reward", statistics.mean(locs["irewbuffer"]), locs["it"])
self.writer.add_scalar("Rnd/weight", self.alg.rnd.weight, locs["it"])
# everything else
self.writer.add_scalar("Train/mean_reward", statistics.mean(locs["rewbuffer"]), locs["it"])
self.writer.add_scalar("Train/mean_episode_length", statistics.mean(locs["lenbuffer"]), locs["it"])
self.writer.add_scalar('Train/mean_step_disc_reward', statistics.mean(locs['step_discrewbuffer']), locs['it'])
if self.logger_type != "wandb": # wandb does not support non-integer x-axis logging
self.writer.add_scalar("Train/mean_reward/time", statistics.mean(locs["rewbuffer"]), self.tot_time)
self.writer.add_scalar(
"Train/mean_episode_length/time", statistics.mean(locs["lenbuffer"]), self.tot_time
)
str = f" \033[1m Learning iteration {locs['it']}/{locs['tot_iter']} \033[0m "
if len(locs["rewbuffer"]) > 0:
log_string = (
f"""{'#' * width}\n"""
f"""{str.center(width, ' ')}\n\n"""
f"""{'Computation:':>{pad}} {fps:.0f} steps/s (collection: {locs[
'collection_time']:.3f}s, learning {locs['learn_time']:.3f}s)\n"""
f"""{'Mean action noise std:':>{pad}} {mean_std.item():.2f}\n"""
f"""{'Step disc reward:':>{pad}} {statistics.mean(locs['step_discrewbuffer']):.2f}\n"""
)
# -- Losses
for key, value in locs["loss_dict"].items():
log_string += f"""{f'Mean {key} loss:':>{pad}} {value:.4f}\n"""
# -- Rewards
if hasattr(self.alg, "rnd") and self.alg.rnd:
log_string += (
f"""{'Mean extrinsic reward:':>{pad}} {statistics.mean(locs['erewbuffer']):.2f}\n"""
f"""{'Mean intrinsic reward:':>{pad}} {statistics.mean(locs['irewbuffer']):.2f}\n"""
)
log_string += f"""{'Mean reward:':>{pad}} {statistics.mean(locs['rewbuffer']):.2f}\n"""
# -- episode info
log_string += f"""{'Mean episode length:':>{pad}} {statistics.mean(locs['lenbuffer']):.2f}\n"""
else:
log_string = (
f"""{'#' * width}\n"""
f"""{str.center(width, ' ')}\n\n"""
f"""{'Computation:':>{pad}} {fps:.0f} steps/s (collection: {locs[
'collection_time']:.3f}s, learning {locs['learn_time']:.3f}s)\n"""
f"""{'Mean action noise std:':>{pad}} {mean_std.item():.2f}\n"""
)
for key, value in locs["loss_dict"].items():
log_string += f"""{f'{key}:':>{pad}} {value:.4f}\n"""
log_string += ep_string
log_string += (
f"""{'-' * width}\n"""
f"""{'Total timesteps:':>{pad}} {self.tot_timesteps}\n"""
f"""{'Iteration time:':>{pad}} {iteration_time:.2f}s\n"""
f"""{'Time elapsed:':>{pad}} {time.strftime("%H:%M:%S", time.gmtime(self.tot_time))}\n"""
f"""{'ETA:':>{pad}} {time.strftime(
"%H:%M:%S",
time.gmtime(
self.tot_time / (locs['it'] - locs['start_iter'] + 1)
* (locs['start_iter'] + locs['num_learning_iterations'] - locs['it'])
)
)}\n"""
)
print(log_string)
def save(self, path: str, infos=None):
# -- Save model
saved_dict = {
"model_state_dict": self.alg.policy.state_dict(),
"optimizer_state_dict": self.alg.optimizer.state_dict(),
"iter": self.current_learning_iteration,
"infos": infos,
}
# -- Save RND model if used
if hasattr(self.alg, "rnd") and self.alg.rnd:
saved_dict["rnd_state_dict"] = self.alg.rnd.state_dict()
saved_dict["rnd_optimizer_state_dict"] = self.alg.rnd_optimizer.state_dict()
torch.save(saved_dict, path)
# upload model to external logging service
if self.logger_type in ["neptune", "wandb"] and not self.disable_logs:
self.writer.save_model(path, self.current_learning_iteration)
def load(self, path: str, load_optimizer: bool = True, map_location: str | None = None):
loaded_dict = torch.load(path, weights_only=False, map_location=map_location)
# -- Load model
resumed_training = self.alg.policy.load_state_dict(loaded_dict["model_state_dict"])
# -- Load RND model if used
if hasattr(self.alg, "rnd") and self.alg.rnd:
self.alg.rnd.load_state_dict(loaded_dict["rnd_state_dict"])
# -- load optimizer if used
if load_optimizer and resumed_training:
# -- algorithm optimizer
self.alg.optimizer.load_state_dict(loaded_dict["optimizer_state_dict"])
# -- RND optimizer if used
if hasattr(self.alg, "rnd") and self.alg.rnd:
self.alg.rnd_optimizer.load_state_dict(loaded_dict["rnd_optimizer_state_dict"])
# -- load current learning iteration
if resumed_training:
self.current_learning_iteration = loaded_dict["iter"]
return loaded_dict["infos"]
def get_inference_policy(self, device=None):
self.eval_mode() # switch to evaluation mode (dropout for example)
if device is not None:
self.alg.policy.to(device)
return self.alg.policy.act_inference
def train_mode(self):
# -- PPO
self.alg.policy.train()
self.alg.discriminator.train()
# -- RND
if hasattr(self.alg, "rnd") and self.alg.rnd:
self.alg.rnd.train()
def eval_mode(self):
# -- PPO
self.alg.policy.eval()
self.alg.discriminator.eval()
# -- RND
if hasattr(self.alg, "rnd") and self.alg.rnd:
self.alg.rnd.eval()
def add_git_repo_to_log(self, repo_file_path):
self.git_status_repos.append(repo_file_path)
"""
Helper functions.
"""
def _configure_multi_gpu(self):
"""Configure multi-gpu training."""
# check if distributed training is enabled
self.gpu_world_size = int(os.getenv("WORLD_SIZE", "1"))
self.is_distributed = self.gpu_world_size > 1
# if not distributed training, set local and global rank to 0 and return
if not self.is_distributed:
self.gpu_local_rank = 0
self.gpu_global_rank = 0
self.multi_gpu_cfg = None
return
# get rank and world size
self.gpu_local_rank = int(os.getenv("LOCAL_RANK", "0"))
self.gpu_global_rank = int(os.getenv("RANK", "0"))
# make a configuration dictionary
self.multi_gpu_cfg = {
"global_rank": self.gpu_global_rank, # rank of the main process
"local_rank": self.gpu_local_rank, # rank of the current process
"world_size": self.gpu_world_size, # total number of processes
}
# check if user has device specified for local rank
if self.device != f"cuda:{self.gpu_local_rank}":
raise ValueError(
f"Device '{self.device}' does not match expected device for local rank '{self.gpu_local_rank}'."
)
# validate multi-gpu configuration
if self.gpu_local_rank >= self.gpu_world_size:
raise ValueError(
f"Local rank '{self.gpu_local_rank}' is greater than or equal to world size '{self.gpu_world_size}'."
)
if self.gpu_global_rank >= self.gpu_world_size:
raise ValueError(
f"Global rank '{self.gpu_global_rank}' is greater than or equal to world size '{self.gpu_world_size}'."
)
# initialize torch distributed
torch.distributed.init_process_group(backend="nccl", rank=self.gpu_global_rank, world_size=self.gpu_world_size)
# set device to the local rank
torch.cuda.set_device(self.gpu_local_rank)
def _construct_algorithm(self, obs) -> AMP_PPO:
"""Construct the actor-critic algorithm."""
# resolve RND config
self.alg_cfg = resolve_rnd_config(self.alg_cfg, obs, self.cfg["obs_groups"], self.env)
# resolve symmetry config
self.alg_cfg = resolve_symmetry_config(self.alg_cfg, self.env)
# resolve deprecated normalization config
if self.cfg.get("empirical_normalization") is not None:
warnings.warn(
"The `empirical_normalization` parameter is deprecated. Please set `actor_obs_normalization` and "
"`critic_obs_normalization` as part of the `policy` configuration instead.",
DeprecationWarning,
)
if self.policy_cfg.get("actor_obs_normalization") is None:
self.policy_cfg["actor_obs_normalization"] = self.cfg["empirical_normalization"]
if self.policy_cfg.get("critic_obs_normalization") is None:
self.policy_cfg["critic_obs_normalization"] = self.cfg["empirical_normalization"]
# initialize the actor-critic
actor_critic_class = eval(self.policy_cfg.pop("class_name"))
actor_critic: ActorCritic | ActorCriticRecurrent = actor_critic_class(
obs, self.cfg["obs_groups"], self.env.num_actions, **self.policy_cfg
).to(self.device)
# initialize the algorithm
alg_class = eval(self.alg_cfg.pop("class_name"))
alg: AMP_PPO = alg_class(actor_critic, self.discriminator, self.amp_data, self.amp_normalizer, self.amp_num_frames, device=self.device, **self.alg_cfg, multi_gpu_cfg=self.multi_gpu_cfg)
# initialize the storage
alg.init_storage(
"rl",
self.env.num_envs,
self.num_steps_per_env,
obs,
[self.env.num_actions],
)
return alg
def _prepare_logging_writer(self):
"""Prepares the logging writers."""
if self.log_dir is not None and self.writer is None and not self.disable_logs:
# Launch either Tensorboard or Neptune & Tensorboard summary writer(s), default: Tensorboard.
self.logger_type = self.cfg.get("logger", "tensorboard")
self.logger_type = self.logger_type.lower()
if self.logger_type == "neptune":
from rsl_rl.utils.neptune_utils import NeptuneSummaryWriter
self.writer = NeptuneSummaryWriter(log_dir=self.log_dir, flush_secs=10, cfg=self.cfg)
self.writer.log_config(self.env.cfg, self.cfg, self.alg_cfg, self.policy_cfg)
elif self.logger_type == "wandb":
from rsl_rl.utils.wandb_utils import WandbSummaryWriter
self.writer = WandbSummaryWriter(log_dir=self.log_dir, flush_secs=10, cfg=self.cfg)
self.writer.log_config(self.env.cfg, self.cfg, self.alg_cfg, self.policy_cfg)
elif self.logger_type == "tensorboard":
from torch.utils.tensorboard import SummaryWriter
self.writer = SummaryWriter(log_dir=self.log_dir, flush_secs=10)
else:
raise ValueError("Logger type not found. Please choose 'neptune', 'wandb' or 'tensorboard'.")

View File

@ -0,0 +1,179 @@
# Copyright (c) 2021-2025, ETH Zurich and NVIDIA CORPORATION
# All rights reserved.
#
# SPDX-License-Identifier: BSD-3-Clause
from __future__ import annotations
import os
import time
import torch
from collections import deque
import rsl_rl
from rsl_rl.algorithms import Distillation
from rsl_rl.env import VecEnv
from rsl_rl.modules import StudentTeacher, StudentTeacherRecurrent
from rsl_rl.runners import OnPolicyRunner
from rsl_rl.utils import resolve_obs_groups, store_code_state
class DistillationRunner(OnPolicyRunner):
"""On-policy runner for training and evaluation of teacher-student training."""
def __init__(self, env: VecEnv, train_cfg: dict, log_dir: str | None = None, device="cpu"):
self.cfg = train_cfg
self.alg_cfg = train_cfg["algorithm"]
self.policy_cfg = train_cfg["policy"]
self.device = device
self.env = env
# check if multi-gpu is enabled
self._configure_multi_gpu()
# store training configuration
self.num_steps_per_env = self.cfg["num_steps_per_env"]
self.save_interval = self.cfg["save_interval"]
# query observations from environment for algorithm construction
obs = self.env.get_observations()
self.cfg["obs_groups"] = resolve_obs_groups(obs, self.cfg["obs_groups"], default_sets=["teacher"])
# create the algorithm
self.alg = self._construct_algorithm(obs)
# Decide whether to disable logging
# We only log from the process with rank 0 (main process)
self.disable_logs = self.is_distributed and self.gpu_global_rank != 0
# Logging
self.log_dir = log_dir
self.writer = None
self.tot_timesteps = 0
self.tot_time = 0
self.current_learning_iteration = 0
self.git_status_repos = [rsl_rl.__file__]
def learn(self, num_learning_iterations: int, init_at_random_ep_len: bool = False): # noqa: C901
# initialize writer
self._prepare_logging_writer()
# check if teacher is loaded
if not self.alg.policy.loaded_teacher:
raise ValueError("Teacher model parameters not loaded. Please load a teacher model to distill.")
# randomize initial episode lengths (for exploration)
if init_at_random_ep_len:
self.env.episode_length_buf = torch.randint_like(
self.env.episode_length_buf, high=int(self.env.max_episode_length)
)
# start learning
obs = self.env.get_observations().to(self.device)
self.train_mode() # switch to train mode (for dropout for example)
# Book keeping
ep_infos = []
rewbuffer = deque(maxlen=100)
lenbuffer = deque(maxlen=100)
cur_reward_sum = torch.zeros(self.env.num_envs, dtype=torch.float, device=self.device)
cur_episode_length = torch.zeros(self.env.num_envs, dtype=torch.float, device=self.device)
# Ensure all parameters are in-synced
if self.is_distributed:
print(f"Synchronizing parameters for rank {self.gpu_global_rank}...")
self.alg.broadcast_parameters()
# Start training
start_iter = self.current_learning_iteration
tot_iter = start_iter + num_learning_iterations
for it in range(start_iter, tot_iter):
start = time.time()
# Rollout
with torch.inference_mode():
for _ in range(self.num_steps_per_env):
# Sample actions
actions = self.alg.act(obs)
# Step the environment
obs, rewards, dones, extras = self.env.step(actions.to(self.env.device))
# Move to device
obs, rewards, dones = (obs.to(self.device), rewards.to(self.device), dones.to(self.device))
# process the step
self.alg.process_env_step(obs, rewards, dones, extras)
# book keeping
if self.log_dir is not None:
if "episode" in extras:
ep_infos.append(extras["episode"])
elif "log" in extras:
ep_infos.append(extras["log"])
# Update rewards
cur_reward_sum += rewards
# Update episode length
cur_episode_length += 1
# Clear data for completed episodes
new_ids = (dones > 0).nonzero(as_tuple=False)
rewbuffer.extend(cur_reward_sum[new_ids][:, 0].cpu().numpy().tolist())
lenbuffer.extend(cur_episode_length[new_ids][:, 0].cpu().numpy().tolist())
cur_reward_sum[new_ids] = 0
cur_episode_length[new_ids] = 0
stop = time.time()
collection_time = stop - start
start = stop
# update policy
loss_dict = self.alg.update()
stop = time.time()
learn_time = stop - start
self.current_learning_iteration = it
# log info
if self.log_dir is not None and not self.disable_logs:
# Log information
self.log(locals())
# Save model
if it % self.save_interval == 0:
self.save(os.path.join(self.log_dir, f"model_{it}.pt"))
# Clear episode infos
ep_infos.clear()
# Save code state
if it == start_iter and not self.disable_logs:
# obtain all the diff files
git_file_paths = store_code_state(self.log_dir, self.git_status_repos)
# if possible store them to wandb
if self.logger_type in ["wandb", "neptune"] and git_file_paths:
for path in git_file_paths:
self.writer.save_file(path)
# Save the final model after training
if self.log_dir is not None and not self.disable_logs:
self.save(os.path.join(self.log_dir, f"model_{self.current_learning_iteration}.pt"))
"""
Helper methods.
"""
def _construct_algorithm(self, obs) -> Distillation:
"""Construct the distillation algorithm."""
# initialize the actor-critic
student_teacher_class = eval(self.policy_cfg.pop("class_name"))
student_teacher: StudentTeacher | StudentTeacherRecurrent = student_teacher_class(
obs, self.cfg["obs_groups"], self.env.num_actions, **self.policy_cfg
).to(self.device)
# initialize the algorithm
alg_class = eval(self.alg_cfg.pop("class_name"))
alg: Distillation = alg_class(
student_teacher, device=self.device, **self.alg_cfg, multi_gpu_cfg=self.multi_gpu_cfg
)
# initialize the storage
alg.init_storage(
"distillation",
self.env.num_envs,
self.num_steps_per_env,
obs,
[self.env.num_actions],
)
return alg

View File

@ -0,0 +1,460 @@
# Copyright (c) 2021-2025, ETH Zurich and NVIDIA CORPORATION
# All rights reserved.
#
# SPDX-License-Identifier: BSD-3-Clause
from __future__ import annotations
import os
import statistics
import time
import torch
import warnings
from collections import deque
import rsl_rl
from rsl_rl.algorithms import PPO
from rsl_rl.env import VecEnv
from rsl_rl.modules import ActorCritic, ActorCriticRecurrent, resolve_rnd_config, resolve_symmetry_config
from rsl_rl.utils import resolve_obs_groups, store_code_state
class OnPolicyRunner:
"""On-policy runner for training and evaluation of actor-critic methods."""
def __init__(self, env: VecEnv, train_cfg: dict, log_dir: str | None = None, device="cpu"):
self.cfg = train_cfg
self.alg_cfg = train_cfg["algorithm"]
self.policy_cfg = train_cfg["policy"]
self.device = device
self.env = env
# check if multi-gpu is enabled
self._configure_multi_gpu()
# store training configuration
self.num_steps_per_env = self.cfg["num_steps_per_env"]
self.save_interval = self.cfg["save_interval"]
# query observations from environment for algorithm construction
obs = self.env.get_observations()
default_sets = ["critic"]
if "rnd_cfg" in self.alg_cfg and self.alg_cfg["rnd_cfg"] is not None:
default_sets.append("rnd_state")
self.cfg["obs_groups"] = resolve_obs_groups(obs, self.cfg["obs_groups"], default_sets)
# create the algorithm
self.alg = self._construct_algorithm(obs)
# Decide whether to disable logging
# We only log from the process with rank 0 (main process)
self.disable_logs = self.is_distributed and self.gpu_global_rank != 0
# Logging
self.log_dir = log_dir
self.writer = None
self.tot_timesteps = 0
self.tot_time = 0
self.current_learning_iteration = 0
self.git_status_repos = [rsl_rl.__file__]
def learn(self, num_learning_iterations: int, init_at_random_ep_len: bool = False): # noqa: C901
# initialize writer
self._prepare_logging_writer()
# randomize initial episode lengths (for exploration)
if init_at_random_ep_len:
self.env.episode_length_buf = torch.randint_like(
self.env.episode_length_buf, high=int(self.env.max_episode_length)
)
# start learning
obs = self.env.get_observations().to(self.device)
self.train_mode() # switch to train mode (for dropout for example)
# Book keeping
ep_infos = []
rewbuffer = deque(maxlen=100)
lenbuffer = deque(maxlen=100)
cur_reward_sum = torch.zeros(self.env.num_envs, dtype=torch.float, device=self.device)
cur_episode_length = torch.zeros(self.env.num_envs, dtype=torch.float, device=self.device)
# create buffers for logging extrinsic and intrinsic rewards
if self.alg.rnd:
erewbuffer = deque(maxlen=100)
irewbuffer = deque(maxlen=100)
cur_ereward_sum = torch.zeros(self.env.num_envs, dtype=torch.float, device=self.device)
cur_ireward_sum = torch.zeros(self.env.num_envs, dtype=torch.float, device=self.device)
# Ensure all parameters are in-synced
if self.is_distributed:
print(f"Synchronizing parameters for rank {self.gpu_global_rank}...")
self.alg.broadcast_parameters()
# Start training
start_iter = self.current_learning_iteration
tot_iter = start_iter + num_learning_iterations
for it in range(start_iter, tot_iter):
start = time.time()
# Rollout
with torch.inference_mode():
for _ in range(self.num_steps_per_env):
# Sample actions
actions = self.alg.act(obs)
# Step the environment
obs, rewards, dones, extras = self.env.step(actions.to(self.env.device))
# Move to device
obs, rewards, dones = (obs.to(self.device), rewards.to(self.device), dones.to(self.device))
# process the step
self.alg.process_env_step(obs, rewards, dones, extras)
# Extract intrinsic rewards (only for logging)
intrinsic_rewards = self.alg.intrinsic_rewards if self.alg.rnd else None
# book keeping
if self.log_dir is not None:
if "episode" in extras:
ep_infos.append(extras["episode"])
elif "log" in extras:
ep_infos.append(extras["log"])
# Update rewards
if self.alg.rnd:
cur_ereward_sum += rewards
cur_ireward_sum += intrinsic_rewards # type: ignore
cur_reward_sum += rewards + intrinsic_rewards
else:
cur_reward_sum += rewards
# Update episode length
cur_episode_length += 1
# Clear data for completed episodes
# -- common
new_ids = (dones > 0).nonzero(as_tuple=False)
rewbuffer.extend(cur_reward_sum[new_ids][:, 0].cpu().numpy().tolist())
lenbuffer.extend(cur_episode_length[new_ids][:, 0].cpu().numpy().tolist())
cur_reward_sum[new_ids] = 0
cur_episode_length[new_ids] = 0
# -- intrinsic and extrinsic rewards
if self.alg.rnd:
erewbuffer.extend(cur_ereward_sum[new_ids][:, 0].cpu().numpy().tolist())
irewbuffer.extend(cur_ireward_sum[new_ids][:, 0].cpu().numpy().tolist())
cur_ereward_sum[new_ids] = 0
cur_ireward_sum[new_ids] = 0
stop = time.time()
collection_time = stop - start
start = stop
# compute returns
self.alg.compute_returns(obs)
# update policy
loss_dict = self.alg.update()
stop = time.time()
learn_time = stop - start
self.current_learning_iteration = it
# log info
if self.log_dir is not None and not self.disable_logs:
# Log information
self.log(locals())
# Save model
if it % self.save_interval == 0:
self.save(os.path.join(self.log_dir, f"model_{it}.pt"))
# Clear episode infos
ep_infos.clear()
# Save code state
if it == start_iter and not self.disable_logs:
# obtain all the diff files
git_file_paths = store_code_state(self.log_dir, self.git_status_repos)
# if possible store them to wandb
if self.logger_type in ["wandb", "neptune"] and git_file_paths:
for path in git_file_paths:
self.writer.save_file(path)
# Save the final model after training
if self.log_dir is not None and not self.disable_logs:
self.save(os.path.join(self.log_dir, f"model_{self.current_learning_iteration}.pt"))
def log(self, locs: dict, width: int = 80, pad: int = 35):
# Compute the collection size
collection_size = self.num_steps_per_env * self.env.num_envs * self.gpu_world_size
# Update total time-steps and time
self.tot_timesteps += collection_size
self.tot_time += locs["collection_time"] + locs["learn_time"]
iteration_time = locs["collection_time"] + locs["learn_time"]
# -- Episode info
ep_string = ""
if locs["ep_infos"]:
for key in locs["ep_infos"][0]:
infotensor = torch.tensor([], device=self.device)
for ep_info in locs["ep_infos"]:
# handle scalar and zero dimensional tensor infos
if key not in ep_info:
continue
if not isinstance(ep_info[key], torch.Tensor):
ep_info[key] = torch.Tensor([ep_info[key]])
if len(ep_info[key].shape) == 0:
ep_info[key] = ep_info[key].unsqueeze(0)
infotensor = torch.cat((infotensor, ep_info[key].to(self.device)))
value = torch.mean(infotensor)
# log to logger and terminal
if "/" in key:
self.writer.add_scalar(key, value, locs["it"])
ep_string += f"""{f'{key}:':>{pad}} {value:.4f}\n"""
else:
self.writer.add_scalar("Episode/" + key, value, locs["it"])
ep_string += f"""{f'Mean episode {key}:':>{pad}} {value:.4f}\n"""
mean_std = self.alg.policy.action_std.mean()
fps = int(collection_size / (locs["collection_time"] + locs["learn_time"]))
# -- Losses
for key, value in locs["loss_dict"].items():
self.writer.add_scalar(f"Loss/{key}", value, locs["it"])
self.writer.add_scalar("Loss/learning_rate", self.alg.learning_rate, locs["it"])
# -- Policy
self.writer.add_scalar("Policy/mean_noise_std", mean_std.item(), locs["it"])
# -- Performance
self.writer.add_scalar("Perf/total_fps", fps, locs["it"])
self.writer.add_scalar("Perf/collection time", locs["collection_time"], locs["it"])
self.writer.add_scalar("Perf/learning_time", locs["learn_time"], locs["it"])
# -- Training
if len(locs["rewbuffer"]) > 0:
# separate logging for intrinsic and extrinsic rewards
if hasattr(self.alg, "rnd") and self.alg.rnd:
self.writer.add_scalar("Rnd/mean_extrinsic_reward", statistics.mean(locs["erewbuffer"]), locs["it"])
self.writer.add_scalar("Rnd/mean_intrinsic_reward", statistics.mean(locs["irewbuffer"]), locs["it"])
self.writer.add_scalar("Rnd/weight", self.alg.rnd.weight, locs["it"])
# everything else
self.writer.add_scalar("Train/mean_reward", statistics.mean(locs["rewbuffer"]), locs["it"])
self.writer.add_scalar("Train/mean_episode_length", statistics.mean(locs["lenbuffer"]), locs["it"])
if self.logger_type != "wandb": # wandb does not support non-integer x-axis logging
self.writer.add_scalar("Train/mean_reward/time", statistics.mean(locs["rewbuffer"]), self.tot_time)
self.writer.add_scalar(
"Train/mean_episode_length/time", statistics.mean(locs["lenbuffer"]), self.tot_time
)
str = f" \033[1m Learning iteration {locs['it']}/{locs['tot_iter']} \033[0m "
if len(locs["rewbuffer"]) > 0:
log_string = (
f"""{'#' * width}\n"""
f"""{str.center(width, ' ')}\n\n"""
f"""{'Computation:':>{pad}} {fps:.0f} steps/s (collection: {locs[
'collection_time']:.3f}s, learning {locs['learn_time']:.3f}s)\n"""
f"""{'Mean action noise std:':>{pad}} {mean_std.item():.2f}\n"""
)
# -- Losses
for key, value in locs["loss_dict"].items():
log_string += f"""{f'Mean {key} loss:':>{pad}} {value:.4f}\n"""
# -- Rewards
if hasattr(self.alg, "rnd") and self.alg.rnd:
log_string += (
f"""{'Mean extrinsic reward:':>{pad}} {statistics.mean(locs['erewbuffer']):.2f}\n"""
f"""{'Mean intrinsic reward:':>{pad}} {statistics.mean(locs['irewbuffer']):.2f}\n"""
)
log_string += f"""{'Mean reward:':>{pad}} {statistics.mean(locs['rewbuffer']):.2f}\n"""
# -- episode info
log_string += f"""{'Mean episode length:':>{pad}} {statistics.mean(locs['lenbuffer']):.2f}\n"""
else:
log_string = (
f"""{'#' * width}\n"""
f"""{str.center(width, ' ')}\n\n"""
f"""{'Computation:':>{pad}} {fps:.0f} steps/s (collection: {locs[
'collection_time']:.3f}s, learning {locs['learn_time']:.3f}s)\n"""
f"""{'Mean action noise std:':>{pad}} {mean_std.item():.2f}\n"""
)
for key, value in locs["loss_dict"].items():
log_string += f"""{f'{key}:':>{pad}} {value:.4f}\n"""
log_string += ep_string
log_string += (
f"""{'-' * width}\n"""
f"""{'Total timesteps:':>{pad}} {self.tot_timesteps}\n"""
f"""{'Iteration time:':>{pad}} {iteration_time:.2f}s\n"""
f"""{'Time elapsed:':>{pad}} {time.strftime("%H:%M:%S", time.gmtime(self.tot_time))}\n"""
f"""{'ETA:':>{pad}} {time.strftime(
"%H:%M:%S",
time.gmtime(
self.tot_time / (locs['it'] - locs['start_iter'] + 1)
* (locs['start_iter'] + locs['num_learning_iterations'] - locs['it'])
)
)}\n"""
)
print(log_string)
def save(self, path: str, infos=None):
# -- Save model
saved_dict = {
"model_state_dict": self.alg.policy.state_dict(),
"optimizer_state_dict": self.alg.optimizer.state_dict(),
"iter": self.current_learning_iteration,
"infos": infos,
}
# -- Save RND model if used
if hasattr(self.alg, "rnd") and self.alg.rnd:
saved_dict["rnd_state_dict"] = self.alg.rnd.state_dict()
saved_dict["rnd_optimizer_state_dict"] = self.alg.rnd_optimizer.state_dict()
torch.save(saved_dict, path)
# upload model to external logging service
if self.logger_type in ["neptune", "wandb"] and not self.disable_logs:
self.writer.save_model(path, self.current_learning_iteration)
def load(self, path: str, load_optimizer: bool = True, map_location: str | None = None):
loaded_dict = torch.load(path, weights_only=False, map_location=map_location)
# -- Load model
resumed_training = self.alg.policy.load_state_dict(loaded_dict["model_state_dict"])
# -- Load RND model if used
if hasattr(self.alg, "rnd") and self.alg.rnd:
self.alg.rnd.load_state_dict(loaded_dict["rnd_state_dict"])
# -- load optimizer if used
if load_optimizer and resumed_training:
# -- algorithm optimizer
self.alg.optimizer.load_state_dict(loaded_dict["optimizer_state_dict"])
# -- RND optimizer if used
if hasattr(self.alg, "rnd") and self.alg.rnd:
self.alg.rnd_optimizer.load_state_dict(loaded_dict["rnd_optimizer_state_dict"])
# -- load current learning iteration
if resumed_training:
self.current_learning_iteration = loaded_dict["iter"]
return loaded_dict["infos"]
def get_inference_policy(self, device=None):
self.eval_mode() # switch to evaluation mode (dropout for example)
if device is not None:
self.alg.policy.to(device)
return self.alg.policy.act_inference
def train_mode(self):
# -- PPO
self.alg.policy.train()
# -- RND
if hasattr(self.alg, "rnd") and self.alg.rnd:
self.alg.rnd.train()
def eval_mode(self):
# -- PPO
self.alg.policy.eval()
# -- RND
if hasattr(self.alg, "rnd") and self.alg.rnd:
self.alg.rnd.eval()
def add_git_repo_to_log(self, repo_file_path):
self.git_status_repos.append(repo_file_path)
"""
Helper functions.
"""
def _configure_multi_gpu(self):
"""Configure multi-gpu training."""
# check if distributed training is enabled
self.gpu_world_size = int(os.getenv("WORLD_SIZE", "1"))
self.is_distributed = self.gpu_world_size > 1
# if not distributed training, set local and global rank to 0 and return
if not self.is_distributed:
self.gpu_local_rank = 0
self.gpu_global_rank = 0
self.multi_gpu_cfg = None
return
# get rank and world size
self.gpu_local_rank = int(os.getenv("LOCAL_RANK", "0"))
self.gpu_global_rank = int(os.getenv("RANK", "0"))
# make a configuration dictionary
self.multi_gpu_cfg = {
"global_rank": self.gpu_global_rank, # rank of the main process
"local_rank": self.gpu_local_rank, # rank of the current process
"world_size": self.gpu_world_size, # total number of processes
}
# check if user has device specified for local rank
if self.device != f"cuda:{self.gpu_local_rank}":
raise ValueError(
f"Device '{self.device}' does not match expected device for local rank '{self.gpu_local_rank}'."
)
# validate multi-gpu configuration
if self.gpu_local_rank >= self.gpu_world_size:
raise ValueError(
f"Local rank '{self.gpu_local_rank}' is greater than or equal to world size '{self.gpu_world_size}'."
)
if self.gpu_global_rank >= self.gpu_world_size:
raise ValueError(
f"Global rank '{self.gpu_global_rank}' is greater than or equal to world size '{self.gpu_world_size}'."
)
# initialize torch distributed
torch.distributed.init_process_group(backend="nccl", rank=self.gpu_global_rank, world_size=self.gpu_world_size)
# set device to the local rank
torch.cuda.set_device(self.gpu_local_rank)
def _construct_algorithm(self, obs) -> PPO:
"""Construct the actor-critic algorithm."""
# resolve RND config
self.alg_cfg = resolve_rnd_config(self.alg_cfg, obs, self.cfg["obs_groups"], self.env)
# resolve symmetry config
self.alg_cfg = resolve_symmetry_config(self.alg_cfg, self.env)
# resolve deprecated normalization config
if self.cfg.get("empirical_normalization") is not None:
warnings.warn(
"The `empirical_normalization` parameter is deprecated. Please set `actor_obs_normalization` and "
"`critic_obs_normalization` as part of the `policy` configuration instead.",
DeprecationWarning,
)
if self.policy_cfg.get("actor_obs_normalization") is None:
self.policy_cfg["actor_obs_normalization"] = self.cfg["empirical_normalization"]
if self.policy_cfg.get("critic_obs_normalization") is None:
self.policy_cfg["critic_obs_normalization"] = self.cfg["empirical_normalization"]
# initialize the actor-critic
actor_critic_class = eval(self.policy_cfg.pop("class_name"))
actor_critic: ActorCritic | ActorCriticRecurrent = actor_critic_class(
obs, self.cfg["obs_groups"], self.env.num_actions, **self.policy_cfg
).to(self.device)
# initialize the algorithm
alg_class = eval(self.alg_cfg.pop("class_name"))
alg: PPO = alg_class(actor_critic, device=self.device, **self.alg_cfg, multi_gpu_cfg=self.multi_gpu_cfg)
# initialize the storage
alg.init_storage(
"rl",
self.env.num_envs,
self.num_steps_per_env,
obs,
[self.env.num_actions],
)
return alg
def _prepare_logging_writer(self):
"""Prepares the logging writers."""
if self.log_dir is not None and self.writer is None and not self.disable_logs:
# Launch either Tensorboard or Neptune & Tensorboard summary writer(s), default: Tensorboard.
self.logger_type = self.cfg.get("logger", "tensorboard")
self.logger_type = self.logger_type.lower()
if self.logger_type == "neptune":
from rsl_rl.utils.neptune_utils import NeptuneSummaryWriter
self.writer = NeptuneSummaryWriter(log_dir=self.log_dir, flush_secs=10, cfg=self.cfg)
self.writer.log_config(self.env.cfg, self.cfg, self.alg_cfg, self.policy_cfg)
elif self.logger_type == "wandb":
from rsl_rl.utils.wandb_utils import WandbSummaryWriter
self.writer = WandbSummaryWriter(log_dir=self.log_dir, flush_secs=10, cfg=self.cfg)
self.writer.log_config(self.env.cfg, self.cfg, self.alg_cfg, self.policy_cfg)
elif self.logger_type == "tensorboard":
from torch.utils.tensorboard import SummaryWriter
self.writer = SummaryWriter(log_dir=self.log_dir, flush_secs=10)
else:
raise ValueError("Logger type not found. Please choose 'neptune', 'wandb' or 'tensorboard'.")

View File

@ -0,0 +1,10 @@
# Copyright (c) 2021-2025, ETH Zurich and NVIDIA CORPORATION
# All rights reserved.
#
# SPDX-License-Identifier: BSD-3-Clause
"""Implementation of transitions storage for RL-agent."""
from .rollout_storage import RolloutStorage
from .replay_buffer_multi import ReplayBufferMulti
__all__ = ["RolloutStorage", "ReplayBufferMulti"]

View File

@ -0,0 +1,38 @@
import torch
import numpy as np
class ReplayBufferMulti:
"""Fixed-size buffer to store experience tuples."""
def __init__(self, obs_dim, buffer_size, num_amp_frames, device):
"""Initialize a ReplayBuffer object.
Arguments:
buffer_size (int): maximum size of buffer
"""
self.states = torch.zeros(buffer_size, num_amp_frames, obs_dim).to(device)
self.num_amp_frames = num_amp_frames
self.buffer_size = buffer_size
self.device = device
self.step = 0
self.num_samples = 0
def insert(self, states):
"""Add new states to memory."""
num_states = states.shape[0]
start_idx = self.step
end_idx = self.step + num_states
if end_idx > self.buffer_size:
self.states[self.step:self.buffer_size] = states[:self.buffer_size - self.step]
self.states[:end_idx - self.buffer_size] = states[self.buffer_size - self.step:]
else:
self.states[start_idx:end_idx] = states
self.num_samples = min(self.buffer_size, max(end_idx, self.num_samples))
self.step = (self.step + num_states) % self.buffer_size
def feed_forward_generator(self, num_mini_batch, mini_batch_size):
for _ in range(num_mini_batch):
sample_idxs = np.random.choice(self.num_samples, size=mini_batch_size)
yield (self.states[sample_idxs].to(self.device))

View File

@ -0,0 +1,260 @@
# Copyright (c) 2021-2025, ETH Zurich and NVIDIA CORPORATION
# All rights reserved.
#
# SPDX-License-Identifier: BSD-3-Clause
from __future__ import annotations
import torch
from tensordict import TensorDict
from rsl_rl.utils import split_and_pad_trajectories
class RolloutStorage:
class Transition:
def __init__(self):
self.observations = None
self.actions = None
self.privileged_actions = None
self.rewards = None
self.dones = None
self.values = None
self.actions_log_prob = None
self.action_mean = None
self.action_sigma = None
self.hidden_states = None
def clear(self):
self.__init__()
def __init__(
self,
training_type,
num_envs,
num_transitions_per_env,
obs,
actions_shape,
device="cpu",
):
# store inputs
self.training_type = training_type
self.device = device
self.num_transitions_per_env = num_transitions_per_env
self.num_envs = num_envs
self.actions_shape = actions_shape
# Core
self.observations = TensorDict(
{key: torch.zeros(num_transitions_per_env, *value.shape, device=device) for key, value in obs.items()},
batch_size=[num_transitions_per_env, num_envs],
device=self.device,
)
self.rewards = torch.zeros(num_transitions_per_env, num_envs, 1, device=self.device)
self.actions = torch.zeros(num_transitions_per_env, num_envs, *actions_shape, device=self.device)
self.dones = torch.zeros(num_transitions_per_env, num_envs, 1, device=self.device).byte()
# for distillation
if training_type == "distillation":
self.privileged_actions = torch.zeros(num_transitions_per_env, num_envs, *actions_shape, device=self.device)
# for reinforcement learning
if training_type == "rl":
self.values = torch.zeros(num_transitions_per_env, num_envs, 1, device=self.device)
self.actions_log_prob = torch.zeros(num_transitions_per_env, num_envs, 1, device=self.device)
self.mu = torch.zeros(num_transitions_per_env, num_envs, *actions_shape, device=self.device)
self.sigma = torch.zeros(num_transitions_per_env, num_envs, *actions_shape, device=self.device)
self.returns = torch.zeros(num_transitions_per_env, num_envs, 1, device=self.device)
self.advantages = torch.zeros(num_transitions_per_env, num_envs, 1, device=self.device)
# For RNN networks
self.saved_hidden_states_a = None
self.saved_hidden_states_c = None
# counter for the number of transitions stored
self.step = 0
def add_transitions(self, transition: Transition):
# check if the transition is valid
if self.step >= self.num_transitions_per_env:
raise OverflowError("Rollout buffer overflow! You should call clear() before adding new transitions.")
# Core
self.observations[self.step].copy_(transition.observations)
self.actions[self.step].copy_(transition.actions)
self.rewards[self.step].copy_(transition.rewards.view(-1, 1))
self.dones[self.step].copy_(transition.dones.view(-1, 1))
# for distillation
if self.training_type == "distillation":
self.privileged_actions[self.step].copy_(transition.privileged_actions)
# for reinforcement learning
if self.training_type == "rl":
self.values[self.step].copy_(transition.values)
self.actions_log_prob[self.step].copy_(transition.actions_log_prob.view(-1, 1))
self.mu[self.step].copy_(transition.action_mean)
self.sigma[self.step].copy_(transition.action_sigma)
# For RNN networks
self._save_hidden_states(transition.hidden_states)
# increment the counter
self.step += 1
def _save_hidden_states(self, hidden_states):
if hidden_states is None or hidden_states == (None, None):
return
# make a tuple out of GRU hidden state sto match the LSTM format
hid_a = hidden_states[0] if isinstance(hidden_states[0], tuple) else (hidden_states[0],)
hid_c = hidden_states[1] if isinstance(hidden_states[1], tuple) else (hidden_states[1],)
# initialize if needed
if self.saved_hidden_states_a is None:
self.saved_hidden_states_a = [
torch.zeros(self.observations.shape[0], *hid_a[i].shape, device=self.device) for i in range(len(hid_a))
]
self.saved_hidden_states_c = [
torch.zeros(self.observations.shape[0], *hid_c[i].shape, device=self.device) for i in range(len(hid_c))
]
# copy the states
for i in range(len(hid_a)):
self.saved_hidden_states_a[i][self.step].copy_(hid_a[i])
self.saved_hidden_states_c[i][self.step].copy_(hid_c[i])
def clear(self):
self.step = 0
def compute_returns(self, last_values, gamma, lam, normalize_advantage: bool = True):
advantage = 0
for step in reversed(range(self.num_transitions_per_env)):
# if we are at the last step, bootstrap the return value
if step == self.num_transitions_per_env - 1:
next_values = last_values
else:
next_values = self.values[step + 1]
# 1 if we are not in a terminal state, 0 otherwise
next_is_not_terminal = 1.0 - self.dones[step].float()
# TD error: r_t + gamma * V(s_{t+1}) - V(s_t)
delta = self.rewards[step] + next_is_not_terminal * gamma * next_values - self.values[step]
# Advantage: A(s_t, a_t) = delta_t + gamma * lambda * A(s_{t+1}, a_{t+1})
advantage = delta + next_is_not_terminal * gamma * lam * advantage
# Return: R_t = A(s_t, a_t) + V(s_t)
self.returns[step] = advantage + self.values[step]
# Compute the advantages
self.advantages = self.returns - self.values
# Normalize the advantages if flag is set
# This is to prevent double normalization (i.e. if per minibatch normalization is used)
if normalize_advantage:
self.advantages = (self.advantages - self.advantages.mean()) / (self.advantages.std() + 1e-8)
# for distillation
def generator(self):
if self.training_type != "distillation":
raise ValueError("This function is only available for distillation training.")
for i in range(self.num_transitions_per_env):
yield self.observations[i], self.actions[i], self.privileged_actions[i], self.dones[i]
# for reinforcement learning with feedforward networks
def mini_batch_generator(self, num_mini_batches, num_epochs=8):
if self.training_type != "rl":
raise ValueError("This function is only available for reinforcement learning training.")
batch_size = self.num_envs * self.num_transitions_per_env
mini_batch_size = batch_size // num_mini_batches
indices = torch.randperm(num_mini_batches * mini_batch_size, requires_grad=False, device=self.device)
# Core
observations = self.observations.flatten(0, 1)
actions = self.actions.flatten(0, 1)
values = self.values.flatten(0, 1)
returns = self.returns.flatten(0, 1)
# For PPO
old_actions_log_prob = self.actions_log_prob.flatten(0, 1)
advantages = self.advantages.flatten(0, 1)
old_mu = self.mu.flatten(0, 1)
old_sigma = self.sigma.flatten(0, 1)
for epoch in range(num_epochs):
for i in range(num_mini_batches):
# Select the indices for the mini-batch
start = i * mini_batch_size
end = (i + 1) * mini_batch_size
batch_idx = indices[start:end]
# Create the mini-batch
# -- Core
obs_batch = observations[batch_idx]
actions_batch = actions[batch_idx]
# -- For PPO
target_values_batch = values[batch_idx]
returns_batch = returns[batch_idx]
old_actions_log_prob_batch = old_actions_log_prob[batch_idx]
advantages_batch = advantages[batch_idx]
old_mu_batch = old_mu[batch_idx]
old_sigma_batch = old_sigma[batch_idx]
# yield the mini-batch
yield obs_batch, actions_batch, target_values_batch, advantages_batch, returns_batch, old_actions_log_prob_batch, old_mu_batch, old_sigma_batch, (
None,
None,
), None
# for reinfrocement learning with recurrent networks
def recurrent_mini_batch_generator(self, num_mini_batches, num_epochs=8):
if self.training_type != "rl":
raise ValueError("This function is only available for reinforcement learning training.")
padded_obs_trajectories, trajectory_masks = split_and_pad_trajectories(self.observations, self.dones)
mini_batch_size = self.num_envs // num_mini_batches
for ep in range(num_epochs):
first_traj = 0
for i in range(num_mini_batches):
start = i * mini_batch_size
stop = (i + 1) * mini_batch_size
dones = self.dones.squeeze(-1)
last_was_done = torch.zeros_like(dones, dtype=torch.bool)
last_was_done[1:] = dones[:-1]
last_was_done[0] = True
trajectories_batch_size = torch.sum(last_was_done[:, start:stop])
last_traj = first_traj + trajectories_batch_size
masks_batch = trajectory_masks[:, first_traj:last_traj]
obs_batch = padded_obs_trajectories[:, first_traj:last_traj]
actions_batch = self.actions[:, start:stop]
old_mu_batch = self.mu[:, start:stop]
old_sigma_batch = self.sigma[:, start:stop]
returns_batch = self.returns[:, start:stop]
advantages_batch = self.advantages[:, start:stop]
values_batch = self.values[:, start:stop]
old_actions_log_prob_batch = self.actions_log_prob[:, start:stop]
# reshape to [num_envs, time, num layers, hidden dim] (original shape: [time, num_layers, num_envs, hidden_dim])
# then take only time steps after dones (flattens num envs and time dimensions),
# take a batch of trajectories and finally reshape back to [num_layers, batch, hidden_dim]
last_was_done = last_was_done.permute(1, 0)
hid_a_batch = [
saved_hidden_states.permute(2, 0, 1, 3)[last_was_done][first_traj:last_traj]
.transpose(1, 0)
.contiguous()
for saved_hidden_states in self.saved_hidden_states_a
]
hid_c_batch = [
saved_hidden_states.permute(2, 0, 1, 3)[last_was_done][first_traj:last_traj]
.transpose(1, 0)
.contiguous()
for saved_hidden_states in self.saved_hidden_states_c
]
# remove the tuple for GRU
hid_a_batch = hid_a_batch[0] if len(hid_a_batch) == 1 else hid_a_batch
hid_c_batch = hid_c_batch[0] if len(hid_c_batch) == 1 else hid_c_batch
yield obs_batch, actions_batch, values_batch, advantages_batch, returns_batch, old_actions_log_prob_batch, old_mu_batch, old_sigma_batch, (
hid_a_batch,
hid_c_batch,
), masks_batch
first_traj = last_traj

View File

@ -0,0 +1,13 @@
# Copyright (c) 2021-2025, ETH Zurich and NVIDIA CORPORATION
# All rights reserved.
#
# SPDX-License-Identifier: BSD-3-Clause
"""Helper functions."""
from .utils import *
from .motion_loader_g1 import G1_AMPLoader
__all__ = [
"G1_AMPLoader",
]

View File

@ -0,0 +1,388 @@
import os
from os.path import join as pjoin
import glob
import json
import logging
import torch
import numpy as np
from pybullet_utils import transformations
from rsl_rl.utils import motion_util
_EPS = np.finfo(float).eps * 4.0
def quaternion_slerp(q0, q1, fraction, spin=0, shortestpath=True):
"""Batch quaternion spherical linear interpolation."""
out = torch.zeros_like(q0)
zero_mask = torch.isclose(fraction, torch.zeros_like(fraction)).squeeze()
ones_mask = torch.isclose(fraction, torch.ones_like(fraction)).squeeze()
out[zero_mask] = q0[zero_mask]
out[ones_mask] = q1[ones_mask]
d = torch.sum(q0 * q1, dim=-1, keepdim=True)
dist_mask = (torch.abs(torch.abs(d) - 1.0) < _EPS).squeeze()
out[dist_mask] = q0[dist_mask]
if shortestpath:
d_old = torch.clone(d)
d = torch.where(d_old < 0, -d, d)
q1 = torch.where(d_old < 0, -q1, q1)
angle = torch.acos(d) + spin * torch.pi
angle_mask = (torch.abs(angle) < _EPS).squeeze()
out[angle_mask] = q0[angle_mask]
final_mask = torch.logical_or(zero_mask, ones_mask)
final_mask = torch.logical_or(final_mask, dist_mask)
final_mask = torch.logical_or(final_mask, angle_mask)
final_mask = torch.logical_not(final_mask)
isin = 1.0 / angle
q0 *= torch.sin((1.0 - fraction) * angle) * isin
q1 *= torch.sin(fraction * angle) * isin
q0 += q1
out[final_mask] = q0[final_mask]
return out
class G1_AMPLoader:
def __init__(
self,
device,
time_between_frames,
motion_files,
preload_transitions=False,
num_preload_transitions=1000000,
num_frames=5,
):
"""Expert dataset provides AMP observations from Dog mocap dataset.
time_between_frames: Amount of time in seconds between transition.
"""
self.device = device
self.time_between_frames = time_between_frames
self.num_frames = num_frames
# Values to store for each trajectory.
self.trajectories = []
self.trajectories_full = []
self.trajectory_names = []
self.trajectory_idxs = []
self.trajectory_lens = [] # Traj length in seconds.
self.trajectory_weights = []
self.trajectory_frame_durations = []
self.trajectory_num_frames = []
self.motion_dir = motion_files
# import ipdb; ipdb.set_trace()
for i, motion_file in enumerate(os.listdir(motion_files)):
self.trajectory_names.append(motion_file)
motion_path = pjoin(motion_files, motion_file)
motion_data = np.load(motion_path, allow_pickle=True)
motion_data_processed = np.zeros((motion_data.shape[0],36))
for f_i in range(motion_data.shape[0]):
motion_data_processed[f_i, :3] = motion_data[f_i, :3] # base pos
motion_data_processed[f_i, 3:7] = motion_data[f_i, 3:7] # base quat (wxyz)
motion_data_processed[f_i, 7:35] = motion_data[f_i, 7:35] # base vel
'''
NOTE The order of motion_data_processed is
base pos 0:3,
base quat 3:7, wxyz
dof pos 7:36, (mujoco joint order)
'''
self.trajectories.append(torch.tensor(
motion_data_processed[:, 7:],
dtype=torch.float32,
device=self.device
))
self.trajectories_full.append(torch.tensor(
motion_data_processed,
dtype=torch.float32,
device=self.device
))
self.trajectory_idxs.append(i)
self.trajectory_weights.append(1 / len(os.listdir(motion_files)))
frame_duration = 1 / 50
self.trajectory_frame_durations.append(frame_duration)
traj_len = (motion_data_processed.shape[0] - 1) * frame_duration # seconds
self.trajectory_lens.append(traj_len)
self.trajectory_num_frames.append(float(motion_data_processed.shape[0]))
print(f"Loaded {traj_len}s. motion from {motion_file}.")
# Trajectory weights are used to sample some trajectories more than others.
self.trajectory_weights = np.array(self.trajectory_weights) / np.sum(self.trajectory_weights)
self.trajectory_frame_durations = np.array(self.trajectory_frame_durations)
self.trajectory_lens = np.array(self.trajectory_lens)
self.trajectory_num_frames = np.array(self.trajectory_num_frames)
# Preload transitions.
self.preload_transitions = preload_transitions
if self.preload_transitions:
print(f'Preloading {num_preload_transitions} transitions')
traj_idxs = self.weighted_traj_idx_sample_batch(num_preload_transitions)
times = self.traj_time_sample_batch(traj_idxs)
self.preloaded_s_prior = self.get_full_frame_at_time_batch(traj_idxs, times - self.time_between_frames)
self.preloaded_s = self.get_full_frame_at_time_batch(traj_idxs, times)
self.preloaded_s_next = self.get_full_frame_at_time_batch(traj_idxs, times + self.time_between_frames)
print(f'Finished preloading')
# 预加载多帧数据
self.preloaded_frames = []
for i in range(self.num_frames):
frame_time = times + (i - (self.num_frames - 2)) * self.time_between_frames
full_frame = self.get_full_frame_at_time_batch(traj_idxs, frame_time)
# 预处理提前提取并连接需要的列7:26 和 29:33避免每次生成时重复切片
processed_frame = torch.cat([
full_frame[:, 7:26],
full_frame[:, 29:33]
], dim=-1)
self.preloaded_frames.append(processed_frame)
print(f'Finished preloading multiple frames')
self.all_trajectories_full = torch.vstack(self.trajectories_full)
def weighted_traj_idx_sample(self):
"""Get traj idx via weighted sampling."""
return np.random.choice(
self.trajectory_idxs, p=self.trajectory_weights)
def weighted_traj_idx_sample_batch(self, size):
"""Batch sample traj idxs."""
return np.random.choice(
self.trajectory_idxs, size=size, p=self.trajectory_weights,
replace=True)
def traj_time_sample(self, traj_idx):
"""Sample random time for traj."""
subst = self.time_between_frames + self.trajectory_frame_durations[traj_idx]
return max(
0, (self.trajectory_lens[traj_idx] * np.random.uniform() - subst))
def traj_time_sample_batch(self, traj_idxs):
"""Sample random time for multiple trajectories."""
subst = self.time_between_frames + self.trajectory_frame_durations[traj_idxs]
time_samples = self.trajectory_lens[traj_idxs] * np.random.uniform(size=len(traj_idxs)) - subst
return np.maximum(np.zeros_like(time_samples), time_samples)
def slerp(self, val0, val1, blend):
return (1.0 - blend) * val0 + blend * val1
def get_trajectory(self, traj_idx):
"""Returns trajectory of AMP observations."""
return self.trajectories_full[traj_idx]
def get_frame_at_time(self, traj_idx, time):
"""Returns frame for the given trajectory at the specified time."""
p = float(time) / self.trajectory_lens[traj_idx]
n = self.trajectories[traj_idx].shape[0]
idx_low, idx_high = int(np.floor(p * n)), int(np.ceil(p * n))
frame_start = self.trajectories[traj_idx][idx_low]
frame_end = self.trajectories[traj_idx][idx_high]
blend = p * n - idx_low
return self.slerp(frame_start, frame_end, blend)
def get_frame_at_time_batch(self, traj_idxs, times):
"""Returns frame for the given trajectory at the specified time."""
p = times / self.trajectory_lens[traj_idxs]
n = self.trajectory_num_frames[traj_idxs]
idx_low, idx_high = np.floor(p * n).astype(np.int32), np.ceil(p * n).astype(np.int32)
all_frame_starts = torch.zeros(len(traj_idxs), self.observation_dim, device=self.device)
all_frame_ends = torch.zeros(len(traj_idxs), self.observation_dim, device=self.device)
for traj_idx in set(traj_idxs):
trajectory = self.trajectories[traj_idx]
traj_mask = traj_idxs == traj_idx
all_frame_starts[traj_mask] = trajectory[idx_low[traj_mask]]
all_frame_ends[traj_mask] = trajectory[idx_high[traj_mask]]
blend = torch.tensor(p * n - idx_low, device=self.device, dtype=torch.float32).unsqueeze(-1)
return self.slerp(all_frame_starts, all_frame_ends, blend)
def get_full_frame_at_time(self, traj_idx, time):
"""Returns full frame for the given trajectory at the specified time."""
p = float(time) / self.trajectory_lens[traj_idx]
n = self.trajectories_full[traj_idx].shape[0]
idx_low, idx_high = int(np.floor(p * n)), int(np.ceil(p * n))
frame_start = self.trajectories_full[traj_idx][idx_low]
frame_end = self.trajectories_full[traj_idx][idx_high]
blend = p * n - idx_low
print(idx_low, idx_high)
return self.blend_frame_pose(frame_start, frame_end, blend)
def get_full_frame_at_time_batch(self, traj_idxs, times):
p = times / self.trajectory_lens[traj_idxs]
n = self.trajectory_num_frames[traj_idxs]
idx_low, idx_high = np.floor(p * n).astype(np.int32), np.ceil(p * n).astype(np.int32)
all_frame_pos_starts = torch.zeros(len(traj_idxs), 3, device=self.device)
all_frame_pos_ends = torch.zeros(len(traj_idxs), 3, device=self.device)
all_frame_rot_starts = torch.zeros(len(traj_idxs), 4, device=self.device)
all_frame_rot_ends = torch.zeros(len(traj_idxs), 4, device=self.device)
all_frame_amp_starts = torch.zeros(len(traj_idxs), 29, device=self.device)
all_frame_amp_ends = torch.zeros(len(traj_idxs), 29, device=self.device)
for traj_idx in set(traj_idxs):
trajectory = self.trajectories_full[traj_idx]
traj_mask = traj_idxs == traj_idx
all_frame_pos_starts[traj_mask] = G1_AMPLoader.get_root_pos_batch(trajectory[idx_low[traj_mask]])
all_frame_pos_ends[traj_mask] = G1_AMPLoader.get_root_pos_batch(trajectory[idx_high[traj_mask]])
all_frame_rot_starts[traj_mask] = G1_AMPLoader.get_root_rot_batch(trajectory[idx_low[traj_mask]])
all_frame_rot_ends[traj_mask] = G1_AMPLoader.get_root_rot_batch(trajectory[idx_high[traj_mask]])
all_frame_amp_starts[traj_mask] = trajectory[idx_low[traj_mask]][:, 7:36] # base vel3+ang3, dof vel23+ang23
all_frame_amp_ends[traj_mask] = trajectory[idx_high[traj_mask]][:, 7:36] # base vel3+ang3, dof vel23+ang23
blend = torch.tensor(p * n - idx_low, device=self.device, dtype=torch.float32).unsqueeze(-1)
pos_blend = self.slerp(all_frame_pos_starts, all_frame_pos_ends, blend)
rot_blend = quaternion_slerp(all_frame_rot_starts, all_frame_rot_ends, blend)
amp_blend = self.slerp(all_frame_amp_starts, all_frame_amp_ends, blend)
return torch.cat([pos_blend, rot_blend, amp_blend], dim=-1)
def get_frame(self):
"""Returns random frame."""
traj_idx = self.weighted_traj_idx_sample()
sampled_time = self.traj_time_sample(traj_idx)
return self.get_frame_at_time(traj_idx, sampled_time)
def get_full_frame(self):
"""Returns random full frame."""
traj_idx = self.weighted_traj_idx_sample()
sampled_time = self.traj_time_sample(traj_idx)
return self.get_full_frame_at_time(traj_idx, sampled_time)
def get_full_frame_batch(self, num_frames):
if self.preload_transitions:
idxs = np.random.choice(
self.preloaded_s.shape[0], size=num_frames)
return self.preloaded_s[idxs]
else:
traj_idxs = self.weighted_traj_idx_sample_batch(num_frames)
times = self.traj_time_sample_batch(traj_idxs)
return self.get_full_frame_at_time_batch(traj_idxs, times)
def blend_frame_pose(self, frame0, frame1, blend):
"""Linearly interpolate between two frames, including orientation.
Args:
frame0: First frame to be blended corresponds to (blend = 0).
frame1: Second frame to be blended corresponds to (blend = 1).
blend: Float between [0, 1], specifying the interpolation between
the two frames.
Returns:
An interpolation of the two frames.
"""
root_pos0, root_pos1 = G1_AMPLoader.get_root_pos(frame0), G1_AMPLoader.get_root_pos(frame1)
root_rot0, root_rot1 = G1_AMPLoader.get_root_rot(frame0), G1_AMPLoader.get_root_rot(frame1)
joints0, joints1 = G1_AMPLoader.get_joint_pose(frame0), G1_AMPLoader.get_joint_pose(frame1)
# tar_toe_pos_0, tar_toe_pos_1 = G1_AMPLoader.get_tar_toe_pos_local(frame0), G1_AMPLoader.get_tar_toe_pos_local(frame1)
linear_vel_0, linear_vel_1 = G1_AMPLoader.get_linear_vel(frame0), G1_AMPLoader.get_linear_vel(frame1)
angular_vel_0, angular_vel_1 = G1_AMPLoader.get_angular_vel(frame0), G1_AMPLoader.get_angular_vel(frame1)
joint_vel_0, joint_vel_1 = G1_AMPLoader.get_joint_vel(frame0), G1_AMPLoader.get_joint_vel(frame1)
blend_root_pos = self.slerp(root_pos0, root_pos1, blend)
blend_root_rot = transformations.quaternion_slerp(root_rot0.cpu().numpy(), root_rot1.cpu().numpy(), blend)
blend_root_rot = torch.tensor(motion_util.standardize_quaternion(blend_root_rot),dtype=torch.float32, device=self.device)
blend_joints = self.slerp(joints0, joints1, blend)
# blend_tar_toe_pos = self.slerp(tar_toe_pos_0, tar_toe_pos_1, blend)
blend_linear_vel = self.slerp(linear_vel_0, linear_vel_1, blend)
blend_angular_vel = self.slerp(angular_vel_0, angular_vel_1, blend)
blend_joints_vel = self.slerp(joint_vel_0, joint_vel_1, blend)
# return
# torch.cat([
# blend_root_pos, blend_root_rot, blend_linear_vel, blend_angular_vel, blend_joints, blend_joints_vel])
return torch.cat([blend_root_pos, blend_root_rot, blend_linear_vel, blend_angular_vel, blend_joints])
def feed_forward_generator_23dof_multi(self, num_mini_batch, mini_batch_size):
"""Generates a batch of AMP transitions."""
# import ipdb; ipdb.set_trace()
for _ in range(num_mini_batch):
if self.preload_transitions:
idxs = np.random.choice(self.preloaded_s.shape[0], size=mini_batch_size)
frames = []
for i in range(self.num_frames):
# 数据已在预加载时预处理,直接索引即可
s = self.preloaded_frames[i][idxs]
frames.append(s)
else:
NotImplementedError('preload transition')
yield torch.stack(frames, dim=1) # [batch, num_frames, 16]
def quaternion_to_euler_array(self, quat):
# Ensure quaternion is in the correct format [x, y, z, w]
x, y, z, w =quat
# Roll (x-axis rotation)
t0 = +2.0 * (w * x + y * z)
t1 = +1.0 - 2.0 * (x * x + y * y)
roll_x = np.arctan2(t0, t1)
# Pitch (y-axis rotation)
t2 = +2.0 * (w * y - z * x)
t2 = np.clip(t2, -1.0, 1.0)
pitch_y = np.arcsin(t2)
# Yaw (z-axis rotation)
t3 = +2.0 * (w * z + x * y)
t4 = +1.0 - 2.0 * (y * y + z * z)
yaw_z = np.arctan2(t3, t4)
# Returns roll, pitch, yaw in a NumPy array in radians
return np.array([roll_x, pitch_y, yaw_z])
def euler_to_quaternion(self, root_rot):
roll, pitch, yaw = root_rot[0], root_rot[1], root_rot[2]
cy = np.cos(yaw * 0.5)
sy = np.sin(yaw * 0.5)
cp = np.cos(pitch * 0.5)
sp = np.sin(pitch * 0.5)
cr = np.cos(roll * 0.5)
sr = np.sin(roll * 0.5)
qw = cy * cp * cr + sy * sp * sr
qx = cy * cp * sr - sy * sp * cr
qy = sy * cp * sr + cy * sp * cr
qz = sy * cp * cr - cy * sp * sr
return np.array([qx, qy, qz, qw])
@property
def observation_dim(self):
"""Size of AMP observations."""
return self.trajectories[0].shape[1] + 1
@property
def num_motions(self):
return len(self.trajectory_names)
@staticmethod
def get_root_pos(pose):
return pose[0:3]
@staticmethod
def get_root_pos_batch(poses):
return poses[:, 0:3]
@staticmethod
def get_root_rot(pose):
return pose[3:7]
@staticmethod
def get_root_rot_batch(poses):
return poses[:, 3:7]
@staticmethod
def get_joint_pose_batch_12dof(poses):
return poses[:, 13:25]
@staticmethod
def get_tar_toe_pos_local(pose):
return pose[G1_AMPLoader.TAR_TOE_POS_LOCAL_START_IDX:G1_AMPLoader.TAR_TOE_POS_LOCAL_END_IDX]
@staticmethod
def get_tar_toe_pos_local_batch(poses):
return poses[:, G1_AMPLoader.TAR_TOE_POS_LOCAL_START_IDX:G1_AMPLoader.TAR_TOE_POS_LOCAL_END_IDX]

View File

@ -0,0 +1,97 @@
# coding=utf-8
# Copyright 2020 The Google Research Authors.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
"""Utility functions for processing motion clips."""
import os
import inspect
currentdir = os.path.dirname(os.path.abspath(inspect.getfile(inspect.currentframe())))
parentdir = os.path.dirname(os.path.dirname(currentdir))
os.sys.path.insert(0, parentdir)
import numpy as np
from rsl_rl.utils import pose3d
# from pybullet_utils import transformations
def standardize_quaternion(q):
"""Returns a quaternion where q.w >= 0 to remove redundancy due to q = -q.
Args:
q: A quaternion to be standardized.
Returns:
A quaternion with q.w >= 0.
"""
if q[-1] < 0:
q = -q
return q
def normalize_rotation_angle(theta):
"""Returns a rotation angle normalized between [-pi, pi].
Args:
theta: angle of rotation (radians).
Returns:
An angle of rotation normalized between [-pi, pi].
"""
norm_theta = theta
if np.abs(norm_theta) > np.pi:
norm_theta = np.fmod(norm_theta, 2 * np.pi)
if norm_theta >= 0:
norm_theta += -2 * np.pi
else:
norm_theta += 2 * np.pi
return norm_theta
def calc_heading(q):
"""Returns the heading of a rotation q, specified as a quaternion.
The heading represents the rotational component of q along the vertical
axis (z axis).
Args:
q: A quaternion that the heading is to be computed from.
Returns:
An angle representing the rotation about the z axis.
"""
ref_dir = np.array([1, 0, 0])
rot_dir = pose3d.QuaternionRotatePoint(ref_dir, q)
heading = np.arctan2(rot_dir[1], rot_dir[0])
return heading
# def calc_heading_rot(q):
# """Return a quaternion representing the heading rotation of q along the vertical axis (z axis).
# Args:
# q: A quaternion that the heading is to be computed from.
# Returns:
# A quaternion representing the rotation about the z axis.
# """
# heading = calc_heading(q)
# q_heading = transformations.quaternion_about_axis(heading, [0, 0, 1])
# return q_heading

View File

@ -0,0 +1,94 @@
# Copyright (c) 2021-2025, ETH Zurich and NVIDIA CORPORATION
# All rights reserved.
#
# SPDX-License-Identifier: BSD-3-Clause
from __future__ import annotations
import os
from dataclasses import asdict
from torch.utils.tensorboard import SummaryWriter
try:
import neptune
except ModuleNotFoundError:
raise ModuleNotFoundError("neptune-client is required to log to Neptune.")
class NeptuneLogger:
def __init__(self, project, token):
self.run = neptune.init_run(project=project, api_token=token)
def store_config(self, env_cfg, runner_cfg, alg_cfg, policy_cfg):
self.run["runner_cfg"] = runner_cfg
self.run["policy_cfg"] = policy_cfg
self.run["alg_cfg"] = alg_cfg
self.run["env_cfg"] = asdict(env_cfg)
class NeptuneSummaryWriter(SummaryWriter):
"""Summary writer for Neptune."""
def __init__(self, log_dir: str, flush_secs: int, cfg):
super().__init__(log_dir, flush_secs)
try:
project = cfg["neptune_project"]
except KeyError:
raise KeyError("Please specify neptune_project in the runner config, e.g. legged_gym.")
try:
token = os.environ["NEPTUNE_API_TOKEN"]
except KeyError:
raise KeyError(
"Neptune api token not found. Please run or add to ~/.bashrc: export NEPTUNE_API_TOKEN=YOUR_API_TOKEN"
)
try:
entity = os.environ["NEPTUNE_USERNAME"]
except KeyError:
raise KeyError(
"Neptune username not found. Please run or add to ~/.bashrc: export NEPTUNE_USERNAME=YOUR_USERNAME"
)
neptune_project = entity + "/" + project
self.neptune_logger = NeptuneLogger(neptune_project, token)
self.name_map = {
"Train/mean_reward/time": "Train/mean_reward_time",
"Train/mean_episode_length/time": "Train/mean_episode_length_time",
}
run_name = os.path.split(log_dir)[-1]
self.neptune_logger.run["log_dir"].log(run_name)
def _map_path(self, path):
if path in self.name_map:
return self.name_map[path]
else:
return path
def add_scalar(self, tag, scalar_value, global_step=None, walltime=None, new_style=False):
super().add_scalar(
tag,
scalar_value,
global_step=global_step,
walltime=walltime,
new_style=new_style,
)
self.neptune_logger.run[self._map_path(tag)].log(scalar_value, step=global_step)
def stop(self):
self.neptune_logger.run.stop()
def log_config(self, env_cfg, runner_cfg, alg_cfg, policy_cfg):
self.neptune_logger.store_config(env_cfg, runner_cfg, alg_cfg, policy_cfg)
def save_model(self, model_path, iter):
self.neptune_logger.run["model/saved_model_" + str(iter)].upload(model_path)
def save_file(self, path, iter=None):
name = path.rsplit("/", 1)[-1].split(".")[0]
self.neptune_logger.run["git_diff/" + name].upload(path)

View File

@ -0,0 +1,283 @@
# coding=utf-8
# Copyright 2020 The Google Research Authors.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
"""Utilities for 3D pose conversion."""
import math
import numpy as np
# from pybullet_utils import transformations
VECTOR3_0 = np.zeros(3, dtype=np.float64)
VECTOR3_1 = np.ones(3, dtype=np.float64)
VECTOR3_X = np.array([1, 0, 0], dtype=np.float64)
VECTOR3_Y = np.array([0, 1, 0], dtype=np.float64)
VECTOR3_Z = np.array([0, 0, 1], dtype=np.float64)
# QUATERNION_IDENTITY is the multiplicative identity 1.0 + 0i + 0j + 0k.
# When interpreted as a rotation, it is the identity rotation.
QUATERNION_IDENTITY = np.array([0.0, 0.0, 0.0, 1.0], dtype=np.float64)
def Vector3RandomNormal(sigma, mu=VECTOR3_0):
"""Returns a random 3D vector from a normal distribution.
Each component is selected independently from a normal distribution.
Args:
sigma: Scale (or stddev) of distribution for all variables.
mu: Mean of distribution for each variable.
Returns:
A 3D vector in a numpy array.
"""
random_v3 = np.random.normal(scale=sigma, size=3) + mu
return random_v3
def Vector3RandomUniform(low=VECTOR3_0, high=VECTOR3_1):
"""Returns a 3D vector selected uniformly from the input box.
Args:
low: The min-value corner of the box.
high: The max-value corner of the box.
Returns:
A 3D vector in a numpy array.
"""
random_x = np.random.uniform(low=low[0], high=high[0])
random_y = np.random.uniform(low=low[1], high=high[1])
random_z = np.random.uniform(low=low[2], high=high[2])
return np.array([random_x, random_y, random_z])
def Vector3RandomUnit():
"""Returns a random 3D vector with unit length.
Generates a 3D vector selected uniformly from the unit sphere.
Returns:
A normalized 3D vector in a numpy array.
"""
longitude = np.random.uniform(low=-math.pi, high=math.pi)
sin_latitude = np.random.uniform(low=-1.0, high=1.0)
cos_latitude = math.sqrt(1.0 - sin_latitude * sin_latitude)
x = math.cos(longitude) * cos_latitude
y = math.sin(longitude) * cos_latitude
z = sin_latitude
return np.array([x, y, z], dtype=np.float64)
def QuaternionNormalize(q):
"""Normalizes the quaternion to length 1.
Divides the quaternion by its magnitude. If the magnitude is too
small, returns the quaternion identity value (1.0).
Args:
q: A quaternion to be normalized.
Raises:
ValueError: If input quaternion has length near zero.
Returns:
A quaternion with magnitude 1 in a numpy array [x, y, z, w].
"""
q_norm = np.linalg.norm(q)
if np.isclose(q_norm, 0.0):
raise ValueError(
'Quaternion may not be zero in QuaternionNormalize: |q| = %f, q = %s' %
(q_norm, q))
return q / q_norm
def QuaternionFromAxisAngle(axis, angle):
"""Returns a quaternion that generates the given axis-angle rotation.
Returns the quaternion: sin(angle/2) * axis + cos(angle/2).
Args:
axis: Axis of rotation, a 3D vector in a numpy array.
angle: The angle of rotation (radians).
Raises:
ValueError: If input axis is not a normalizable 3D vector.
Returns:
A unit quaternion in a numpy array.
"""
if len(axis) != 3:
raise ValueError('Axis vector should have three components: %s' % axis)
axis_norm = np.linalg.norm(axis)
if np.isclose(axis_norm, 0.0):
raise ValueError('Axis vector may not have zero length: |v| = %f, v = %s' %
(axis_norm, axis))
half_angle = angle * 0.5
q = np.zeros(4, dtype=np.float64)
q[0:3] = axis
q[0:3] *= math.sin(half_angle) / axis_norm
q[3] = math.cos(half_angle)
return q
def QuaternionToAxisAngle(quat, default_axis=VECTOR3_Z, direction_axis=None):
"""Calculates axis and angle of rotation performed by a quaternion.
Calculates the axis and angle of the rotation performed by the quaternion.
The quaternion should have four values and be normalized.
Args:
quat: Unit quaternion in a numpy array.
default_axis: 3D vector axis used if the rotation is near to zero. Without
this default, small rotations would result in an exception. It is
reasonable to use a default axis for tiny rotations, because zero angle
rotations about any axis are equivalent.
direction_axis: Used to disambiguate rotation directions. If the
direction_axis is specified, the axis of the rotation will be chosen such
that its inner product with the direction_axis is non-negative.
Raises:
ValueError: If quat is not a normalized quaternion.
Returns:
axis: Axis of rotation.
angle: Angle in radians.
"""
if len(quat) != 4:
raise ValueError(
'Quaternion should have four components [x, y, z, w]: %s' % quat)
if not np.isclose(1.0, np.linalg.norm(quat)):
raise ValueError('Quaternion should have unit length: |q| = %f, q = %s' %
(np.linalg.norm(quat), quat))
axis = quat[:3].copy()
axis_norm = np.linalg.norm(axis)
min_axis_norm = 1e-8
if axis_norm < min_axis_norm:
axis = default_axis
if len(default_axis) != 3:
raise ValueError('Axis vector should have three components: %s' % axis)
if not np.isclose(np.linalg.norm(axis), 1.0):
raise ValueError('Axis vector should have unit length: |v| = %f, v = %s' %
(np.linalg.norm(axis), axis))
else:
axis /= axis_norm
sin_half_angle = axis_norm
if direction_axis is not None and np.inner(axis, direction_axis) < 0:
sin_half_angle = -sin_half_angle
axis = -axis
cos_half_angle = quat[3]
half_angle = math.atan2(sin_half_angle, cos_half_angle)
angle = half_angle * 2
return axis, angle
def QuaternionRandomRotation(max_angle=math.pi):
"""Creates a random small rotation around a random axis.
Generates a small rotation with the axis vector selected uniformly
from the unit sphere and an angle selected from a uniform
distribution over [0, max_angle].
If the max_angle is not specified, the rotation should be selected
uniformly over all possible rotation angles.
Args:
max_angle: The maximum angle of rotation (radians).
Returns:
A unit quaternion in a numpy array.
"""
angle = np.random.uniform(low=0, high=max_angle)
axis = Vector3RandomUnit()
return QuaternionFromAxisAngle(axis, angle)
# def QuaternionRotatePoint(point, quat):
# """Performs a rotation by quaternion.
# Rotate the point by the quaternion using quaternion multiplication,
# (q * p * q^-1), without constructing the rotation matrix.
# Args:
# point: The point to be rotated.
# quat: The rotation represented as a quaternion [x, y, z, w].
# Returns:
# A 3D vector in a numpy array.
# """
# q_point = np.array([point[0], point[1], point[2], 0.0])
# quat_inverse = transformations.quaternion_inverse(quat)
# q_point_rotated = transformations.quaternion_multiply(
# transformations.quaternion_multiply(quat, q_point), quat_inverse)
# return q_point_rotated[:3]
def IsRotationMatrix(m):
"""Returns true if the 3x3 submatrix represents a rotation.
Args:
m: A transformation matrix.
Raises:
ValueError: If input is not a matrix of size at least 3x3.
Returns:
True if the 3x3 submatrix is a rotation (orthogonal).
"""
if len(m.shape) != 2 or m.shape[0] < 3 or m.shape[1] < 3:
raise ValueError('Matrix should be 3x3 or 4x4: %s\n %s' % (m.shape, m))
rot = m[:3, :3]
eye = np.matmul(rot, np.transpose(rot))
return np.isclose(eye, np.identity(3), atol=1e-4).all()
# def ZAxisAlignedRobotPoseTool(robot_pose_tool):
# """Returns the current gripper pose rotated for alignment with the z-axis.
# Args:
# robot_pose_tool: a pose3d.Pose3d() instance.
# Returns:
# An instance of pose.Transform representing the current gripper pose
# rotated for alignment with the z-axis.
# """
# # Align the current pose to the z-axis.
# robot_pose_tool.quaternion = transformations.quaternion_multiply(
# RotationBetween(
# robot_pose_tool.matrix4x4[0:3, 0:3].dot(np.array([0, 0, 1])),
# np.array([0.0, 0.0, -1.0])), robot_pose_tool.quaternion)
# return robot_pose_tool
# def RotationBetween(a_translation_b, a_translation_c):
# """Computes the rotation from one vector to another.
# The computed rotation has the property that:
# a_translation_c = a_rotation_b_to_c * a_translation_b
# Args:
# a_translation_b: vec3, vector to rotate from
# a_translation_c: vec3, vector to rotate to
# Returns:
# a_rotation_b_to_c: new Orientation
# """
# rotation = rotation3.Rotation3.rotation_between(
# a_translation_b, a_translation_c, err_msg='RotationBetween')
# return rotation.quaternion.xyzw

View File

@ -0,0 +1,360 @@
# Copyright (c) 2021-2025, ETH Zurich and NVIDIA CORPORATION
# All rights reserved.
#
# SPDX-License-Identifier: BSD-3-Clause
from __future__ import annotations
import git
import importlib
import os
import pathlib
import torch
import warnings
from tensordict import TensorDict
from typing import Callable
import numpy as np
class RunningMeanStd:
def __init__(self, epsilon: float = 1e-4, shape: Tuple[int, ...] = ()):
"""
Calculates the running mean and std of a data stream
https://en.wikipedia.org/wiki/Algorithms_for_calculating_variance#Parallel_algorithm
:param epsilon: helps with arithmetic issues
:param shape: the shape of the data stream's output
"""
self.mean = np.zeros(shape, np.float64)
self.var = np.ones(shape, np.float64)
self.count = epsilon
def update(self, arr: np.ndarray) -> None:
batch_mean = np.mean(arr, axis=0)
batch_var = np.var(arr, axis=0)
batch_count = arr.shape[0]
self.update_from_moments(batch_mean, batch_var, batch_count)
def update_from_moments(self, batch_mean: np.ndarray, batch_var: np.ndarray, batch_count: int) -> None:
delta = batch_mean - self.mean
tot_count = self.count + batch_count
new_mean = self.mean + delta * batch_count / tot_count
m_a = self.var * self.count
m_b = batch_var * batch_count
m_2 = m_a + m_b + np.square(delta) * self.count * batch_count / (self.count + batch_count)
new_var = m_2 / (self.count + batch_count)
new_count = batch_count + self.count
self.mean = new_mean
self.var = new_var
self.count = new_count
class Normalizer(RunningMeanStd):
def __init__(self, input_dim, epsilon=1e-4, clip_obs=10.0):
super().__init__(shape=input_dim)
self.epsilon = epsilon
self.clip_obs = clip_obs
def normalize(self, input):
return np.clip((input - self.mean) / np.sqrt(self.var + self.epsilon), -self.clip_obs, self.clip_obs)
def normalize_torch(self, input, device):
mean_torch = torch.tensor(self.mean, device=device, dtype=torch.float32)
std_torch = torch.sqrt(torch.tensor(self.var + self.epsilon, device=device, dtype=torch.float32))
return torch.clamp((input - mean_torch) / std_torch, -self.clip_obs, self.clip_obs)
def update_normalizer(self, rollouts, expert_loader):
policy_data_generator = rollouts.feed_forward_generator_amp(None, mini_batch_size=expert_loader.batch_size)
expert_data_generator = expert_loader.dataset.feed_forward_generator_amp(expert_loader.batch_size)
for expert_batch, policy_batch in zip(expert_data_generator, policy_data_generator):
self.update(torch.vstack(tuple(policy_batch) + tuple(expert_batch)).cpu().numpy())
def resolve_nn_activation(act_name: str) -> torch.nn.Module:
"""Resolves the activation function from the name.
Args:
act_name: The name of the activation function.
Returns:
The activation function.
Raises:
ValueError: If the activation function is not found.
"""
act_dict = {
"elu": torch.nn.ELU(),
"selu": torch.nn.SELU(),
"relu": torch.nn.ReLU(),
"crelu": torch.nn.CELU(),
"lrelu": torch.nn.LeakyReLU(),
"tanh": torch.nn.Tanh(),
"sigmoid": torch.nn.Sigmoid(),
"softplus": torch.nn.Softplus(),
"gelu": torch.nn.GELU(),
"swish": torch.nn.SiLU(),
"mish": torch.nn.Mish(),
"identity": torch.nn.Identity(),
}
act_name = act_name.lower()
if act_name in act_dict:
return act_dict[act_name]
else:
raise ValueError(f"Invalid activation function '{act_name}'. Valid activations are: {list(act_dict.keys())}")
def resolve_optimizer(optimizer_name: str) -> torch.optim.Optimizer:
"""Resolves the optimizer from the name.
Args:
optimizer_name: The name of the optimizer.
Returns:
The optimizer.
Raises:
ValueError: If the optimizer is not found.
"""
optimizer_dict = {
"adam": torch.optim.Adam,
"adamw": torch.optim.AdamW,
"sgd": torch.optim.SGD,
"rmsprop": torch.optim.RMSprop,
}
optimizer_name = optimizer_name.lower()
if optimizer_name in optimizer_dict:
return optimizer_dict[optimizer_name]
else:
raise ValueError(f"Invalid optimizer '{optimizer_name}'. Valid optimizers are: {list(optimizer_dict.keys())}")
def split_and_pad_trajectories(
tensor: torch.Tensor | TensorDict, dones: torch.Tensor
) -> tuple[torch.Tensor | TensorDict, torch.Tensor]:
"""Splits trajectories at done indices. Then concatenates them and pads with zeros up to the length of the longest
trajectory. Returns masks corresponding to valid parts of the trajectories.
Example:
Input: [[a1, a2, a3, a4 | a5, a6],
[b1, b2 | b3, b4, b5 | b6]]
Output:[[a1, a2, a3, a4], | [[True, True, True, True],
[a5, a6, 0, 0], | [True, True, False, False],
[b1, b2, 0, 0], | [True, True, False, False],
[b3, b4, b5, 0], | [True, True, True, False],
[b6, 0, 0, 0]] | [True, False, False, False]]
Assumes that the input has the following order of dimensions: [time, number of envs, additional dimensions]
"""
dones = dones.clone()
dones[-1] = 1
# Permute the buffers to have order (num_envs, num_transitions_per_env, ...), for correct reshaping
flat_dones = dones.transpose(1, 0).reshape(-1, 1)
# Get length of trajectory by counting the number of successive not done elements
done_indices = torch.cat((flat_dones.new_tensor([-1], dtype=torch.int64), flat_dones.nonzero()[:, 0]))
trajectory_lengths = done_indices[1:] - done_indices[:-1]
trajectory_lengths_list = trajectory_lengths.tolist()
# Extract the individual trajectories
if isinstance(tensor, TensorDict):
padded_trajectories = {}
for k, v in tensor.items():
# split the tensor into trajectories
trajectories = torch.split(v.transpose(1, 0).flatten(0, 1), trajectory_lengths_list)
# add at least one full length trajectory
trajectories = trajectories + (torch.zeros(v.shape[0], *v.shape[2:], device=v.device),)
# pad the trajectories to the length of the longest trajectory
padded_trajectories[k] = torch.nn.utils.rnn.pad_sequence(trajectories)
# remove the added tensor
padded_trajectories[k] = padded_trajectories[k][:, :-1]
padded_trajectories = TensorDict(
padded_trajectories, batch_size=[tensor.batch_size[0], len(trajectory_lengths_list)]
)
else:
# split the tensor into trajectories
trajectories = torch.split(tensor.transpose(1, 0).flatten(0, 1), trajectory_lengths_list)
# add at least one full length trajectory
trajectories = trajectories + (torch.zeros(tensor.shape[0], *tensor.shape[2:], device=tensor.device),)
# pad the trajectories to the length of the longest trajectory
padded_trajectories = torch.nn.utils.rnn.pad_sequence(trajectories)
# remove the added tensor
padded_trajectories = padded_trajectories[:, :-1]
# create masks for the valid parts of the trajectories
trajectory_masks = trajectory_lengths > torch.arange(0, tensor.shape[0], device=tensor.device).unsqueeze(1)
return padded_trajectories, trajectory_masks
def unpad_trajectories(trajectories, masks):
"""Does the inverse operation of split_and_pad_trajectories()"""
# Need to transpose before and after the masking to have proper reshaping
return (
trajectories.transpose(1, 0)[masks.transpose(1, 0)]
.view(-1, trajectories.shape[0], trajectories.shape[-1])
.transpose(1, 0)
)
def store_code_state(logdir, repositories) -> list:
git_log_dir = os.path.join(logdir, "git")
os.makedirs(git_log_dir, exist_ok=True)
file_paths = []
for repository_file_path in repositories:
try:
repo = git.Repo(repository_file_path, search_parent_directories=True)
t = repo.head.commit.tree
except Exception:
print(f"Could not find git repository in {repository_file_path}. Skipping.")
# skip if not a git repository
continue
# get the name of the repository
repo_name = pathlib.Path(repo.working_dir).name
diff_file_name = os.path.join(git_log_dir, f"{repo_name}.diff")
# check if the diff file already exists
if os.path.isfile(diff_file_name):
continue
# write the diff file
print(f"Storing git diff for '{repo_name}' in: {diff_file_name}")
with open(diff_file_name, "x", encoding="utf-8") as f:
content = f"--- git status ---\n{repo.git.status()} \n\n\n--- git diff ---\n{repo.git.diff(t)}"
f.write(content)
# add the file path to the list of files to be uploaded
file_paths.append(diff_file_name)
return file_paths
def string_to_callable(name: str) -> Callable:
"""Resolves the module and function names to return the function.
Args:
name: The function name. The format should be 'module:attribute_name'.
Raises:
ValueError: When the resolved attribute is not a function.
ValueError: When unable to resolve the attribute.
Returns:
The function loaded from the module.
"""
try:
mod_name, attr_name = name.split(":")
mod = importlib.import_module(mod_name)
callable_object = getattr(mod, attr_name)
# check if attribute is callable
if callable(callable_object):
return callable_object
else:
raise ValueError(f"The imported object is not callable: '{name}'")
except AttributeError as e:
msg = (
"We could not interpret the entry as a callable object. The format of input should be"
f" 'module:attribute_name'\nWhile processing input '{name}', received the error:\n {e}."
)
raise ValueError(msg)
def resolve_obs_groups(
obs: TensorDict, obs_groups: dict[str, list[str]], default_sets: list[str]
) -> dict[str, list[str]]:
"""Validates the observation configuration and defaults missing observation sets.
The input is an observation dictionary `obs` containing observation groups and a configuration dictionary
`obs_groups` where the keys are the observation sets and the values are lists of observation groups.
The configuration dictionary could for example look like:
{
"policy": ["group_1", "group_2"],
"critic": ["group_1", "group_3"]
}
This means that the 'policy' observation set will contain the observations "group_1" and "group_2" and the
'critic' observation set will contain the observations "group_1" and "group_3". This function will check that all
the observations in the 'policy' and 'critic' observation sets are present in the observation dictionary from the
environment.
Additionally, if one of the `default_sets`, e.g. "critic", is not present in the configuration dictionary,
this function will:
1. Check if a group with the same name exists in the observations and assign this group to the observation set.
2. If 1. fails, it will assign the observations from the 'policy' observation set to the default observation set.
Args:
obs: Observations from the environment in the form of a dictionary.
obs_groups: Observation sets configuration.
default_sets: Reserved observation set names used by the algorithm (besides 'policy').
If not provided in 'obs_groups', a default behavior gets triggered.
Returns:
The resolved observation groups.
Raises:
ValueError: If any observation set is an empty list.
ValueError: If any observation set contains an observation term that is not present in the observations.
"""
# check if policy observation set exists
if "policy" not in obs_groups.keys():
if "policy" in obs:
obs_groups["policy"] = ["policy"]
warnings.warn(
"The observation configuration dictionary 'obs_groups' must contain the 'policy' key."
" As an observation group with the name 'policy' was found, this is assumed to be the observation set."
" Consider adding the 'policy' key to the 'obs_groups' dictionary for clarity."
" This behavior will be removed in a future version."
)
else:
raise ValueError(
"The observation configuration dictionary 'obs_groups' must contain the 'policy' key."
f" Found keys: {list(obs_groups.keys())}"
)
# check all observation sets for valid observation groups
for set_name, groups in obs_groups.items():
# check if the list is empty
if len(groups) == 0:
msg = f"The '{set_name}' key in the 'obs_groups' dictionary can not be an empty list."
if set_name in default_sets:
if set_name not in obs:
msg += " Consider removing the key to default to the observations used for the 'policy' set."
else:
msg += (
f" Consider removing the key to default to the observation '{set_name}' from the environment."
)
raise ValueError(msg)
# check groups exist inside the observations from the environment
for group in groups:
if group not in obs:
raise ValueError(
f"Observation '{group}' in observation set '{set_name}' not found in the observations from the"
f" environment. Available observations from the environment: {list(obs.keys())}"
)
# fill missing observation sets
for default_set_name in default_sets:
if default_set_name not in obs_groups.keys():
if default_set_name in obs:
obs_groups[default_set_name] = [default_set_name]
warnings.warn(
f"The observation configuration dictionary 'obs_groups' must contain the '{default_set_name}' key."
f" As an observation group with the name '{default_set_name}' was found, this is assumed to be the"
f" observation set. Consider adding the '{default_set_name}' key to the 'obs_groups' dictionary for"
" clarity. This behavior will be removed in a future version."
)
else:
obs_groups[default_set_name] = obs_groups["policy"].copy()
warnings.warn(
f"The observation configuration dictionary 'obs_groups' must contain the '{default_set_name}' key."
f" As the configuration for '{default_set_name}' is missing, the observations from the 'policy' set"
f" are used. Consider adding the '{default_set_name}' key to the 'obs_groups' dictionary for"
" clarity. This behavior will be removed in a future version."
)
# print the final parsed observation sets
print("-" * 80)
print("Resolved observation sets: ")
for set_name, groups in obs_groups.items():
print("\t", set_name, ": ", groups)
print("-" * 80)
return obs_groups

View File

@ -0,0 +1,87 @@
# Copyright (c) 2021-2025, ETH Zurich and NVIDIA CORPORATION
# All rights reserved.
#
# SPDX-License-Identifier: BSD-3-Clause
from __future__ import annotations
import os
from dataclasses import asdict
from torch.utils.tensorboard import SummaryWriter
try:
import wandb
except ModuleNotFoundError:
raise ModuleNotFoundError("Wandb is required to log to Weights and Biases.")
class WandbSummaryWriter(SummaryWriter):
"""Summary writer for Weights and Biases."""
def __init__(self, log_dir: str, flush_secs: int, cfg):
super().__init__(log_dir, flush_secs)
# Get the run name
run_name = os.path.split(log_dir)[-1]
try:
project = cfg["wandb_project"]
except KeyError:
raise KeyError("Please specify wandb_project in the runner config, e.g. legged_gym.")
try:
entity = os.environ["WANDB_USERNAME"]
except KeyError:
entity = None
# Initialize wandb
wandb.init(project=project, entity=entity, name=run_name)
# Add log directory to wandb
wandb.config.update({"log_dir": log_dir})
self.name_map = {
"Train/mean_reward/time": "Train/mean_reward_time",
"Train/mean_episode_length/time": "Train/mean_episode_length_time",
}
def store_config(self, env_cfg, runner_cfg, alg_cfg, policy_cfg):
wandb.config.update({"runner_cfg": runner_cfg})
wandb.config.update({"policy_cfg": policy_cfg})
wandb.config.update({"alg_cfg": alg_cfg})
try:
wandb.config.update({"env_cfg": env_cfg.to_dict()})
except Exception:
wandb.config.update({"env_cfg": asdict(env_cfg)})
def add_scalar(self, tag, scalar_value, global_step=None, walltime=None, new_style=False):
super().add_scalar(
tag,
scalar_value,
global_step=global_step,
walltime=walltime,
new_style=new_style,
)
wandb.log({self._map_path(tag): scalar_value}, step=global_step)
def stop(self):
wandb.finish()
def log_config(self, env_cfg, runner_cfg, alg_cfg, policy_cfg):
self.store_config(env_cfg, runner_cfg, alg_cfg, policy_cfg)
def save_model(self, model_path, iter):
wandb.save(model_path, base_path=os.path.dirname(model_path))
def save_file(self, path, iter=None):
wandb.save(path, base_path=os.path.dirname(path))
"""
Private methods.
"""
def _map_path(self, path):
if path in self.name_map:
return self.name_map[path]
else:
return path

View File

@ -0,0 +1,10 @@
# Copyright (c) 2021-2025, ETH Zurich and NVIDIA CORPORATION
# All rights reserved.
#
# SPDX-License-Identifier: BSD-3-Clause
"""Submodule defining the environment definitions."""
from .vec_env import VecEnv
__all__ = ["VecEnv"]

View File

@ -0,0 +1,113 @@
# Copyright (c) 2021-2025, ETH Zurich and NVIDIA CORPORATION
# All rights reserved.
#
# SPDX-License-Identifier: BSD-3-Clause
from __future__ import annotations
import torch
from abc import ABC, abstractmethod
from tensordict import TensorDict
class VecEnv(ABC):
"""Abstract class for a vectorized environment.
The vectorized environment is a collection of environments that are synchronized. This means that
the same type of action is applied to all environments and the same type of observation is returned from all
environments.
"""
num_envs: int
"""Number of environments."""
num_actions: int
"""Number of actions."""
max_episode_length: int | torch.Tensor
max_episode_length_s: float
"""Maximum episode length.
The maximum episode length can be a scalar or a tensor. If it is a scalar, it is the same for all environments.
If it is a tensor, it is the maximum episode length for each environment. This is useful for dynamic episode
lengths.
"""
episode_length_buf: torch.Tensor
"""Buffer for current episode lengths."""
device: torch.device | str
"""Device to use."""
cfg: dict | object
"""Configuration object."""
reset_env_ids: torch.Tensor | None = None
contact_phase: torch.Tensor | None = None
"""
Operations.
"""
@abstractmethod
def get_observations(self) -> TensorDict:
"""Return the current observations.
Returns:
observations (TensorDict): Observations from the environment.
"""
raise NotImplementedError
@abstractmethod
def get_amp_observations(self) -> TensorDict:
"""Return the current AMP observations.
Returns:
observations (TensorDict): Observations from the environment.
"""
raise NotImplementedError
@abstractmethod
def step(self, actions: torch.Tensor) -> tuple[TensorDict, torch.Tensor, torch.Tensor, dict]:
"""Apply input action to the environment.
Args:
actions (torch.Tensor): Input actions to apply. Shape: (num_envs, num_actions)
Returns:
observations (TensorDict): Observations from the environment.
rewards (torch.Tensor): Rewards from the environment. Shape: (num_envs,)
dones (torch.Tensor): Done flags from the environment. Shape: (num_envs,)
extras (dict): Extra information from the environment.
Observations:
The observations TensorDict usually contains multiple observation groups. The `obs_groups`
dictionary of the runner configuration specifies which observation groups are used for which
purpose, i.e., it maps the available observation groups to observation sets. The observation sets
(keys of the `obs_groups` dictionary) currently used by rsl_rl are:
- "policy": Specified observation groups are used as input to the actor/student network.
- "critic": Specified observation groups are used as input to the critic network.
- "teacher": Specified observation groups are used as input to the teacher network.
- "rnd_state": Specified observation groups are used as input to the RND network.
Incomplete or incorrect configurations are handled in the `resolve_obs_groups()` function in
`rsl_rl/utils/utils.py`.
Extras:
The extras dictionary includes metrics such as the episode reward, episode length, etc. The following
dictionary keys are used by rsl_rl:
- "time_outs" (torch.Tensor): Timeouts for the environments. These correspond to terminations that
happen due to time limits and not due to the environment reaching a terminal state. This is useful
for environments that have a fixed episode length.
- "log" (dict[str, float | torch.Tensor]): Additional information for logging and debugging purposes.
The key should be a string and start with "/" for namespacing. The value can be a scalar or a
tensor. If it is a tensor, the mean of the tensor is used for logging.
"""
raise NotImplementedError

View File

@ -0,0 +1,21 @@
# Copyright (c) 2021-2025, ETH Zurich and NVIDIA CORPORATION
# All rights reserved.
#
# SPDX-License-Identifier: BSD-3-Clause
"""Definitions for neural-network components for RL-agents."""
from .actor_critic import ActorCritic
from .actor_critic_recurrent import ActorCriticRecurrent
from .rnd import *
from .student_teacher import StudentTeacher
from .student_teacher_recurrent import StudentTeacherRecurrent
from .symmetry import *
from .discriminator_multi import DiscriminatorMulti
__all__ = [
"ActorCritic",
"ActorCriticRecurrent",
"StudentTeacher",
"StudentTeacherRecurrent",
"DiscriminatorMulti",
]

View File

@ -0,0 +1,195 @@
# Copyright (c) 2021-2025, ETH Zurich and NVIDIA CORPORATION
# All rights reserved.
#
# SPDX-License-Identifier: BSD-3-Clause
from __future__ import annotations
import torch
import torch.nn as nn
from torch.distributions import Normal
from rsl_rl.networks import MLP, EmpiricalNormalization
class ActorCritic(nn.Module):
is_recurrent = False
def __init__(
self,
obs,
obs_groups,
num_actions,
actor_obs_normalization=False,
critic_obs_normalization=False,
actor_hidden_dims=[256, 256, 256],
critic_hidden_dims=[256, 256, 256],
activation="elu",
init_noise_std=1.0,
noise_std_type: str = "scalar",
state_dependent_std=False,
**kwargs,
):
if kwargs:
print(
"ActorCritic.__init__ got unexpected arguments, which will be ignored: "
+ str([key for key in kwargs.keys()])
)
super().__init__()
# get the observation dimensions
self.obs_groups = obs_groups
num_actor_obs = 0
for obs_group in obs_groups["policy"]:
assert len(obs[obs_group].shape) == 2, "The ActorCritic module only supports 1D observations."
num_actor_obs += obs[obs_group].shape[-1]
num_critic_obs = 0
for obs_group in obs_groups["critic"]:
assert len(obs[obs_group].shape) == 2, "The ActorCritic module only supports 1D observations."
num_critic_obs += obs[obs_group].shape[-1]
self.state_dependent_std = state_dependent_std
# actor
if self.state_dependent_std:
self.actor = MLP(num_actor_obs, [2, num_actions], actor_hidden_dims, activation)
else:
self.actor = MLP(num_actor_obs, num_actions, actor_hidden_dims, activation)
# actor observation normalization
self.actor_obs_normalization = actor_obs_normalization
if actor_obs_normalization:
self.actor_obs_normalizer = EmpiricalNormalization(num_actor_obs)
else:
self.actor_obs_normalizer = torch.nn.Identity()
print(f"Actor MLP: {self.actor}")
# critic
self.critic = MLP(num_critic_obs, 1, critic_hidden_dims, activation)
# critic observation normalization
self.critic_obs_normalization = critic_obs_normalization
if critic_obs_normalization:
self.critic_obs_normalizer = EmpiricalNormalization(num_critic_obs)
else:
self.critic_obs_normalizer = torch.nn.Identity()
print(f"Critic MLP: {self.critic}")
# Action noise
self.noise_std_type = noise_std_type
if self.state_dependent_std:
torch.nn.init.zeros_(self.actor[-2].weight[num_actions:])
if self.noise_std_type == "scalar":
torch.nn.init.constant_(self.actor[-2].bias[num_actions:], init_noise_std)
elif self.noise_std_type == "log":
torch.nn.init.constant_(
self.actor[-2].bias[num_actions:], torch.log(torch.tensor(init_noise_std + 1e-7))
)
else:
raise ValueError(f"Unknown standard deviation type: {self.noise_std_type}. Should be 'scalar' or 'log'")
else:
if self.noise_std_type == "scalar":
self.std = nn.Parameter(init_noise_std * torch.ones(num_actions))
elif self.noise_std_type == "log":
self.log_std = nn.Parameter(torch.log(init_noise_std * torch.ones(num_actions)))
else:
raise ValueError(f"Unknown standard deviation type: {self.noise_std_type}. Should be 'scalar' or 'log'")
# Action distribution (populated in update_distribution)
self.distribution = None
# disable args validation for speedup
Normal.set_default_validate_args(False)
def reset(self, dones=None):
pass
def forward(self):
raise NotImplementedError
@property
def action_mean(self):
return self.distribution.mean
@property
def action_std(self):
return self.distribution.stddev
@property
def entropy(self):
return self.distribution.entropy().sum(dim=-1)
def update_distribution(self, obs):
if self.state_dependent_std:
# compute mean and standard deviation
mean_and_std = self.actor(obs)
if self.noise_std_type == "scalar":
mean, std = torch.unbind(mean_and_std, dim=-2)
elif self.noise_std_type == "log":
mean, log_std = torch.unbind(mean_and_std, dim=-2)
std = torch.exp(log_std)
else:
raise ValueError(f"Unknown standard deviation type: {self.noise_std_type}. Should be 'scalar' or 'log'")
else:
# compute mean
mean = self.actor(obs)
# compute standard deviation
if self.noise_std_type == "scalar":
std = self.std.expand_as(mean)
elif self.noise_std_type == "log":
std = torch.exp(self.log_std).expand_as(mean)
else:
raise ValueError(f"Unknown standard deviation type: {self.noise_std_type}. Should be 'scalar' or 'log'")
# create distribution
self.distribution = Normal(mean, std)
def act(self, obs, **kwargs):
obs = self.get_actor_obs(obs)
obs = self.actor_obs_normalizer(obs)
self.update_distribution(obs)
return self.distribution.sample()
def act_inference(self, obs):
obs = self.get_actor_obs(obs)
obs = self.actor_obs_normalizer(obs)
return self.actor(obs)
def evaluate(self, obs, **kwargs):
obs = self.get_critic_obs(obs)
obs = self.critic_obs_normalizer(obs)
return self.critic(obs)
def get_actor_obs(self, obs):
obs_list = []
for obs_group in self.obs_groups["policy"]:
obs_list.append(obs[obs_group])
return torch.cat(obs_list, dim=-1)
def get_critic_obs(self, obs):
obs_list = []
for obs_group in self.obs_groups["critic"]:
obs_list.append(obs[obs_group])
return torch.cat(obs_list, dim=-1)
def get_actions_log_prob(self, actions):
return self.distribution.log_prob(actions).sum(dim=-1)
def update_normalization(self, obs):
if self.actor_obs_normalization:
actor_obs = self.get_actor_obs(obs)
self.actor_obs_normalizer.update(actor_obs)
if self.critic_obs_normalization:
critic_obs = self.get_critic_obs(obs)
self.critic_obs_normalizer.update(critic_obs)
def load_state_dict(self, state_dict, strict=True):
"""Load the parameters of the actor-critic model.
Args:
state_dict (dict): State dictionary of the model.
strict (bool): Whether to strictly enforce that the keys in state_dict match the keys returned by this
module's state_dict() function.
Returns:
bool: Whether this training resumes a previous training. This flag is used by the `load()` function of
`OnPolicyRunner` to determine how to load further parameters (relevant for, e.g., distillation).
"""
super().load_state_dict(state_dict, strict=strict)
return True # training resumes

View File

@ -0,0 +1,218 @@
# Copyright (c) 2021-2025, ETH Zurich and NVIDIA CORPORATION
# All rights reserved.
#
# SPDX-License-Identifier: BSD-3-Clause
from __future__ import annotations
import torch
import torch.nn as nn
import warnings
from torch.distributions import Normal
from rsl_rl.networks import MLP, EmpiricalNormalization, Memory
class ActorCriticRecurrent(nn.Module):
is_recurrent = True
def __init__(
self,
obs,
obs_groups,
num_actions,
actor_obs_normalization=False,
critic_obs_normalization=False,
actor_hidden_dims=[256, 256, 256],
critic_hidden_dims=[256, 256, 256],
activation="elu",
init_noise_std=1.0,
noise_std_type: str = "scalar",
state_dependent_std=False,
rnn_type="lstm",
rnn_hidden_dim=256,
rnn_num_layers=1,
**kwargs,
):
if "rnn_hidden_size" in kwargs:
warnings.warn(
"The argument `rnn_hidden_size` is deprecated and will be removed in a future version. "
"Please use `rnn_hidden_dim` instead.",
DeprecationWarning,
)
if rnn_hidden_dim == 256: # Only override if the new argument is at its default
rnn_hidden_dim = kwargs.pop("rnn_hidden_size")
if kwargs:
print(
"ActorCriticRecurrent.__init__ got unexpected arguments, which will be ignored: " + str(kwargs.keys()),
)
super().__init__()
# get the observation dimensions
self.obs_groups = obs_groups
num_actor_obs = 0
for obs_group in obs_groups["policy"]:
assert len(obs[obs_group].shape) == 2, "The ActorCriticRecurrent module only supports 1D observations."
num_actor_obs += obs[obs_group].shape[-1]
num_critic_obs = 0
for obs_group in obs_groups["critic"]:
assert len(obs[obs_group].shape) == 2, "The ActorCriticRecurrent module only supports 1D observations."
num_critic_obs += obs[obs_group].shape[-1]
self.state_dependent_std = state_dependent_std
# actor
self.memory_a = Memory(num_actor_obs, type=rnn_type, num_layers=rnn_num_layers, hidden_size=rnn_hidden_dim)
if self.state_dependent_std:
self.actor = MLP(rnn_hidden_dim, [2, num_actions], actor_hidden_dims, activation)
else:
self.actor = MLP(rnn_hidden_dim, num_actions, actor_hidden_dims, activation)
# actor observation normalization
self.actor_obs_normalization = actor_obs_normalization
if actor_obs_normalization:
self.actor_obs_normalizer = EmpiricalNormalization(num_actor_obs)
else:
self.actor_obs_normalizer = torch.nn.Identity()
print(f"Actor RNN: {self.memory_a}")
print(f"Actor MLP: {self.actor}")
# critic
self.memory_c = Memory(num_critic_obs, type=rnn_type, num_layers=rnn_num_layers, hidden_size=rnn_hidden_dim)
self.critic = MLP(rnn_hidden_dim, 1, critic_hidden_dims, activation)
# critic observation normalization
self.critic_obs_normalization = critic_obs_normalization
if critic_obs_normalization:
self.critic_obs_normalizer = EmpiricalNormalization(num_critic_obs)
else:
self.critic_obs_normalizer = torch.nn.Identity()
print(f"Critic RNN: {self.memory_c}")
print(f"Critic MLP: {self.critic}")
# Action noise
self.noise_std_type = noise_std_type
if self.state_dependent_std:
torch.nn.init.zeros_(self.actor[-2].weight[num_actions:])
if self.noise_std_type == "scalar":
torch.nn.init.constant_(self.actor[-2].bias[num_actions:], init_noise_std)
elif self.noise_std_type == "log":
torch.nn.init.constant_(
self.actor[-2].bias[num_actions:], torch.log(torch.tensor(init_noise_std + 1e-7))
)
else:
raise ValueError(f"Unknown standard deviation type: {self.noise_std_type}. Should be 'scalar' or 'log'")
else:
if self.noise_std_type == "scalar":
self.std = nn.Parameter(init_noise_std * torch.ones(num_actions))
elif self.noise_std_type == "log":
self.log_std = nn.Parameter(torch.log(init_noise_std * torch.ones(num_actions)))
else:
raise ValueError(f"Unknown standard deviation type: {self.noise_std_type}. Should be 'scalar' or 'log'")
# Action distribution (populated in update_distribution)
self.distribution = None
# disable args validation for speedup
Normal.set_default_validate_args(False)
@property
def action_mean(self):
return self.distribution.mean
@property
def action_std(self):
return self.distribution.stddev
@property
def entropy(self):
return self.distribution.entropy().sum(dim=-1)
def reset(self, dones=None):
self.memory_a.reset(dones)
self.memory_c.reset(dones)
def forward(self):
raise NotImplementedError
def update_distribution(self, obs):
if self.state_dependent_std:
# compute mean and standard deviation
mean_and_std = self.actor(obs)
if self.noise_std_type == "scalar":
mean, std = torch.unbind(mean_and_std, dim=-2)
elif self.noise_std_type == "log":
mean, log_std = torch.unbind(mean_and_std, dim=-2)
std = torch.exp(log_std)
else:
raise ValueError(f"Unknown standard deviation type: {self.noise_std_type}. Should be 'scalar' or 'log'")
else:
# compute mean
mean = self.actor(obs)
# compute standard deviation
if self.noise_std_type == "scalar":
std = self.std.expand_as(mean)
elif self.noise_std_type == "log":
std = torch.exp(self.log_std).expand_as(mean)
else:
raise ValueError(f"Unknown standard deviation type: {self.noise_std_type}. Should be 'scalar' or 'log'")
# create distribution
self.distribution = Normal(mean, std)
def act(self, obs, masks=None, hidden_states=None):
obs = self.get_actor_obs(obs)
obs = self.actor_obs_normalizer(obs)
out_mem = self.memory_a(obs, masks, hidden_states).squeeze(0)
self.update_distribution(out_mem)
return self.distribution.sample()
def act_inference(self, obs):
obs = self.get_actor_obs(obs)
obs = self.actor_obs_normalizer(obs)
out_mem = self.memory_a(obs).squeeze(0)
return self.actor(out_mem)
def evaluate(self, obs, masks=None, hidden_states=None):
obs = self.get_critic_obs(obs)
obs = self.critic_obs_normalizer(obs)
out_mem = self.memory_c(obs, masks, hidden_states).squeeze(0)
return self.critic(out_mem)
def get_actor_obs(self, obs):
obs_list = []
for obs_group in self.obs_groups["policy"]:
obs_list.append(obs[obs_group])
return torch.cat(obs_list, dim=-1)
def get_critic_obs(self, obs):
obs_list = []
for obs_group in self.obs_groups["critic"]:
obs_list.append(obs[obs_group])
return torch.cat(obs_list, dim=-1)
def get_actions_log_prob(self, actions):
return self.distribution.log_prob(actions).sum(dim=-1)
def get_hidden_states(self):
return self.memory_a.hidden_states, self.memory_c.hidden_states
def update_normalization(self, obs):
if self.actor_obs_normalization:
actor_obs = self.get_actor_obs(obs)
self.actor_obs_normalizer.update(actor_obs)
if self.critic_obs_normalization:
critic_obs = self.get_critic_obs(obs)
self.critic_obs_normalizer.update(critic_obs)
def load_state_dict(self, state_dict, strict=True):
"""Load the parameters of the actor-critic model.
Args:
state_dict (dict): State dictionary of the model.
strict (bool): Whether to strictly enforce that the keys in state_dict match the keys returned by this
module's state_dict() function.
Returns:
bool: Whether this training resumes a previous training. This flag is used by the `load()` function of
`OnPolicyRunner` to determine how to load further parameters (relevant for, e.g., distillation).
"""
super().load_state_dict(state_dict, strict=strict)
return True

View File

@ -0,0 +1,102 @@
import torch
import torch.nn as nn
from torch import autograd
import torch.nn.utils.spectral_norm as spectral_norm
class DiscriminatorMulti(nn.Module):
def __init__(
self, state_dim, amp_reward_coef, hidden_layer_sizes, device,
num_frames=2, task_reward_lerp=0.0, use_lerp=True):
super(DiscriminatorMulti, self).__init__()
self.device = device
self.state_dim = state_dim
self.use_lerp = use_lerp
self.num_frames = num_frames # 存储帧数参数
self.amp_reward_coef = amp_reward_coef
amp_layers = []
curr_in_dim = state_dim * num_frames
for hidden_dim in hidden_layer_sizes:
amp_layers.append(spectral_norm(nn.Linear(curr_in_dim, hidden_dim)))
amp_layers.append(nn.ReLU())
curr_in_dim = hidden_dim
self.trunk = nn.Sequential(*amp_layers).to(device)
self.amp_linear = spectral_norm(nn.Linear(hidden_layer_sizes[-1], 1)).to(device)
self.trunk.train()
self.amp_linear.train()
self.task_reward_lerp = task_reward_lerp
def forward(self, x):
h = self.trunk(x)
d = self.amp_linear(h)
return d
def compute_grad_pen(self,
expert_states, # 改为接收多帧状态列表
lambda_=10):
# 将多帧状态沿最后一个维度拼接
expert_data = expert_states.flatten(1)
expert_data.requires_grad = True
disc = self.amp_linear(self.trunk(expert_data))
ones = torch.ones(disc.size(), device=disc.device)
grad = autograd.grad(
outputs=disc, inputs=expert_data,
grad_outputs=ones, create_graph=True,
retain_graph=True, only_inputs=True)[0]
# Enforce that the grad norm approaches 0.
grad_pen = lambda_ * (grad.norm(2, dim=1) - 0).pow(2).mean()
return grad_pen
def get_disc_weights(self):
weights = []
for m in self.trunk.modules():
if isinstance(m, nn.Linear):
weights.append(torch.flatten(m.weight))
weights.append(torch.flatten(self.amp_linear.weight))
return weights
def get_disc_logit_weights(self):
return torch.flatten(self.amp_linear.weight)
def predict_amp_reward(
self, states, # 改为接收多帧状态列表
task_reward, normalizer=None):
"""
states: torch.Tensor, shape=(num_envs, num_frames, state_dim)
task_reward: torch.Tensor, shape=(num_envs, 1)
"""
# import ipdb; ipdb.set_trace()
with torch.no_grad():
self.eval()
if normalizer is not None:
# 对每一帧状态进行归一化
states = normalizer.normalize_torch(states, self.device)
# 拼接多帧状态
state_cat = states.flatten(1)
d = self.amp_linear(self.trunk(state_cat))
disc_reward = self.amp_reward_coef * torch.clamp(1 - (1/4) * torch.square(d - 1), min=0)
if self.use_lerp:
if self.task_reward_lerp > 0:
reward = self._lerp_reward(disc_reward, task_reward.unsqueeze(-1))
self.train()
return reward.squeeze(), d, disc_reward.squeeze() * (1.0 - self.task_reward_lerp)
else:
disc_reward *= 0.02
reward = task_reward.unsqueeze(-1) + disc_reward
self.train()
return reward.squeeze(), d, disc_reward.squeeze()
def _lerp_reward(self, disc_r, task_r):
r = (1.0 - self.task_reward_lerp) * disc_r + self.task_reward_lerp * task_r
return r

View File

@ -0,0 +1,209 @@
# Copyright (c) 2021-2025, ETH Zurich and NVIDIA CORPORATION
# All rights reserved.
#
# SPDX-License-Identifier: BSD-3-Clause
from __future__ import annotations
import torch
import torch.nn as nn
from rsl_rl.networks import MLP, EmpiricalDiscountedVariationNormalization, EmpiricalNormalization
class RandomNetworkDistillation(nn.Module):
"""Implementation of Random Network Distillation (RND) [1]
References:
.. [1] Burda, Yuri, et al. "Exploration by random network distillation." arXiv preprint arXiv:1810.12894 (2018).
"""
def __init__(
self,
num_states: int,
obs_groups: dict,
num_outputs: int,
predictor_hidden_dims: list[int],
target_hidden_dims: list[int],
activation: str = "elu",
weight: float = 0.0,
state_normalization: bool = False,
reward_normalization: bool = False,
device: str = "cpu",
weight_schedule: dict | None = None,
):
"""Initialize the RND module.
- If :attr:`state_normalization` is True, then the input state is normalized using an Empirical Normalization layer.
- If :attr:`reward_normalization` is True, then the intrinsic reward is normalized using an Empirical Discounted
Variation Normalization layer.
.. note::
If the hidden dimensions are -1 in the predictor and target networks configuration, then the number of states
is used as the hidden dimension.
Args:
num_states: Number of states/inputs to the predictor and target networks.
num_outputs: Number of outputs (embedding size) of the predictor and target networks.
predictor_hidden_dims: List of hidden dimensions of the predictor network.
target_hidden_dims: List of hidden dimensions of the target network.
activation: Activation function. Defaults to "elu".
weight: Scaling factor of the intrinsic reward. Defaults to 0.0.
state_normalization: Whether to normalize the input state. Defaults to False.
reward_normalization: Whether to normalize the intrinsic reward. Defaults to False.
device: Device to use. Defaults to "cpu".
weight_schedule: The type of schedule to use for the RND weight parameter.
Defaults to None, in which case the weight parameter is constant.
It is a dictionary with the following keys:
- "mode": The type of schedule to use for the RND weight parameter.
- "constant": Constant weight schedule.
- "step": Step weight schedule.
- "linear": Linear weight schedule.
For the "step" weight schedule, the following parameters are required:
- "final_step": The step at which the weight parameter is set to the final value.
- "final_value": The final value of the weight parameter.
For the "linear" weight schedule, the following parameters are required:
- "initial_step": The step at which the weight parameter is set to the initial value.
- "final_step": The step at which the weight parameter is set to the final value.
- "final_value": The final value of the weight parameter.
"""
# initialize parent class
super().__init__()
# Store parameters
self.num_states = num_states
self.obs_groups = obs_groups
self.num_outputs = num_outputs
self.initial_weight = weight
self.device = device
self.state_normalization = state_normalization
self.reward_normalization = reward_normalization
# Normalization of input gates
if state_normalization:
self.state_normalizer = EmpiricalNormalization(shape=[self.num_states], until=1.0e8).to(self.device)
else:
self.state_normalizer = torch.nn.Identity()
# Normalization of intrinsic reward
if reward_normalization:
self.reward_normalizer = EmpiricalDiscountedVariationNormalization(shape=[], until=1.0e8).to(self.device)
else:
self.reward_normalizer = torch.nn.Identity()
# counter for the number of updates
self.update_counter = 0
# resolve weight schedule
if weight_schedule is not None:
self.weight_scheduler_params = weight_schedule
self.weight_scheduler = getattr(self, f"_{weight_schedule['mode']}_weight_schedule")
else:
self.weight_scheduler = None
# Create network architecture
self.predictor = MLP(num_states, num_outputs, predictor_hidden_dims, activation).to(self.device)
self.target = MLP(num_states, num_outputs, target_hidden_dims, activation).to(self.device)
# make target network not trainable
self.target.eval()
def get_intrinsic_reward(self, obs) -> torch.Tensor:
# Note: the counter is updated number of env steps per learning iteration
self.update_counter += 1
# Extract the rnd state from the observation
rnd_state = self.get_rnd_state(obs)
rnd_state = self.state_normalizer(rnd_state)
# Obtain the embedding of the rnd state from the target and predictor networks
target_embedding = self.target(rnd_state).detach()
predictor_embedding = self.predictor(rnd_state).detach()
# Compute the intrinsic reward as the distance between the embeddings
intrinsic_reward = torch.linalg.norm(target_embedding - predictor_embedding, dim=1)
# Normalize intrinsic reward
intrinsic_reward = self.reward_normalizer(intrinsic_reward)
# Check the weight schedule
if self.weight_scheduler is not None:
self.weight = self.weight_scheduler(step=self.update_counter, **self.weight_scheduler_params)
else:
self.weight = self.initial_weight
# Scale intrinsic reward
intrinsic_reward *= self.weight
return intrinsic_reward
def forward(self, *args, **kwargs):
raise RuntimeError("Forward method is not implemented. Use get_intrinsic_reward instead.")
def train(self, mode: bool = True):
# sets module into training mode
self.predictor.train(mode)
if self.state_normalization:
self.state_normalizer.train(mode)
if self.reward_normalization:
self.reward_normalizer.train(mode)
return self
def eval(self):
return self.train(False)
def get_rnd_state(self, obs):
obs_list = []
for obs_group in self.obs_groups["rnd_state"]:
obs_list.append(obs[obs_group])
return torch.cat(obs_list, dim=-1)
def update_normalization(self, obs):
# Normalize the state
if self.state_normalization:
rnd_state = self.get_rnd_state(obs)
self.state_normalizer.update(rnd_state)
"""
Different weight schedules.
"""
def _constant_weight_schedule(self, step: int, **kwargs):
return self.initial_weight
def _step_weight_schedule(self, step: int, final_step: int, final_value: float, **kwargs):
return self.initial_weight if step < final_step else final_value
def _linear_weight_schedule(self, step: int, initial_step: int, final_step: int, final_value: float, **kwargs):
if step < initial_step:
return self.initial_weight
elif step > final_step:
return final_value
else:
return self.initial_weight + (final_value - self.initial_weight) * (step - initial_step) / (
final_step - initial_step
)
def resolve_rnd_config(alg_cfg, obs, obs_groups, env):
"""Resolve the RND configuration.
Args:
alg_cfg: The algorithm configuration dictionary.
obs: The observation dictionary.
obs_groups: The observation groups dictionary.
env: The environment.
Returns:
The resolved algorithm configuration dictionary.
"""
# resolve dimension of rnd gated state
if "rnd_cfg" in alg_cfg and alg_cfg["rnd_cfg"] is not None:
# get dimension of rnd gated state
num_rnd_state = 0
for obs_group in obs_groups["rnd_state"]:
assert len(obs[obs_group].shape) == 2, "The RND module only supports 1D observations."
num_rnd_state += obs[obs_group].shape[-1]
# add rnd gated state to config
alg_cfg["rnd_cfg"]["num_states"] = num_rnd_state
alg_cfg["rnd_cfg"]["obs_groups"] = obs_groups
# scale down the rnd weight with timestep
alg_cfg["rnd_cfg"]["weight"] *= env.unwrapped.step_dt
return alg_cfg

View File

@ -0,0 +1,206 @@
# Copyright (c) 2021-2025, ETH Zurich and NVIDIA CORPORATION
# All rights reserved.
#
# SPDX-License-Identifier: BSD-3-Clause
from __future__ import annotations
import torch
import torch.nn as nn
from torch.distributions import Normal
from rsl_rl.networks import MLP, EmpiricalNormalization
class StudentTeacher(nn.Module):
is_recurrent = False
def __init__(
self,
obs,
obs_groups,
num_actions,
student_obs_normalization=False,
teacher_obs_normalization=False,
student_hidden_dims=[256, 256, 256],
teacher_hidden_dims=[256, 256, 256],
activation="elu",
init_noise_std=0.1,
noise_std_type: str = "scalar",
**kwargs,
):
if kwargs:
print(
"StudentTeacher.__init__ got unexpected arguments, which will be ignored: "
+ str([key for key in kwargs.keys()])
)
super().__init__()
self.loaded_teacher = False # indicates if teacher has been loaded
# get the observation dimensions
self.obs_groups = obs_groups
num_student_obs = 0
for obs_group in obs_groups["policy"]:
assert len(obs[obs_group].shape) == 2, "The StudentTeacher module only supports 1D observations."
num_student_obs += obs[obs_group].shape[-1]
num_teacher_obs = 0
for obs_group in obs_groups["teacher"]:
assert len(obs[obs_group].shape) == 2, "The StudentTeacher module only supports 1D observations."
num_teacher_obs += obs[obs_group].shape[-1]
# student
self.student = MLP(num_student_obs, num_actions, student_hidden_dims, activation)
# student observation normalization
self.student_obs_normalization = student_obs_normalization
if student_obs_normalization:
self.student_obs_normalizer = EmpiricalNormalization(num_student_obs)
else:
self.student_obs_normalizer = torch.nn.Identity()
print(f"Student MLP: {self.student}")
# teacher
self.teacher = MLP(num_teacher_obs, num_actions, teacher_hidden_dims, activation)
self.teacher.eval()
# teacher observation normalization
self.teacher_obs_normalization = teacher_obs_normalization
if teacher_obs_normalization:
self.teacher_obs_normalizer = EmpiricalNormalization(num_teacher_obs)
else:
self.teacher_obs_normalizer = torch.nn.Identity()
print(f"Teacher MLP: {self.teacher}")
# action noise
self.noise_std_type = noise_std_type
if self.noise_std_type == "scalar":
self.std = nn.Parameter(init_noise_std * torch.ones(num_actions))
elif self.noise_std_type == "log":
self.log_std = nn.Parameter(torch.log(init_noise_std * torch.ones(num_actions)))
else:
raise ValueError(f"Unknown standard deviation type: {self.noise_std_type}. Should be 'scalar' or 'log'")
# action distribution (populated in update_distribution)
self.distribution = None
# disable args validation for speedup
Normal.set_default_validate_args(False)
def reset(self, dones=None, hidden_states=None):
pass
def forward(self):
raise NotImplementedError
@property
def action_mean(self):
return self.distribution.mean
@property
def action_std(self):
return self.distribution.stddev
@property
def entropy(self):
return self.distribution.entropy().sum(dim=-1)
def update_distribution(self, obs):
# compute mean
mean = self.student(obs)
# compute standard deviation
if self.noise_std_type == "scalar":
std = self.std.expand_as(mean)
elif self.noise_std_type == "log":
std = torch.exp(self.log_std).expand_as(mean)
else:
raise ValueError(f"Unknown standard deviation type: {self.noise_std_type}. Should be 'scalar' or 'log'")
# create distribution
self.distribution = Normal(mean, std)
def act(self, obs):
obs = self.get_student_obs(obs)
obs = self.student_obs_normalizer(obs)
self.update_distribution(obs)
return self.distribution.sample()
def act_inference(self, obs):
obs = self.get_student_obs(obs)
obs = self.student_obs_normalizer(obs)
return self.student(obs)
def evaluate(self, obs):
obs = self.get_teacher_obs(obs)
obs = self.teacher_obs_normalizer(obs)
with torch.no_grad():
return self.teacher(obs)
def get_student_obs(self, obs):
obs_list = []
for obs_group in self.obs_groups["policy"]:
obs_list.append(obs[obs_group])
return torch.cat(obs_list, dim=-1)
def get_teacher_obs(self, obs):
obs_list = []
for obs_group in self.obs_groups["teacher"]:
obs_list.append(obs[obs_group])
return torch.cat(obs_list, dim=-1)
def get_hidden_states(self):
return None
def detach_hidden_states(self, dones=None):
pass
def train(self, mode=True):
super().train(mode)
# make sure teacher is in eval mode
self.teacher.eval()
self.teacher_obs_normalizer.eval()
def update_normalization(self, obs):
if self.student_obs_normalization:
student_obs = self.get_student_obs(obs)
self.student_obs_normalizer.update(student_obs)
def load_state_dict(self, state_dict, strict=True):
"""Load the parameters of the student and teacher networks.
Args:
state_dict (dict): State dictionary of the model.
strict (bool): Whether to strictly enforce that the keys in state_dict match the keys returned by this
module's state_dict() function.
Returns:
bool: Whether this training resumes a previous training. This flag is used by the `load()` function of
`OnPolicyRunner` to determine how to load further parameters.
"""
# check if state_dict contains teacher and student or just teacher parameters
if any("actor" in key for key in state_dict.keys()): # loading parameters from rl training
# rename keys to match teacher and remove critic parameters
teacher_state_dict = {}
teacher_obs_normalizer_state_dict = {}
for key, value in state_dict.items():
if "actor." in key:
teacher_state_dict[key.replace("actor.", "")] = value
if "actor_obs_normalizer." in key:
teacher_obs_normalizer_state_dict[key.replace("actor_obs_normalizer.", "")] = value
self.teacher.load_state_dict(teacher_state_dict, strict=strict)
self.teacher_obs_normalizer.load_state_dict(teacher_obs_normalizer_state_dict, strict=strict)
# set flag for successfully loading the parameters
self.loaded_teacher = True
self.teacher.eval()
self.teacher_obs_normalizer.eval()
return False # training does not resume
elif any("student" in key for key in state_dict.keys()): # loading parameters from distillation training
super().load_state_dict(state_dict, strict=strict)
# set flag for successfully loading the parameters
self.loaded_teacher = True
self.teacher.eval()
self.teacher_obs_normalizer.eval()
return True # training resumes
else:
raise ValueError("state_dict does not contain student or teacher parameters")

View File

@ -0,0 +1,249 @@
# Copyright (c) 2021-2025, ETH Zurich and NVIDIA CORPORATION
# All rights reserved.
#
# SPDX-License-Identifier: BSD-3-Clause
from __future__ import annotations
import torch
import torch.nn as nn
import warnings
from torch.distributions import Normal
from rsl_rl.networks import MLP, EmpiricalNormalization, Memory
class StudentTeacherRecurrent(nn.Module):
is_recurrent = True
def __init__(
self,
obs,
obs_groups,
num_actions,
student_obs_normalization=False,
teacher_obs_normalization=False,
student_hidden_dims=[256, 256, 256],
teacher_hidden_dims=[256, 256, 256],
activation="elu",
init_noise_std=0.1,
noise_std_type: str = "scalar",
rnn_type="lstm",
rnn_hidden_dim=256,
rnn_num_layers=1,
teacher_recurrent=False,
**kwargs,
):
if "rnn_hidden_size" in kwargs:
warnings.warn(
"The argument `rnn_hidden_size` is deprecated and will be removed in a future version. "
"Please use `rnn_hidden_dim` instead.",
DeprecationWarning,
)
if rnn_hidden_dim == 256: # Only override if the new argument is at its default
rnn_hidden_dim = kwargs.pop("rnn_hidden_size")
if kwargs:
print(
"StudentTeacherRecurrent.__init__ got unexpected arguments, which will be ignored: "
+ str(kwargs.keys()),
)
super().__init__()
self.loaded_teacher = False # indicates if teacher has been loaded
self.teacher_recurrent = teacher_recurrent # indicates if teacher is recurrent too
# get the observation dimensions
self.obs_groups = obs_groups
num_student_obs = 0
for obs_group in obs_groups["policy"]:
assert len(obs[obs_group].shape) == 2, "The StudentTeacher module only supports 1D observations."
num_student_obs += obs[obs_group].shape[-1]
num_teacher_obs = 0
for obs_group in obs_groups["teacher"]:
assert len(obs[obs_group].shape) == 2, "The StudentTeacher module only supports 1D observations."
num_teacher_obs += obs[obs_group].shape[-1]
# student
self.memory_s = Memory(num_student_obs, type=rnn_type, num_layers=rnn_num_layers, hidden_size=rnn_hidden_dim)
self.student = MLP(rnn_hidden_dim, num_actions, student_hidden_dims, activation)
# student observation normalization
self.student_obs_normalization = student_obs_normalization
if student_obs_normalization:
self.student_obs_normalizer = EmpiricalNormalization(num_student_obs)
else:
self.student_obs_normalizer = torch.nn.Identity()
print(f"Student RNN: {self.memory_s}")
print(f"Student MLP: {self.student}")
# teacher
if self.teacher_recurrent:
self.memory_t = Memory(
num_teacher_obs, type=rnn_type, num_layers=rnn_num_layers, hidden_size=rnn_hidden_dim
)
num_teacher_obs = rnn_hidden_dim
self.teacher = MLP(num_teacher_obs, num_actions, teacher_hidden_dims, activation)
# teacher observation normalization
self.teacher_obs_normalization = teacher_obs_normalization
if teacher_obs_normalization:
self.teacher_obs_normalizer = EmpiricalNormalization(num_teacher_obs)
else:
self.teacher_obs_normalizer = torch.nn.Identity()
if self.teacher_recurrent:
print(f"Teacher RNN: {self.memory_t}")
print(f"Teacher MLP: {self.teacher}")
# action noise
self.noise_std_type = noise_std_type
if self.noise_std_type == "scalar":
self.std = nn.Parameter(init_noise_std * torch.ones(num_actions))
elif self.noise_std_type == "log":
self.log_std = nn.Parameter(torch.log(init_noise_std * torch.ones(num_actions)))
else:
raise ValueError(f"Unknown standard deviation type: {self.noise_std_type}. Should be 'scalar' or 'log'")
# action distribution (populated in update_distribution)
self.distribution = None
# disable args validation for speedup
Normal.set_default_validate_args(False)
def reset(self, dones=None, hidden_states=None):
if hidden_states is None:
hidden_states = (None, None)
self.memory_s.reset(dones, hidden_states[0])
if self.teacher_recurrent:
self.memory_t.reset(dones, hidden_states[1])
def forward(self):
raise NotImplementedError
@property
def action_mean(self):
return self.distribution.mean
@property
def action_std(self):
return self.distribution.stddev
@property
def entropy(self):
return self.distribution.entropy().sum(dim=-1)
def update_distribution(self, obs):
# compute mean
mean = self.student(obs)
# compute standard deviation
if self.noise_std_type == "scalar":
std = self.std.expand_as(mean)
elif self.noise_std_type == "log":
std = torch.exp(self.log_std).expand_as(mean)
else:
raise ValueError(f"Unknown standard deviation type: {self.noise_std_type}. Should be 'scalar' or 'log'")
# create distribution
self.distribution = Normal(mean, std)
def act(self, obs):
obs = self.get_student_obs(obs)
obs = self.student_obs_normalizer(obs)
out_mem = self.memory_s(obs).squeeze(0)
self.update_distribution(out_mem)
return self.distribution.sample()
def act_inference(self, obs):
obs = self.get_student_obs(obs)
obs = self.student_obs_normalizer(obs)
out_mem = self.memory_s(obs).squeeze(0)
return self.student(out_mem)
def evaluate(self, obs):
obs = self.get_teacher_obs(obs)
obs = self.teacher_obs_normalizer(obs)
with torch.no_grad():
if self.teacher_recurrent:
self.memory_t.eval()
obs = self.memory_t(obs).squeeze(0)
return self.teacher(obs)
def get_student_obs(self, obs):
obs_list = []
for obs_group in self.obs_groups["policy"]:
obs_list.append(obs[obs_group])
return torch.cat(obs_list, dim=-1)
def get_teacher_obs(self, obs):
obs_list = []
for obs_group in self.obs_groups["teacher"]:
obs_list.append(obs[obs_group])
return torch.cat(obs_list, dim=-1)
def get_hidden_states(self):
if self.teacher_recurrent:
return self.memory_s.hidden_states, self.memory_t.hidden_states
else:
return self.memory_s.hidden_states, None
def detach_hidden_states(self, dones=None):
self.memory_s.detach_hidden_states(dones)
if self.teacher_recurrent:
self.memory_t.detach_hidden_states(dones)
def train(self, mode=True):
super().train(mode)
# make sure teacher is in eval mode
self.teacher.eval()
self.teacher_obs_normalizer.eval()
def update_normalization(self, obs):
if self.student_obs_normalization:
student_obs = self.get_student_obs(obs)
self.student_obs_normalizer.update(student_obs)
def load_state_dict(self, state_dict, strict=True):
"""Load the parameters of the student and teacher networks.
Args:
state_dict (dict): State dictionary of the model.
strict (bool): Whether to strictly enforce that the keys in state_dict match the keys returned by this
module's state_dict() function.
Returns:
bool: Whether this training resumes a previous training. This flag is used by the `load()` function of
`OnPolicyRunner` to determine how to load further parameters.
"""
# check if state_dict contains teacher and student or just teacher parameters
if any("actor" in key for key in state_dict.keys()): # loading parameters from rl training
# rename keys to match teacher and remove critic parameters
teacher_state_dict = {}
teacher_obs_normalizer_state_dict = {}
for key, value in state_dict.items():
if "actor." in key:
teacher_state_dict[key.replace("actor.", "")] = value
if "actor_obs_normalizer." in key:
teacher_obs_normalizer_state_dict[key.replace("actor_obs_normalizer.", "")] = value
self.teacher.load_state_dict(teacher_state_dict, strict=strict)
self.teacher_obs_normalizer.load_state_dict(teacher_obs_normalizer_state_dict, strict=strict)
# also load recurrent memory if teacher is recurrent
if self.teacher_recurrent:
memory_t_state_dict = {}
for key, value in state_dict.items():
if "memory_a." in key:
memory_t_state_dict[key.replace("memory_a.", "")] = value
self.memory_t.load_state_dict(memory_t_state_dict, strict=strict)
# set flag for successfully loading the parameters
self.loaded_teacher = True
self.teacher.eval()
self.teacher_obs_normalizer.eval()
return False # training does not resume
elif any("student" in key for key in state_dict.keys()): # loading parameters from distillation training
super().load_state_dict(state_dict, strict=strict)
# set flag for successfully loading the parameters
self.loaded_teacher = True
self.teacher.eval()
self.teacher_obs_normalizer.eval()
return True # training resumes
else:
raise ValueError("state_dict does not contain student or teacher parameters")

View File

@ -0,0 +1,24 @@
# Copyright (c) 2021-2025, ETH Zurich and NVIDIA CORPORATION
# All rights reserved.
#
# SPDX-License-Identifier: BSD-3-Clause
from __future__ import annotations
def resolve_symmetry_config(alg_cfg, env):
"""Resolve the symmetry configuration.
Args:
alg_cfg: The algorithm configuration dictionary.
env: The environment.
Returns:
The resolved algorithm configuration dictionary.
"""
# if using symmetry then pass the environment config object
if "symmetry_cfg" in alg_cfg and alg_cfg["symmetry_cfg"] is not None:
# this is used by the symmetry function for handling different observation terms
alg_cfg["symmetry_cfg"]["_env"] = env
return alg_cfg

View File

@ -0,0 +1,10 @@
# Copyright (c) 2021-2025, ETH Zurich and NVIDIA CORPORATION
# All rights reserved.
#
# SPDX-License-Identifier: BSD-3-Clause
"""Definitions for components of modules."""
from .memory import Memory
from .mlp import MLP
from .normalization import EmpiricalDiscountedVariationNormalization, EmpiricalNormalization

View File

@ -0,0 +1,70 @@
# Copyright (c) 2021-2025, ETH Zurich and NVIDIA CORPORATION
# All rights reserved.
#
# SPDX-License-Identifier: BSD-3-Clause
from __future__ import annotations
import torch.nn as nn
from rsl_rl.utils import unpad_trajectories
class Memory(nn.Module):
"""Memory module for recurrent networks.
This module is used to store the hidden states of the policy.
Currently only supports GRU and LSTM.
"""
def __init__(self, input_size, type="lstm", num_layers=1, hidden_size=256):
super().__init__()
# RNN
rnn_cls = nn.GRU if type.lower() == "gru" else nn.LSTM
self.rnn = rnn_cls(input_size=input_size, hidden_size=hidden_size, num_layers=num_layers)
self.hidden_states = None
def forward(self, input, masks=None, hidden_states=None):
batch_mode = masks is not None
if batch_mode:
# batch mode: needs saved hidden states
if hidden_states is None:
raise ValueError("Hidden states not passed to memory module during policy update")
out, _ = self.rnn(input, hidden_states)
out = unpad_trajectories(out, masks)
else:
# inference/distillation mode: uses hidden states of last step
out, self.hidden_states = self.rnn(input.unsqueeze(0), self.hidden_states)
return out
def reset(self, dones=None, hidden_states=None):
if dones is None: # reset all hidden states
if hidden_states is None:
self.hidden_states = None
else:
self.hidden_states = hidden_states
elif self.hidden_states is not None: # reset hidden states of done environments
if hidden_states is None:
if isinstance(self.hidden_states, tuple): # tuple in case of LSTM
for hidden_state in self.hidden_states:
hidden_state[..., dones == 1, :] = 0.0
else:
self.hidden_states[..., dones == 1, :] = 0.0
else:
NotImplementedError(
"Resetting hidden states of done environments with custom hidden states is not implemented"
)
def detach_hidden_states(self, dones=None):
if self.hidden_states is not None:
if dones is None: # detach all hidden states
if isinstance(self.hidden_states, tuple): # tuple in case of LSTM
self.hidden_states = tuple(hidden_state.detach() for hidden_state in self.hidden_states)
else:
self.hidden_states = self.hidden_states.detach()
else: # detach hidden states of done environments
if isinstance(self.hidden_states, tuple): # tuple in case of LSTM
for hidden_state in self.hidden_states:
hidden_state[..., dones == 1, :] = hidden_state[..., dones == 1, :].detach()
else:
self.hidden_states[..., dones == 1, :] = self.hidden_states[..., dones == 1, :].detach()

View File

@ -0,0 +1,120 @@
# Copyright (c) 2021-2025, ETH Zurich and NVIDIA CORPORATION
# All rights reserved.
#
# SPDX-License-Identifier: BSD-3-Clause
from __future__ import annotations
import torch
import torch.nn as nn
from functools import reduce
from rsl_rl.utils import resolve_nn_activation
class MLP(nn.Sequential):
"""Multi-layer perceptron.
The MLP network is a sequence of linear layers and activation functions. The
last layer is a linear layer that outputs the desired dimension unless the
last activation function is specified.
It provides additional conveniences:
- If the hidden dimensions have a value of ``-1``, the dimension is inferred
from the input dimension.
- If the output dimension is a tuple, the output is reshaped to the desired
shape.
"""
def __init__(
self,
input_dim: int,
output_dim: int | tuple[int] | list[int],
hidden_dims: tuple[int] | list[int],
activation: str = "elu",
last_activation: str | None = None,
):
"""Initialize the MLP.
Args:
input_dim: Dimension of the input.
output_dim: Dimension of the output.
hidden_dims: Dimensions of the hidden layers. A value of ``-1`` indicates
that the dimension should be inferred from the input dimension.
activation: Activation function. Defaults to "elu".
last_activation: Activation function of the last layer. Defaults to None,
in which case the last layer is linear.
"""
super().__init__()
# resolve activation functions
activation_mod = resolve_nn_activation(activation)
last_activation_mod = resolve_nn_activation(last_activation) if last_activation is not None else None
# resolve number of hidden dims if they are -1
hidden_dims_processed = [input_dim if dim == -1 else dim for dim in hidden_dims]
# create layers sequentially
layers = []
layers.append(nn.Linear(input_dim, hidden_dims_processed[0]))
layers.append(activation_mod)
for layer_index in range(len(hidden_dims_processed) - 1):
layers.append(nn.Linear(hidden_dims_processed[layer_index], hidden_dims_processed[layer_index + 1]))
layers.append(activation_mod)
# add last layer
if isinstance(output_dim, int):
layers.append(nn.Linear(hidden_dims_processed[-1], output_dim))
else:
# compute the total output dimension
total_out_dim = reduce(lambda x, y: x * y, output_dim)
# add a layer to reshape the output to the desired shape
layers.append(nn.Linear(hidden_dims_processed[-1], total_out_dim))
layers.append(nn.Unflatten(dim=-1, unflattened_size=output_dim))
# add last activation function if specified
if last_activation_mod is not None:
layers.append(last_activation_mod)
# register the layers
for idx, layer in enumerate(layers):
self.add_module(f"{idx}", layer)
def init_weights(self, scales: float | tuple[float]):
"""Initialize the weights of the MLP.
Args:
scales: Scale factor for the weights.
"""
def get_scale(idx) -> float:
"""Get the scale factor for the weights of the MLP.
Args:
idx: Index of the layer.
"""
return scales[idx] if isinstance(scales, (list, tuple)) else scales
# initialize the weights
for idx, module in enumerate(self):
if isinstance(module, nn.Linear):
nn.init.orthogonal_(module.weight, gain=get_scale(idx))
nn.init.zeros_(module.bias)
def forward(self, x: torch.Tensor) -> torch.Tensor:
"""Forward pass of the MLP.
Args:
x: Input tensor.
"""
for layer in self:
x = layer(x)
return x
def reset(self, dones=None, hidden_states=None):
pass
def detach_hidden_states(self, dones=None):
pass

View File

@ -0,0 +1,130 @@
# Copyright (c) 2021-2025, ETH Zurich and NVIDIA CORPORATION
# All rights reserved.
#
# SPDX-License-Identifier: BSD-3-Clause
# Copyright (c) 2020 Preferred Networks, Inc.
from __future__ import annotations
import torch
from torch import nn
class EmpiricalNormalization(nn.Module):
"""Normalize mean and variance of values based on empirical values."""
def __init__(self, shape, eps=1e-2, until=None):
"""Initialize EmpiricalNormalization module.
Args:
shape (int or tuple of int): Shape of input values except batch axis.
eps (float): Small value for stability.
until (int or None): If this arg is specified, the module learns input values until the sum of batch sizes
exceeds it.
Note: The normalization parameters are computed over the whole batch, not for each environment separately.
"""
super().__init__()
self.eps = eps
self.until = until
self.register_buffer("_mean", torch.zeros(shape).unsqueeze(0))
self.register_buffer("_var", torch.ones(shape).unsqueeze(0))
self.register_buffer("_std", torch.ones(shape).unsqueeze(0))
self.register_buffer("count", torch.tensor(0, dtype=torch.long))
@property
def mean(self):
return self._mean.squeeze(0).clone()
@property
def std(self):
return self._std.squeeze(0).clone()
def forward(self, x):
"""Normalize mean and variance of values based on empirical values."""
return (x - self._mean) / (self._std + self.eps)
@torch.jit.unused
def update(self, x):
"""Learn input values without computing the output values of them"""
if not self.training:
return
if self.until is not None and self.count >= self.until:
return
count_x = x.shape[0]
self.count += count_x
rate = count_x / self.count
var_x = torch.var(x, dim=0, unbiased=False, keepdim=True)
mean_x = torch.mean(x, dim=0, keepdim=True)
delta_mean = mean_x - self._mean
self._mean += rate * delta_mean
self._var += rate * (var_x - self._var + delta_mean * (mean_x - self._mean))
self._std = torch.sqrt(self._var)
@torch.jit.unused
def inverse(self, y):
"""De-normalize values based on empirical values."""
return y * (self._std + self.eps) + self._mean
class EmpiricalDiscountedVariationNormalization(nn.Module):
"""Reward normalization from Pathak's large scale study on PPO.
Reward normalization. Since the reward function is non-stationary, it is useful to normalize
the scale of the rewards so that the value function can learn quickly. We did this by dividing
the rewards by a running estimate of the standard deviation of the sum of discounted rewards.
"""
def __init__(self, shape, eps=1e-2, gamma=0.99, until=None):
super().__init__()
self.emp_norm = EmpiricalNormalization(shape, eps, until)
self.disc_avg = _DiscountedAverage(gamma)
def forward(self, rew):
if self.training:
# update discounted rewards
avg = self.disc_avg.update(rew)
# update moments from discounted rewards
self.emp_norm.update(avg)
# normalize rewards with the empirical std
if self.emp_norm._std > 0:
return rew / self.emp_norm._std
else:
return rew
"""
Helper class.
"""
class _DiscountedAverage:
r"""Discounted average of rewards.
The discounted average is defined as:
.. math::
\bar{R}_t = \gamma \bar{R}_{t-1} + r_t
Args:
gamma (float): Discount factor.
"""
def __init__(self, gamma):
self.avg = None
self.gamma = gamma
def update(self, rew: torch.Tensor) -> torch.Tensor:
if self.avg is None:
self.avg = rew
else:
self.avg = self.avg * self.gamma + rew
return self.avg

View File

@ -0,0 +1,12 @@
# Copyright (c) 2021-2025, ETH Zurich and NVIDIA CORPORATION
# All rights reserved.
#
# SPDX-License-Identifier: BSD-3-Clause
"""Implementation of runners for environment-agent interaction."""
from .on_policy_runner import OnPolicyRunner # isort:skip
from .distillation_runner import DistillationRunner
from .amp_on_policy_runner import AMPOnPolicyRunner
__all__ = ["OnPolicyRunner", "DistillationRunner", "AMPOnPolicyRunner"]

View File

@ -0,0 +1,521 @@
# Copyright (c) 2021-2025, ETH Zurich and NVIDIA CORPORATION
# All rights reserved.
#
# SPDX-License-Identifier: BSD-3-Clause
from __future__ import annotations
import os
import statistics
import time
import torch
import warnings
from collections import deque
import rsl_rl
from rsl_rl.algorithms import AMP_PPO
from rsl_rl.env import VecEnv
from rsl_rl.modules import ActorCritic, ActorCriticRecurrent,DiscriminatorMulti, resolve_rnd_config, resolve_symmetry_config
from rsl_rl.utils import resolve_obs_groups, store_code_state, Normalizer, G1_AMPLoader
class AMPOnPolicyRunner:
"""On-policy runner for training and evaluation of actor-critic methods."""
def __init__(self, env: VecEnv, train_cfg: dict, log_dir: str | None = None, device="cpu"):
self.cfg = train_cfg
self.alg_cfg = train_cfg["algorithm"]
self.policy_cfg = train_cfg["policy"]
self.device = device
self.env = env
# check if multi-gpu is enabled
self._configure_multi_gpu()
# store training configuration
self.num_steps_per_env = self.cfg["num_steps_per_env"]
self.save_interval = self.cfg["save_interval"]
# query observations from environment for algorithm construction
obs = self.env.get_observations()
default_sets = ["critic"]
if "rnd_cfg" in self.alg_cfg and self.alg_cfg["rnd_cfg"] is not None:
default_sets.append("rnd_state")
self.cfg["obs_groups"] = resolve_obs_groups(obs, self.cfg["obs_groups"], default_sets)
self.amp_data = G1_AMPLoader(
device,
time_between_frames=1/50.0,
preload_transitions=True,
num_preload_transitions=train_cfg["amp_num_preload_transitions"],
motion_files=train_cfg["amp_motion_files"],
num_frames=train_cfg['amp_num_frames']
)
self.amp_observation_dim = self.amp_data.observation_dim if self.cfg["amp_num_obs"] == 0 else self.cfg["amp_num_obs"]
self.amp_num_frames = 0 if self.cfg["amp_num_frames"] == 0 else self.cfg["amp_num_frames"]
self.amp_normalizer = Normalizer(self.amp_observation_dim)
self.discriminator = DiscriminatorMulti(
self.amp_observation_dim,
train_cfg["amp_reward_coef"],
train_cfg["amp_discr_hidden_dims"],
device,
train_cfg["amp_num_frames"],
train_cfg["amp_task_reward_lerp"],
train_cfg['use_lerp'],
).to(self.device)
# create the algorithm
self.alg = self._construct_algorithm(obs)
# Decide whether to disable logging
# We only log from the process with rank 0 (main process)
self.disable_logs = self.is_distributed and self.gpu_global_rank != 0
# Logging
self.log_dir = log_dir
self.writer = None
self.tot_timesteps = 0
self.tot_time = 0
self.current_learning_iteration = 0
self.git_status_repos = [rsl_rl.__file__]
def learn(self, num_learning_iterations: int, init_at_random_ep_len: bool = False): # noqa: C901
# initialize writer
self._prepare_logging_writer()
# randomize initial episode lengths (for exploration)
if init_at_random_ep_len:
self.env.episode_length_buf = torch.randint_like(
self.env.episode_length_buf, high=int(self.env.max_episode_length)
)
# start learning
obs = self.env.get_observations().to(self.device)
self.train_mode() # switch to train mode (for dropout for example)
amp_obs = self.env.get_amp_observations()
amp_obs = amp_obs.to(self.device)
if self.amp_num_frames != 0:
self.amp_obs_frames = torch.zeros(size=(self.env.num_envs, self.amp_num_frames, self.amp_observation_dim), device=self.device)
self.amp_obs_frames = torch.concat((self.amp_obs_frames[:, 1:], amp_obs.unsqueeze(1)), dim=1)
# Book keeping
ep_infos = []
rewbuffer = deque(maxlen=100)
lenbuffer = deque(maxlen=100)
step_discrewbuffer = deque(maxlen=100)
cur_reward_sum = torch.zeros(self.env.num_envs, dtype=torch.float, device=self.device)
cur_episode_length = torch.zeros(self.env.num_envs, dtype=torch.float, device=self.device)
cur_single_step_disc_rew = torch.zeros(self.env.num_envs, dtype=torch.float, device=self.device)
# create buffers for logging extrinsic and intrinsic rewards
if self.alg.rnd:
erewbuffer = deque(maxlen=100)
irewbuffer = deque(maxlen=100)
cur_ereward_sum = torch.zeros(self.env.num_envs, dtype=torch.float, device=self.device)
cur_ireward_sum = torch.zeros(self.env.num_envs, dtype=torch.float, device=self.device)
# Ensure all parameters are in-synced
if self.is_distributed:
print(f"Synchronizing parameters for rank {self.gpu_global_rank}...")
self.alg.broadcast_parameters()
# Start training
start_iter = self.current_learning_iteration
tot_iter = start_iter + num_learning_iterations
for it in range(start_iter, tot_iter):
start = time.time()
# Rollout
with torch.inference_mode():
for _ in range(self.num_steps_per_env):
# Sample actions
actions = self.alg.act(obs,amp_obs)
# Step the environment
obs, rewards, dones, extras = self.env.step(actions.to(self.env.device))
# Move to device
obs, rewards, dones = (obs.to(self.device), rewards.to(self.device), dones.to(self.device))
next_amp_obs = self.env.get_amp_observations()
next_amp_obs = next_amp_obs.to(self.device)
next_amp_obs_with_term = torch.clone(next_amp_obs)
reset_env_ids = self.env.reset_env_ids
terminal_amp_states = self.env.get_amp_observations()[reset_env_ids]
next_amp_obs_with_term[reset_env_ids] = terminal_amp_states
self.amp_obs_frames = torch.concat((self.amp_obs_frames[:, 1:], next_amp_obs_with_term.unsqueeze(1)), dim=1)
amp_reward = torch.zeros(self.env.num_envs, device=obs.device)
mask = self.env.contact_phase[:, 0] == 1.0
if mask.any():
rewards[mask], logit, disc_reward = self.alg.discriminator.predict_amp_reward(
self.amp_obs_frames[mask], rewards[mask], normalizer=self.alg.amp_normalizer
)
amp_reward[mask] += disc_reward
# process the step
self.alg.process_env_step(obs, rewards, dones, extras, next_amp_obs_with_term, self.amp_obs_frames)
self.amp_obs_frames[reset_env_ids] = 0
amp_obs = torch.clone(next_amp_obs)
# Extract intrinsic rewards (only for logging)
intrinsic_rewards = self.alg.intrinsic_rewards if self.alg.rnd else None
# book keeping
if self.log_dir is not None:
if "episode" in extras:
ep_infos.append(extras["episode"])
elif "log" in extras:
ep_infos.append(extras["log"])
# Update rewards
if self.alg.rnd:
cur_ereward_sum += rewards
cur_ireward_sum += intrinsic_rewards # type: ignore
cur_reward_sum += rewards + intrinsic_rewards
else:
cur_reward_sum += rewards
# Update episode length
cur_episode_length += 1
cur_single_step_disc_rew += amp_reward
# Clear data for completed episodes
# -- common
new_ids = (dones > 0).nonzero(as_tuple=False)
rewbuffer.extend(cur_reward_sum[new_ids][:, 0].cpu().numpy().tolist())
lenbuffer.extend(cur_episode_length[new_ids][:, 0].cpu().numpy().tolist())
cur_reward_sum[new_ids] = 0
cur_episode_length[new_ids] = 0
to_extend_disc = (cur_single_step_disc_rew[new_ids] / self.env.max_episode_length_s)[:, 0].cpu().numpy()
step_discrewbuffer.extend(to_extend_disc.tolist())
cur_single_step_disc_rew[new_ids] = 0
# -- intrinsic and extrinsic rewards
if self.alg.rnd:
erewbuffer.extend(cur_ereward_sum[new_ids][:, 0].cpu().numpy().tolist())
irewbuffer.extend(cur_ireward_sum[new_ids][:, 0].cpu().numpy().tolist())
cur_ereward_sum[new_ids] = 0
cur_ireward_sum[new_ids] = 0
stop = time.time()
collection_time = stop - start
start = stop
# compute returns
self.alg.compute_returns(obs)
# update policy
loss_dict = self.alg.update()
stop = time.time()
learn_time = stop - start
self.current_learning_iteration = it
# log info
if self.log_dir is not None and not self.disable_logs:
# Log information
self.log(locals())
# Save model
if it % self.save_interval == 0:
self.save(os.path.join(self.log_dir, f"model_{it}.pt"))
# Clear episode infos
ep_infos.clear()
# Save code state
if it == start_iter and not self.disable_logs:
# obtain all the diff files
git_file_paths = store_code_state(self.log_dir, self.git_status_repos)
# if possible store them to wandb
if self.logger_type in ["wandb", "neptune"] and git_file_paths:
for path in git_file_paths:
self.writer.save_file(path)
# Save the final model after training
if self.log_dir is not None and not self.disable_logs:
self.save(os.path.join(self.log_dir, f"model_{self.current_learning_iteration}.pt"))
def log(self, locs: dict, width: int = 80, pad: int = 35):
# Compute the collection size
collection_size = self.num_steps_per_env * self.env.num_envs * self.gpu_world_size
# Update total time-steps and time
self.tot_timesteps += collection_size
self.tot_time += locs["collection_time"] + locs["learn_time"]
iteration_time = locs["collection_time"] + locs["learn_time"]
# -- Episode info
ep_string = ""
if locs["ep_infos"]:
for key in locs["ep_infos"][0]:
infotensor = torch.tensor([], device=self.device)
for ep_info in locs["ep_infos"]:
# handle scalar and zero dimensional tensor infos
if key not in ep_info:
continue
if not isinstance(ep_info[key], torch.Tensor):
ep_info[key] = torch.Tensor([ep_info[key]])
if len(ep_info[key].shape) == 0:
ep_info[key] = ep_info[key].unsqueeze(0)
infotensor = torch.cat((infotensor, ep_info[key].to(self.device)))
value = torch.mean(infotensor)
# log to logger and terminal
if "/" in key:
self.writer.add_scalar(key, value, locs["it"])
ep_string += f"""{f'{key}:':>{pad}} {value:.4f}\n"""
else:
self.writer.add_scalar("Episode/" + key, value, locs["it"])
ep_string += f"""{f'Mean episode {key}:':>{pad}} {value:.4f}\n"""
mean_std = self.alg.policy.action_std.mean()
fps = int(collection_size / (locs["collection_time"] + locs["learn_time"]))
# -- Losses
for key, value in locs["loss_dict"].items():
self.writer.add_scalar(f"Loss/{key}", value, locs["it"])
self.writer.add_scalar("Loss/learning_rate", self.alg.learning_rate, locs["it"])
# -- Policy
self.writer.add_scalar("Policy/mean_noise_std", mean_std.item(), locs["it"])
# -- Performance
self.writer.add_scalar("Perf/total_fps", fps, locs["it"])
self.writer.add_scalar("Perf/collection time", locs["collection_time"], locs["it"])
self.writer.add_scalar("Perf/learning_time", locs["learn_time"], locs["it"])
# -- Training
if len(locs["rewbuffer"]) > 0:
# separate logging for intrinsic and extrinsic rewards
if hasattr(self.alg, "rnd") and self.alg.rnd:
self.writer.add_scalar("Rnd/mean_extrinsic_reward", statistics.mean(locs["erewbuffer"]), locs["it"])
self.writer.add_scalar("Rnd/mean_intrinsic_reward", statistics.mean(locs["irewbuffer"]), locs["it"])
self.writer.add_scalar("Rnd/weight", self.alg.rnd.weight, locs["it"])
# everything else
self.writer.add_scalar("Train/mean_reward", statistics.mean(locs["rewbuffer"]), locs["it"])
self.writer.add_scalar("Train/mean_episode_length", statistics.mean(locs["lenbuffer"]), locs["it"])
self.writer.add_scalar('Train/mean_step_disc_reward', statistics.mean(locs['step_discrewbuffer']), locs['it'])
if self.logger_type != "wandb": # wandb does not support non-integer x-axis logging
self.writer.add_scalar("Train/mean_reward/time", statistics.mean(locs["rewbuffer"]), self.tot_time)
self.writer.add_scalar(
"Train/mean_episode_length/time", statistics.mean(locs["lenbuffer"]), self.tot_time
)
str = f" \033[1m Learning iteration {locs['it']}/{locs['tot_iter']} \033[0m "
if len(locs["rewbuffer"]) > 0:
log_string = (
f"""{'#' * width}\n"""
f"""{str.center(width, ' ')}\n\n"""
f"""{'Computation:':>{pad}} {fps:.0f} steps/s (collection: {locs[
'collection_time']:.3f}s, learning {locs['learn_time']:.3f}s)\n"""
f"""{'Mean action noise std:':>{pad}} {mean_std.item():.2f}\n"""
f"""{'Step disc reward:':>{pad}} {statistics.mean(locs['step_discrewbuffer']):.2f}\n"""
)
# -- Losses
for key, value in locs["loss_dict"].items():
log_string += f"""{f'Mean {key} loss:':>{pad}} {value:.4f}\n"""
# -- Rewards
if hasattr(self.alg, "rnd") and self.alg.rnd:
log_string += (
f"""{'Mean extrinsic reward:':>{pad}} {statistics.mean(locs['erewbuffer']):.2f}\n"""
f"""{'Mean intrinsic reward:':>{pad}} {statistics.mean(locs['irewbuffer']):.2f}\n"""
)
log_string += f"""{'Mean reward:':>{pad}} {statistics.mean(locs['rewbuffer']):.2f}\n"""
# -- episode info
log_string += f"""{'Mean episode length:':>{pad}} {statistics.mean(locs['lenbuffer']):.2f}\n"""
else:
log_string = (
f"""{'#' * width}\n"""
f"""{str.center(width, ' ')}\n\n"""
f"""{'Computation:':>{pad}} {fps:.0f} steps/s (collection: {locs[
'collection_time']:.3f}s, learning {locs['learn_time']:.3f}s)\n"""
f"""{'Mean action noise std:':>{pad}} {mean_std.item():.2f}\n"""
)
for key, value in locs["loss_dict"].items():
log_string += f"""{f'{key}:':>{pad}} {value:.4f}\n"""
log_string += ep_string
log_string += (
f"""{'-' * width}\n"""
f"""{'Total timesteps:':>{pad}} {self.tot_timesteps}\n"""
f"""{'Iteration time:':>{pad}} {iteration_time:.2f}s\n"""
f"""{'Time elapsed:':>{pad}} {time.strftime("%H:%M:%S", time.gmtime(self.tot_time))}\n"""
f"""{'ETA:':>{pad}} {time.strftime(
"%H:%M:%S",
time.gmtime(
self.tot_time / (locs['it'] - locs['start_iter'] + 1)
* (locs['start_iter'] + locs['num_learning_iterations'] - locs['it'])
)
)}\n"""
)
print(log_string)
def save(self, path: str, infos=None):
# -- Save model
saved_dict = {
"model_state_dict": self.alg.policy.state_dict(),
"optimizer_state_dict": self.alg.optimizer.state_dict(),
"iter": self.current_learning_iteration,
"infos": infos,
}
# -- Save RND model if used
if hasattr(self.alg, "rnd") and self.alg.rnd:
saved_dict["rnd_state_dict"] = self.alg.rnd.state_dict()
saved_dict["rnd_optimizer_state_dict"] = self.alg.rnd_optimizer.state_dict()
torch.save(saved_dict, path)
# upload model to external logging service
if self.logger_type in ["neptune", "wandb"] and not self.disable_logs:
self.writer.save_model(path, self.current_learning_iteration)
def load(self, path: str, load_optimizer: bool = True, map_location: str | None = None):
loaded_dict = torch.load(path, weights_only=False, map_location=map_location)
# -- Load model
resumed_training = self.alg.policy.load_state_dict(loaded_dict["model_state_dict"])
# -- Load RND model if used
if hasattr(self.alg, "rnd") and self.alg.rnd:
self.alg.rnd.load_state_dict(loaded_dict["rnd_state_dict"])
# -- load optimizer if used
if load_optimizer and resumed_training:
# -- algorithm optimizer
self.alg.optimizer.load_state_dict(loaded_dict["optimizer_state_dict"])
# -- RND optimizer if used
if hasattr(self.alg, "rnd") and self.alg.rnd:
self.alg.rnd_optimizer.load_state_dict(loaded_dict["rnd_optimizer_state_dict"])
# -- load current learning iteration
if resumed_training:
self.current_learning_iteration = loaded_dict["iter"]
return loaded_dict["infos"]
def get_inference_policy(self, device=None):
self.eval_mode() # switch to evaluation mode (dropout for example)
if device is not None:
self.alg.policy.to(device)
return self.alg.policy.act_inference
def train_mode(self):
# -- PPO
self.alg.policy.train()
self.alg.discriminator.train()
# -- RND
if hasattr(self.alg, "rnd") and self.alg.rnd:
self.alg.rnd.train()
def eval_mode(self):
# -- PPO
self.alg.policy.eval()
self.alg.discriminator.eval()
# -- RND
if hasattr(self.alg, "rnd") and self.alg.rnd:
self.alg.rnd.eval()
def add_git_repo_to_log(self, repo_file_path):
self.git_status_repos.append(repo_file_path)
"""
Helper functions.
"""
def _configure_multi_gpu(self):
"""Configure multi-gpu training."""
# check if distributed training is enabled
self.gpu_world_size = int(os.getenv("WORLD_SIZE", "1"))
self.is_distributed = self.gpu_world_size > 1
# if not distributed training, set local and global rank to 0 and return
if not self.is_distributed:
self.gpu_local_rank = 0
self.gpu_global_rank = 0
self.multi_gpu_cfg = None
return
# get rank and world size
self.gpu_local_rank = int(os.getenv("LOCAL_RANK", "0"))
self.gpu_global_rank = int(os.getenv("RANK", "0"))
# make a configuration dictionary
self.multi_gpu_cfg = {
"global_rank": self.gpu_global_rank, # rank of the main process
"local_rank": self.gpu_local_rank, # rank of the current process
"world_size": self.gpu_world_size, # total number of processes
}
# check if user has device specified for local rank
if self.device != f"cuda:{self.gpu_local_rank}":
raise ValueError(
f"Device '{self.device}' does not match expected device for local rank '{self.gpu_local_rank}'."
)
# validate multi-gpu configuration
if self.gpu_local_rank >= self.gpu_world_size:
raise ValueError(
f"Local rank '{self.gpu_local_rank}' is greater than or equal to world size '{self.gpu_world_size}'."
)
if self.gpu_global_rank >= self.gpu_world_size:
raise ValueError(
f"Global rank '{self.gpu_global_rank}' is greater than or equal to world size '{self.gpu_world_size}'."
)
# initialize torch distributed
torch.distributed.init_process_group(backend="nccl", rank=self.gpu_global_rank, world_size=self.gpu_world_size)
# set device to the local rank
torch.cuda.set_device(self.gpu_local_rank)
def _construct_algorithm(self, obs) -> AMP_PPO:
"""Construct the actor-critic algorithm."""
# resolve RND config
self.alg_cfg = resolve_rnd_config(self.alg_cfg, obs, self.cfg["obs_groups"], self.env)
# resolve symmetry config
self.alg_cfg = resolve_symmetry_config(self.alg_cfg, self.env)
# resolve deprecated normalization config
if self.cfg.get("empirical_normalization") is not None:
warnings.warn(
"The `empirical_normalization` parameter is deprecated. Please set `actor_obs_normalization` and "
"`critic_obs_normalization` as part of the `policy` configuration instead.",
DeprecationWarning,
)
if self.policy_cfg.get("actor_obs_normalization") is None:
self.policy_cfg["actor_obs_normalization"] = self.cfg["empirical_normalization"]
if self.policy_cfg.get("critic_obs_normalization") is None:
self.policy_cfg["critic_obs_normalization"] = self.cfg["empirical_normalization"]
# initialize the actor-critic
actor_critic_class = eval(self.policy_cfg.pop("class_name"))
actor_critic: ActorCritic | ActorCriticRecurrent = actor_critic_class(
obs, self.cfg["obs_groups"], self.env.num_actions, **self.policy_cfg
).to(self.device)
# initialize the algorithm
alg_class = eval(self.alg_cfg.pop("class_name"))
alg: AMP_PPO = alg_class(actor_critic, self.discriminator, self.amp_data, self.amp_normalizer, self.amp_num_frames, device=self.device, **self.alg_cfg, multi_gpu_cfg=self.multi_gpu_cfg)
# initialize the storage
alg.init_storage(
"rl",
self.env.num_envs,
self.num_steps_per_env,
obs,
[self.env.num_actions],
)
return alg
def _prepare_logging_writer(self):
"""Prepares the logging writers."""
if self.log_dir is not None and self.writer is None and not self.disable_logs:
# Launch either Tensorboard or Neptune & Tensorboard summary writer(s), default: Tensorboard.
self.logger_type = self.cfg.get("logger", "tensorboard")
self.logger_type = self.logger_type.lower()
if self.logger_type == "neptune":
from rsl_rl.utils.neptune_utils import NeptuneSummaryWriter
self.writer = NeptuneSummaryWriter(log_dir=self.log_dir, flush_secs=10, cfg=self.cfg)
self.writer.log_config(self.env.cfg, self.cfg, self.alg_cfg, self.policy_cfg)
elif self.logger_type == "wandb":
from rsl_rl.utils.wandb_utils import WandbSummaryWriter
self.writer = WandbSummaryWriter(log_dir=self.log_dir, flush_secs=10, cfg=self.cfg)
self.writer.log_config(self.env.cfg, self.cfg, self.alg_cfg, self.policy_cfg)
elif self.logger_type == "tensorboard":
from torch.utils.tensorboard import SummaryWriter
self.writer = SummaryWriter(log_dir=self.log_dir, flush_secs=10)
else:
raise ValueError("Logger type not found. Please choose 'neptune', 'wandb' or 'tensorboard'.")

View File

@ -0,0 +1,179 @@
# Copyright (c) 2021-2025, ETH Zurich and NVIDIA CORPORATION
# All rights reserved.
#
# SPDX-License-Identifier: BSD-3-Clause
from __future__ import annotations
import os
import time
import torch
from collections import deque
import rsl_rl
from rsl_rl.algorithms import Distillation
from rsl_rl.env import VecEnv
from rsl_rl.modules import StudentTeacher, StudentTeacherRecurrent
from rsl_rl.runners import OnPolicyRunner
from rsl_rl.utils import resolve_obs_groups, store_code_state
class DistillationRunner(OnPolicyRunner):
"""On-policy runner for training and evaluation of teacher-student training."""
def __init__(self, env: VecEnv, train_cfg: dict, log_dir: str | None = None, device="cpu"):
self.cfg = train_cfg
self.alg_cfg = train_cfg["algorithm"]
self.policy_cfg = train_cfg["policy"]
self.device = device
self.env = env
# check if multi-gpu is enabled
self._configure_multi_gpu()
# store training configuration
self.num_steps_per_env = self.cfg["num_steps_per_env"]
self.save_interval = self.cfg["save_interval"]
# query observations from environment for algorithm construction
obs = self.env.get_observations()
self.cfg["obs_groups"] = resolve_obs_groups(obs, self.cfg["obs_groups"], default_sets=["teacher"])
# create the algorithm
self.alg = self._construct_algorithm(obs)
# Decide whether to disable logging
# We only log from the process with rank 0 (main process)
self.disable_logs = self.is_distributed and self.gpu_global_rank != 0
# Logging
self.log_dir = log_dir
self.writer = None
self.tot_timesteps = 0
self.tot_time = 0
self.current_learning_iteration = 0
self.git_status_repos = [rsl_rl.__file__]
def learn(self, num_learning_iterations: int, init_at_random_ep_len: bool = False): # noqa: C901
# initialize writer
self._prepare_logging_writer()
# check if teacher is loaded
if not self.alg.policy.loaded_teacher:
raise ValueError("Teacher model parameters not loaded. Please load a teacher model to distill.")
# randomize initial episode lengths (for exploration)
if init_at_random_ep_len:
self.env.episode_length_buf = torch.randint_like(
self.env.episode_length_buf, high=int(self.env.max_episode_length)
)
# start learning
obs = self.env.get_observations().to(self.device)
self.train_mode() # switch to train mode (for dropout for example)
# Book keeping
ep_infos = []
rewbuffer = deque(maxlen=100)
lenbuffer = deque(maxlen=100)
cur_reward_sum = torch.zeros(self.env.num_envs, dtype=torch.float, device=self.device)
cur_episode_length = torch.zeros(self.env.num_envs, dtype=torch.float, device=self.device)
# Ensure all parameters are in-synced
if self.is_distributed:
print(f"Synchronizing parameters for rank {self.gpu_global_rank}...")
self.alg.broadcast_parameters()
# Start training
start_iter = self.current_learning_iteration
tot_iter = start_iter + num_learning_iterations
for it in range(start_iter, tot_iter):
start = time.time()
# Rollout
with torch.inference_mode():
for _ in range(self.num_steps_per_env):
# Sample actions
actions = self.alg.act(obs)
# Step the environment
obs, rewards, dones, extras = self.env.step(actions.to(self.env.device))
# Move to device
obs, rewards, dones = (obs.to(self.device), rewards.to(self.device), dones.to(self.device))
# process the step
self.alg.process_env_step(obs, rewards, dones, extras)
# book keeping
if self.log_dir is not None:
if "episode" in extras:
ep_infos.append(extras["episode"])
elif "log" in extras:
ep_infos.append(extras["log"])
# Update rewards
cur_reward_sum += rewards
# Update episode length
cur_episode_length += 1
# Clear data for completed episodes
new_ids = (dones > 0).nonzero(as_tuple=False)
rewbuffer.extend(cur_reward_sum[new_ids][:, 0].cpu().numpy().tolist())
lenbuffer.extend(cur_episode_length[new_ids][:, 0].cpu().numpy().tolist())
cur_reward_sum[new_ids] = 0
cur_episode_length[new_ids] = 0
stop = time.time()
collection_time = stop - start
start = stop
# update policy
loss_dict = self.alg.update()
stop = time.time()
learn_time = stop - start
self.current_learning_iteration = it
# log info
if self.log_dir is not None and not self.disable_logs:
# Log information
self.log(locals())
# Save model
if it % self.save_interval == 0:
self.save(os.path.join(self.log_dir, f"model_{it}.pt"))
# Clear episode infos
ep_infos.clear()
# Save code state
if it == start_iter and not self.disable_logs:
# obtain all the diff files
git_file_paths = store_code_state(self.log_dir, self.git_status_repos)
# if possible store them to wandb
if self.logger_type in ["wandb", "neptune"] and git_file_paths:
for path in git_file_paths:
self.writer.save_file(path)
# Save the final model after training
if self.log_dir is not None and not self.disable_logs:
self.save(os.path.join(self.log_dir, f"model_{self.current_learning_iteration}.pt"))
"""
Helper methods.
"""
def _construct_algorithm(self, obs) -> Distillation:
"""Construct the distillation algorithm."""
# initialize the actor-critic
student_teacher_class = eval(self.policy_cfg.pop("class_name"))
student_teacher: StudentTeacher | StudentTeacherRecurrent = student_teacher_class(
obs, self.cfg["obs_groups"], self.env.num_actions, **self.policy_cfg
).to(self.device)
# initialize the algorithm
alg_class = eval(self.alg_cfg.pop("class_name"))
alg: Distillation = alg_class(
student_teacher, device=self.device, **self.alg_cfg, multi_gpu_cfg=self.multi_gpu_cfg
)
# initialize the storage
alg.init_storage(
"distillation",
self.env.num_envs,
self.num_steps_per_env,
obs,
[self.env.num_actions],
)
return alg

View File

@ -0,0 +1,460 @@
# Copyright (c) 2021-2025, ETH Zurich and NVIDIA CORPORATION
# All rights reserved.
#
# SPDX-License-Identifier: BSD-3-Clause
from __future__ import annotations
import os
import statistics
import time
import torch
import warnings
from collections import deque
import rsl_rl
from rsl_rl.algorithms import PPO
from rsl_rl.env import VecEnv
from rsl_rl.modules import ActorCritic, ActorCriticRecurrent, resolve_rnd_config, resolve_symmetry_config
from rsl_rl.utils import resolve_obs_groups, store_code_state
class OnPolicyRunner:
"""On-policy runner for training and evaluation of actor-critic methods."""
def __init__(self, env: VecEnv, train_cfg: dict, log_dir: str | None = None, device="cpu"):
self.cfg = train_cfg
self.alg_cfg = train_cfg["algorithm"]
self.policy_cfg = train_cfg["policy"]
self.device = device
self.env = env
# check if multi-gpu is enabled
self._configure_multi_gpu()
# store training configuration
self.num_steps_per_env = self.cfg["num_steps_per_env"]
self.save_interval = self.cfg["save_interval"]
# query observations from environment for algorithm construction
obs = self.env.get_observations()
default_sets = ["critic"]
if "rnd_cfg" in self.alg_cfg and self.alg_cfg["rnd_cfg"] is not None:
default_sets.append("rnd_state")
self.cfg["obs_groups"] = resolve_obs_groups(obs, self.cfg["obs_groups"], default_sets)
# create the algorithm
self.alg = self._construct_algorithm(obs)
# Decide whether to disable logging
# We only log from the process with rank 0 (main process)
self.disable_logs = self.is_distributed and self.gpu_global_rank != 0
# Logging
self.log_dir = log_dir
self.writer = None
self.tot_timesteps = 0
self.tot_time = 0
self.current_learning_iteration = 0
self.git_status_repos = [rsl_rl.__file__]
def learn(self, num_learning_iterations: int, init_at_random_ep_len: bool = False): # noqa: C901
# initialize writer
self._prepare_logging_writer()
# randomize initial episode lengths (for exploration)
if init_at_random_ep_len:
self.env.episode_length_buf = torch.randint_like(
self.env.episode_length_buf, high=int(self.env.max_episode_length)
)
# start learning
obs = self.env.get_observations().to(self.device)
self.train_mode() # switch to train mode (for dropout for example)
# Book keeping
ep_infos = []
rewbuffer = deque(maxlen=100)
lenbuffer = deque(maxlen=100)
cur_reward_sum = torch.zeros(self.env.num_envs, dtype=torch.float, device=self.device)
cur_episode_length = torch.zeros(self.env.num_envs, dtype=torch.float, device=self.device)
# create buffers for logging extrinsic and intrinsic rewards
if self.alg.rnd:
erewbuffer = deque(maxlen=100)
irewbuffer = deque(maxlen=100)
cur_ereward_sum = torch.zeros(self.env.num_envs, dtype=torch.float, device=self.device)
cur_ireward_sum = torch.zeros(self.env.num_envs, dtype=torch.float, device=self.device)
# Ensure all parameters are in-synced
if self.is_distributed:
print(f"Synchronizing parameters for rank {self.gpu_global_rank}...")
self.alg.broadcast_parameters()
# Start training
start_iter = self.current_learning_iteration
tot_iter = start_iter + num_learning_iterations
for it in range(start_iter, tot_iter):
start = time.time()
# Rollout
with torch.inference_mode():
for _ in range(self.num_steps_per_env):
# Sample actions
actions = self.alg.act(obs)
# Step the environment
obs, rewards, dones, extras = self.env.step(actions.to(self.env.device))
# Move to device
obs, rewards, dones = (obs.to(self.device), rewards.to(self.device), dones.to(self.device))
# process the step
self.alg.process_env_step(obs, rewards, dones, extras)
# Extract intrinsic rewards (only for logging)
intrinsic_rewards = self.alg.intrinsic_rewards if self.alg.rnd else None
# book keeping
if self.log_dir is not None:
if "episode" in extras:
ep_infos.append(extras["episode"])
elif "log" in extras:
ep_infos.append(extras["log"])
# Update rewards
if self.alg.rnd:
cur_ereward_sum += rewards
cur_ireward_sum += intrinsic_rewards # type: ignore
cur_reward_sum += rewards + intrinsic_rewards
else:
cur_reward_sum += rewards
# Update episode length
cur_episode_length += 1
# Clear data for completed episodes
# -- common
new_ids = (dones > 0).nonzero(as_tuple=False)
rewbuffer.extend(cur_reward_sum[new_ids][:, 0].cpu().numpy().tolist())
lenbuffer.extend(cur_episode_length[new_ids][:, 0].cpu().numpy().tolist())
cur_reward_sum[new_ids] = 0
cur_episode_length[new_ids] = 0
# -- intrinsic and extrinsic rewards
if self.alg.rnd:
erewbuffer.extend(cur_ereward_sum[new_ids][:, 0].cpu().numpy().tolist())
irewbuffer.extend(cur_ireward_sum[new_ids][:, 0].cpu().numpy().tolist())
cur_ereward_sum[new_ids] = 0
cur_ireward_sum[new_ids] = 0
stop = time.time()
collection_time = stop - start
start = stop
# compute returns
self.alg.compute_returns(obs)
# update policy
loss_dict = self.alg.update()
stop = time.time()
learn_time = stop - start
self.current_learning_iteration = it
# log info
if self.log_dir is not None and not self.disable_logs:
# Log information
self.log(locals())
# Save model
if it % self.save_interval == 0:
self.save(os.path.join(self.log_dir, f"model_{it}.pt"))
# Clear episode infos
ep_infos.clear()
# Save code state
if it == start_iter and not self.disable_logs:
# obtain all the diff files
git_file_paths = store_code_state(self.log_dir, self.git_status_repos)
# if possible store them to wandb
if self.logger_type in ["wandb", "neptune"] and git_file_paths:
for path in git_file_paths:
self.writer.save_file(path)
# Save the final model after training
if self.log_dir is not None and not self.disable_logs:
self.save(os.path.join(self.log_dir, f"model_{self.current_learning_iteration}.pt"))
def log(self, locs: dict, width: int = 80, pad: int = 35):
# Compute the collection size
collection_size = self.num_steps_per_env * self.env.num_envs * self.gpu_world_size
# Update total time-steps and time
self.tot_timesteps += collection_size
self.tot_time += locs["collection_time"] + locs["learn_time"]
iteration_time = locs["collection_time"] + locs["learn_time"]
# -- Episode info
ep_string = ""
if locs["ep_infos"]:
for key in locs["ep_infos"][0]:
infotensor = torch.tensor([], device=self.device)
for ep_info in locs["ep_infos"]:
# handle scalar and zero dimensional tensor infos
if key not in ep_info:
continue
if not isinstance(ep_info[key], torch.Tensor):
ep_info[key] = torch.Tensor([ep_info[key]])
if len(ep_info[key].shape) == 0:
ep_info[key] = ep_info[key].unsqueeze(0)
infotensor = torch.cat((infotensor, ep_info[key].to(self.device)))
value = torch.mean(infotensor)
# log to logger and terminal
if "/" in key:
self.writer.add_scalar(key, value, locs["it"])
ep_string += f"""{f'{key}:':>{pad}} {value:.4f}\n"""
else:
self.writer.add_scalar("Episode/" + key, value, locs["it"])
ep_string += f"""{f'Mean episode {key}:':>{pad}} {value:.4f}\n"""
mean_std = self.alg.policy.action_std.mean()
fps = int(collection_size / (locs["collection_time"] + locs["learn_time"]))
# -- Losses
for key, value in locs["loss_dict"].items():
self.writer.add_scalar(f"Loss/{key}", value, locs["it"])
self.writer.add_scalar("Loss/learning_rate", self.alg.learning_rate, locs["it"])
# -- Policy
self.writer.add_scalar("Policy/mean_noise_std", mean_std.item(), locs["it"])
# -- Performance
self.writer.add_scalar("Perf/total_fps", fps, locs["it"])
self.writer.add_scalar("Perf/collection time", locs["collection_time"], locs["it"])
self.writer.add_scalar("Perf/learning_time", locs["learn_time"], locs["it"])
# -- Training
if len(locs["rewbuffer"]) > 0:
# separate logging for intrinsic and extrinsic rewards
if hasattr(self.alg, "rnd") and self.alg.rnd:
self.writer.add_scalar("Rnd/mean_extrinsic_reward", statistics.mean(locs["erewbuffer"]), locs["it"])
self.writer.add_scalar("Rnd/mean_intrinsic_reward", statistics.mean(locs["irewbuffer"]), locs["it"])
self.writer.add_scalar("Rnd/weight", self.alg.rnd.weight, locs["it"])
# everything else
self.writer.add_scalar("Train/mean_reward", statistics.mean(locs["rewbuffer"]), locs["it"])
self.writer.add_scalar("Train/mean_episode_length", statistics.mean(locs["lenbuffer"]), locs["it"])
if self.logger_type != "wandb": # wandb does not support non-integer x-axis logging
self.writer.add_scalar("Train/mean_reward/time", statistics.mean(locs["rewbuffer"]), self.tot_time)
self.writer.add_scalar(
"Train/mean_episode_length/time", statistics.mean(locs["lenbuffer"]), self.tot_time
)
str = f" \033[1m Learning iteration {locs['it']}/{locs['tot_iter']} \033[0m "
if len(locs["rewbuffer"]) > 0:
log_string = (
f"""{'#' * width}\n"""
f"""{str.center(width, ' ')}\n\n"""
f"""{'Computation:':>{pad}} {fps:.0f} steps/s (collection: {locs[
'collection_time']:.3f}s, learning {locs['learn_time']:.3f}s)\n"""
f"""{'Mean action noise std:':>{pad}} {mean_std.item():.2f}\n"""
)
# -- Losses
for key, value in locs["loss_dict"].items():
log_string += f"""{f'Mean {key} loss:':>{pad}} {value:.4f}\n"""
# -- Rewards
if hasattr(self.alg, "rnd") and self.alg.rnd:
log_string += (
f"""{'Mean extrinsic reward:':>{pad}} {statistics.mean(locs['erewbuffer']):.2f}\n"""
f"""{'Mean intrinsic reward:':>{pad}} {statistics.mean(locs['irewbuffer']):.2f}\n"""
)
log_string += f"""{'Mean reward:':>{pad}} {statistics.mean(locs['rewbuffer']):.2f}\n"""
# -- episode info
log_string += f"""{'Mean episode length:':>{pad}} {statistics.mean(locs['lenbuffer']):.2f}\n"""
else:
log_string = (
f"""{'#' * width}\n"""
f"""{str.center(width, ' ')}\n\n"""
f"""{'Computation:':>{pad}} {fps:.0f} steps/s (collection: {locs[
'collection_time']:.3f}s, learning {locs['learn_time']:.3f}s)\n"""
f"""{'Mean action noise std:':>{pad}} {mean_std.item():.2f}\n"""
)
for key, value in locs["loss_dict"].items():
log_string += f"""{f'{key}:':>{pad}} {value:.4f}\n"""
log_string += ep_string
log_string += (
f"""{'-' * width}\n"""
f"""{'Total timesteps:':>{pad}} {self.tot_timesteps}\n"""
f"""{'Iteration time:':>{pad}} {iteration_time:.2f}s\n"""
f"""{'Time elapsed:':>{pad}} {time.strftime("%H:%M:%S", time.gmtime(self.tot_time))}\n"""
f"""{'ETA:':>{pad}} {time.strftime(
"%H:%M:%S",
time.gmtime(
self.tot_time / (locs['it'] - locs['start_iter'] + 1)
* (locs['start_iter'] + locs['num_learning_iterations'] - locs['it'])
)
)}\n"""
)
print(log_string)
def save(self, path: str, infos=None):
# -- Save model
saved_dict = {
"model_state_dict": self.alg.policy.state_dict(),
"optimizer_state_dict": self.alg.optimizer.state_dict(),
"iter": self.current_learning_iteration,
"infos": infos,
}
# -- Save RND model if used
if hasattr(self.alg, "rnd") and self.alg.rnd:
saved_dict["rnd_state_dict"] = self.alg.rnd.state_dict()
saved_dict["rnd_optimizer_state_dict"] = self.alg.rnd_optimizer.state_dict()
torch.save(saved_dict, path)
# upload model to external logging service
if self.logger_type in ["neptune", "wandb"] and not self.disable_logs:
self.writer.save_model(path, self.current_learning_iteration)
def load(self, path: str, load_optimizer: bool = True, map_location: str | None = None):
loaded_dict = torch.load(path, weights_only=False, map_location=map_location)
# -- Load model
resumed_training = self.alg.policy.load_state_dict(loaded_dict["model_state_dict"])
# -- Load RND model if used
if hasattr(self.alg, "rnd") and self.alg.rnd:
self.alg.rnd.load_state_dict(loaded_dict["rnd_state_dict"])
# -- load optimizer if used
if load_optimizer and resumed_training:
# -- algorithm optimizer
self.alg.optimizer.load_state_dict(loaded_dict["optimizer_state_dict"])
# -- RND optimizer if used
if hasattr(self.alg, "rnd") and self.alg.rnd:
self.alg.rnd_optimizer.load_state_dict(loaded_dict["rnd_optimizer_state_dict"])
# -- load current learning iteration
if resumed_training:
self.current_learning_iteration = loaded_dict["iter"]
return loaded_dict["infos"]
def get_inference_policy(self, device=None):
self.eval_mode() # switch to evaluation mode (dropout for example)
if device is not None:
self.alg.policy.to(device)
return self.alg.policy.act_inference
def train_mode(self):
# -- PPO
self.alg.policy.train()
# -- RND
if hasattr(self.alg, "rnd") and self.alg.rnd:
self.alg.rnd.train()
def eval_mode(self):
# -- PPO
self.alg.policy.eval()
# -- RND
if hasattr(self.alg, "rnd") and self.alg.rnd:
self.alg.rnd.eval()
def add_git_repo_to_log(self, repo_file_path):
self.git_status_repos.append(repo_file_path)
"""
Helper functions.
"""
def _configure_multi_gpu(self):
"""Configure multi-gpu training."""
# check if distributed training is enabled
self.gpu_world_size = int(os.getenv("WORLD_SIZE", "1"))
self.is_distributed = self.gpu_world_size > 1
# if not distributed training, set local and global rank to 0 and return
if not self.is_distributed:
self.gpu_local_rank = 0
self.gpu_global_rank = 0
self.multi_gpu_cfg = None
return
# get rank and world size
self.gpu_local_rank = int(os.getenv("LOCAL_RANK", "0"))
self.gpu_global_rank = int(os.getenv("RANK", "0"))
# make a configuration dictionary
self.multi_gpu_cfg = {
"global_rank": self.gpu_global_rank, # rank of the main process
"local_rank": self.gpu_local_rank, # rank of the current process
"world_size": self.gpu_world_size, # total number of processes
}
# check if user has device specified for local rank
if self.device != f"cuda:{self.gpu_local_rank}":
raise ValueError(
f"Device '{self.device}' does not match expected device for local rank '{self.gpu_local_rank}'."
)
# validate multi-gpu configuration
if self.gpu_local_rank >= self.gpu_world_size:
raise ValueError(
f"Local rank '{self.gpu_local_rank}' is greater than or equal to world size '{self.gpu_world_size}'."
)
if self.gpu_global_rank >= self.gpu_world_size:
raise ValueError(
f"Global rank '{self.gpu_global_rank}' is greater than or equal to world size '{self.gpu_world_size}'."
)
# initialize torch distributed
torch.distributed.init_process_group(backend="nccl", rank=self.gpu_global_rank, world_size=self.gpu_world_size)
# set device to the local rank
torch.cuda.set_device(self.gpu_local_rank)
def _construct_algorithm(self, obs) -> PPO:
"""Construct the actor-critic algorithm."""
# resolve RND config
self.alg_cfg = resolve_rnd_config(self.alg_cfg, obs, self.cfg["obs_groups"], self.env)
# resolve symmetry config
self.alg_cfg = resolve_symmetry_config(self.alg_cfg, self.env)
# resolve deprecated normalization config
if self.cfg.get("empirical_normalization") is not None:
warnings.warn(
"The `empirical_normalization` parameter is deprecated. Please set `actor_obs_normalization` and "
"`critic_obs_normalization` as part of the `policy` configuration instead.",
DeprecationWarning,
)
if self.policy_cfg.get("actor_obs_normalization") is None:
self.policy_cfg["actor_obs_normalization"] = self.cfg["empirical_normalization"]
if self.policy_cfg.get("critic_obs_normalization") is None:
self.policy_cfg["critic_obs_normalization"] = self.cfg["empirical_normalization"]
# initialize the actor-critic
actor_critic_class = eval(self.policy_cfg.pop("class_name"))
actor_critic: ActorCritic | ActorCriticRecurrent = actor_critic_class(
obs, self.cfg["obs_groups"], self.env.num_actions, **self.policy_cfg
).to(self.device)
# initialize the algorithm
alg_class = eval(self.alg_cfg.pop("class_name"))
alg: PPO = alg_class(actor_critic, device=self.device, **self.alg_cfg, multi_gpu_cfg=self.multi_gpu_cfg)
# initialize the storage
alg.init_storage(
"rl",
self.env.num_envs,
self.num_steps_per_env,
obs,
[self.env.num_actions],
)
return alg
def _prepare_logging_writer(self):
"""Prepares the logging writers."""
if self.log_dir is not None and self.writer is None and not self.disable_logs:
# Launch either Tensorboard or Neptune & Tensorboard summary writer(s), default: Tensorboard.
self.logger_type = self.cfg.get("logger", "tensorboard")
self.logger_type = self.logger_type.lower()
if self.logger_type == "neptune":
from rsl_rl.utils.neptune_utils import NeptuneSummaryWriter
self.writer = NeptuneSummaryWriter(log_dir=self.log_dir, flush_secs=10, cfg=self.cfg)
self.writer.log_config(self.env.cfg, self.cfg, self.alg_cfg, self.policy_cfg)
elif self.logger_type == "wandb":
from rsl_rl.utils.wandb_utils import WandbSummaryWriter
self.writer = WandbSummaryWriter(log_dir=self.log_dir, flush_secs=10, cfg=self.cfg)
self.writer.log_config(self.env.cfg, self.cfg, self.alg_cfg, self.policy_cfg)
elif self.logger_type == "tensorboard":
from torch.utils.tensorboard import SummaryWriter
self.writer = SummaryWriter(log_dir=self.log_dir, flush_secs=10)
else:
raise ValueError("Logger type not found. Please choose 'neptune', 'wandb' or 'tensorboard'.")

View File

@ -0,0 +1,10 @@
# Copyright (c) 2021-2025, ETH Zurich and NVIDIA CORPORATION
# All rights reserved.
#
# SPDX-License-Identifier: BSD-3-Clause
"""Implementation of transitions storage for RL-agent."""
from .rollout_storage import RolloutStorage
from .replay_buffer_multi import ReplayBufferMulti
__all__ = ["RolloutStorage", "ReplayBufferMulti"]

View File

@ -0,0 +1,38 @@
import torch
import numpy as np
class ReplayBufferMulti:
"""Fixed-size buffer to store experience tuples."""
def __init__(self, obs_dim, buffer_size, num_amp_frames, device):
"""Initialize a ReplayBuffer object.
Arguments:
buffer_size (int): maximum size of buffer
"""
self.states = torch.zeros(buffer_size, num_amp_frames, obs_dim).to(device)
self.num_amp_frames = num_amp_frames
self.buffer_size = buffer_size
self.device = device
self.step = 0
self.num_samples = 0
def insert(self, states):
"""Add new states to memory."""
num_states = states.shape[0]
start_idx = self.step
end_idx = self.step + num_states
if end_idx > self.buffer_size:
self.states[self.step:self.buffer_size] = states[:self.buffer_size - self.step]
self.states[:end_idx - self.buffer_size] = states[self.buffer_size - self.step:]
else:
self.states[start_idx:end_idx] = states
self.num_samples = min(self.buffer_size, max(end_idx, self.num_samples))
self.step = (self.step + num_states) % self.buffer_size
def feed_forward_generator(self, num_mini_batch, mini_batch_size):
for _ in range(num_mini_batch):
sample_idxs = np.random.choice(self.num_samples, size=mini_batch_size)
yield (self.states[sample_idxs].to(self.device))

View File

@ -0,0 +1,260 @@
# Copyright (c) 2021-2025, ETH Zurich and NVIDIA CORPORATION
# All rights reserved.
#
# SPDX-License-Identifier: BSD-3-Clause
from __future__ import annotations
import torch
from tensordict import TensorDict
from rsl_rl.utils import split_and_pad_trajectories
class RolloutStorage:
class Transition:
def __init__(self):
self.observations = None
self.actions = None
self.privileged_actions = None
self.rewards = None
self.dones = None
self.values = None
self.actions_log_prob = None
self.action_mean = None
self.action_sigma = None
self.hidden_states = None
def clear(self):
self.__init__()
def __init__(
self,
training_type,
num_envs,
num_transitions_per_env,
obs,
actions_shape,
device="cpu",
):
# store inputs
self.training_type = training_type
self.device = device
self.num_transitions_per_env = num_transitions_per_env
self.num_envs = num_envs
self.actions_shape = actions_shape
# Core
self.observations = TensorDict(
{key: torch.zeros(num_transitions_per_env, *value.shape, device=device) for key, value in obs.items()},
batch_size=[num_transitions_per_env, num_envs],
device=self.device,
)
self.rewards = torch.zeros(num_transitions_per_env, num_envs, 1, device=self.device)
self.actions = torch.zeros(num_transitions_per_env, num_envs, *actions_shape, device=self.device)
self.dones = torch.zeros(num_transitions_per_env, num_envs, 1, device=self.device).byte()
# for distillation
if training_type == "distillation":
self.privileged_actions = torch.zeros(num_transitions_per_env, num_envs, *actions_shape, device=self.device)
# for reinforcement learning
if training_type == "rl":
self.values = torch.zeros(num_transitions_per_env, num_envs, 1, device=self.device)
self.actions_log_prob = torch.zeros(num_transitions_per_env, num_envs, 1, device=self.device)
self.mu = torch.zeros(num_transitions_per_env, num_envs, *actions_shape, device=self.device)
self.sigma = torch.zeros(num_transitions_per_env, num_envs, *actions_shape, device=self.device)
self.returns = torch.zeros(num_transitions_per_env, num_envs, 1, device=self.device)
self.advantages = torch.zeros(num_transitions_per_env, num_envs, 1, device=self.device)
# For RNN networks
self.saved_hidden_states_a = None
self.saved_hidden_states_c = None
# counter for the number of transitions stored
self.step = 0
def add_transitions(self, transition: Transition):
# check if the transition is valid
if self.step >= self.num_transitions_per_env:
raise OverflowError("Rollout buffer overflow! You should call clear() before adding new transitions.")
# Core
self.observations[self.step].copy_(transition.observations)
self.actions[self.step].copy_(transition.actions)
self.rewards[self.step].copy_(transition.rewards.view(-1, 1))
self.dones[self.step].copy_(transition.dones.view(-1, 1))
# for distillation
if self.training_type == "distillation":
self.privileged_actions[self.step].copy_(transition.privileged_actions)
# for reinforcement learning
if self.training_type == "rl":
self.values[self.step].copy_(transition.values)
self.actions_log_prob[self.step].copy_(transition.actions_log_prob.view(-1, 1))
self.mu[self.step].copy_(transition.action_mean)
self.sigma[self.step].copy_(transition.action_sigma)
# For RNN networks
self._save_hidden_states(transition.hidden_states)
# increment the counter
self.step += 1
def _save_hidden_states(self, hidden_states):
if hidden_states is None or hidden_states == (None, None):
return
# make a tuple out of GRU hidden state sto match the LSTM format
hid_a = hidden_states[0] if isinstance(hidden_states[0], tuple) else (hidden_states[0],)
hid_c = hidden_states[1] if isinstance(hidden_states[1], tuple) else (hidden_states[1],)
# initialize if needed
if self.saved_hidden_states_a is None:
self.saved_hidden_states_a = [
torch.zeros(self.observations.shape[0], *hid_a[i].shape, device=self.device) for i in range(len(hid_a))
]
self.saved_hidden_states_c = [
torch.zeros(self.observations.shape[0], *hid_c[i].shape, device=self.device) for i in range(len(hid_c))
]
# copy the states
for i in range(len(hid_a)):
self.saved_hidden_states_a[i][self.step].copy_(hid_a[i])
self.saved_hidden_states_c[i][self.step].copy_(hid_c[i])
def clear(self):
self.step = 0
def compute_returns(self, last_values, gamma, lam, normalize_advantage: bool = True):
advantage = 0
for step in reversed(range(self.num_transitions_per_env)):
# if we are at the last step, bootstrap the return value
if step == self.num_transitions_per_env - 1:
next_values = last_values
else:
next_values = self.values[step + 1]
# 1 if we are not in a terminal state, 0 otherwise
next_is_not_terminal = 1.0 - self.dones[step].float()
# TD error: r_t + gamma * V(s_{t+1}) - V(s_t)
delta = self.rewards[step] + next_is_not_terminal * gamma * next_values - self.values[step]
# Advantage: A(s_t, a_t) = delta_t + gamma * lambda * A(s_{t+1}, a_{t+1})
advantage = delta + next_is_not_terminal * gamma * lam * advantage
# Return: R_t = A(s_t, a_t) + V(s_t)
self.returns[step] = advantage + self.values[step]
# Compute the advantages
self.advantages = self.returns - self.values
# Normalize the advantages if flag is set
# This is to prevent double normalization (i.e. if per minibatch normalization is used)
if normalize_advantage:
self.advantages = (self.advantages - self.advantages.mean()) / (self.advantages.std() + 1e-8)
# for distillation
def generator(self):
if self.training_type != "distillation":
raise ValueError("This function is only available for distillation training.")
for i in range(self.num_transitions_per_env):
yield self.observations[i], self.actions[i], self.privileged_actions[i], self.dones[i]
# for reinforcement learning with feedforward networks
def mini_batch_generator(self, num_mini_batches, num_epochs=8):
if self.training_type != "rl":
raise ValueError("This function is only available for reinforcement learning training.")
batch_size = self.num_envs * self.num_transitions_per_env
mini_batch_size = batch_size // num_mini_batches
indices = torch.randperm(num_mini_batches * mini_batch_size, requires_grad=False, device=self.device)
# Core
observations = self.observations.flatten(0, 1)
actions = self.actions.flatten(0, 1)
values = self.values.flatten(0, 1)
returns = self.returns.flatten(0, 1)
# For PPO
old_actions_log_prob = self.actions_log_prob.flatten(0, 1)
advantages = self.advantages.flatten(0, 1)
old_mu = self.mu.flatten(0, 1)
old_sigma = self.sigma.flatten(0, 1)
for epoch in range(num_epochs):
for i in range(num_mini_batches):
# Select the indices for the mini-batch
start = i * mini_batch_size
end = (i + 1) * mini_batch_size
batch_idx = indices[start:end]
# Create the mini-batch
# -- Core
obs_batch = observations[batch_idx]
actions_batch = actions[batch_idx]
# -- For PPO
target_values_batch = values[batch_idx]
returns_batch = returns[batch_idx]
old_actions_log_prob_batch = old_actions_log_prob[batch_idx]
advantages_batch = advantages[batch_idx]
old_mu_batch = old_mu[batch_idx]
old_sigma_batch = old_sigma[batch_idx]
# yield the mini-batch
yield obs_batch, actions_batch, target_values_batch, advantages_batch, returns_batch, old_actions_log_prob_batch, old_mu_batch, old_sigma_batch, (
None,
None,
), None
# for reinfrocement learning with recurrent networks
def recurrent_mini_batch_generator(self, num_mini_batches, num_epochs=8):
if self.training_type != "rl":
raise ValueError("This function is only available for reinforcement learning training.")
padded_obs_trajectories, trajectory_masks = split_and_pad_trajectories(self.observations, self.dones)
mini_batch_size = self.num_envs // num_mini_batches
for ep in range(num_epochs):
first_traj = 0
for i in range(num_mini_batches):
start = i * mini_batch_size
stop = (i + 1) * mini_batch_size
dones = self.dones.squeeze(-1)
last_was_done = torch.zeros_like(dones, dtype=torch.bool)
last_was_done[1:] = dones[:-1]
last_was_done[0] = True
trajectories_batch_size = torch.sum(last_was_done[:, start:stop])
last_traj = first_traj + trajectories_batch_size
masks_batch = trajectory_masks[:, first_traj:last_traj]
obs_batch = padded_obs_trajectories[:, first_traj:last_traj]
actions_batch = self.actions[:, start:stop]
old_mu_batch = self.mu[:, start:stop]
old_sigma_batch = self.sigma[:, start:stop]
returns_batch = self.returns[:, start:stop]
advantages_batch = self.advantages[:, start:stop]
values_batch = self.values[:, start:stop]
old_actions_log_prob_batch = self.actions_log_prob[:, start:stop]
# reshape to [num_envs, time, num layers, hidden dim] (original shape: [time, num_layers, num_envs, hidden_dim])
# then take only time steps after dones (flattens num envs and time dimensions),
# take a batch of trajectories and finally reshape back to [num_layers, batch, hidden_dim]
last_was_done = last_was_done.permute(1, 0)
hid_a_batch = [
saved_hidden_states.permute(2, 0, 1, 3)[last_was_done][first_traj:last_traj]
.transpose(1, 0)
.contiguous()
for saved_hidden_states in self.saved_hidden_states_a
]
hid_c_batch = [
saved_hidden_states.permute(2, 0, 1, 3)[last_was_done][first_traj:last_traj]
.transpose(1, 0)
.contiguous()
for saved_hidden_states in self.saved_hidden_states_c
]
# remove the tuple for GRU
hid_a_batch = hid_a_batch[0] if len(hid_a_batch) == 1 else hid_a_batch
hid_c_batch = hid_c_batch[0] if len(hid_c_batch) == 1 else hid_c_batch
yield obs_batch, actions_batch, values_batch, advantages_batch, returns_batch, old_actions_log_prob_batch, old_mu_batch, old_sigma_batch, (
hid_a_batch,
hid_c_batch,
), masks_batch
first_traj = last_traj

View File

@ -0,0 +1,13 @@
# Copyright (c) 2021-2025, ETH Zurich and NVIDIA CORPORATION
# All rights reserved.
#
# SPDX-License-Identifier: BSD-3-Clause
"""Helper functions."""
from .utils import *
from .motion_loader_g1 import G1_AMPLoader
__all__ = [
"G1_AMPLoader",
]

View File

@ -0,0 +1,388 @@
import os
from os.path import join as pjoin
import glob
import json
import logging
import torch
import numpy as np
from pybullet_utils import transformations
from rsl_rl.utils import motion_util
_EPS = np.finfo(float).eps * 4.0
def quaternion_slerp(q0, q1, fraction, spin=0, shortestpath=True):
"""Batch quaternion spherical linear interpolation."""
out = torch.zeros_like(q0)
zero_mask = torch.isclose(fraction, torch.zeros_like(fraction)).squeeze()
ones_mask = torch.isclose(fraction, torch.ones_like(fraction)).squeeze()
out[zero_mask] = q0[zero_mask]
out[ones_mask] = q1[ones_mask]
d = torch.sum(q0 * q1, dim=-1, keepdim=True)
dist_mask = (torch.abs(torch.abs(d) - 1.0) < _EPS).squeeze()
out[dist_mask] = q0[dist_mask]
if shortestpath:
d_old = torch.clone(d)
d = torch.where(d_old < 0, -d, d)
q1 = torch.where(d_old < 0, -q1, q1)
angle = torch.acos(d) + spin * torch.pi
angle_mask = (torch.abs(angle) < _EPS).squeeze()
out[angle_mask] = q0[angle_mask]
final_mask = torch.logical_or(zero_mask, ones_mask)
final_mask = torch.logical_or(final_mask, dist_mask)
final_mask = torch.logical_or(final_mask, angle_mask)
final_mask = torch.logical_not(final_mask)
isin = 1.0 / angle
q0 *= torch.sin((1.0 - fraction) * angle) * isin
q1 *= torch.sin(fraction * angle) * isin
q0 += q1
out[final_mask] = q0[final_mask]
return out
class G1_AMPLoader:
def __init__(
self,
device,
time_between_frames,
motion_files,
preload_transitions=False,
num_preload_transitions=1000000,
num_frames=5,
):
"""Expert dataset provides AMP observations from Dog mocap dataset.
time_between_frames: Amount of time in seconds between transition.
"""
self.device = device
self.time_between_frames = time_between_frames
self.num_frames = num_frames
# Values to store for each trajectory.
self.trajectories = []
self.trajectories_full = []
self.trajectory_names = []
self.trajectory_idxs = []
self.trajectory_lens = [] # Traj length in seconds.
self.trajectory_weights = []
self.trajectory_frame_durations = []
self.trajectory_num_frames = []
self.motion_dir = motion_files
# import ipdb; ipdb.set_trace()
for i, motion_file in enumerate(os.listdir(motion_files)):
self.trajectory_names.append(motion_file)
motion_path = pjoin(motion_files, motion_file)
motion_data = np.load(motion_path, allow_pickle=True)
motion_data_processed = np.zeros((motion_data.shape[0],36))
for f_i in range(motion_data.shape[0]):
motion_data_processed[f_i, :3] = motion_data[f_i, :3] # base pos
motion_data_processed[f_i, 3:7] = motion_data[f_i, 3:7] # base quat (wxyz)
motion_data_processed[f_i, 7:35] = motion_data[f_i, 7:35] # base vel
'''
NOTE The order of motion_data_processed is
base pos 0:3,
base quat 3:7, wxyz
dof pos 7:36, (mujoco joint order)
'''
self.trajectories.append(torch.tensor(
motion_data_processed[:, 7:],
dtype=torch.float32,
device=self.device
))
self.trajectories_full.append(torch.tensor(
motion_data_processed,
dtype=torch.float32,
device=self.device
))
self.trajectory_idxs.append(i)
self.trajectory_weights.append(1 / len(os.listdir(motion_files)))
frame_duration = 1 / 50
self.trajectory_frame_durations.append(frame_duration)
traj_len = (motion_data_processed.shape[0] - 1) * frame_duration # seconds
self.trajectory_lens.append(traj_len)
self.trajectory_num_frames.append(float(motion_data_processed.shape[0]))
print(f"Loaded {traj_len}s. motion from {motion_file}.")
# Trajectory weights are used to sample some trajectories more than others.
self.trajectory_weights = np.array(self.trajectory_weights) / np.sum(self.trajectory_weights)
self.trajectory_frame_durations = np.array(self.trajectory_frame_durations)
self.trajectory_lens = np.array(self.trajectory_lens)
self.trajectory_num_frames = np.array(self.trajectory_num_frames)
# Preload transitions.
self.preload_transitions = preload_transitions
if self.preload_transitions:
print(f'Preloading {num_preload_transitions} transitions')
traj_idxs = self.weighted_traj_idx_sample_batch(num_preload_transitions)
times = self.traj_time_sample_batch(traj_idxs)
self.preloaded_s_prior = self.get_full_frame_at_time_batch(traj_idxs, times - self.time_between_frames)
self.preloaded_s = self.get_full_frame_at_time_batch(traj_idxs, times)
self.preloaded_s_next = self.get_full_frame_at_time_batch(traj_idxs, times + self.time_between_frames)
print(f'Finished preloading')
# 预加载多帧数据
self.preloaded_frames = []
for i in range(self.num_frames):
frame_time = times + (i - (self.num_frames - 2)) * self.time_between_frames
full_frame = self.get_full_frame_at_time_batch(traj_idxs, frame_time)
# 预处理提前提取并连接需要的列7:26 和 29:33避免每次生成时重复切片
processed_frame = torch.cat([
full_frame[:, 7:26],
full_frame[:, 29:33]
], dim=-1)
self.preloaded_frames.append(processed_frame)
print(f'Finished preloading multiple frames')
self.all_trajectories_full = torch.vstack(self.trajectories_full)
def weighted_traj_idx_sample(self):
"""Get traj idx via weighted sampling."""
return np.random.choice(
self.trajectory_idxs, p=self.trajectory_weights)
def weighted_traj_idx_sample_batch(self, size):
"""Batch sample traj idxs."""
return np.random.choice(
self.trajectory_idxs, size=size, p=self.trajectory_weights,
replace=True)
def traj_time_sample(self, traj_idx):
"""Sample random time for traj."""
subst = self.time_between_frames + self.trajectory_frame_durations[traj_idx]
return max(
0, (self.trajectory_lens[traj_idx] * np.random.uniform() - subst))
def traj_time_sample_batch(self, traj_idxs):
"""Sample random time for multiple trajectories."""
subst = self.time_between_frames + self.trajectory_frame_durations[traj_idxs]
time_samples = self.trajectory_lens[traj_idxs] * np.random.uniform(size=len(traj_idxs)) - subst
return np.maximum(np.zeros_like(time_samples), time_samples)
def slerp(self, val0, val1, blend):
return (1.0 - blend) * val0 + blend * val1
def get_trajectory(self, traj_idx):
"""Returns trajectory of AMP observations."""
return self.trajectories_full[traj_idx]
def get_frame_at_time(self, traj_idx, time):
"""Returns frame for the given trajectory at the specified time."""
p = float(time) / self.trajectory_lens[traj_idx]
n = self.trajectories[traj_idx].shape[0]
idx_low, idx_high = int(np.floor(p * n)), int(np.ceil(p * n))
frame_start = self.trajectories[traj_idx][idx_low]
frame_end = self.trajectories[traj_idx][idx_high]
blend = p * n - idx_low
return self.slerp(frame_start, frame_end, blend)
def get_frame_at_time_batch(self, traj_idxs, times):
"""Returns frame for the given trajectory at the specified time."""
p = times / self.trajectory_lens[traj_idxs]
n = self.trajectory_num_frames[traj_idxs]
idx_low, idx_high = np.floor(p * n).astype(np.int32), np.ceil(p * n).astype(np.int32)
all_frame_starts = torch.zeros(len(traj_idxs), self.observation_dim, device=self.device)
all_frame_ends = torch.zeros(len(traj_idxs), self.observation_dim, device=self.device)
for traj_idx in set(traj_idxs):
trajectory = self.trajectories[traj_idx]
traj_mask = traj_idxs == traj_idx
all_frame_starts[traj_mask] = trajectory[idx_low[traj_mask]]
all_frame_ends[traj_mask] = trajectory[idx_high[traj_mask]]
blend = torch.tensor(p * n - idx_low, device=self.device, dtype=torch.float32).unsqueeze(-1)
return self.slerp(all_frame_starts, all_frame_ends, blend)
def get_full_frame_at_time(self, traj_idx, time):
"""Returns full frame for the given trajectory at the specified time."""
p = float(time) / self.trajectory_lens[traj_idx]
n = self.trajectories_full[traj_idx].shape[0]
idx_low, idx_high = int(np.floor(p * n)), int(np.ceil(p * n))
frame_start = self.trajectories_full[traj_idx][idx_low]
frame_end = self.trajectories_full[traj_idx][idx_high]
blend = p * n - idx_low
print(idx_low, idx_high)
return self.blend_frame_pose(frame_start, frame_end, blend)
def get_full_frame_at_time_batch(self, traj_idxs, times):
p = times / self.trajectory_lens[traj_idxs]
n = self.trajectory_num_frames[traj_idxs]
idx_low, idx_high = np.floor(p * n).astype(np.int32), np.ceil(p * n).astype(np.int32)
all_frame_pos_starts = torch.zeros(len(traj_idxs), 3, device=self.device)
all_frame_pos_ends = torch.zeros(len(traj_idxs), 3, device=self.device)
all_frame_rot_starts = torch.zeros(len(traj_idxs), 4, device=self.device)
all_frame_rot_ends = torch.zeros(len(traj_idxs), 4, device=self.device)
all_frame_amp_starts = torch.zeros(len(traj_idxs), 29, device=self.device)
all_frame_amp_ends = torch.zeros(len(traj_idxs), 29, device=self.device)
for traj_idx in set(traj_idxs):
trajectory = self.trajectories_full[traj_idx]
traj_mask = traj_idxs == traj_idx
all_frame_pos_starts[traj_mask] = G1_AMPLoader.get_root_pos_batch(trajectory[idx_low[traj_mask]])
all_frame_pos_ends[traj_mask] = G1_AMPLoader.get_root_pos_batch(trajectory[idx_high[traj_mask]])
all_frame_rot_starts[traj_mask] = G1_AMPLoader.get_root_rot_batch(trajectory[idx_low[traj_mask]])
all_frame_rot_ends[traj_mask] = G1_AMPLoader.get_root_rot_batch(trajectory[idx_high[traj_mask]])
all_frame_amp_starts[traj_mask] = trajectory[idx_low[traj_mask]][:, 7:36] # base vel3+ang3, dof vel23+ang23
all_frame_amp_ends[traj_mask] = trajectory[idx_high[traj_mask]][:, 7:36] # base vel3+ang3, dof vel23+ang23
blend = torch.tensor(p * n - idx_low, device=self.device, dtype=torch.float32).unsqueeze(-1)
pos_blend = self.slerp(all_frame_pos_starts, all_frame_pos_ends, blend)
rot_blend = quaternion_slerp(all_frame_rot_starts, all_frame_rot_ends, blend)
amp_blend = self.slerp(all_frame_amp_starts, all_frame_amp_ends, blend)
return torch.cat([pos_blend, rot_blend, amp_blend], dim=-1)
def get_frame(self):
"""Returns random frame."""
traj_idx = self.weighted_traj_idx_sample()
sampled_time = self.traj_time_sample(traj_idx)
return self.get_frame_at_time(traj_idx, sampled_time)
def get_full_frame(self):
"""Returns random full frame."""
traj_idx = self.weighted_traj_idx_sample()
sampled_time = self.traj_time_sample(traj_idx)
return self.get_full_frame_at_time(traj_idx, sampled_time)
def get_full_frame_batch(self, num_frames):
if self.preload_transitions:
idxs = np.random.choice(
self.preloaded_s.shape[0], size=num_frames)
return self.preloaded_s[idxs]
else:
traj_idxs = self.weighted_traj_idx_sample_batch(num_frames)
times = self.traj_time_sample_batch(traj_idxs)
return self.get_full_frame_at_time_batch(traj_idxs, times)
def blend_frame_pose(self, frame0, frame1, blend):
"""Linearly interpolate between two frames, including orientation.
Args:
frame0: First frame to be blended corresponds to (blend = 0).
frame1: Second frame to be blended corresponds to (blend = 1).
blend: Float between [0, 1], specifying the interpolation between
the two frames.
Returns:
An interpolation of the two frames.
"""
root_pos0, root_pos1 = G1_AMPLoader.get_root_pos(frame0), G1_AMPLoader.get_root_pos(frame1)
root_rot0, root_rot1 = G1_AMPLoader.get_root_rot(frame0), G1_AMPLoader.get_root_rot(frame1)
joints0, joints1 = G1_AMPLoader.get_joint_pose(frame0), G1_AMPLoader.get_joint_pose(frame1)
# tar_toe_pos_0, tar_toe_pos_1 = G1_AMPLoader.get_tar_toe_pos_local(frame0), G1_AMPLoader.get_tar_toe_pos_local(frame1)
linear_vel_0, linear_vel_1 = G1_AMPLoader.get_linear_vel(frame0), G1_AMPLoader.get_linear_vel(frame1)
angular_vel_0, angular_vel_1 = G1_AMPLoader.get_angular_vel(frame0), G1_AMPLoader.get_angular_vel(frame1)
joint_vel_0, joint_vel_1 = G1_AMPLoader.get_joint_vel(frame0), G1_AMPLoader.get_joint_vel(frame1)
blend_root_pos = self.slerp(root_pos0, root_pos1, blend)
blend_root_rot = transformations.quaternion_slerp(root_rot0.cpu().numpy(), root_rot1.cpu().numpy(), blend)
blend_root_rot = torch.tensor(motion_util.standardize_quaternion(blend_root_rot),dtype=torch.float32, device=self.device)
blend_joints = self.slerp(joints0, joints1, blend)
# blend_tar_toe_pos = self.slerp(tar_toe_pos_0, tar_toe_pos_1, blend)
blend_linear_vel = self.slerp(linear_vel_0, linear_vel_1, blend)
blend_angular_vel = self.slerp(angular_vel_0, angular_vel_1, blend)
blend_joints_vel = self.slerp(joint_vel_0, joint_vel_1, blend)
# return
# torch.cat([
# blend_root_pos, blend_root_rot, blend_linear_vel, blend_angular_vel, blend_joints, blend_joints_vel])
return torch.cat([blend_root_pos, blend_root_rot, blend_linear_vel, blend_angular_vel, blend_joints])
def feed_forward_generator_23dof_multi(self, num_mini_batch, mini_batch_size):
"""Generates a batch of AMP transitions."""
# import ipdb; ipdb.set_trace()
for _ in range(num_mini_batch):
if self.preload_transitions:
idxs = np.random.choice(self.preloaded_s.shape[0], size=mini_batch_size)
frames = []
for i in range(self.num_frames):
# 数据已在预加载时预处理,直接索引即可
s = self.preloaded_frames[i][idxs]
frames.append(s)
else:
NotImplementedError('preload transition')
yield torch.stack(frames, dim=1) # [batch, num_frames, 16]
def quaternion_to_euler_array(self, quat):
# Ensure quaternion is in the correct format [x, y, z, w]
x, y, z, w =quat
# Roll (x-axis rotation)
t0 = +2.0 * (w * x + y * z)
t1 = +1.0 - 2.0 * (x * x + y * y)
roll_x = np.arctan2(t0, t1)
# Pitch (y-axis rotation)
t2 = +2.0 * (w * y - z * x)
t2 = np.clip(t2, -1.0, 1.0)
pitch_y = np.arcsin(t2)
# Yaw (z-axis rotation)
t3 = +2.0 * (w * z + x * y)
t4 = +1.0 - 2.0 * (y * y + z * z)
yaw_z = np.arctan2(t3, t4)
# Returns roll, pitch, yaw in a NumPy array in radians
return np.array([roll_x, pitch_y, yaw_z])
def euler_to_quaternion(self, root_rot):
roll, pitch, yaw = root_rot[0], root_rot[1], root_rot[2]
cy = np.cos(yaw * 0.5)
sy = np.sin(yaw * 0.5)
cp = np.cos(pitch * 0.5)
sp = np.sin(pitch * 0.5)
cr = np.cos(roll * 0.5)
sr = np.sin(roll * 0.5)
qw = cy * cp * cr + sy * sp * sr
qx = cy * cp * sr - sy * sp * cr
qy = sy * cp * sr + cy * sp * cr
qz = sy * cp * cr - cy * sp * sr
return np.array([qx, qy, qz, qw])
@property
def observation_dim(self):
"""Size of AMP observations."""
return self.trajectories[0].shape[1] + 1
@property
def num_motions(self):
return len(self.trajectory_names)
@staticmethod
def get_root_pos(pose):
return pose[0:3]
@staticmethod
def get_root_pos_batch(poses):
return poses[:, 0:3]
@staticmethod
def get_root_rot(pose):
return pose[3:7]
@staticmethod
def get_root_rot_batch(poses):
return poses[:, 3:7]
@staticmethod
def get_joint_pose_batch_12dof(poses):
return poses[:, 13:25]
@staticmethod
def get_tar_toe_pos_local(pose):
return pose[G1_AMPLoader.TAR_TOE_POS_LOCAL_START_IDX:G1_AMPLoader.TAR_TOE_POS_LOCAL_END_IDX]
@staticmethod
def get_tar_toe_pos_local_batch(poses):
return poses[:, G1_AMPLoader.TAR_TOE_POS_LOCAL_START_IDX:G1_AMPLoader.TAR_TOE_POS_LOCAL_END_IDX]

View File

@ -0,0 +1,97 @@
# coding=utf-8
# Copyright 2020 The Google Research Authors.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
"""Utility functions for processing motion clips."""
import os
import inspect
currentdir = os.path.dirname(os.path.abspath(inspect.getfile(inspect.currentframe())))
parentdir = os.path.dirname(os.path.dirname(currentdir))
os.sys.path.insert(0, parentdir)
import numpy as np
from rsl_rl.utils import pose3d
# from pybullet_utils import transformations
def standardize_quaternion(q):
"""Returns a quaternion where q.w >= 0 to remove redundancy due to q = -q.
Args:
q: A quaternion to be standardized.
Returns:
A quaternion with q.w >= 0.
"""
if q[-1] < 0:
q = -q
return q
def normalize_rotation_angle(theta):
"""Returns a rotation angle normalized between [-pi, pi].
Args:
theta: angle of rotation (radians).
Returns:
An angle of rotation normalized between [-pi, pi].
"""
norm_theta = theta
if np.abs(norm_theta) > np.pi:
norm_theta = np.fmod(norm_theta, 2 * np.pi)
if norm_theta >= 0:
norm_theta += -2 * np.pi
else:
norm_theta += 2 * np.pi
return norm_theta
def calc_heading(q):
"""Returns the heading of a rotation q, specified as a quaternion.
The heading represents the rotational component of q along the vertical
axis (z axis).
Args:
q: A quaternion that the heading is to be computed from.
Returns:
An angle representing the rotation about the z axis.
"""
ref_dir = np.array([1, 0, 0])
rot_dir = pose3d.QuaternionRotatePoint(ref_dir, q)
heading = np.arctan2(rot_dir[1], rot_dir[0])
return heading
# def calc_heading_rot(q):
# """Return a quaternion representing the heading rotation of q along the vertical axis (z axis).
# Args:
# q: A quaternion that the heading is to be computed from.
# Returns:
# A quaternion representing the rotation about the z axis.
# """
# heading = calc_heading(q)
# q_heading = transformations.quaternion_about_axis(heading, [0, 0, 1])
# return q_heading

Some files were not shown because too many files have changed in this diff Show More