Upload folder using huggingface_hub
Browse filesThis view is limited to 50 files because it contains too many changes. See raw diff
- .gitignore +246 -0
- COPYRIGHT +4 -0
- LICENSE +190 -0
- README.md +274 -0
- docker/Dockerfile +71 -0
- docs/DATASET_TO_GEAR_AND_TRAIN.md +471 -0
- docs/DROID_CONVERSION.md +62 -0
- docs/WAN22_BACKBONE.md +167 -0
- docs/data-preparation.md +194 -0
- docs/training-guide.md +222 -0
- eval_utils/policy_client.py +93 -0
- eval_utils/policy_server.py +130 -0
- eval_utils/run_sim_eval.py +219 -0
- eval_utils/serve_dreamzero_wan22.py +391 -0
- groot/__init__.py +0 -0
- groot/control/__init__.py +0 -0
- groot/control/tensorrt_utils.py +852 -0
- groot/vla/__init__.py +0 -0
- groot/vla/common/__init__.py +1 -0
- groot/vla/common/utils/__init__.py +3 -0
- groot/vla/common/utils/data_structure/__init__.py +2 -0
- groot/vla/common/utils/data_structure/shape_utils.py +283 -0
- groot/vla/common/utils/data_structure/tree_utils.py +219 -0
- groot/vla/common/utils/io/__init__.py +6 -0
- groot/vla/common/utils/io/config_utils.py +260 -0
- groot/vla/common/utils/io/file_utils.py +707 -0
- groot/vla/common/utils/io/hdf5_utils.py +84 -0
- groot/vla/common/utils/io/json_utils.py +270 -0
- groot/vla/common/utils/io/print_utils.py +362 -0
- groot/vla/common/utils/io/termcolor.py +186 -0
- groot/vla/common/utils/misc/__init__.py +5 -0
- groot/vla/common/utils/misc/array_tensor_utils.py +372 -0
- groot/vla/common/utils/misc/functional_utils.py +635 -0
- groot/vla/common/utils/misc/image_utils.py +225 -0
- groot/vla/common/utils/misc/misc_utils.py +261 -0
- groot/vla/common/utils/misc/torch_utils.py +748 -0
- groot/vla/common/utils/misc/video_utils.py +487 -0
- groot/vla/configs/conf.yaml +135 -0
- groot/vla/configs/data/dreamzero/agibot_relative.yaml +54 -0
- groot/vla/configs/data/dreamzero/base_48_wan_fine_aug_relative.yaml +439 -0
- groot/vla/configs/data/dreamzero/droid_relative.yaml +49 -0
- groot/vla/configs/data/dreamzero/droid_relative_wan22.yaml +51 -0
- groot/vla/configs/data/dreamzero/libero.yaml +85 -0
- groot/vla/configs/data/dreamzero/libero_sim_relative.yaml +58 -0
- groot/vla/configs/data/dreamzero/libero_sim_relative.yaml.bak +54 -0
- groot/vla/configs/data/dreamzero/manifeel.yaml +86 -0
- groot/vla/configs/data/dreamzero/robotwin.yaml +35 -0
- groot/vla/configs/data/dreamzero/yam_relative.yaml +52 -0
- groot/vla/configs/deepspeed/zero2.json +28 -0
- groot/vla/configs/deepspeed/zero2_offload.json +32 -0
.gitignore
ADDED
|
@@ -0,0 +1,246 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# Byte-compiled / optimized / DLL files
|
| 2 |
+
__pycache__/
|
| 3 |
+
*.py[cod]
|
| 4 |
+
*$py.class
|
| 5 |
+
wandb/
|
| 6 |
+
|
| 7 |
+
# C extensions
|
| 8 |
+
|
| 9 |
+
|
| 10 |
+
# Distribution / packaging
|
| 11 |
+
.Python
|
| 12 |
+
build/
|
| 13 |
+
develop-eggs/
|
| 14 |
+
dist/
|
| 15 |
+
downloads/
|
| 16 |
+
eggs/
|
| 17 |
+
.eggs/
|
| 18 |
+
sdist/
|
| 19 |
+
var/
|
| 20 |
+
wheels/
|
| 21 |
+
share/python-wheels/
|
| 22 |
+
*.egg-info/
|
| 23 |
+
.installed.cfg
|
| 24 |
+
*.egg
|
| 25 |
+
*.swp
|
| 26 |
+
*.swo
|
| 27 |
+
gear_working_dir/
|
| 28 |
+
MANIFEST
|
| 29 |
+
|
| 30 |
+
# PyInstaller
|
| 31 |
+
# Usually these files are written by a python script from a template
|
| 32 |
+
# before PyInstaller builds the exe, so as to inject date/other infos into it.
|
| 33 |
+
*.manifest
|
| 34 |
+
*.spec
|
| 35 |
+
|
| 36 |
+
# Installer logs
|
| 37 |
+
pip-log.txt
|
| 38 |
+
pip-delete-this-directory.txt
|
| 39 |
+
|
| 40 |
+
# Unit test / coverage reports
|
| 41 |
+
htmlcov/
|
| 42 |
+
.tox/
|
| 43 |
+
.nox/
|
| 44 |
+
.coverage
|
| 45 |
+
.coverage.*
|
| 46 |
+
.cache
|
| 47 |
+
nosetests.xml
|
| 48 |
+
coverage.xml
|
| 49 |
+
*.cover
|
| 50 |
+
*.py,cover
|
| 51 |
+
.hypothesis/
|
| 52 |
+
.pytest_cache/
|
| 53 |
+
cover/
|
| 54 |
+
/groot/control/wbc_checkpoints/
|
| 55 |
+
# Translations
|
| 56 |
+
*.mo
|
| 57 |
+
*.pot
|
| 58 |
+
|
| 59 |
+
# Django stuff:
|
| 60 |
+
*.log
|
| 61 |
+
local_settings.py
|
| 62 |
+
db.sqlite3
|
| 63 |
+
db.sqlite3-journal
|
| 64 |
+
|
| 65 |
+
# Flask stuff:
|
| 66 |
+
instance/
|
| 67 |
+
.webassets-cache
|
| 68 |
+
|
| 69 |
+
.vscode/
|
| 70 |
+
.cursor/
|
| 71 |
+
.cursor*
|
| 72 |
+
.claude/
|
| 73 |
+
CLAUDE.md
|
| 74 |
+
|
| 75 |
+
# Scrapy stuff:
|
| 76 |
+
.scrapy
|
| 77 |
+
|
| 78 |
+
# Sphinx documentation
|
| 79 |
+
docs/_build/
|
| 80 |
+
|
| 81 |
+
# PyBuilder
|
| 82 |
+
.pybuilder/
|
| 83 |
+
target/
|
| 84 |
+
|
| 85 |
+
# Jupyter Notebook
|
| 86 |
+
.ipynb_checkpoints
|
| 87 |
+
|
| 88 |
+
# IPython
|
| 89 |
+
profile_default/
|
| 90 |
+
ipython_config.py
|
| 91 |
+
|
| 92 |
+
# pyenv
|
| 93 |
+
# For a library or package, you might want to ignore these files since the code is
|
| 94 |
+
# intended to run in multiple environments; otherwise, check them in:
|
| 95 |
+
# .python-version
|
| 96 |
+
|
| 97 |
+
# pipenv
|
| 98 |
+
# According to pypa/pipenv#598, it is recommended to include Pipfile.lock in version control.
|
| 99 |
+
# However, in case of collaboration, if having platform-specific dependencies or dependencies
|
| 100 |
+
# having no cross-platform support, pipenv may install dependencies that don't work, or not
|
| 101 |
+
# install all needed dependencies.
|
| 102 |
+
#Pipfile.lock
|
| 103 |
+
|
| 104 |
+
# UV
|
| 105 |
+
# Similar to Pipfile.lock, it is generally recommended to include uv.lock in version control.
|
| 106 |
+
# This is especially recommended for binary packages to ensure reproducibility, and is more
|
| 107 |
+
# commonly ignored for libraries.
|
| 108 |
+
#uv.lock
|
| 109 |
+
|
| 110 |
+
# poetry
|
| 111 |
+
# Similar to Pipfile.lock, it is generally recommended to include poetry.lock in version control.
|
| 112 |
+
# This is especially recommended for binary packages to ensure reproducibility, and is more
|
| 113 |
+
# commonly ignored for libraries.
|
| 114 |
+
# https://python-poetry.org/docs/basic-usage/#commit-your-poetrylock-file-to-version-control
|
| 115 |
+
#poetry.lock
|
| 116 |
+
|
| 117 |
+
# pdm
|
| 118 |
+
# Similar to Pipfile.lock, it is generally recommended to include pdm.lock in version control.
|
| 119 |
+
#pdm.lock
|
| 120 |
+
# pdm stores project-wide configurations in .pdm.toml, but it is recommended to not include it
|
| 121 |
+
# in version control.
|
| 122 |
+
# https://pdm.fming.dev/latest/usage/project/#working-with-version-control
|
| 123 |
+
.pdm.toml
|
| 124 |
+
.pdm-python
|
| 125 |
+
.pdm-build/
|
| 126 |
+
|
| 127 |
+
# PEP 582; used by e.g. github.com/David-OConnor/pyflow and github.com/pdm-project/pdm
|
| 128 |
+
__pypackages__/
|
| 129 |
+
|
| 130 |
+
# Celery stuff
|
| 131 |
+
celerybeat-schedule
|
| 132 |
+
celerybeat.pid
|
| 133 |
+
|
| 134 |
+
# SageMath parsed files
|
| 135 |
+
*.sage.py
|
| 136 |
+
|
| 137 |
+
# Spyder project settings
|
| 138 |
+
.spyderproject
|
| 139 |
+
.spyproject
|
| 140 |
+
|
| 141 |
+
# Rope project settings
|
| 142 |
+
.ropeproject
|
| 143 |
+
|
| 144 |
+
# mkdocs documentation
|
| 145 |
+
/site
|
| 146 |
+
|
| 147 |
+
# mypy
|
| 148 |
+
.mypy_cache/
|
| 149 |
+
.dmypy.json
|
| 150 |
+
dmypy.json
|
| 151 |
+
|
| 152 |
+
# Pyre type checker
|
| 153 |
+
.pyre/
|
| 154 |
+
|
| 155 |
+
# pytype static type analyzer
|
| 156 |
+
.pytype/
|
| 157 |
+
|
| 158 |
+
# Cython debug symbols
|
| 159 |
+
cython_debug/
|
| 160 |
+
# IDE
|
| 161 |
+
.idea/
|
| 162 |
+
.vscode/
|
| 163 |
+
|
| 164 |
+
# log
|
| 165 |
+
outputs/
|
| 166 |
+
logs/
|
| 167 |
+
*logs_rl*
|
| 168 |
+
!external_dependencies/OpenHomie/HomieRL/legged_gym/logs/exported/policies/*.onnx
|
| 169 |
+
|
| 170 |
+
# Ruff stuff:
|
| 171 |
+
.ruff_cache/
|
| 172 |
+
|
| 173 |
+
# PyPI configuration file
|
| 174 |
+
.pypirc
|
| 175 |
+
|
| 176 |
+
outputs/
|
| 177 |
+
osmo/
|
| 178 |
+
source/
|
| 179 |
+
groot/rl/data/
|
| 180 |
+
|
| 181 |
+
.DS_Store
|
| 182 |
+
sim_dependencies/
|
| 183 |
+
.venv/*
|
| 184 |
+
|
| 185 |
+
|
| 186 |
+
# Mujoco
|
| 187 |
+
MUJOCO_LOG.TXT
|
| 188 |
+
playground/
|
| 189 |
+
|
| 190 |
+
*.code-workspace
|
| 191 |
+
batch_*.sh
|
| 192 |
+
logs_*/
|
| 193 |
+
runs/
|
| 194 |
+
out/
|
| 195 |
+
recordings/
|
| 196 |
+
pyrightconfig.json
|
| 197 |
+
*.npz
|
| 198 |
+
*.nbize.py
|
| 199 |
+
|
| 200 |
+
# Git worktrees
|
| 201 |
+
/wt
|
| 202 |
+
|
| 203 |
+
# Gear working directory
|
| 204 |
+
gear_working_dir/
|
| 205 |
+
|
| 206 |
+
episode_data/
|
| 207 |
+
|
| 208 |
+
# third party packages
|
| 209 |
+
third_parties/
|
| 210 |
+
|
| 211 |
+
# redis related
|
| 212 |
+
*.rdb
|
| 213 |
+
download_model.py
|
| 214 |
+
.dockerignore
|
| 215 |
+
|
| 216 |
+
# Sysid files
|
| 217 |
+
plots/
|
| 218 |
+
sysid_data/
|
| 219 |
+
isaac_data/
|
| 220 |
+
|
| 221 |
+
# External dependencies
|
| 222 |
+
external_dependencies/ws_lidar_slam_ros2/build
|
| 223 |
+
external_dependencies/ws_lidar_slam_ros2/install
|
| 224 |
+
external_dependencies/ws_lidar_slam_ros2/log
|
| 225 |
+
external_dependencies/ws_slam/
|
| 226 |
+
external_dependencies/Livox-SDK2/
|
| 227 |
+
external_dependencies/mfm
|
| 228 |
+
# external_dependencies/genmo
|
| 229 |
+
external_dependencies/whole_body_tracking
|
| 230 |
+
inputs/
|
| 231 |
+
outputs/
|
| 232 |
+
|
| 233 |
+
groot/dexmg/grootrobosuite/docs/images/
|
| 234 |
+
external_dependencies/mfm/
|
| 235 |
+
dump.rdb
|
| 236 |
+
|
| 237 |
+
# Open loop eval plots
|
| 238 |
+
offline_open_loop_eval_plots/
|
| 239 |
+
|
| 240 |
+
# yam eval data
|
| 241 |
+
groot/control/envs/yam/data/eval/*
|
| 242 |
+
/models
|
| 243 |
+
/data
|
| 244 |
+
|
| 245 |
+
checkpoints/
|
| 246 |
+
video_pred_output*/
|
COPYRIGHT
ADDED
|
@@ -0,0 +1,4 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
Copyright (c) 2025 NVIDIA Corporation. All rights reserved.
|
| 2 |
+
|
| 3 |
+
Licensed under the Apache License, Version 2.0.
|
| 4 |
+
See LICENSE for the full license text.
|
LICENSE
ADDED
|
@@ -0,0 +1,190 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
Apache License
|
| 2 |
+
Version 2.0, January 2004
|
| 3 |
+
http://www.apache.org/licenses/
|
| 4 |
+
|
| 5 |
+
TERMS AND CONDITIONS FOR USE, REPRODUCTION, AND DISTRIBUTION
|
| 6 |
+
|
| 7 |
+
1. Definitions.
|
| 8 |
+
|
| 9 |
+
"License" shall mean the terms and conditions for use, reproduction,
|
| 10 |
+
and distribution as defined by Sections 1 through 9 of this document.
|
| 11 |
+
|
| 12 |
+
"Licensor" shall mean the copyright owner or entity authorized by
|
| 13 |
+
the copyright owner that is granting the License.
|
| 14 |
+
|
| 15 |
+
"Legal Entity" shall mean the union of the acting entity and all
|
| 16 |
+
other entities that control, are controlled by, or are under common
|
| 17 |
+
control with that entity. For the purposes of this definition,
|
| 18 |
+
"control" means (i) the power, direct or indirect, to cause the
|
| 19 |
+
direction or management of such entity, whether by contract or
|
| 20 |
+
otherwise, or (ii) ownership of fifty percent (50%) or more of the
|
| 21 |
+
outstanding shares, or (iii) beneficial ownership of such entity.
|
| 22 |
+
|
| 23 |
+
"You" (or "Your") shall mean an individual or Legal Entity
|
| 24 |
+
exercising permissions granted by this License.
|
| 25 |
+
|
| 26 |
+
"Source" form shall mean the preferred form for making modifications,
|
| 27 |
+
including but not limited to software source code, documentation
|
| 28 |
+
source, and configuration files.
|
| 29 |
+
|
| 30 |
+
"Object" form shall mean any form resulting from mechanical
|
| 31 |
+
transformation or translation of a Source form, including but
|
| 32 |
+
not limited to compiled object code, generated documentation,
|
| 33 |
+
and conversions to other media types.
|
| 34 |
+
|
| 35 |
+
"Work" shall mean the work of authorship, whether in Source or
|
| 36 |
+
Object form, made available under the License, as indicated by a
|
| 37 |
+
copyright notice that is included in or attached to the work
|
| 38 |
+
(an example is provided in the Appendix below).
|
| 39 |
+
|
| 40 |
+
"Derivative Works" shall mean any work, whether in Source or Object
|
| 41 |
+
form, that is based on (or derived from) the Work and for which the
|
| 42 |
+
editorial revisions, annotations, elaborations, or other modifications
|
| 43 |
+
represent, as a whole, an original work of authorship. For the purposes
|
| 44 |
+
of this License, Derivative Works shall not include works that remain
|
| 45 |
+
separable from, or merely link (or bind by name) to the interfaces of,
|
| 46 |
+
the Work and Derivative Works thereof.
|
| 47 |
+
|
| 48 |
+
"Contribution" shall mean any work of authorship, including
|
| 49 |
+
the original version of the Work and any modifications or additions
|
| 50 |
+
to that Work or Derivative Works thereof, that is intentionally
|
| 51 |
+
submitted to the Licensor for inclusion in the Work by the copyright owner
|
| 52 |
+
or by an individual or Legal Entity authorized to submit on behalf of
|
| 53 |
+
the copyright owner. For the purposes of this definition, "submitted"
|
| 54 |
+
means any form of electronic, verbal, or written communication sent
|
| 55 |
+
to the Licensor or its representatives, including but not limited to
|
| 56 |
+
communication on electronic mailing lists, source code control systems,
|
| 57 |
+
and issue tracking systems that are managed by, or on behalf of, the
|
| 58 |
+
Licensor for the purpose of discussing and improving the Work, but
|
| 59 |
+
excluding communication that is conspicuously marked or otherwise
|
| 60 |
+
designated in writing by the copyright owner as "Not a Contribution."
|
| 61 |
+
|
| 62 |
+
"Contributor" shall mean Licensor and any individual or Legal Entity
|
| 63 |
+
on behalf of whom a Contribution has been received by the Licensor and
|
| 64 |
+
subsequently incorporated within the Work.
|
| 65 |
+
|
| 66 |
+
2. Grant of Copyright License. Subject to the terms and conditions of
|
| 67 |
+
this License, each Contributor hereby grants to You a perpetual,
|
| 68 |
+
worldwide, non-exclusive, no-charge, royalty-free, irrevocable
|
| 69 |
+
copyright license to reproduce, prepare Derivative Works of,
|
| 70 |
+
publicly display, publicly perform, sublicense, and distribute the
|
| 71 |
+
Work and such Derivative Works in Source or Object form.
|
| 72 |
+
|
| 73 |
+
3. Grant of Patent License. Subject to the terms and conditions of
|
| 74 |
+
this License, each Contributor hereby grants to You a perpetual,
|
| 75 |
+
worldwide, non-exclusive, no-charge, royalty-free, irrevocable
|
| 76 |
+
(except as stated in this section) patent license to make, have made,
|
| 77 |
+
use, offer to sell, sell, import, and otherwise transfer the Work,
|
| 78 |
+
where such license applies only to those patent claims licensable
|
| 79 |
+
by such Contributor that are necessarily infringed by their
|
| 80 |
+
Contribution(s) alone or by combination of their Contribution(s)
|
| 81 |
+
with the Work to which such Contribution(s) was submitted. If You
|
| 82 |
+
institute patent litigation against any entity (including a
|
| 83 |
+
cross-claim or counterclaim in a lawsuit) alleging that the Work
|
| 84 |
+
or a Contribution incorporated within the Work constitutes direct
|
| 85 |
+
or contributory patent infringement, then any patent licenses
|
| 86 |
+
granted to You under this License for that Work shall terminate
|
| 87 |
+
as of the date such litigation is filed.
|
| 88 |
+
|
| 89 |
+
4. Redistribution. You may reproduce and distribute copies of the
|
| 90 |
+
Work or Derivative Works thereof in any medium, with or without
|
| 91 |
+
modifications, and in Source or Object form, provided that You
|
| 92 |
+
meet the following conditions:
|
| 93 |
+
|
| 94 |
+
(a) You must give any other recipients of the Work or
|
| 95 |
+
Derivative Works a copy of this License; and
|
| 96 |
+
|
| 97 |
+
(b) You must cause any modified files to carry prominent notices
|
| 98 |
+
stating that You changed the files; and
|
| 99 |
+
|
| 100 |
+
(c) You must retain, in the Source form of any Derivative Works
|
| 101 |
+
that You distribute, all copyright, patent, trademark, and
|
| 102 |
+
attribution notices from the Source form of the Work,
|
| 103 |
+
excluding those notices that do not pertain to any part of
|
| 104 |
+
the Derivative Works; and
|
| 105 |
+
|
| 106 |
+
(d) If the Work includes a "NOTICE" text file as part of its
|
| 107 |
+
distribution, then any Derivative Works that You distribute must
|
| 108 |
+
include a readable copy of the attribution notices contained
|
| 109 |
+
within such NOTICE file, excluding any notices that do not
|
| 110 |
+
pertain to any part of the Derivative Works, in at least one
|
| 111 |
+
of the following places: within a NOTICE text file distributed
|
| 112 |
+
as part of the Derivative Works; within the Source form or
|
| 113 |
+
documentation, if provided along with the Derivative Works; or,
|
| 114 |
+
within a display generated by the Derivative Works, if and
|
| 115 |
+
wherever such third-party notices normally appear. The contents
|
| 116 |
+
of the NOTICE file are for informational purposes only and
|
| 117 |
+
do not modify the License. You may add Your own attribution
|
| 118 |
+
notices within Derivative Works that You distribute, alongside
|
| 119 |
+
or as an addendum to the NOTICE text from the Work, provided
|
| 120 |
+
that such additional attribution notices cannot be construed
|
| 121 |
+
as modifying the License.
|
| 122 |
+
|
| 123 |
+
You may add Your own copyright statement to Your modifications and
|
| 124 |
+
may provide additional or different license terms and conditions
|
| 125 |
+
for use, reproduction, or distribution of Your modifications, or
|
| 126 |
+
for any such Derivative Works as a whole, provided Your use,
|
| 127 |
+
reproduction, and distribution of the Work otherwise complies with
|
| 128 |
+
the conditions stated in this License.
|
| 129 |
+
|
| 130 |
+
5. Submission of Contributions. Unless You explicitly state otherwise,
|
| 131 |
+
any Contribution intentionally submitted for inclusion in the Work
|
| 132 |
+
by You to the Licensor shall be under the terms and conditions of
|
| 133 |
+
this License, without any additional terms or conditions.
|
| 134 |
+
Notwithstanding the above, nothing herein shall supersede or modify
|
| 135 |
+
the terms of any separate license agreement you may have executed
|
| 136 |
+
with Licensor regarding such Contributions.
|
| 137 |
+
|
| 138 |
+
6. Trademarks. This License does not grant permission to use the trade
|
| 139 |
+
names, trademarks, service marks, or product names of the Licensor,
|
| 140 |
+
except as required for reasonable and customary use in describing the
|
| 141 |
+
origin of the Work and reproducing the content of the NOTICE file.
|
| 142 |
+
|
| 143 |
+
7. Disclaimer of Warranty. Unless required by applicable law or
|
| 144 |
+
agreed to in writing, Licensor provides the Work (and each
|
| 145 |
+
Contributor provides its Contributions) on an "AS IS" BASIS,
|
| 146 |
+
WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or
|
| 147 |
+
implied, including, without limitation, any warranties or conditions
|
| 148 |
+
of TITLE, NON-INFRINGEMENT, MERCHANTABILITY, or FITNESS FOR A
|
| 149 |
+
PARTICULAR PURPOSE. You are solely responsible for determining the
|
| 150 |
+
appropriateness of using or redistributing the Work and assume any
|
| 151 |
+
risks associated with Your exercise of permissions under this License.
|
| 152 |
+
|
| 153 |
+
8. Limitation of Liability. In no event and under no legal theory,
|
| 154 |
+
whether in tort (including negligence), contract, or otherwise,
|
| 155 |
+
unless required by applicable law (such as deliberate and grossly
|
| 156 |
+
negligent acts) or agreed to in writing, shall any Contributor be
|
| 157 |
+
liable to You for damages, including any direct, indirect, special,
|
| 158 |
+
incidental, or consequential damages of any character arising as a
|
| 159 |
+
result of this License or out of the use or inability to use the
|
| 160 |
+
Work (including but not limited to damages for loss of goodwill,
|
| 161 |
+
work stoppage, computer failure or malfunction, or any and all
|
| 162 |
+
other commercial damages or losses), even if such Contributor
|
| 163 |
+
has been advised of the possibility of such damages.
|
| 164 |
+
|
| 165 |
+
9. Accepting Warranty or Additional Liability. While redistributing
|
| 166 |
+
the Work or Derivative Works thereof, You may choose to offer,
|
| 167 |
+
and charge a fee for, acceptance of support, warranty, indemnity,
|
| 168 |
+
or other liability obligations and/or rights consistent with this
|
| 169 |
+
License. However, in accepting such obligations, You may act only
|
| 170 |
+
on Your own behalf and on Your sole responsibility, not on behalf
|
| 171 |
+
of any other Contributor, and only if You agree to indemnify,
|
| 172 |
+
defend, and hold each Contributor harmless for any liability
|
| 173 |
+
incurred by, or claims asserted against, such Contributor by reason
|
| 174 |
+
of your accepting any such warranty or additional liability.
|
| 175 |
+
|
| 176 |
+
END OF TERMS AND CONDITIONS
|
| 177 |
+
|
| 178 |
+
Copyright 2025 NVIDIA Corporation
|
| 179 |
+
|
| 180 |
+
Licensed under the Apache License, Version 2.0 (the "License");
|
| 181 |
+
you may not use this file except in compliance with the License.
|
| 182 |
+
You may obtain a copy of the License at
|
| 183 |
+
|
| 184 |
+
http://www.apache.org/licenses/LICENSE-2.0
|
| 185 |
+
|
| 186 |
+
Unless required by applicable law or agreed to in writing, software
|
| 187 |
+
distributed under the License is distributed on an "AS IS" BASIS,
|
| 188 |
+
WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
| 189 |
+
See the License for the specific language governing permissions and
|
| 190 |
+
limitations under the License.
|
README.md
ADDED
|
@@ -0,0 +1,274 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# NVIDIA DreamZero: World Action Models Are Zero-Shot Policies
|
| 2 |
+
A research project from [NVIDIA GEAR Lab](https://research.nvidia.com/labs/gear/).
|
| 3 |
+
|
| 4 |
+
[](https://www.nvidia.com) [](LICENSE) [](https://arxiv.org/abs/2602.15922)
|
| 5 |
+
|
| 6 |
+
[[Project Page](https://dreamzero0.github.io/)] [[Paper](https://arxiv.org/abs/2602.15922)]
|
| 7 |
+
|
| 8 |
+
DreamZero is a World Action Model that jointly predicts actions and videos, achieving strong zero-shot performance on unseen tasks. This release package contains everything needed to load a pretrained DreamZero model and run distributed inference via a WebSocket server.
|
| 9 |
+
|
| 10 |
+
## News
|
| 11 |
+
|
| 12 |
+
- **02/27:** DreamZero is **#1 on both [MolmoSpaces]([https://huggingface.co/spaces/ai2-adapt/MolmoSpaces](https://molmospaces.allen.ai/leaderboard)) and [RoboArena]([https://robo-arena.github.io/](https://robo-arena.github.io/leaderboard))**! DreamZero-DROID is trained *from scratch* using only the DROID dataset — no pretraining on large-scale robot data, unlike competing VLAs. This demonstrates the strength of video-model backbones for generalist robot policies (VAMs/WAMs).
|
| 13 |
+
- **02/27:** Released **DreamZero-AgiBot checkpoint** and **post-training code** for efficient few-shot adaptation. Post-train on just ~30 minutes of play data for your specific robot, and see the robot do basic language following and pick-and-place (see YAM experiments in our paper for more detail).
|
| 14 |
+
- **02/20:** Released the **full training codebase, preprocessed dataset, and guide for new embodiments** to replicate the DreamZero-DROID checkpoint and train on your own robot. See [Adding a New Embodiment to DreamZero](docs/DATASET_TO_GEAR_AND_TRAIN.md) for a step-by-step walkthrough.
|
| 15 |
+
|
| 16 |
+
## Features
|
| 17 |
+
|
| 18 |
+
**Available Now**
|
| 19 |
+
- Pretrained DreamZero-DROID model checkpoint [[Huggingface](https://huggingface.co/GEAR-Dreams/DreamZero-DROID)]
|
| 20 |
+
- Pretrained DreamZero-AgiBot checkpoint (for post-training on new embodiments) [[Huggingface](https://huggingface.co/GEAR-Dreams/DreamZero-AgiBot)]
|
| 21 |
+
- Distributed WebSocket inference server (GB200, H100)
|
| 22 |
+
- DiT caching for optimized inference (~0.6s on GB200, ~3s on H100)
|
| 23 |
+
- DROID simulation evaluation support
|
| 24 |
+
- [RoboArena](https://robo-arena.github.io/) integration (DROID real)
|
| 25 |
+
- Video generation and saving (MP4)
|
| 26 |
+
- LoRA and full fine-tuning training scripts
|
| 27 |
+
- Training on new embodiments (AgiBot, YAM) — see [guide](docs/DATASET_TO_GEAR_AND_TRAIN.md)
|
| 28 |
+
|
| 29 |
+
**Coming Soon**
|
| 30 |
+
- [PolaRiS](https://polaris-evals.github.io/) simulation environment support
|
| 31 |
+
- [Genie 3.0](https://arxiv.org/abs/2601.02078) sim environment support for DreamZero-AgiBot
|
| 32 |
+
|
| 33 |
+
## Testing Out DreamZero in Simulation with API
|
| 34 |
+
We provide an inference script that directly evaluates a hosted DreamZero-DROID policy on [`sim_evals`](https://github.com/arhanjain/sim-evals). To test out the policy, first request access to the API via this form [link](https://forms.gle/zCj5zjDvHsoeuMXU7). Then, follow these instructions to install [`sim_evals`](https://github.com/arhanjain/sim-evals) and launch evaluation.
|
| 35 |
+
|
| 36 |
+
```bash
|
| 37 |
+
# Clone repository
|
| 38 |
+
git clone --recurse-submodules https://github.com/arhanjain/sim-evals.git
|
| 39 |
+
cd sim-evals
|
| 40 |
+
|
| 41 |
+
# Install uv
|
| 42 |
+
curl -LsSf https://astral.sh/uv/install.sh | sh
|
| 43 |
+
|
| 44 |
+
# Activate uv environment
|
| 45 |
+
uv sync
|
| 46 |
+
source .venv/bin/activate
|
| 47 |
+
|
| 48 |
+
# [Optional] update pytorch versions
|
| 49 |
+
pip install torch==2.9.1 torchvision==0.24.1 torchaudio==2.9.1 --index-url https://download.pytorch.org/whl/cu129
|
| 50 |
+
|
| 51 |
+
# Download assets (may need to export HF_TOKEN=<YOUR_HUGGINGFACE_TOKEN> first)
|
| 52 |
+
uvx hf download owhan/DROID-sim-environments --repo-type dataset --local-dir assets
|
| 53 |
+
|
| 54 |
+
# Run eval script
|
| 55 |
+
cd ..
|
| 56 |
+
python eval_utils/run_sim_eval.py --host <API_HOST> --port <API_PORT>
|
| 57 |
+
```
|
| 58 |
+
|
| 59 |
+
The outputs are saved in `runs` directory.
|
| 60 |
+
|
| 61 |
+
|
| 62 |
+
## Quick Start
|
| 63 |
+
|
| 64 |
+
### Prerequisites
|
| 65 |
+
|
| 66 |
+
- **Python**: 3.11
|
| 67 |
+
- **Hardware**: Multi-GPU setup (tested on GB200, H100)
|
| 68 |
+
- Minimum: 2 GPUs for distributed inference
|
| 69 |
+
- **CUDA**: Compatible GPU with CUDA 12.9+
|
| 70 |
+
|
| 71 |
+
### Installation
|
| 72 |
+
|
| 73 |
+
1. **Create conda environment:**
|
| 74 |
+
```bash
|
| 75 |
+
conda create -n dreamzero python=3.11
|
| 76 |
+
conda activate dreamzero
|
| 77 |
+
```
|
| 78 |
+
|
| 79 |
+
2. **Install dependencies (PyTorch 2.8+ with CUDA 12.9+):**
|
| 80 |
+
```bash
|
| 81 |
+
pip install -e . --extra-index-url https://download.pytorch.org/whl/cu129
|
| 82 |
+
```
|
| 83 |
+
|
| 84 |
+
3. **Install flash attention:**
|
| 85 |
+
```bash
|
| 86 |
+
MAX_JOBS=8 pip install --no-build-isolation flash-attn
|
| 87 |
+
```
|
| 88 |
+
|
| 89 |
+
4. **[GB200 ONLY, SKIP FOR H100] Install Transformer Engine:**
|
| 90 |
+
```bash
|
| 91 |
+
pip install --no-build-isolation transformer_engine[pytorch]
|
| 92 |
+
```
|
| 93 |
+
|
| 94 |
+
5. **[GB200 ONLY FOR TENSORRT, SKIP FOR H100] Install Tensorrt:**
|
| 95 |
+
```bash
|
| 96 |
+
pip install tensorrt==10.13.2.6 tensorrt_cu13==10.13.2.6 tensorrt_cu13_libs==10.13.2.6 tensorrt_cu13_bindings==10.13.2.6 --no-deps
|
| 97 |
+
pip install transformer_engine==2.10.0 transformer_engine_cu12==2.10.0 transformer_engine_torch==2.10.0
|
| 98 |
+
```
|
| 99 |
+
|
| 100 |
+
## Downloading Pretrained Checkpoints
|
| 101 |
+
|
| 102 |
+
### DreamZero-DROID (for inference)
|
| 103 |
+
|
| 104 |
+
We release a 14B pretrained DROID checkpoint on [Huggingface](https://huggingface.co/GEAR-Dreams/DreamZero-DROID). To download the checkpoint, run
|
| 105 |
+
|
| 106 |
+
```bash
|
| 107 |
+
hf download GEAR-Dreams/DreamZero-DROID --repo-type model --local-dir <path/to/checkpoint>
|
| 108 |
+
```
|
| 109 |
+
|
| 110 |
+
### DreamZero-AgiBot (for fine-tuning on new embodiments)
|
| 111 |
+
|
| 112 |
+
To fine-tune DreamZero on a new embodiment (e.g. YAM, AgiBot), download the pretrained [DreamZero-AgiBot](https://huggingface.co/GEAR-Dreams/DreamZero-AgiBot) checkpoint (~45GB) to `./checkpoints/DreamZero-AgiBot`:
|
| 113 |
+
|
| 114 |
+
```bash
|
| 115 |
+
git clone https://huggingface.co/GEAR-Dreams/DreamZero-AgiBot ./checkpoints/DreamZero-AgiBot
|
| 116 |
+
```
|
| 117 |
+
|
| 118 |
+
Or with the Hugging Face CLI:
|
| 119 |
+
|
| 120 |
+
```bash
|
| 121 |
+
hf download GEAR-Dreams/DreamZero-AgiBot --repo-type model --local-dir ./checkpoints/DreamZero-AgiBot
|
| 122 |
+
```
|
| 123 |
+
|
| 124 |
+
The YAM and AgiBot training scripts use `pretrained_model_path=./checkpoints/DreamZero-AgiBot` by default. See the [new embodiment guide](docs/DATASET_TO_GEAR_AND_TRAIN.md) for usage.
|
| 125 |
+
|
| 126 |
+
## Running the Inference Server
|
| 127 |
+
|
| 128 |
+
### Command Overview
|
| 129 |
+
|
| 130 |
+
The inference server uses PyTorch distributed training utilities to parallelize the model across multiple GPUs:
|
| 131 |
+
|
| 132 |
+
```bash
|
| 133 |
+
CUDA_VISIBLE_DEVICES=0,1 python -m torch.distributed.run --standalone --nproc_per_node=2 socket_test_optimized_AR.py --port 5000 --enable-dit-cache --model-path <path/to/checkpoint>
|
| 134 |
+
```
|
| 135 |
+
|
| 136 |
+
(Optional only for GB200) Tensorrt enables faster generation
|
| 137 |
+
```bash
|
| 138 |
+
export LOAD_TRT_ENGINE=<path/to/checkpoint>/tensorrt/wan/WanModel_nvfp4.trt
|
| 139 |
+
export DYNAMIC_CACHE_SCHEDULE=true
|
| 140 |
+
CUDA_VISIBLE_DEVICES=0,1 python -m torch.distributed.run --standalone --nproc_per_node=2 /mnt/aws-lfs-02/shared/seonghyeony/dreamzero/socket_test_optimized_AR.py --port 8000 --enable-dit-cache --model-path <path/to/checkpoint>
|
| 141 |
+
```
|
| 142 |
+
To verify the server is working, run a test client. The first few inferences will take a few minutes to warm up. After warming up, inference takes ~0.6s on GB200 and ~3s on H100.
|
| 143 |
+
|
| 144 |
+
```
|
| 145 |
+
python test_client_AR.py --port 5000
|
| 146 |
+
```
|
| 147 |
+
|
| 148 |
+
### Command-line Arguments
|
| 149 |
+
|
| 150 |
+
- `--port`: Port number for the WebSocket server (default: 8000)
|
| 151 |
+
- `--model-path`: Path to the pretrained model checkpoint directory
|
| 152 |
+
- `--enable-dit-cache`: Enable caching in DiT layers for faster inference (recommended)
|
| 153 |
+
- `--max-chunk-size`: Override max_chunk_size for inference (optional)
|
| 154 |
+
- `--timeout-seconds`: Server timeout in seconds (default: 50000)
|
| 155 |
+
- `--index`: Index for output directory naming (default: 0)
|
| 156 |
+
|
| 157 |
+
|
| 158 |
+
### Output
|
| 159 |
+
|
| 160 |
+
The server saves:
|
| 161 |
+
- **Videos**: Generated video predictions as MP4 files in `{model_path}/real_world_eval_gen_{date}_{index}/{checkpoint_name}/`
|
| 162 |
+
- **Input observations**: Saved per message in `{output_dir}/inputs/{msg_index}_{timestamp}/`
|
| 163 |
+
|
| 164 |
+
|
| 165 |
+
## Training
|
| 166 |
+
|
| 167 |
+
> **Training on a new embodiment?** See [Adding a New Embodiment to DreamZero](docs/DATASET_TO_GEAR_AND_TRAIN.md) for a complete guide on converting your dataset, configuring modalities, and launching training. <em>Make sure to align the 3 camera view order to ensure positive transfer.</em>
|
| 168 |
+
|
| 169 |
+
### Downloading Pretrained Base Model Weights
|
| 170 |
+
|
| 171 |
+
DreamZero is built on top of [Wan2.1-I2V-14B-480P](https://huggingface.co/Wan-AI/Wan2.1-I2V-14B-480P) and uses the [umt5-xxl](https://huggingface.co/google/umt5-xxl) tokenizer. Download both before training:
|
| 172 |
+
|
| 173 |
+
```bash
|
| 174 |
+
pip install "huggingface_hub[cli]"
|
| 175 |
+
|
| 176 |
+
# You may need to set your HuggingFace token:
|
| 177 |
+
# export HF_TOKEN=<YOUR_HUGGINGFACE_TOKEN>
|
| 178 |
+
|
| 179 |
+
# Download Wan2.1 model weights (~28GB)
|
| 180 |
+
hf download Wan-AI/Wan2.1-I2V-14B-480P --local-dir ./checkpoints/Wan2.1-I2V-14B-480P
|
| 181 |
+
|
| 182 |
+
# Download umt5-xxl tokenizer
|
| 183 |
+
hf download google/umt5-xxl --local-dir ./checkpoints/umt5-xxl
|
| 184 |
+
```
|
| 185 |
+
|
| 186 |
+
> **Note:** The training script will auto-download these if they are not found at the configured paths, but pre-downloading is recommended to avoid delays at launch.
|
| 187 |
+
|
| 188 |
+
### DROID Dataset
|
| 189 |
+
|
| 190 |
+
We release the preprocessed DROID dataset used to train DreamZero on HuggingFace: [GEAR-Dreams/DreamZero-DROID-Data](https://huggingface.co/datasets/GEAR-Dreams/DreamZero-DROID-Data).
|
| 191 |
+
|
| 192 |
+
This dataset is derived from the [DROID 1.0.1](https://droid-dataset.github.io/) dataset with the following modifications:
|
| 193 |
+
- Converted from RLDS/TFDS format to [LeRobot](https://github.com/huggingface/lerobot) v2.0 format
|
| 194 |
+
- Idle frames removed using [Physical Intelligence's idle frame detector](https://github.com/Physical-Intelligence/openpi/blob/main/examples/droid/README_train.md#data-filtering) (`droid_sample_ranges_v1_0_1.json`)
|
| 195 |
+
- Episodes without language annotations are filtered out
|
| 196 |
+
- Successful episodes only (episodes with non-zero reward)
|
| 197 |
+
- 3 camera views: `exterior_image_1_left`, `exterior_image_2_left`, `wrist_image_left`
|
| 198 |
+
|
| 199 |
+
**To download the preprocessed dataset (~131GB):**
|
| 200 |
+
|
| 201 |
+
```bash
|
| 202 |
+
huggingface-cli download GEAR-Dreams/DreamZero-DROID-Data --repo-type dataset --local-dir ./data/droid_lerobot
|
| 203 |
+
```
|
| 204 |
+
|
| 205 |
+
If you want to reproduce the dataset conversion from raw DROID 1.0.1 yourself (or modify the filtering), see [docs/DROID_CONVERSION.md](docs/DROID_CONVERSION.md).
|
| 206 |
+
|
| 207 |
+
### Running Training
|
| 208 |
+
|
| 209 |
+
```bash
|
| 210 |
+
# Configure paths (override defaults as needed)
|
| 211 |
+
export DROID_DATA_ROOT="./data/droid_lerobot"
|
| 212 |
+
export OUTPUT_DIR="./checkpoints/dreamzero_droid"
|
| 213 |
+
export NUM_GPUS=4
|
| 214 |
+
|
| 215 |
+
# Point to your downloaded model weights (if not using default paths)
|
| 216 |
+
export WAN_CKPT_DIR="./checkpoints/Wan2.1-I2V-14B-480P"
|
| 217 |
+
export TOKENIZER_DIR="./checkpoints/umt5-xxl"
|
| 218 |
+
|
| 219 |
+
# Launch training
|
| 220 |
+
bash scripts/train/droid_training.sh
|
| 221 |
+
```
|
| 222 |
+
|
| 223 |
+
**Using Wan2.2-TI2V-5B backbone (5B params, lower VRAM):** To train with the smaller Wan2.2-TI2V-5B model instead of Wan2.1-I2V-14B, see [docs/WAN22_BACKBONE.md](docs/WAN22_BACKBONE.md) and run `bash scripts/train/droid_training_wan22.sh`.
|
| 224 |
+
|
| 225 |
+
### Training Configuration
|
| 226 |
+
|
| 227 |
+
The training script uses Hydra for configuration and DeepSpeed ZeRO Stage 2 for distributed training. Key defaults:
|
| 228 |
+
|
| 229 |
+
| Parameter | Default | Description |
|
| 230 |
+
|---|---|---|
|
| 231 |
+
| `NUM_GPUS` | 4 | Number of GPUs |
|
| 232 |
+
| `per_device_train_batch_size` | 1 | Batch size per GPU |
|
| 233 |
+
| `learning_rate` | 1e-5 | Learning rate |
|
| 234 |
+
| `max_steps` | 10 | Max training steps (increase for full training) |
|
| 235 |
+
| `warmup_ratio` | 0.05 | Warmup ratio |
|
| 236 |
+
| `weight_decay` | 1e-5 | Weight decay |
|
| 237 |
+
| `image_resolution_width` | 320 | Image width |
|
| 238 |
+
| `image_resolution_height` | 176 | Image height |
|
| 239 |
+
| `num_frames` | 33 | Number of video frames |
|
| 240 |
+
| `action_horizon` | 24 | Action prediction horizon |
|
| 241 |
+
| `save_lora_only` | true | Only save LoRA weights |
|
| 242 |
+
| `bf16` | true | Use bfloat16 precision |
|
| 243 |
+
|
| 244 |
+
> **Note:** `max_steps=10` is set for a quick sanity check. For full training, increase this to your desired number of steps and configure `save_steps` / `save_strategy` accordingly.
|
| 245 |
+
|
| 246 |
+
|
| 247 |
+
## Citation
|
| 248 |
+
|
| 249 |
+
If you use DreamZero in your research, please cite:
|
| 250 |
+
|
| 251 |
+
```bibtex
|
| 252 |
+
@misc{ye2026worldactionmodelszeroshot,
|
| 253 |
+
title={World Action Models are Zero-shot Policies},
|
| 254 |
+
author={Seonghyeon Ye and Yunhao Ge and Kaiyuan Zheng and Shenyuan Gao and Sihyun Yu and George Kurian and Suneel Indupuru and You Liang Tan and Chuning Zhu and Jiannan Xiang and Ayaan Malik and Kyungmin Lee and William Liang and Nadun Ranawaka and Jiasheng Gu and Yinzhen Xu and Guanzhi Wang and Fengyuan Hu and Avnish Narayan and Johan Bjorck and Jing Wang and Gwanghyun Kim and Dantong Niu and Ruijie Zheng and Yuqi Xie and Jimmy Wu and Qi Wang and Ryan Julian and Danfei Xu and Yilun Du and Yevgen Chebotar and Scott Reed and Jan Kautz and Yuke Zhu and Linxi "Jim" Fan and Joel Jang},
|
| 255 |
+
year={2026},
|
| 256 |
+
eprint={2602.15922},
|
| 257 |
+
archivePrefix={arXiv},
|
| 258 |
+
primaryClass={cs.RO},
|
| 259 |
+
url={https://arxiv.org/abs/2602.15922},
|
| 260 |
+
}
|
| 261 |
+
```
|
| 262 |
+
|
| 263 |
+
## License
|
| 264 |
+
|
| 265 |
+
This project is licensed under the [Apache License 2.0](LICENSE).
|
| 266 |
+
|
| 267 |
+
## Support
|
| 268 |
+
|
| 269 |
+
For issues and questions:
|
| 270 |
+
- Check the troubleshooting section above
|
| 271 |
+
- Review server logs for detailed error messages
|
| 272 |
+
- Verify your checkpoint is compatible with this release
|
| 273 |
+
|
| 274 |
+
[](https://star-history.com/#dreamzero0/dreamzero&Date)
|
docker/Dockerfile
ADDED
|
@@ -0,0 +1,71 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# DreamZero 训练 Docker 镜像
|
| 2 |
+
#
|
| 3 |
+
# 构建:
|
| 4 |
+
# docker build -t dreamzero:latest -f docker/Dockerfile .
|
| 5 |
+
#
|
| 6 |
+
# 运行 (单节点 8 GPU):
|
| 7 |
+
# docker run --gpus all --shm-size=64g \
|
| 8 |
+
# -v /path/to/checkpoints:/workspace/dreamzero/checkpoints \
|
| 9 |
+
# -v /path/to/data:/workspace/dreamzero/data \
|
| 10 |
+
# -v /path/to/output:/workspace/dreamzero/output \
|
| 11 |
+
# dreamzero:latest \
|
| 12 |
+
# data=dreamzero/libero \
|
| 13 |
+
# train_architecture=full \
|
| 14 |
+
# max_steps=200000 \
|
| 15 |
+
# output_dir=./output/libero_full
|
| 16 |
+
|
| 17 |
+
FROM nvidia/cuda:12.4.1-devel-ubuntu22.04
|
| 18 |
+
|
| 19 |
+
LABEL description="DreamZero: Video Diffusion Model for Scalable Robot Learning"
|
| 20 |
+
LABEL maintainer="DreamZero Team"
|
| 21 |
+
|
| 22 |
+
# ============ 系统依赖 ============
|
| 23 |
+
RUN apt-get update && apt-get install -y --no-install-recommends \
|
| 24 |
+
python3.11 \
|
| 25 |
+
python3.11-dev \
|
| 26 |
+
python3.11-distutils \
|
| 27 |
+
python3-pip \
|
| 28 |
+
git \
|
| 29 |
+
ffmpeg \
|
| 30 |
+
libsm6 \
|
| 31 |
+
libxext6 \
|
| 32 |
+
libglib2.0-0 \
|
| 33 |
+
libxrender-dev \
|
| 34 |
+
libgl1-mesa-glx \
|
| 35 |
+
wget \
|
| 36 |
+
curl \
|
| 37 |
+
&& rm -rf /var/lib/apt/lists/*
|
| 38 |
+
|
| 39 |
+
# 设置 python3.11 为默认
|
| 40 |
+
RUN update-alternatives --install /usr/bin/python python /usr/bin/python3.11 1
|
| 41 |
+
RUN update-alternatives --install /usr/bin/python3 python3 /usr/bin/python3.11 1
|
| 42 |
+
|
| 43 |
+
# ============ 安装 PyTorch ============
|
| 44 |
+
RUN pip install --no-cache-dir \
|
| 45 |
+
torch==2.4.0 \
|
| 46 |
+
torchvision==0.19.0 \
|
| 47 |
+
--index-url https://download.pytorch.org/whl/cu124
|
| 48 |
+
|
| 49 |
+
# ============ 复制代码 ============
|
| 50 |
+
WORKDIR /workspace/dreamzero
|
| 51 |
+
COPY groot/ ./groot/
|
| 52 |
+
COPY scripts/ ./scripts/
|
| 53 |
+
COPY pyproject.toml requirements.txt ./
|
| 54 |
+
|
| 55 |
+
# ============ 安装依赖 ============
|
| 56 |
+
RUN pip install --no-cache-dir -e ".[train]"
|
| 57 |
+
|
| 58 |
+
# ============ NCCL 配置(多节点优化) ============
|
| 59 |
+
ENV NCCL_DEBUG=WARN
|
| 60 |
+
ENV NCCL_IB_DISABLE=0
|
| 61 |
+
ENV NCCL_SOCKET_IFNAME=^docker0,lo
|
| 62 |
+
ENV NCCL_IB_TIMEOUT=22
|
| 63 |
+
ENV NCCL_IB_RETRY_CNT=4
|
| 64 |
+
|
| 65 |
+
# ============ 中国用户镜像 ============
|
| 66 |
+
ENV HF_ENDPOINT=https://hf-mirror.com
|
| 67 |
+
|
| 68 |
+
# ============ 默认入口 ============
|
| 69 |
+
ENTRYPOINT ["torchrun", "--standalone", "--nproc_per_node", "8", \
|
| 70 |
+
"/workspace/dreamzero/groot/vla/experiment/experiment.py"]
|
| 71 |
+
CMD ["--help"]
|
docs/DATASET_TO_GEAR_AND_TRAIN.md
ADDED
|
@@ -0,0 +1,471 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# Adding a New Embodiment to DreamZero
|
| 2 |
+
|
| 3 |
+
How to take a LeRobot v2 dataset for a new robot, convert it to GEAR format, define its modality config, and train a DreamZero policy.
|
| 4 |
+
|
| 5 |
+
Throughout this guide, replace `<EMBODIMENT>` with your robot's name (e.g. `myrobot`, `franka`, `aloha`).
|
| 6 |
+
|
| 7 |
+
---
|
| 8 |
+
|
| 9 |
+
## Overview
|
| 10 |
+
|
| 11 |
+
```
|
| 12 |
+
Step 1 Convert LeRobot v2 dataset → GEAR metadata
|
| 13 |
+
Step 2 Register the embodiment tag
|
| 14 |
+
Step 3 Add modality config + transforms to base YAML
|
| 15 |
+
Step 4 Create a dataset YAML
|
| 16 |
+
Step 5 Create a training script
|
| 17 |
+
Step 6 Train
|
| 18 |
+
```
|
| 19 |
+
|
| 20 |
+
---
|
| 21 |
+
|
| 22 |
+
## Step 1: Convert Dataset to GEAR Format
|
| 23 |
+
|
| 24 |
+
The converter reads a LeRobot v2 dataset and generates the metadata files DreamZero needs. It does **not** modify your parquet files or videos — it only writes to `meta/`.
|
| 25 |
+
|
| 26 |
+
### Expected input structure
|
| 27 |
+
|
| 28 |
+
```
|
| 29 |
+
your_dataset/
|
| 30 |
+
├── data/
|
| 31 |
+
│ └── chunk-000/
|
| 32 |
+
│ ├── episode_000000.parquet
|
| 33 |
+
│ └── ...
|
| 34 |
+
├── videos/
|
| 35 |
+
│ └── chunk-000/
|
| 36 |
+
│ ├── observation.images.cam0/
|
| 37 |
+
│ │ ├── episode_000000.mp4
|
| 38 |
+
│ │ └── ...
|
| 39 |
+
│ └── observation.images.cam1/
|
| 40 |
+
│ └── ...
|
| 41 |
+
└── meta/
|
| 42 |
+
└── info.json # must contain: features, total_episodes, fps
|
| 43 |
+
```
|
| 44 |
+
|
| 45 |
+
### Run the converter
|
| 46 |
+
|
| 47 |
+
```bash
|
| 48 |
+
python scripts/data/convert_lerobot_to_gear.py \
|
| 49 |
+
--dataset-path /path/to/your_dataset \
|
| 50 |
+
--embodiment-tag <EMBODIMENT> \
|
| 51 |
+
--state-keys '{"joint_pos": [0, 6], "gripper_pos": [6, 7]}' \
|
| 52 |
+
--action-keys '{"joint_pos": [0, 6], "gripper_pos": [6, 7]}' \
|
| 53 |
+
--relative-action-keys joint_pos gripper_pos \
|
| 54 |
+
--task-key annotation.task
|
| 55 |
+
```
|
| 56 |
+
|
| 57 |
+
`--state-keys` and `--action-keys` tell the converter how to split a packed vector column into named sub-keys. The JSON maps sub-key name → `[start_index, end_index]`. Omit these flags to let the converter auto-detect.
|
| 58 |
+
|
| 59 |
+
### Arguments
|
| 60 |
+
|
| 61 |
+
| Argument | Default | Description |
|
| 62 |
+
|---|---|---|
|
| 63 |
+
| `--dataset-path` | *(required)* | Path to the LeRobot v2 dataset |
|
| 64 |
+
| `--output-path` | *(in-place)* | Write to a different directory instead of in-place |
|
| 65 |
+
| `--embodiment-tag` | `xdof` | Tag for `meta/embodiment.json`; must match the key you use in Step 3 |
|
| 66 |
+
| `--state-keys` | *(auto)* | JSON: sub-key name → `[start, end]` index range |
|
| 67 |
+
| `--action-keys` | *(auto)* | JSON: sub-key name → `[start, end]` index range |
|
| 68 |
+
| `--relative-action-keys` | *(none)* | Sub-key names to compute relative action stats for |
|
| 69 |
+
| `--task-key` | *(auto)* | Column name for language/task annotations |
|
| 70 |
+
| `--fps` | *(from info.json)* | Override dataset FPS |
|
| 71 |
+
| `--action-horizon` | `24` | Horizon for relative stats computation |
|
| 72 |
+
| `--force` | `false` | Overwrite existing metadata files |
|
| 73 |
+
|
| 74 |
+
### Generated files
|
| 75 |
+
|
| 76 |
+
The converter creates these under `meta/`:
|
| 77 |
+
|
| 78 |
+
| File | Contents |
|
| 79 |
+
|---|---|
|
| 80 |
+
| `modality.json` | Maps state, action, video, and annotation keys with index ranges and dtypes |
|
| 81 |
+
| `embodiment.json` | `{"embodiment_tag": "<EMBODIMENT>"}` |
|
| 82 |
+
| `stats.json` | Per-feature statistics (mean, std, min, max, q01, q99) |
|
| 83 |
+
| `relative_stats_dreamzero.json` | Relative action statistics (action − reference state) |
|
| 84 |
+
| `tasks.jsonl` | Unique task descriptions |
|
| 85 |
+
| `episodes.jsonl` | Per-episode metadata (index, tasks, length) |
|
| 86 |
+
|
| 87 |
+
---
|
| 88 |
+
|
| 89 |
+
## Step 2: Register the Embodiment Tag
|
| 90 |
+
|
| 91 |
+
1. Add to the enum in `groot/vla/data/schema/embodiment_tags.py`:
|
| 92 |
+
|
| 93 |
+
```python
|
| 94 |
+
class EmbodimentTag(str, Enum):
|
| 95 |
+
...
|
| 96 |
+
MY_ROBOT = "<EMBODIMENT>"
|
| 97 |
+
```
|
| 98 |
+
|
| 99 |
+
2. Add to `VALID_EMBODIMENT_TAGS` in `scripts/data/convert_lerobot_to_gear.py` (if you want the converter to accept the tag without `--force`):
|
| 100 |
+
|
| 101 |
+
```python
|
| 102 |
+
VALID_EMBODIMENT_TAGS = [
|
| 103 |
+
...,
|
| 104 |
+
"<EMBODIMENT>",
|
| 105 |
+
]
|
| 106 |
+
```
|
| 107 |
+
|
| 108 |
+
---
|
| 109 |
+
|
| 110 |
+
## Step 3: Add Modality Config and Transforms
|
| 111 |
+
|
| 112 |
+
Edit `groot/vla/configs/data/dreamzero/base_48_wan_fine_aug_relative.yaml`.
|
| 113 |
+
|
| 114 |
+
### 3a. Understanding modality
|
| 115 |
+
|
| 116 |
+
**Modality** connects your dataset columns to the training pipeline. The `modality.json` from Step 1 contains entries like:
|
| 117 |
+
|
| 118 |
+
```json
|
| 119 |
+
{
|
| 120 |
+
"state": {
|
| 121 |
+
"joint_pos": {"original_key": "observation.state", "start": 0, "end": 6},
|
| 122 |
+
"gripper_pos": {"original_key": "observation.state", "start": 6, "end": 7}
|
| 123 |
+
},
|
| 124 |
+
"action": {
|
| 125 |
+
"joint_pos": {"original_key": "action", "start": 0, "end": 6},
|
| 126 |
+
"gripper_pos": {"original_key": "action", "start": 6, "end": 7}
|
| 127 |
+
},
|
| 128 |
+
"video": {
|
| 129 |
+
"cam0": {"original_key": "observation.images.cam0"}
|
| 130 |
+
},
|
| 131 |
+
"annotation": {
|
| 132 |
+
"task": {"original_key": "annotation.task"}
|
| 133 |
+
}
|
| 134 |
+
}
|
| 135 |
+
```
|
| 136 |
+
|
| 137 |
+
The YAML config must reference **exactly these key names** with type prefixes:
|
| 138 |
+
|
| 139 |
+
| Modality | YAML key format | Example |
|
| 140 |
+
|---|---|---|
|
| 141 |
+
| State | `state.<name>` | `state.joint_pos` |
|
| 142 |
+
| Action | `action.<name>` | `action.joint_pos` |
|
| 143 |
+
| Video | `video.<name>` | `video.cam0` |
|
| 144 |
+
| Language | `annotation.<name>` | `annotation.task` |
|
| 145 |
+
|
| 146 |
+
If the YAML keys don't match `modality.json`, training will fail with missing-key errors.
|
| 147 |
+
|
| 148 |
+
### 3b. Add `modality_config_<EMBODIMENT>`
|
| 149 |
+
|
| 150 |
+
```yaml
|
| 151 |
+
modality_config_<EMBODIMENT>:
|
| 152 |
+
video:
|
| 153 |
+
_target_: groot.vla.data.dataset.ModalityConfig
|
| 154 |
+
delta_indices: [0, 1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 12, 13, 14, 15, 16, 17, 18, 19, 20, 21, 22, 23, 24]
|
| 155 |
+
eval_delta_indices: [0]
|
| 156 |
+
modality_keys: # one entry per camera, matching modality.json
|
| 157 |
+
- video.cam0
|
| 158 |
+
- video.cam1
|
| 159 |
+
- video.cam2
|
| 160 |
+
state:
|
| 161 |
+
_target_: groot.vla.data.dataset.ModalityConfig
|
| 162 |
+
delta_indices: [0]
|
| 163 |
+
modality_keys: # matching modality.json state keys
|
| 164 |
+
- state.joint_pos
|
| 165 |
+
- state.gripper_pos
|
| 166 |
+
action:
|
| 167 |
+
_target_: groot.vla.data.dataset.ModalityConfig
|
| 168 |
+
delta_indices: [0, 1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 12, 13, 14, 15, 16, 17, 18, 19, 20, 21, 22, 23]
|
| 169 |
+
modality_keys: # matching modality.json action keys
|
| 170 |
+
- action.joint_pos
|
| 171 |
+
- action.gripper_pos
|
| 172 |
+
language:
|
| 173 |
+
_target_: groot.vla.data.dataset.ModalityConfig
|
| 174 |
+
delta_indices: [0]
|
| 175 |
+
modality_keys:
|
| 176 |
+
- annotation.task
|
| 177 |
+
```
|
| 178 |
+
|
| 179 |
+
**`delta_indices` explained:**
|
| 180 |
+
|
| 181 |
+
- **Video** — frame offsets to sample (25 entries = 25 frames from the trajectory).
|
| 182 |
+
- **State / Language** — `[0]` = current timestep only.
|
| 183 |
+
- **Action** — future offsets (24 entries = 24-step action chunk).
|
| 184 |
+
|
| 185 |
+
Adjust these to match your `num_frames` and `action_horizon` training settings.
|
| 186 |
+
|
| 187 |
+
### 3c. Add `transform_<EMBODIMENT>`
|
| 188 |
+
|
| 189 |
+
```yaml
|
| 190 |
+
transform_<EMBODIMENT>:
|
| 191 |
+
_target_: groot.vla.data.transform.ComposedModalityTransform
|
| 192 |
+
transforms:
|
| 193 |
+
# Video
|
| 194 |
+
- <<: *totensor_cfg
|
| 195 |
+
apply_to: ${modality_config_<EMBODIMENT>.video.modality_keys}
|
| 196 |
+
- <<: *crop_cfg
|
| 197 |
+
apply_to: ${modality_config_<EMBODIMENT>.video.modality_keys}
|
| 198 |
+
- <<: *resize_cfg
|
| 199 |
+
apply_to: ${modality_config_<EMBODIMENT>.video.modality_keys}
|
| 200 |
+
- <<: *color_jitter_cfg
|
| 201 |
+
apply_to: ${modality_config_<EMBODIMENT>.video.modality_keys}
|
| 202 |
+
- <<: *to_numpy_cfg
|
| 203 |
+
apply_to: ${modality_config_<EMBODIMENT>.video.modality_keys}
|
| 204 |
+
|
| 205 |
+
# State
|
| 206 |
+
- _target_: groot.vla.data.transform.StateActionToTensor
|
| 207 |
+
apply_to: ${modality_config_<EMBODIMENT>.state.modality_keys}
|
| 208 |
+
- _target_: groot.vla.data.transform.StateActionTransform
|
| 209 |
+
apply_to: ${modality_config_<EMBODIMENT>.state.modality_keys}
|
| 210 |
+
normalization_modes:
|
| 211 |
+
state.joint_pos: q99 # every state key needs a normalization mode
|
| 212 |
+
state.gripper_pos: q99
|
| 213 |
+
|
| 214 |
+
# Action
|
| 215 |
+
- _target_: groot.vla.data.transform.StateActionToTensor
|
| 216 |
+
apply_to: ${modality_config_<EMBODIMENT>.action.modality_keys}
|
| 217 |
+
- _target_: groot.vla.data.transform.StateActionTransform
|
| 218 |
+
apply_to: ${modality_config_<EMBODIMENT>.action.modality_keys}
|
| 219 |
+
normalization_modes:
|
| 220 |
+
action.joint_pos: q99 # every action key needs a normalization mode
|
| 221 |
+
action.gripper_pos: q99
|
| 222 |
+
|
| 223 |
+
# Concat
|
| 224 |
+
- _target_: groot.vla.data.transform.ConcatTransform
|
| 225 |
+
video_concat_order: ${modality_config_<EMBODIMENT>.video.modality_keys}
|
| 226 |
+
state_concat_order: ${modality_config_<EMBODIMENT>.state.modality_keys}
|
| 227 |
+
action_concat_order: ${modality_config_<EMBODIMENT>.action.modality_keys}
|
| 228 |
+
|
| 229 |
+
# Model-specific (required, don't change)
|
| 230 |
+
- ${model_specific_transform}
|
| 231 |
+
```
|
| 232 |
+
|
| 233 |
+
Every state and action key **must** appear in `normalization_modes`. The strategy is typically `q99`.
|
| 234 |
+
|
| 235 |
+
### 3d. Register in the global maps
|
| 236 |
+
|
| 237 |
+
Add your embodiment to each of these four maps (at the bottom of the base YAML):
|
| 238 |
+
|
| 239 |
+
```yaml
|
| 240 |
+
modality_configs:
|
| 241 |
+
...
|
| 242 |
+
<EMBODIMENT>: ${modality_config_<EMBODIMENT>}
|
| 243 |
+
|
| 244 |
+
transforms:
|
| 245 |
+
...
|
| 246 |
+
<EMBODIMENT>: ${transform_<EMBODIMENT>}
|
| 247 |
+
|
| 248 |
+
metadata_versions:
|
| 249 |
+
...
|
| 250 |
+
<EMBODIMENT>: '0221'
|
| 251 |
+
|
| 252 |
+
fps:
|
| 253 |
+
...
|
| 254 |
+
<EMBODIMENT>: 30 # set to your dataset's FPS
|
| 255 |
+
```
|
| 256 |
+
|
| 257 |
+
---
|
| 258 |
+
|
| 259 |
+
## Step 4: Create a Dataset YAML
|
| 260 |
+
|
| 261 |
+
Create `groot/vla/configs/data/dreamzero/<EMBODIMENT>_relative.yaml`:
|
| 262 |
+
|
| 263 |
+
```yaml
|
| 264 |
+
# @package _global_
|
| 265 |
+
|
| 266 |
+
defaults:
|
| 267 |
+
- dreamzero/base_48_wan_fine_aug_relative
|
| 268 |
+
- _self_
|
| 269 |
+
|
| 270 |
+
max_state_dim: 64
|
| 271 |
+
use_global_metadata: false
|
| 272 |
+
relative_action: true
|
| 273 |
+
relative_action_per_horizon: false
|
| 274 |
+
relative_action_keys:
|
| 275 |
+
- joint_pos # sub-key names (without state./action. prefix)
|
| 276 |
+
- gripper_pos # that should use relative actions
|
| 277 |
+
max_chunk_size: 5
|
| 278 |
+
dataset_shard_sampling_rate: 0.1
|
| 279 |
+
mixture_dataset_cls: groot.vla.data.dataset.lerobot_sharded.ShardedLeRobotMixtureDataset.from_mixture_spec
|
| 280 |
+
single_dataset_cls: groot.vla.data.dataset.lerobot_sharded.ShardedLeRobotSubLangSingleActionChunkDatasetDROID
|
| 281 |
+
|
| 282 |
+
<EMBODIMENT>_data_root: ??? # set via CLI or env var
|
| 283 |
+
|
| 284 |
+
train_dataset:
|
| 285 |
+
_target_: ${mixture_dataset_cls}
|
| 286 |
+
_convert_: object
|
| 287 |
+
mixture_spec:
|
| 288 |
+
- dataset_path:
|
| 289 |
+
<EMBODIMENT>: # must match key in modality_configs/transforms
|
| 290 |
+
- ${<EMBODIMENT>_data_root}
|
| 291 |
+
dataset_weight: 1.0
|
| 292 |
+
distribute_weights: true
|
| 293 |
+
|
| 294 |
+
dataset_class: ${single_dataset_cls}
|
| 295 |
+
all_modality_configs: ${modality_configs}
|
| 296 |
+
all_transforms: ${transforms}
|
| 297 |
+
metadata_versions: ${metadata_versions}
|
| 298 |
+
fps: ${fps}
|
| 299 |
+
dataset_kwargs:
|
| 300 |
+
video_backend: decord
|
| 301 |
+
use_global_metadata: ${use_global_metadata}
|
| 302 |
+
max_chunk_size: ${max_chunk_size}
|
| 303 |
+
relative_action: ${relative_action}
|
| 304 |
+
relative_action_keys: ${relative_action_keys}
|
| 305 |
+
relative_action_per_horizon: ${relative_action_per_horizon}
|
| 306 |
+
mixture_kwargs:
|
| 307 |
+
training: true
|
| 308 |
+
balance_dataset_weights: false
|
| 309 |
+
seed: 42
|
| 310 |
+
shard_sampling_rate: ${dataset_shard_sampling_rate}
|
| 311 |
+
```
|
| 312 |
+
|
| 313 |
+
The critical things to get right:
|
| 314 |
+
|
| 315 |
+
- `<EMBODIMENT>` in `mixture_spec.dataset_path` must match the key in `modality_configs` and `transforms`.
|
| 316 |
+
- `relative_action_keys` lists the sub-key names (without `state.`/`action.` prefix) that exist in **both** state and action modalities.
|
| 317 |
+
|
| 318 |
+
---
|
| 319 |
+
|
| 320 |
+
## Step 5: Create a Training Script
|
| 321 |
+
|
| 322 |
+
Create `scripts/train/<EMBODIMENT>_training.sh`:
|
| 323 |
+
|
| 324 |
+
```bash
|
| 325 |
+
#!/bin/bash
|
| 326 |
+
export HYDRA_FULL_ERROR=1
|
| 327 |
+
|
| 328 |
+
# ============ CONFIGURATION ============
|
| 329 |
+
DATA_ROOT=${DATA_ROOT:?"Set DATA_ROOT to your GEAR-converted dataset"}
|
| 330 |
+
OUTPUT_DIR=${OUTPUT_DIR:-"./checkpoints/dreamzero_<EMBODIMENT>_lora"}
|
| 331 |
+
|
| 332 |
+
if [ -z "${NUM_GPUS:-}" ]; then
|
| 333 |
+
NUM_GPUS=$(nvidia-smi -L 2>/dev/null | wc -l)
|
| 334 |
+
fi
|
| 335 |
+
NUM_GPUS=${NUM_GPUS:-8}
|
| 336 |
+
|
| 337 |
+
WAN_CKPT_DIR=${WAN_CKPT_DIR:-"./checkpoints/Wan2.1-I2V-14B-480P"}
|
| 338 |
+
TOKENIZER_DIR=${TOKENIZER_DIR:-"./checkpoints/umt5-xxl"}
|
| 339 |
+
# =======================================
|
| 340 |
+
|
| 341 |
+
# Auto-download weights if missing
|
| 342 |
+
if [ ! -d "$WAN_CKPT_DIR" ] || [ -z "$(ls -A "$WAN_CKPT_DIR" 2>/dev/null)" ]; then
|
| 343 |
+
huggingface-cli download Wan-AI/Wan2.1-I2V-14B-480P --local-dir "$WAN_CKPT_DIR"
|
| 344 |
+
fi
|
| 345 |
+
if [ ! -d "$TOKENIZER_DIR" ] || [ -z "$(ls -A "$TOKENIZER_DIR" 2>/dev/null)" ]; then
|
| 346 |
+
huggingface-cli download google/umt5-xxl --local-dir "$TOKENIZER_DIR"
|
| 347 |
+
fi
|
| 348 |
+
|
| 349 |
+
if [ ! -d "$DATA_ROOT" ]; then
|
| 350 |
+
echo "ERROR: Dataset not found at $DATA_ROOT"
|
| 351 |
+
exit 1
|
| 352 |
+
fi
|
| 353 |
+
if [ ! -f "$DATA_ROOT/meta/embodiment.json" ]; then
|
| 354 |
+
echo "ERROR: meta/embodiment.json missing — run convert_lerobot_to_gear.py first"
|
| 355 |
+
exit 1
|
| 356 |
+
fi
|
| 357 |
+
|
| 358 |
+
torchrun --nproc_per_node $NUM_GPUS --standalone \
|
| 359 |
+
groot/vla/experiment/experiment.py \
|
| 360 |
+
report_to=wandb \
|
| 361 |
+
data=dreamzero/<EMBODIMENT>_relative \
|
| 362 |
+
wandb_project=dreamzero \
|
| 363 |
+
train_architecture=lora \
|
| 364 |
+
num_frames=33 \
|
| 365 |
+
action_horizon=24 \
|
| 366 |
+
num_views=3 \
|
| 367 |
+
model=dreamzero/vla \
|
| 368 |
+
model/dreamzero/action_head=wan_flow_matching_action_tf \
|
| 369 |
+
model/dreamzero/transform=dreamzero_cotrain \
|
| 370 |
+
num_frame_per_block=2 \
|
| 371 |
+
num_action_per_block=24 \
|
| 372 |
+
num_state_per_block=1 \
|
| 373 |
+
seed=42 \
|
| 374 |
+
training_args.learning_rate=1e-5 \
|
| 375 |
+
training_args.deepspeed="groot/vla/configs/deepspeed/zero2.json" \
|
| 376 |
+
save_steps=10000 \
|
| 377 |
+
training_args.warmup_ratio=0.05 \
|
| 378 |
+
output_dir=$OUTPUT_DIR \
|
| 379 |
+
per_device_train_batch_size=4 \
|
| 380 |
+
max_steps=100000 \
|
| 381 |
+
weight_decay=1e-5 \
|
| 382 |
+
save_total_limit=10 \
|
| 383 |
+
upload_checkpoints=false \
|
| 384 |
+
bf16=true \
|
| 385 |
+
tf32=true \
|
| 386 |
+
eval_bf16=true \
|
| 387 |
+
dataloader_pin_memory=false \
|
| 388 |
+
dataloader_num_workers=1 \
|
| 389 |
+
image_resolution_width=320 \
|
| 390 |
+
image_resolution_height=176 \
|
| 391 |
+
save_lora_only=true \
|
| 392 |
+
max_chunk_size=4 \
|
| 393 |
+
frame_seqlen=880 \
|
| 394 |
+
save_strategy=steps \
|
| 395 |
+
<EMBODIMENT>_data_root=$DATA_ROOT \
|
| 396 |
+
dit_version=$WAN_CKPT_DIR \
|
| 397 |
+
text_encoder_pretrained_path=$WAN_CKPT_DIR/models_t5_umt5-xxl-enc-bf16.pth \
|
| 398 |
+
image_encoder_pretrained_path=$WAN_CKPT_DIR/models_clip_open-clip-xlm-roberta-large-vit-huge-14.pth \
|
| 399 |
+
vae_pretrained_path=$WAN_CKPT_DIR/Wan2.1_VAE.pth \
|
| 400 |
+
tokenizer_path=$TOKENIZER_DIR \
|
| 401 |
+
pretrained_model_path=./checkpoints/DreamZero-AgiBot \
|
| 402 |
+
++action_head_cfg.config.skip_component_loading=true \
|
| 403 |
+
++action_head_cfg.config.defer_lora_injection=true
|
| 404 |
+
```
|
| 405 |
+
|
| 406 |
+
### Key parameters to adjust per embodiment
|
| 407 |
+
|
| 408 |
+
| Parameter | Default | When to change |
|
| 409 |
+
|---|---|---|
|
| 410 |
+
| `num_views` | `3` | Number of cameras your robot has |
|
| 411 |
+
| `action_horizon` | `24` | Must match the number of action `delta_indices` |
|
| 412 |
+
| `num_frames` | `33` | Must be `len(video delta_indices) + num_frame_per_block * (blocks - 1)` |
|
| 413 |
+
| `image_resolution_width` | `320` | Match your camera resolution (or desired resize) |
|
| 414 |
+
| `image_resolution_height` | `176` | Match your camera resolution (or desired resize) |
|
| 415 |
+
| `max_steps` | `100000` | Scale with dataset size |
|
| 416 |
+
| `per_device_train_batch_size` | `4` | Adjust for GPU memory |
|
| 417 |
+
|
| 418 |
+
---
|
| 419 |
+
|
| 420 |
+
## Step 6: Train
|
| 421 |
+
|
| 422 |
+
### Download the pretrained checkpoint
|
| 423 |
+
|
| 424 |
+
The training scripts load from a pretrained DreamZero checkpoint for LoRA fine-tuning. Download [DreamZero-AgiBot](https://huggingface.co/GEAR-Dreams/DreamZero-AgiBot) (~45GB) to `./checkpoints/DreamZero-AgiBot`:
|
| 425 |
+
|
| 426 |
+
```bash
|
| 427 |
+
git clone https://huggingface.co/GEAR-Dreams/DreamZero-AgiBot ./checkpoints/DreamZero-AgiBot
|
| 428 |
+
```
|
| 429 |
+
|
| 430 |
+
Or with the Hugging Face CLI:
|
| 431 |
+
|
| 432 |
+
```bash
|
| 433 |
+
hf download GEAR-Dreams/DreamZero-AgiBot --repo-type model --local-dir ./checkpoints/DreamZero-AgiBot
|
| 434 |
+
```
|
| 435 |
+
|
| 436 |
+
### Launch training
|
| 437 |
+
|
| 438 |
+
```bash
|
| 439 |
+
DATA_ROOT=/path/to/your_dataset bash scripts/train/<EMBODIMENT>_training.sh
|
| 440 |
+
|
| 441 |
+
# With overrides:
|
| 442 |
+
DATA_ROOT=/path/to/your_dataset OUTPUT_DIR=./checkpoints/run1 NUM_GPUS=4 \
|
| 443 |
+
bash scripts/train/<EMBODIMENT>_training.sh
|
| 444 |
+
```
|
| 445 |
+
|
| 446 |
+
---
|
| 447 |
+
|
| 448 |
+
## Pre-Training Checklist
|
| 449 |
+
|
| 450 |
+
- [ ] `meta/embodiment.json` exists and has the correct tag
|
| 451 |
+
- [ ] `meta/modality.json` state/action/video/annotation keys are populated
|
| 452 |
+
- [ ] `meta/stats.json` and `meta/relative_stats_dreamzero.json` exist
|
| 453 |
+
- [ ] `meta/tasks.jsonl` and `meta/episodes.jsonl` exist
|
| 454 |
+
- [ ] Embodiment tag in `embodiment.json` matches the key in `modality_configs` / `transforms` / `metadata_versions` / `fps`
|
| 455 |
+
- [ ] YAML `modality_keys` match `modality.json` keys exactly (with `state.`/`action.`/`video.`/`annotation.` prefix)
|
| 456 |
+
- [ ] Every state and action key appears in `normalization_modes` in the transform block
|
| 457 |
+
- [ ] `relative_action_keys` are sub-key names that exist in both state and action
|
| 458 |
+
- [ ] Wan2.1-I2V-14B-480P and umt5-xxl weights are available
|
| 459 |
+
- [ ] DreamZero-AgiBot checkpoint is downloaded to `./checkpoints/DreamZero-AgiBot`
|
| 460 |
+
|
| 461 |
+
---
|
| 462 |
+
|
| 463 |
+
## Quick Reference: Existing Embodiments
|
| 464 |
+
|
| 465 |
+
| Embodiment | Data Config | Layout |
|
| 466 |
+
|---|---|---|
|
| 467 |
+
| `oxe_droid` | `droid_relative.yaml` | 3 cameras, joint_position + gripper_position |
|
| 468 |
+
| `agibot` | `agibot_relative.yaml` | 3 cameras, 6 state keys, 7 action keys |
|
| 469 |
+
| `yam` | `yam_relative.yaml` | 3 cameras (top/left/right), bimanual left/right joint_pos + gripper_pos |
|
| 470 |
+
|
| 471 |
+
Use these as concrete examples when building your own config.
|
docs/DROID_CONVERSION.md
ADDED
|
@@ -0,0 +1,62 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# Converting DROID from Scratch
|
| 2 |
+
|
| 3 |
+
If you want to reproduce the DreamZero DROID dataset conversion yourself (or modify the filtering), follow the steps below. This requires the raw DROID 1.0.1 dataset in RLDS format and the idle filter ranges JSON.
|
| 4 |
+
|
| 5 |
+
> **Most users should skip this** and simply download the preprocessed dataset:
|
| 6 |
+
> ```bash
|
| 7 |
+
> huggingface-cli download GEAR-Dreams/DreamZero-DROID-Data --repo-type dataset --local-dir ./data/droid_lerobot
|
| 8 |
+
> ```
|
| 9 |
+
|
| 10 |
+
## Step 1: Install conversion dependencies
|
| 11 |
+
|
| 12 |
+
```bash
|
| 13 |
+
pip install tensorflow tensorflow-datasets polars av
|
| 14 |
+
```
|
| 15 |
+
|
| 16 |
+
## Step 2: Download the raw DROID 1.0.1 dataset
|
| 17 |
+
|
| 18 |
+
This requires `gsutil` ([Google Cloud CLI](https://cloud.google.com/storage/docs/gsutil_install)). The full dataset is ~1.7TB.
|
| 19 |
+
|
| 20 |
+
```bash
|
| 21 |
+
gsutil -m cp -r gs://gresearch/robotics/droid/1.0.1 ./data/droid/1.0.1
|
| 22 |
+
```
|
| 23 |
+
|
| 24 |
+
> **Important:** Use version 1.0.1, not 1.0.0. Version 1.0.1 contains the complete set of language annotations (~75k episodes).
|
| 25 |
+
|
| 26 |
+
## Step 3: Download the idle filter ranges
|
| 27 |
+
|
| 28 |
+
This JSON file maps each episode to the frame ranges that should be kept (non-idle frames). It was originally computed by [Physical Intelligence](https://github.com/Physical-Intelligence/openpi) for training pi0-DROID models.
|
| 29 |
+
|
| 30 |
+
```bash
|
| 31 |
+
gsutil cp gs://openpi-assets/droid/droid_sample_ranges_v1_0_1.json ./data/keep_ranges.json
|
| 32 |
+
```
|
| 33 |
+
|
| 34 |
+
## Step 4: Run the conversion
|
| 35 |
+
|
| 36 |
+
```bash
|
| 37 |
+
python scripts/data/convert_droid.py \
|
| 38 |
+
./data/droid/1.0.1 \
|
| 39 |
+
./data/droid_lerobot \
|
| 40 |
+
--keep-ranges-path ./data/keep_ranges.json \
|
| 41 |
+
--filter-failed \
|
| 42 |
+
-n 16
|
| 43 |
+
```
|
| 44 |
+
|
| 45 |
+
For a quick test with a small subset:
|
| 46 |
+
```bash
|
| 47 |
+
python scripts/data/convert_droid.py \
|
| 48 |
+
./data/droid/1.0.1 \
|
| 49 |
+
./data/droid_lerobot_test \
|
| 50 |
+
--keep-ranges-path ./data/keep_ranges.json \
|
| 51 |
+
--filter-failed \
|
| 52 |
+
--first-n 5 \
|
| 53 |
+
-n 4
|
| 54 |
+
```
|
| 55 |
+
|
| 56 |
+
## Script reference
|
| 57 |
+
|
| 58 |
+
See [`scripts/data/convert_droid.py`](scripts/data/convert_droid.py) for full usage:
|
| 59 |
+
|
| 60 |
+
```
|
| 61 |
+
python scripts/data/convert_droid.py --help
|
| 62 |
+
```
|
docs/WAN22_BACKBONE.md
ADDED
|
@@ -0,0 +1,167 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# Training DreamZero with Wan2.2-TI2V-5B Backbone
|
| 2 |
+
|
| 3 |
+
This guide explains how to train DreamZero on the DROID dataset using **Wan2.2-TI2V-5B** as the backbone instead of the default Wan2.1-I2V-14B.
|
| 4 |
+
|
| 5 |
+
## Architecture Differences
|
| 6 |
+
|
| 7 |
+
| Component | Wan2.1-I2V-14B | Wan2.2-TI2V-5B |
|
| 8 |
+
|-----------|-----------------|----------------|
|
| 9 |
+
| DiT dim | 5120 | 3072 |
|
| 10 |
+
| DiT layers | 32 | 30 |
|
| 11 |
+
| DiT heads | 16 | 24 |
|
| 12 |
+
| FFN dim | 13824 | 14336 |
|
| 13 |
+
| VAE latent channels | 16 | 48 |
|
| 14 |
+
| VAE spatial stride | 8× | 16× |
|
| 15 |
+
| Model type | i2v | ti2v |
|
| 16 |
+
|
| 17 |
+
**FFN** = Feed-Forward Network: the two-layer MLP in each transformer block (Linear → GELU → Linear). FFN dim is the intermediate hidden size (e.g. 14336 for 5B).
|
| 18 |
+
|
| 19 |
+
DreamZero uses a **CausalWanModel** wrapper that extends the base Wan architecture with **action/state registers** for robot policy learning. The same `CausalWanModel` class supports both Wan2.1 and Wan2.2 backbones via configuration—no new class is required. The config switches the architecture parameters (dim, in_dim, out_dim, etc.) and uses `WanVideoVAE38` for the 48-channel Wan2.2 VAE.
|
| 20 |
+
|
| 21 |
+
**What action/state registers do:** The DiT sees a single sequence `[video_tokens | action_register]` where the action register is encoded action and state features (one chunk per block). All tokens share the same transformer (with causal masking and RoPE). The model learns to predict **video noise** (dynamics) and **action noise** (policy): the action-register slice is decoded by `action_decoder` to produce action noise predictions. So the model is conditioned on current state and (noisy) actions and learns to denoise both video and actions for closed-loop policy learning.
|
| 22 |
+
|
| 23 |
+
## Causal masking, RoPE, and sequence layout
|
| 24 |
+
|
| 25 |
+
### Causal masking
|
| 26 |
+
|
| 27 |
+
In attention, **causal masking** means each position can only attend to **past and current** positions (no future). So token at index `i` can see keys at indices `j ≤ i`. That keeps the model autoregressive: it never uses future video frames or future actions when predicting the current step. In CausalWanModel the masking is **blockwise**: the first frame attends to itself; each later block of video frames can attend to the first frame plus previous (and optionally current) blocks. Action and state tokens have their own causal pattern so each action chunk only sees past video and past actions/state. This matches policy learning where you condition on observed history and predict the next action chunk.
|
| 28 |
+
|
| 29 |
+
### RoPE (Rotary Position Embeddings)
|
| 30 |
+
|
| 31 |
+
**RoPE** encodes position by rotating query and key vectors in a complex plane with position-dependent angles. Unlike adding a position vector, RoPE makes attention scores depend on the *relative* position of query and key, which generalizes better to longer sequences. In CausalWanModel:
|
| 32 |
+
|
| 33 |
+
- **Video tokens** use **3D RoPE**: separate frequency components for frame index (time), height, and width of the patch grid. So each token knows its (t, h, w) in the video.
|
| 34 |
+
- **Action and state tokens** use **1D RoPE**: a single position index along the sequence (frame/block index). So the model knows the temporal order of action chunks and state.
|
| 35 |
+
|
| 36 |
+
Freqs are built in `_create_freqs()` from the patch grid size (F, H, W) and concatenated with separate 1D freqs for the action register.
|
| 37 |
+
|
| 38 |
+
### Tokens, blocks, and chunks
|
| 39 |
+
|
| 40 |
+
- **Token**: The smallest unit the transformer sees. After **patch_embedding** (stride 1×2×2 on the latent), one frame yields a 2D grid of tokens; the total per frame is **frame_seqlen** (e.g. 50 for 160×320). So one **token** = one patch (e.g. 1×2×2 in latent space).
|
| 41 |
+
|
| 42 |
+
- **Block (image block)**: A group of consecutive **frames**, not tokens. **num_frame_per_block** (e.g. 2) frames form one “image block.” So with 33 frames you get multiple blocks. **num_image_blocks** = `(num_frames - 1) // num_frame_per_block`. Blocks are used for blockwise causal attention and to align video with action/state.
|
| 43 |
+
|
| 44 |
+
- **Chunk**: In policy terms, an **action chunk** is the sequence of actions the policy outputs for one block (e.g. **num_action_per_block** = 24 actions per block). The **action register** in the DiT has one chunk per image block: for each block there are `num_action_per_block` action tokens and `num_state_per_block` state tokens. So the register length is `num_image_blocks * (num_action_per_block + num_state_per_block)`. “Chunk” and “block” are often used together: one video block corresponds to one action chunk (and one state token) in the register.
|
| 45 |
+
|
| 46 |
+
Summary: **tokens** = patch-level units (50 per frame); **blocks** = groups of frames (e.g. 2 frames per block); **chunks** = per-block action (and state) outputs that are packed into the action register.
|
| 47 |
+
|
| 48 |
+
## Inference: blocks, chunks, and closed-loop
|
| 49 |
+
|
| 50 |
+
### How blocks and chunks are used when predicting actions
|
| 51 |
+
|
| 52 |
+
At **inference**, the model predicts **one action chunk** per call, conditioned on **one block** of video (and current state):
|
| 53 |
+
|
| 54 |
+
1. **Input**: A short video of the current block — e.g. `num_frame_per_block` frames (e.g. 2) — plus current **state** and (during the denoising loop) **noisy actions** for the chunk being predicted. The first time in a trajectory, the “context” is the first frame (and optionally a warm-up pass with no action to fill the KV cache).
|
| 55 |
+
|
| 56 |
+
2. **DiT input**: The sequence is `[video_tokens for this block | action_register]`. The action register holds encoded **noisy** actions and **state** for this block only (one chunk). So the DiT sees: “this block of video + this chunk’s noisy actions and state.”
|
| 57 |
+
|
| 58 |
+
3. **KV cache**: To keep inference causal and efficient, the model uses a **KV cache** over previous blocks. So for the *next* block, the cache already contains keys/values for earlier frames; the DiT only runs on the **new** block’s tokens plus the new action register. `current_start_frame` tells the DiT which block we’re on so RoPE and cache indexing are correct.
|
| 59 |
+
|
| 60 |
+
4. **Output**: The DiT predicts **video noise** and **action noise**. The action noise is decoded by `action_decoder` into a prediction for the **current chunk**. The scheduler then updates the noisy action toward clean; after `num_inference_steps` denoising steps you get **one clean action chunk** (e.g. 24 actions).
|
| 61 |
+
|
| 62 |
+
So: **one block of frames** (and state) in → **one action chunk** out. Blocks and chunks are aligned: one image block ↔ one action chunk in the register.
|
| 63 |
+
|
| 64 |
+
### How DreamZero does closed-loop inference
|
| 65 |
+
|
| 66 |
+
Closed-loop execution reuses the same block/chunk logic in a loop:
|
| 67 |
+
|
| 68 |
+
1. **Observe**: Robot has current observation (e.g. image history + state). The policy is called with this observation (e.g. via `lazy_joint_video_action` or `lazy_joint_video_action_causal`).
|
| 69 |
+
|
| 70 |
+
2. **Predict**: The action head runs the diffusion loop for the **current block**: it encodes the observed frames to latent, runs the DiT (with KV cache and `current_start_frame`) for each denoising step, and returns one denoised **action chunk** (e.g. 24 actions).
|
| 71 |
+
|
| 72 |
+
3. **Execute**: The robot **executes** that chunk (e.g. 24 steps at 5 Hz → ~4.8 s). No new model call during execution.
|
| 73 |
+
|
| 74 |
+
4. **Repeat**: After execution, new observation is available. The policy is called again with the new video (e.g. last N frames). If the task/language is unchanged, `current_start_frame` is incremented by `num_frame_per_block` and the KV cache is reused; the DiT only processes the **new** block and predicts the **next** action chunk. If the task or language changes (or the cache is full), the cache and `current_start_frame` are reset.
|
| 75 |
+
|
| 76 |
+
So closed-loop = **repeated “one block in → one chunk out”** with KV cache across steps so the model never re-processes past frames.
|
| 77 |
+
|
| 78 |
+
### What changes when you swap the backbone to 5B
|
| 79 |
+
|
| 80 |
+
The **inference algorithm and API stay the same** for 14B vs 5B:
|
| 81 |
+
|
| 82 |
+
- Same **block/chunk layout**: `num_frame_per_block`, `num_action_per_block`, `num_state_per_block` (and thus one block → one chunk) are defined by config and data; they do not depend on which backbone (14B vs 5B) you use.
|
| 83 |
+
- Same **closed-loop flow**: `lazy_joint_video_action`, KV cache, `current_start_frame`, and the denoising loop are in the **action head** and are shared. The policy still calls the same methods (`get_action`, `lazy_joint_video_action`, etc.).
|
| 84 |
+
- Same **backbone role**: The backbone only produces conditioning (e.g. text embeddings). The action head owns the DiT, VAE, and action/state encoders. So “swapping to 5B” means swapping the **action head config** (and checkpoints) to the Wan22 5B DiT + VAE38 + 160×320; the high-level inference path (backbone → action_head → one chunk) is unchanged.
|
| 85 |
+
|
| 86 |
+
What **does** change with 5B:
|
| 87 |
+
|
| 88 |
+
- **DiT size and layout**: 5B uses a smaller DiT (dim 3072, 30 layers, 24 heads), **frame_seqlen = 50** (for 160×320), and **no** first-frame latent concat (`concat_first_frame_latent=False`). First frame is conditioned via **CLIP** in the context, not as extra channel in the latent.
|
| 89 |
+
- **VAE and resolution**: 5B uses **WanVideoVAE38** (48 channels, 16× spatial) and **160×320** video. So latent is 10×20; tokens per frame = 50.
|
| 90 |
+
- **Conditioning**: 5B uses CLIP image embedding for the first frame in the context; 14B can concatenate the first-frame latent to the DiT input. The action head handles this inside the same `_forward_inference` / `_forward_blocks`; no change to the external inference API.
|
| 91 |
+
|
| 92 |
+
So: **blocks and chunks** are used the same way for predicting actions at inference; **closed-loop** is the same loop of “observe → predict one chunk → execute → repeat” with KV cache; **swapping to 5B** keeps that flow and only changes the internal model (DiT/VAE) and resolution/conditioning.
|
| 93 |
+
|
| 94 |
+
## Prerequisites
|
| 95 |
+
|
| 96 |
+
1. **Wan2.2-TI2V-5B** weights:
|
| 97 |
+
```bash
|
| 98 |
+
huggingface-cli download Wan-AI/Wan2.2-TI2V-5B --local-dir ./checkpoints/Wan2.2-TI2V-5B
|
| 99 |
+
```
|
| 100 |
+
Or clone from [Wan2.2 GitHub](https://github.com/Wan-Video/Wan2.2) and follow their download instructions.
|
| 101 |
+
|
| 102 |
+
2. **Image encoder (CLIP)**: Wan2.2-TI2V-5B does not include the CLIP image encoder. Use the one from Wan2.1:
|
| 103 |
+
```bash
|
| 104 |
+
huggingface-cli download Wan-AI/Wan2.1-I2V-14B-480P --local-dir ./checkpoints/Wan2.1-I2V-14B-480P
|
| 105 |
+
```
|
| 106 |
+
Only `models_clip_open-clip-xlm-roberta-large-vit-huge-14.pth` is needed.
|
| 107 |
+
|
| 108 |
+
3. **DROID dataset** in LeRobot format:
|
| 109 |
+
```bash
|
| 110 |
+
huggingface-cli download GEAR-Dreams/DreamZero-DROID-Data --repo-type dataset --local-dir ./data/droid_lerobot
|
| 111 |
+
```
|
| 112 |
+
|
| 113 |
+
## Quick Start
|
| 114 |
+
|
| 115 |
+
```bash
|
| 116 |
+
# Set paths (optional - defaults shown)
|
| 117 |
+
export WAN22_CKPT_DIR=./checkpoints/Wan2.2-TI2V-5B
|
| 118 |
+
export IMAGE_ENCODER_DIR=./checkpoints/Wan2.1-I2V-14B-480P # for CLIP only
|
| 119 |
+
export DROID_DATA_ROOT=./data/droid_lerobot
|
| 120 |
+
|
| 121 |
+
# Run training
|
| 122 |
+
bash scripts/train/droid_training_wan22.sh
|
| 123 |
+
```
|
| 124 |
+
|
| 125 |
+
## Configuration Details
|
| 126 |
+
|
| 127 |
+
The Wan2.2 config (`wan_flow_matching_action_tf_wan22.yaml`) overrides:
|
| 128 |
+
|
| 129 |
+
- **model/dreamzero/action_head**: `wan_flow_matching_action_tf_wan22`
|
| 130 |
+
- **diffusion_model_cfg**: Wan2.2 architecture (dim=3072, in_dim=48, out_dim=48, etc.)
|
| 131 |
+
- **vae_cfg**: `WanVideoVAE38` (48-channel Wan2.2 VAE)
|
| 132 |
+
- **frame_seqlen**: 50 (patch output per frame)
|
| 133 |
+
- **target_video_height / target_video_width**: 160 and 320 so latent spatial size is **even** (10×20 after VAE38 16×), avoiding a dynamics-loss crop. Previously 176×320 gave latent 11×20 (odd height); we use **160×320** (H×W) so both latent dimensions are even after the DiT’s stride-(1,2,2) patch embedding.
|
| 134 |
+
|
| 135 |
+
For other resolutions, `frame_seqlen` must match patch output per frame; use H and W divisible by 32 for even latent:
|
| 136 |
+
- 160×320 (H×W): latent 10×20 → 50
|
| 137 |
+
- 176×320: latent 11×20 → 50 (odd H; loss uses crop)
|
| 138 |
+
- 640×352: 220
|
| 139 |
+
|
| 140 |
+
## Using with Custom Training Scripts
|
| 141 |
+
|
| 142 |
+
To use Wan2.2 in your own training script, add:
|
| 143 |
+
|
| 144 |
+
```bash
|
| 145 |
+
model/dreamzero/action_head=wan_flow_matching_action_tf_wan22 \
|
| 146 |
+
dit_version=$WAN22_CKPT_DIR \
|
| 147 |
+
text_encoder_pretrained_path=$WAN22_CKPT_DIR/models_t5_umt5-xxl-enc-bf16.pth \
|
| 148 |
+
image_encoder_pretrained_path=$IMAGE_ENCODER_DIR/models_clip_open-clip-xlm-roberta-large-vit-huge-14.pth \
|
| 149 |
+
vae_pretrained_path=$WAN22_CKPT_DIR/Wan2.2_VAE.pth
|
| 150 |
+
```
|
| 151 |
+
(Do not pass `frame_seqlen`; the Wan22 config uses 50.)
|
| 152 |
+
|
| 153 |
+
## File Layout
|
| 154 |
+
|
| 155 |
+
```
|
| 156 |
+
dreamzero/
|
| 157 |
+
├── groot/vla/configs/model/dreamzero/action_head/
|
| 158 |
+
│ ├── wan_flow_matching_action_tf.yaml # Wan2.1 (default)
|
| 159 |
+
│ └── wan_flow_matching_action_tf_wan22.yaml # Wan2.2-TI2V-5B
|
| 160 |
+
├── scripts/train/
|
| 161 |
+
│ ├── droid_training.sh # Wan2.1 backbone
|
| 162 |
+
│ └── droid_training_wan22.sh # Wan2.2 backbone
|
| 163 |
+
└── docs/
|
| 164 |
+
└── WAN22_BACKBONE.md # This file
|
| 165 |
+
```
|
| 166 |
+
|
| 167 |
+
The action head (`wan_flow_matching_action_tf.py`) automatically detects Wan2.2 vs Wan2.1 based on `in_dim` (48 vs 16) and `vae.z_dim` (48 vs 16), and loads the correct checkpoint files from the appropriate HuggingFace repos when local paths are not found.
|
docs/data-preparation.md
ADDED
|
@@ -0,0 +1,194 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# DreamZero 数据准备指南
|
| 2 |
+
|
| 3 |
+
---
|
| 4 |
+
|
| 5 |
+
## 1. 数据格式要求
|
| 6 |
+
|
| 7 |
+
DreamZero 期望每个训练样本包含以下字段:
|
| 8 |
+
|
| 9 |
+
| 字段 | 类型 | 形状 | 说明 |
|
| 10 |
+
|------|------|------|------|
|
| 11 |
+
| `video` | uint8 numpy | [T, V, H, W, 3] | T=帧数, V=视角数 |
|
| 12 |
+
| `state` | float32 numpy | [state_horizon, max_state_dim] | 机器人状态(需 padding) |
|
| 13 |
+
| `action` | float32 numpy | [action_horizon, max_action_dim] | 动作序列(需 padding) |
|
| 14 |
+
| `language` | str | - | 任务描述 |
|
| 15 |
+
|
| 16 |
+
### 参数约束
|
| 17 |
+
|
| 18 |
+
```
|
| 19 |
+
action_horizon / (lat_T - 1) = num_action_per_block / num_frame_per_block
|
| 20 |
+
(lat_T - 1) / state_horizon = num_frame_per_block / num_state_per_block
|
| 21 |
+
lat_T = num_frames // 4 # Wan2.2 VAE 4x 时间下采样
|
| 22 |
+
```
|
| 23 |
+
|
| 24 |
+
**标准参数**(已验证):
|
| 25 |
+
- `num_frames=12, action_horizon=12, state_horizon=1`
|
| 26 |
+
- `num_frame_per_block=2, num_action_per_block=12, num_state_per_block=1`
|
| 27 |
+
- `max_state_dim=44, max_action_dim=32`
|
| 28 |
+
|
| 29 |
+
---
|
| 30 |
+
|
| 31 |
+
## 2. State/Action Padding
|
| 32 |
+
|
| 33 |
+
State 和 Action 统一 padding 到固定维度:
|
| 34 |
+
|
| 35 |
+
```python
|
| 36 |
+
import numpy as np
|
| 37 |
+
|
| 38 |
+
MAX_STATE_DIM = 44
|
| 39 |
+
MAX_ACTION_DIM = 32
|
| 40 |
+
|
| 41 |
+
def pad_state(state: np.ndarray) -> np.ndarray:
|
| 42 |
+
"""Pad state to MAX_STATE_DIM."""
|
| 43 |
+
d = state.shape[-1]
|
| 44 |
+
padded = np.zeros((MAX_STATE_DIM,), dtype=np.float32)
|
| 45 |
+
padded[:d] = state.astype(np.float32)
|
| 46 |
+
return padded
|
| 47 |
+
|
| 48 |
+
def pad_action(action_chunk: np.ndarray) -> np.ndarray:
|
| 49 |
+
"""Pad action to [horizon, MAX_ACTION_DIM]."""
|
| 50 |
+
d = action_chunk.shape[-1]
|
| 51 |
+
padded = np.zeros((*action_chunk.shape[:-1], MAX_ACTION_DIM), dtype=np.float32)
|
| 52 |
+
padded[..., :d] = action_chunk.astype(np.float32)
|
| 53 |
+
return padded
|
| 54 |
+
```
|
| 55 |
+
|
| 56 |
+
---
|
| 57 |
+
|
| 58 |
+
## 3. 添加新数据集步骤
|
| 59 |
+
|
| 60 |
+
### 3.1 创建 Dataset 类
|
| 61 |
+
|
| 62 |
+
```python
|
| 63 |
+
# groot/vla/data/dataset/my_dataset.py
|
| 64 |
+
from pathlib import Path
|
| 65 |
+
import numpy as np
|
| 66 |
+
from groot.vla.model.dreamzero.transform.dreamzero_cotrain import DreamTransform
|
| 67 |
+
|
| 68 |
+
class MyDataset:
|
| 69 |
+
def __init__(self, dataset_dir, num_frames=12, action_horizon=12,
|
| 70 |
+
state_horizon=1, num_views=1, ...):
|
| 71 |
+
# 1. 扫描数据文件
|
| 72 |
+
# 2. 构建 episode 列表
|
| 73 |
+
# 3. 创建 DreamTransform 实例
|
| 74 |
+
|
| 75 |
+
self.transform = DreamTransform(
|
| 76 |
+
default_instruction="Perform the task.",
|
| 77 |
+
max_state_dim=44, max_action_dim=32,
|
| 78 |
+
state_horizon=state_horizon,
|
| 79 |
+
action_horizon=action_horizon,
|
| 80 |
+
num_views=num_views,
|
| 81 |
+
embodiment_tag_mapping={"my_robot": 17},
|
| 82 |
+
tokenizer_path="/path/to/umt5-xxl",
|
| 83 |
+
)
|
| 84 |
+
# 必须设置 metadata
|
| 85 |
+
self.transform.set_metadata(self.merged_metadata["my_robot"])
|
| 86 |
+
self.transform.train()
|
| 87 |
+
|
| 88 |
+
def __getitem__(self, idx):
|
| 89 |
+
# 返回格式:
|
| 90 |
+
return dict(self.transform({
|
| 91 |
+
"video": video, # [T, V, H, W, C] uint8
|
| 92 |
+
"state": state, # [T, D] float32
|
| 93 |
+
"action": action, # [T, D] float32
|
| 94 |
+
"annotation.human.action.task_description": text,
|
| 95 |
+
}))
|
| 96 |
+
```
|
| 97 |
+
|
| 98 |
+
### 3.2 创建 Hydra 配置
|
| 99 |
+
|
| 100 |
+
```yaml
|
| 101 |
+
# groot/vla/configs/data/dreamzero/my_dataset.yaml
|
| 102 |
+
defaults:
|
| 103 |
+
- dreamzero/base_48_wan_fine_aug_relative
|
| 104 |
+
- _self_
|
| 105 |
+
|
| 106 |
+
my_dataset_dir: /path/to/data
|
| 107 |
+
|
| 108 |
+
num_frames: 12
|
| 109 |
+
action_horizon: 12
|
| 110 |
+
state_horizon: 1
|
| 111 |
+
num_views: 1
|
| 112 |
+
num_frame_per_block: 2
|
| 113 |
+
num_action_per_block: 12
|
| 114 |
+
num_state_per_block: 1
|
| 115 |
+
max_state_dim: 44
|
| 116 |
+
max_action_dim: 32
|
| 117 |
+
max_chunk_size: 4
|
| 118 |
+
image_resolution_height: 160
|
| 119 |
+
image_resolution_width: 320
|
| 120 |
+
frame_seqlen: 50
|
| 121 |
+
|
| 122 |
+
train_dataset:
|
| 123 |
+
_target_: groot.vla.data.dataset.my_dataset.MyDataset
|
| 124 |
+
_convert_: object
|
| 125 |
+
dataset_dir: ${my_dataset_dir}
|
| 126 |
+
num_frames: ${num_frames}
|
| 127 |
+
action_horizon: ${action_horizon}
|
| 128 |
+
state_horizon: ${state_horizon}
|
| 129 |
+
num_views: ${num_views}
|
| 130 |
+
video_height: ${image_resolution_height}
|
| 131 |
+
video_width: ${image_resolution_width}
|
| 132 |
+
```
|
| 133 |
+
|
| 134 |
+
### 3.3 训练
|
| 135 |
+
|
| 136 |
+
```bash
|
| 137 |
+
torchrun ... experiment.py \
|
| 138 |
+
data=dreamzero/my_dataset \
|
| 139 |
+
num_frames=12 action_horizon=12 num_views=1
|
| 140 |
+
```
|
| 141 |
+
|
| 142 |
+
---
|
| 143 |
+
|
| 144 |
+
## 4. Embodiment Tag
|
| 145 |
+
|
| 146 |
+
每个数据集需要注册一个 embodiment tag(用于识别机器人种类):
|
| 147 |
+
|
| 148 |
+
```python
|
| 149 |
+
# 内置 tag 映射 (部分)
|
| 150 |
+
embodiment_tag_mapping = {
|
| 151 |
+
"oxe_droid": 17, # DROID 机器人
|
| 152 |
+
"libero": 18, # LIBERO
|
| 153 |
+
"panda": 19, # Franka Panda
|
| 154 |
+
"manifeel": 20, # ManiFeel
|
| 155 |
+
"robotwin": 21, # RoboTwin
|
| 156 |
+
}
|
| 157 |
+
```
|
| 158 |
+
|
| 159 |
+
如果不确定用哪个 tag,使用 `"oxe_droid": 17`。
|
| 160 |
+
|
| 161 |
+
---
|
| 162 |
+
|
| 163 |
+
## 5. 数据验证
|
| 164 |
+
|
| 165 |
+
训练前验证数据格式:
|
| 166 |
+
|
| 167 |
+
```python
|
| 168 |
+
# 快速验证
|
| 169 |
+
ds = MyDataset(dataset_dir="/path/to/data", max_episodes=3)
|
| 170 |
+
sample = ds[0]
|
| 171 |
+
print(sample.keys())
|
| 172 |
+
for k, v in sample.items():
|
| 173 |
+
if hasattr(v, 'shape'):
|
| 174 |
+
print(f" {k}: {v.shape}, {v.dtype}")
|
| 175 |
+
|
| 176 |
+
# 验证 DreamTransform 输出
|
| 177 |
+
transformed = dict(ds.transform({
|
| 178 |
+
"video": sample["video"],
|
| 179 |
+
"state": sample["state"],
|
| 180 |
+
"action": sample["action"],
|
| 181 |
+
"annotation.human.action.task_description": "test",
|
| 182 |
+
}))
|
| 183 |
+
```
|
| 184 |
+
|
| 185 |
+
---
|
| 186 |
+
|
| 187 |
+
## 6. 常见数据问题
|
| 188 |
+
|
| 189 |
+
| 问题 | 原因 | 解决 |
|
| 190 |
+
|------|------|------|
|
| 191 |
+
| `lat_T = 0` | num_frames < 4 | 设置 num_frames >= 4 |
|
| 192 |
+
| `reshape error` | num_frame_per_block 不对齐 | 确保 (lat_T-1) % num_frame_per_block == 0 |
|
| 193 |
+
| `InterpolationKeyError` | config 缺少字段 | 检查 YAML 包含所有 `${...}` 引用 |
|
| 194 |
+
| VAE dtype error | bf16 vs float32 不匹配 | VAE 加载时用 `dtype=self.dtype` |
|
docs/training-guide.md
ADDED
|
@@ -0,0 +1,222 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# DreamZero 训练指南
|
| 2 |
+
|
| 3 |
+
---
|
| 4 |
+
|
| 5 |
+
## 1. 超参调优
|
| 6 |
+
|
| 7 |
+
### 全量 SFT 推荐参数
|
| 8 |
+
|
| 9 |
+
| 参数 | 推荐值 | 说明 |
|
| 10 |
+
|------|--------|------|
|
| 11 |
+
| `learning_rate` | 1e-5 | 全量微调,比 LoRA 小 10x |
|
| 12 |
+
| `weight_decay` | 1e-5 | 轻微权重衰减 |
|
| 13 |
+
| `warmup_ratio` | 0.05 | 5% 步数线性预热 |
|
| 14 |
+
| `lr_scheduler_type` | cosine | 余弦退火 |
|
| 15 |
+
| `adam_beta1` | 0.95 | Adam beta1 |
|
| 16 |
+
| `adam_beta2` | 0.999 | Adam beta2 |
|
| 17 |
+
|
| 18 |
+
### LoRA 推荐参数
|
| 19 |
+
|
| 20 |
+
| 参数 | 推荐值 |
|
| 21 |
+
|------|--------|
|
| 22 |
+
| `learning_rate` | 1e-4 |
|
| 23 |
+
| LoRA rank | 64 |
|
| 24 |
+
| LoRA alpha | 128 |
|
| 25 |
+
| LoRA target modules | q_proj, k_proj, v_proj, o_proj |
|
| 26 |
+
|
| 27 |
+
### 学习率调整策略
|
| 28 |
+
|
| 29 |
+
- **从零开始训练**: `learning_rate=1e-5`, `warmup_ratio=0.05`
|
| 30 |
+
- **从 checkpoint resume**: 学习率通常重置为初始值或减半
|
| 31 |
+
- **Loss 震荡**: 减小学习率至 5e-6
|
| 32 |
+
- **Loss 不降**: 检查数据预处理是否正确
|
| 33 |
+
|
| 34 |
+
---
|
| 35 |
+
|
| 36 |
+
## 2. Batch Size 和 Global Batch Size
|
| 37 |
+
|
| 38 |
+
```bash
|
| 39 |
+
# global_batch_size = per_device_train_batch_size * num_gpus * gradient_accumulation_steps
|
| 40 |
+
# 32 GPU 训练: 推荐 global_batch_size=128
|
| 41 |
+
global_batch_size=128
|
| 42 |
+
per_device_train_batch_size=1
|
| 43 |
+
# gradient_accumulation_steps = 128 / (1 * 32) = 4
|
| 44 |
+
```
|
| 45 |
+
|
| 46 |
+
---
|
| 47 |
+
|
| 48 |
+
## 3. 显存优化决策树
|
| 49 |
+
|
| 50 |
+
```
|
| 51 |
+
训练时 OOM?
|
| 52 |
+
├── 使用 8-bit Adam? → 未使用 → 加 optim=adamw_bnb_8bit
|
| 53 |
+
├── 开 gradient_checkpointing? → 未开 → 加 model.gradient_checkpointing=true
|
| 54 |
+
├── 可降低视频分辨率? → 是 → 降低 image_resolution_width/height
|
| 55 |
+
├── 可减少帧数? → 是 → 减少 num_frames
|
| 56 |
+
└── 以上都不行 → 使用 ZeRO-3 + CPU offload
|
| 57 |
+
```
|
| 58 |
+
|
| 59 |
+
### 各优化节约的显存
|
| 60 |
+
|
| 61 |
+
| 优化手段 | 显存节约 | 计算开销 |
|
| 62 |
+
|---------|---------|---------|
|
| 63 |
+
| 8-bit Adam | ~50% 优化器显存 | 无(更慢的更新) |
|
| 64 |
+
| Gradient Checkpointing | ~30% 激活显存 | ~15% 计算时间 |
|
| 65 |
+
| ZeRO-3 (vs ZeRO-2) | ~20% 总显存 | ~5% 通信开销 |
|
| 66 |
+
| CPU Offload | 额外 ~30% | ~20% 时间 |
|
| 67 |
+
|
| 68 |
+
---
|
| 69 |
+
|
| 70 |
+
## 4. Checkpoint 管理
|
| 71 |
+
|
| 72 |
+
### 保存
|
| 73 |
+
|
| 74 |
+
```bash
|
| 75 |
+
save_steps=2000 # 每 2000 步保存
|
| 76 |
+
save_total_limit=4 # 保留最近 4 个 + 最优
|
| 77 |
+
save_strategy=steps # 按步数保存
|
| 78 |
+
```
|
| 79 |
+
|
| 80 |
+
### Resume
|
| 81 |
+
|
| 82 |
+
```bash
|
| 83 |
+
# 从指定 checkpoint 恢复
|
| 84 |
+
torchrun ... experiment.py \
|
| 85 |
+
++training_args.resume_from_checkpoint=./output/libero_full/checkpoint-50000
|
| 86 |
+
```
|
| 87 |
+
|
| 88 |
+
**注意**: ZeRO-3 checkpoint 按 rank 分片保存,恢复时 GPU 数量必须一致。
|
| 89 |
+
|
| 90 |
+
### ZeRO-3 Checkpoint 合并
|
| 91 |
+
|
| 92 |
+
如果需要在不同 GPU 数之间迁移 checkpoint:
|
| 93 |
+
|
| 94 |
+
```bash
|
| 95 |
+
# 使用 DeepSpeed 的 zero_to_fp32.py
|
| 96 |
+
python zero_to_fp32.py \
|
| 97 |
+
--checkpoint_dir ./checkpoint-50000 \
|
| 98 |
+
--output_file ./pytorch_model.bin
|
| 99 |
+
```
|
| 100 |
+
|
| 101 |
+
---
|
| 102 |
+
|
| 103 |
+
## 5. 多节点训练注意事项
|
| 104 |
+
|
| 105 |
+
### NCCL 配置
|
| 106 |
+
|
| 107 |
+
```bash
|
| 108 |
+
export NCCL_DEBUG=WARN # 减少日志量
|
| 109 |
+
export NCCL_IB_DISABLE=0 # 启用 InfiniBand
|
| 110 |
+
export NCCL_IB_TIMEOUT=22 # IB 超时
|
| 111 |
+
export NCCL_IB_RETRY_CNT=4 # IB 重试次数
|
| 112 |
+
export NCCL_SOCKET_IFNAME=^docker0,lo # 排除 docker/loopback
|
| 113 |
+
```
|
| 114 |
+
|
| 115 |
+
### 常见问题
|
| 116 |
+
|
| 117 |
+
| 问题 | 原因 | 解决 |
|
| 118 |
+
|------|------|------|
|
| 119 |
+
| NCCL timeout | 网络配置错误 | 检查 NCCL_IB_*、NCCL_SOCKET_IFNAME |
|
| 120 |
+
| 训练挂起 | DataLoader worker 死锁 | 设置 `dataloader_pin_memory=false` |
|
| 121 |
+
| Loss 不一致 | 数据加载不统一 | 检查 `seed` 和 DataLoader shuffle |
|
| 122 |
+
| OOM | 显存不足 | 参考显存优化决策树 |
|
| 123 |
+
|
| 124 |
+
---
|
| 125 |
+
|
| 126 |
+
## 6. DataLoader 优化
|
| 127 |
+
|
| 128 |
+
| 场景 | num_workers | pin_memory | prefetch_factor |
|
| 129 |
+
|------|-------------|------------|-----------------|
|
| 130 |
+
| 单 GPU | 4 | false | 2 |
|
| 131 |
+
| 单节点 8 GPU | 8 | true | 4 |
|
| 132 |
+
| 多节点 | 8 | true | 4 |
|
| 133 |
+
|
| 134 |
+
**注意**: `dataloader_pin_memory=true` 在 decord 视频后端可能导致问题,如遇训练挂起可设为 false。
|
| 135 |
+
|
| 136 |
+
---
|
| 137 |
+
|
| 138 |
+
## 7. 实验跟踪(Wandb)
|
| 139 |
+
|
| 140 |
+
```bash
|
| 141 |
+
# 配置 Wandb
|
| 142 |
+
export WANDB_API_KEY=your_key
|
| 143 |
+
export WANDB_ENTITY=your_entity
|
| 144 |
+
|
| 145 |
+
# 启动训练
|
| 146 |
+
torchrun ... experiment.py \
|
| 147 |
+
report_to=wandb \
|
| 148 |
+
wandb_project=dreamzero-sft \
|
| 149 |
+
...
|
| 150 |
+
```
|
| 151 |
+
|
| 152 |
+
Wandb 记录内容:
|
| 153 |
+
- `loss`: 总损失
|
| 154 |
+
- `dynamics_loss_avg`: 视频预测损失
|
| 155 |
+
- `action_loss_avg`: 动作预测损失
|
| 156 |
+
- `learning_rate`: 学习率
|
| 157 |
+
- `grad_norm`: 梯度范数
|
| 158 |
+
- `training_step_time`: 每步时间
|
| 159 |
+
- `model_forward_time`: 前向时间
|
| 160 |
+
|
| 161 |
+
---
|
| 162 |
+
|
| 163 |
+
## 8. 从 LoRA 切换到全量 SFT
|
| 164 |
+
|
| 165 |
+
```bash
|
| 166 |
+
# LoRA
|
| 167 |
+
train_architecture=lora
|
| 168 |
+
optim=adamw_torch
|
| 169 |
+
learning_rate=1e-4
|
| 170 |
+
|
| 171 |
+
# 全量 SFT
|
| 172 |
+
train_architecture=full
|
| 173 |
+
optim=adamw_bnb_8bit # 必须!否则 5B 优化器状态 ~40GB
|
| 174 |
+
learning_rate=1e-5
|
| 175 |
+
```
|
| 176 |
+
|
| 177 |
+
---
|
| 178 |
+
|
| 179 |
+
## 9. 基准 Loss 参考
|
| 180 |
+
|
| 181 |
+
| Benchmark | Step 1 | Step 5 | Step 100 | Step 1k | Step 10k | Step 100k |
|
| 182 |
+
|-----------|--------|--------|----------|---------|----------|-----------|
|
| 183 |
+
| LIBERO | ~2-3 | ~1.0-1.5 | ~0.8-1.0 | ~0.6-0.8 | ~0.5-0.7 | ~0.35-0.5 |
|
| 184 |
+
| ManiFeel | ~1.5-2.5 | ~0.8-1.2 | ~0.6-0.8 | ~0.5-0.6 | ~0.35-0.5 | ~0.25-0.35 |
|
| 185 |
+
| RoboTwin | ~3-5 | ~1.5-2.0 | ~1.0-1.5 | ~0.8-1.2 | ~0.6-0.9 | ~0.4-0.6 |
|
| 186 |
+
|
| 187 |
+
Loss 持续高于参考范围的可能原因:
|
| 188 |
+
- 学习率不合适
|
| 189 |
+
- 数据预��理不一致(state/action 未正确 padding)
|
| 190 |
+
- 视频帧率与 state 频率不对齐
|
| 191 |
+
|
| 192 |
+
---
|
| 193 |
+
|
| 194 |
+
## 10. 故障排除
|
| 195 |
+
|
| 196 |
+
### "CUDA out of memory"
|
| 197 |
+
|
| 198 |
+
```bash
|
| 199 |
+
# 添加这些参数
|
| 200 |
+
per_device_train_batch_size=1
|
| 201 |
+
model.gradient_checkpointing=true
|
| 202 |
+
optim=adamw_bnb_8bit
|
| 203 |
+
training_args.deepspeed="groot/vla/configs/deepspeed/zero2_offload.json"
|
| 204 |
+
```
|
| 205 |
+
|
| 206 |
+
### "RuntimeError: Expected all tensors to be on the same device"
|
| 207 |
+
|
| 208 |
+
可能是 VAE 或模型设备不匹配。检查 `experiment.py` 中的设备分配。
|
| 209 |
+
|
| 210 |
+
### 训练突然变慢
|
| 211 |
+
|
| 212 |
+
- 检查 CPU 内存(dataloader_workers 可能导致内存交换)
|
| 213 |
+
- 检查磁盘 IO(数据是否在 SSD 上)
|
| 214 |
+
- 减少 dataloader_num_workers
|
| 215 |
+
|
| 216 |
+
### NCCL 初始化失败
|
| 217 |
+
|
| 218 |
+
```bash
|
| 219 |
+
export NCCL_DEBUG=INFO # 查看详细 NCCL 日志
|
| 220 |
+
export NCCL_IB_DISABLE=1 # 如果 IB 不可用
|
| 221 |
+
export NCCL_SOCKET_IFNAME=eth0 # 指定网络接口
|
| 222 |
+
```
|
eval_utils/policy_client.py
ADDED
|
@@ -0,0 +1,93 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""Client for communicating with a policy server.
|
| 2 |
+
|
| 3 |
+
Adapted from https://github.com/robo-arena/roboarena/
|
| 4 |
+
|
| 5 |
+
"""
|
| 6 |
+
|
| 7 |
+
import logging
|
| 8 |
+
import time
|
| 9 |
+
from typing import Dict, Tuple
|
| 10 |
+
|
| 11 |
+
import websockets.sync.client
|
| 12 |
+
from typing_extensions import override
|
| 13 |
+
|
| 14 |
+
from openpi_client.base_policy import BasePolicy
|
| 15 |
+
from openpi_client import msgpack_numpy
|
| 16 |
+
|
| 17 |
+
# The websockets library by default sends a ping every 20 seconds and
|
| 18 |
+
# expects a pong response within 20 seconds. However, the sever may not
|
| 19 |
+
# send a pong response immediately if it is busy processing a request.
|
| 20 |
+
# Increase the ping interval and timeout so that the client can wait
|
| 21 |
+
# for a longer time before closing the connection.
|
| 22 |
+
PING_INTERVAL_SECS = 60
|
| 23 |
+
PING_TIMEOUT_SECS = 600
|
| 24 |
+
|
| 25 |
+
class WebsocketClientPolicy(BasePolicy):
|
| 26 |
+
"""Implements the Policy interface by communicating with a server over websocket.
|
| 27 |
+
|
| 28 |
+
See WebsocketPolicyServer for a corresponding server implementation.
|
| 29 |
+
"""
|
| 30 |
+
|
| 31 |
+
def __init__(self, host: str = "0.0.0.0", port: int = 8000) -> None:
|
| 32 |
+
self._uri = f"ws://{host}:{port}"
|
| 33 |
+
self._packer = msgpack_numpy.Packer()
|
| 34 |
+
self._ws, self._server_metadata = self._wait_for_server()
|
| 35 |
+
|
| 36 |
+
def get_server_metadata(self) -> Dict:
|
| 37 |
+
return self._server_metadata
|
| 38 |
+
|
| 39 |
+
def _wait_for_server(self) -> Tuple[websockets.sync.client.ClientConnection, Dict]:
|
| 40 |
+
logging.info(f"Waiting for server at {self._uri}...")
|
| 41 |
+
try:
|
| 42 |
+
conn = websockets.sync.client.connect(
|
| 43 |
+
self._uri,
|
| 44 |
+
compression=None,
|
| 45 |
+
max_size=None,
|
| 46 |
+
ping_interval=PING_INTERVAL_SECS,
|
| 47 |
+
ping_timeout=PING_TIMEOUT_SECS,
|
| 48 |
+
)
|
| 49 |
+
metadata = msgpack_numpy.unpackb(conn.recv())
|
| 50 |
+
return conn, metadata
|
| 51 |
+
except:
|
| 52 |
+
logging.info("Connection to server with ws:// failed. Trying wss:// ...")
|
| 53 |
+
|
| 54 |
+
self._uri = "wss://" + self._uri.split("//")[1]
|
| 55 |
+
conn = websockets.sync.client.connect(
|
| 56 |
+
self._uri,
|
| 57 |
+
compression=None,
|
| 58 |
+
max_size=None,
|
| 59 |
+
ping_interval=PING_INTERVAL_SECS,
|
| 60 |
+
ping_timeout=PING_TIMEOUT_SECS,
|
| 61 |
+
)
|
| 62 |
+
metadata = msgpack_numpy.unpackb(conn.recv())
|
| 63 |
+
return conn, metadata
|
| 64 |
+
|
| 65 |
+
@override
|
| 66 |
+
def infer(self, obs: Dict) -> Dict: # noqa: UP006
|
| 67 |
+
# Notify server that we're calling the infer endpoint (as opposed to the reset endpoint)
|
| 68 |
+
obs["endpoint"] = "infer"
|
| 69 |
+
|
| 70 |
+
data = self._packer.pack(obs)
|
| 71 |
+
self._ws.send(data)
|
| 72 |
+
response = self._ws.recv()
|
| 73 |
+
if isinstance(response, str):
|
| 74 |
+
# we're expecting bytes; if the server sends a string, it's an error.
|
| 75 |
+
raise RuntimeError(f"Error in inference server:\n{response}")
|
| 76 |
+
return msgpack_numpy.unpackb(response)
|
| 77 |
+
|
| 78 |
+
@override
|
| 79 |
+
def reset(self, reset_info: Dict) -> None:
|
| 80 |
+
# Notify server that we're calling the reset endpoint (as opposed to the infer endpoint)
|
| 81 |
+
reset_info["endpoint"] = "reset"
|
| 82 |
+
|
| 83 |
+
data = self._packer.pack(reset_info)
|
| 84 |
+
self._ws.send(data)
|
| 85 |
+
response = self._ws.recv()
|
| 86 |
+
return response
|
| 87 |
+
|
| 88 |
+
if __name__ == "__main__":
|
| 89 |
+
logging.basicConfig(level=logging.INFO)
|
| 90 |
+
client = WebsocketClientPolicy()
|
| 91 |
+
actions = client.infer({})
|
| 92 |
+
print(f"Actions received: {actions}")
|
| 93 |
+
client.reset({})
|
eval_utils/policy_server.py
ADDED
|
@@ -0,0 +1,130 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""Server for serving a policy over websockets.
|
| 2 |
+
|
| 3 |
+
Adapted from https://github.com/robo-arena/roboarena/
|
| 4 |
+
|
| 5 |
+
"""
|
| 6 |
+
|
| 7 |
+
|
| 8 |
+
import asyncio
|
| 9 |
+
import dataclasses
|
| 10 |
+
import logging
|
| 11 |
+
import traceback
|
| 12 |
+
|
| 13 |
+
from openpi_client.base_policy import BasePolicy
|
| 14 |
+
from openpi_client import msgpack_numpy
|
| 15 |
+
import websockets.asyncio.server
|
| 16 |
+
import websockets.frames
|
| 17 |
+
|
| 18 |
+
|
| 19 |
+
@dataclasses.dataclass
|
| 20 |
+
class PolicyServerConfig:
|
| 21 |
+
# Resolution that images get resized to client-side, None means no resizing.
|
| 22 |
+
# It's beneficial to resize images to the desired resolution client-side for faster communication.
|
| 23 |
+
image_resolution: tuple[int, int] | None = (224, 224)
|
| 24 |
+
# Whether or not wrist camera image(s) should be sent.
|
| 25 |
+
needs_wrist_camera: bool = True
|
| 26 |
+
# Number of external cameras to send.
|
| 27 |
+
n_external_cameras: int = 1 # can be in [0, 1, 2]
|
| 28 |
+
# Whether or not stereo camera image(s) should be sent.
|
| 29 |
+
needs_stereo_camera: bool = False
|
| 30 |
+
# Whether or not the unique eval session id should be sent (e.g. for policies that want to keep track of history).
|
| 31 |
+
needs_session_id: bool = False
|
| 32 |
+
# Which action space to use.
|
| 33 |
+
action_space: str = "joint_position" # can be in ["joint_position", "joint_velocity", "cartesian_position", "cartesian_velocity"]
|
| 34 |
+
|
| 35 |
+
|
| 36 |
+
class WebsocketPolicyServer:
|
| 37 |
+
"""
|
| 38 |
+
Serves a policy using the websocket protocol.
|
| 39 |
+
|
| 40 |
+
Interface:
|
| 41 |
+
Observation:
|
| 42 |
+
- observation/wrist_image_left: (H, W, 3) if needs_wrist_camera is True
|
| 43 |
+
- observation/wrist_image_right: (H, W, 3) if needs_wrist_camera is True and needs_stereo_camera is True
|
| 44 |
+
- observation/exterior_image_{i}_left: (H, W, 3) if n_external_cameras >= 1
|
| 45 |
+
- observation/exterior_image_{i}_right: (H, W, 3) if needs_stereo_camera is True
|
| 46 |
+
- session_id: (1,) if needs_session_id is True
|
| 47 |
+
- observation/joint_position: (7,)
|
| 48 |
+
- observation/cartesian_position: (6,)
|
| 49 |
+
- observation/gripper_position: (1,)
|
| 50 |
+
- prompt: str, the natural language task instruction for the policy
|
| 51 |
+
|
| 52 |
+
Action:
|
| 53 |
+
- action: (N, 8,) or (N, 7,): either 7 movement actions (for joint action spaces) or 6 (for cartesian) plus one dimension for gripper position
|
| 54 |
+
--> all N actions will get executed on the robot before the server is queried again
|
| 55 |
+
|
| 56 |
+
"""
|
| 57 |
+
|
| 58 |
+
def __init__(
|
| 59 |
+
self,
|
| 60 |
+
policy: BasePolicy,
|
| 61 |
+
server_config: PolicyServerConfig,
|
| 62 |
+
host: str = "0.0.0.0",
|
| 63 |
+
port: int = 8000,
|
| 64 |
+
) -> None:
|
| 65 |
+
self._policy = policy
|
| 66 |
+
self._server_config = server_config
|
| 67 |
+
self._host = host
|
| 68 |
+
self._port = port
|
| 69 |
+
logging.getLogger("websockets.server").setLevel(logging.INFO)
|
| 70 |
+
|
| 71 |
+
def serve_forever(self) -> None:
|
| 72 |
+
asyncio.run(self.run())
|
| 73 |
+
|
| 74 |
+
async def run(self):
|
| 75 |
+
async with websockets.asyncio.server.serve(
|
| 76 |
+
self._handler,
|
| 77 |
+
self._host,
|
| 78 |
+
self._port,
|
| 79 |
+
compression=None,
|
| 80 |
+
max_size=None,
|
| 81 |
+
) as server:
|
| 82 |
+
await server.serve_forever()
|
| 83 |
+
|
| 84 |
+
async def _handler(self, websocket: websockets.asyncio.server.ServerConnection):
|
| 85 |
+
logging.info(f"Connection from {websocket.remote_address} opened")
|
| 86 |
+
packer = msgpack_numpy.Packer()
|
| 87 |
+
|
| 88 |
+
# Send server config to client to configure what gets sent to server.
|
| 89 |
+
await websocket.send(packer.pack(dataclasses.asdict(self._server_config)))
|
| 90 |
+
|
| 91 |
+
while True:
|
| 92 |
+
try:
|
| 93 |
+
obs = msgpack_numpy.unpackb(await websocket.recv())
|
| 94 |
+
|
| 95 |
+
endpoint = obs["endpoint"]
|
| 96 |
+
del obs["endpoint"]
|
| 97 |
+
if endpoint == "reset":
|
| 98 |
+
self._policy.reset(obs)
|
| 99 |
+
to_return = "reset successful"
|
| 100 |
+
else:
|
| 101 |
+
action = self._policy.infer(obs)
|
| 102 |
+
to_return = packer.pack(action)
|
| 103 |
+
await websocket.send(to_return)
|
| 104 |
+
except websockets.ConnectionClosed:
|
| 105 |
+
logging.info(f"Connection from {websocket.remote_address} closed")
|
| 106 |
+
break
|
| 107 |
+
except Exception:
|
| 108 |
+
await websocket.send(traceback.format_exc())
|
| 109 |
+
await websocket.close(
|
| 110 |
+
code=websockets.frames.CloseCode.INTERNAL_ERROR,
|
| 111 |
+
reason="Internal server error. Traceback included in previous frame.",
|
| 112 |
+
)
|
| 113 |
+
raise
|
| 114 |
+
|
| 115 |
+
|
| 116 |
+
if __name__ == "__main__":
|
| 117 |
+
import numpy as np
|
| 118 |
+
|
| 119 |
+
class DummyPolicy(BasePolicy):
|
| 120 |
+
def infer(self, obs):
|
| 121 |
+
return np.zeros((1, 8), dtype=np.float32)
|
| 122 |
+
|
| 123 |
+
def reset(self, reset_info):
|
| 124 |
+
pass
|
| 125 |
+
|
| 126 |
+
logging.basicConfig(level=logging.INFO)
|
| 127 |
+
policy = DummyPolicy()
|
| 128 |
+
server = WebsocketPolicyServer(policy, PolicyServerConfig())
|
| 129 |
+
server.serve_forever()
|
| 130 |
+
|
eval_utils/run_sim_eval.py
ADDED
|
@@ -0,0 +1,219 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""
|
| 2 |
+
Example script for running 10 rollouts of a DROID policy on the example environment.
|
| 3 |
+
|
| 4 |
+
Usage:
|
| 5 |
+
|
| 6 |
+
First, make sure you download the simulation assets and unpack them into the root directory of this package.
|
| 7 |
+
|
| 8 |
+
Then, in a separate terminal, launch the policy server on localhost:8000
|
| 9 |
+
-- make sure to set XLA_PYTHON_CLIENT_MEM_FRACTION to avoid JAX hogging all the GPU memory.
|
| 10 |
+
|
| 11 |
+
For example, to launch a pi0-FAST-DROID policy (with joint position control),
|
| 12 |
+
run the command below in a separate terminal from the openpi "karl/droid_policies" branch:
|
| 13 |
+
|
| 14 |
+
XLA_PYTHON_CLIENT_MEM_FRACTION=0.5 uv run scripts/serve_policy.py policy:checkpoint --policy.config=pi0_fast_droid_jointpos --policy.dir=s3://openpi-assets-simeval/pi0_fast_droid_jointpos
|
| 15 |
+
|
| 16 |
+
Finally, run the evaluation script:
|
| 17 |
+
|
| 18 |
+
python run_eval.py --episodes 10 --headless
|
| 19 |
+
"""
|
| 20 |
+
|
| 21 |
+
import uuid
|
| 22 |
+
|
| 23 |
+
import tyro
|
| 24 |
+
import argparse
|
| 25 |
+
import gymnasium as gym
|
| 26 |
+
import torch
|
| 27 |
+
import cv2
|
| 28 |
+
import mediapy
|
| 29 |
+
import numpy as np
|
| 30 |
+
from datetime import datetime
|
| 31 |
+
from pathlib import Path
|
| 32 |
+
from PIL import Image
|
| 33 |
+
from tqdm import tqdm
|
| 34 |
+
|
| 35 |
+
from openpi_client import image_tools
|
| 36 |
+
from sim_evals.inference.abstract_client import InferenceClient
|
| 37 |
+
from policy_client import WebsocketClientPolicy
|
| 38 |
+
|
| 39 |
+
|
| 40 |
+
class DreamZeroJointPosClient(InferenceClient):
|
| 41 |
+
def __init__(self,
|
| 42 |
+
remote_host:str = "localhost",
|
| 43 |
+
remote_port:int = 6000,
|
| 44 |
+
open_loop_horizon:int = 8,
|
| 45 |
+
) -> None:
|
| 46 |
+
self.client = WebsocketClientPolicy(remote_host, remote_port)
|
| 47 |
+
self.open_loop_horizon = open_loop_horizon
|
| 48 |
+
self.actions_from_chunk_completed = 0
|
| 49 |
+
self.pred_action_chunk = None
|
| 50 |
+
self.session_id = str(uuid.uuid4())
|
| 51 |
+
|
| 52 |
+
def visualize(self, request: dict):
|
| 53 |
+
"""
|
| 54 |
+
Return the camera views how the model sees it
|
| 55 |
+
"""
|
| 56 |
+
curr_obs = self._extract_observation(request)
|
| 57 |
+
right_img = image_tools.resize_with_pad(curr_obs["right_image"], 224, 224)
|
| 58 |
+
wrist_img = image_tools.resize_with_pad(curr_obs["wrist_image"], 224, 224)
|
| 59 |
+
left_img = image_tools.resize_with_pad(curr_obs["left_image"], 224, 224)
|
| 60 |
+
combined = np.concatenate([right_img, wrist_img, left_img], axis=1)
|
| 61 |
+
return combined
|
| 62 |
+
|
| 63 |
+
def reset(self):
|
| 64 |
+
self.actions_from_chunk_completed = 0
|
| 65 |
+
self.pred_action_chunk = None
|
| 66 |
+
self.session_id = str(uuid.uuid4())
|
| 67 |
+
|
| 68 |
+
def infer(self, obs: dict, instruction: str) -> dict:
|
| 69 |
+
"""
|
| 70 |
+
Infer the next action from the policy in a server-client setup
|
| 71 |
+
"""
|
| 72 |
+
curr_obs = self._extract_observation(obs)
|
| 73 |
+
if (
|
| 74 |
+
self.actions_from_chunk_completed == 0
|
| 75 |
+
or self.actions_from_chunk_completed >= self.open_loop_horizon
|
| 76 |
+
):
|
| 77 |
+
self.actions_from_chunk_completed = 0
|
| 78 |
+
request_data = {
|
| 79 |
+
"observation/exterior_image_0_left": image_tools.resize_with_pad(curr_obs["right_image"], 180, 320),
|
| 80 |
+
"observation/exterior_image_1_left": image_tools.resize_with_pad(curr_obs["left_image"], 180, 320),
|
| 81 |
+
"observation/wrist_image_left": image_tools.resize_with_pad(curr_obs["wrist_image"], 180, 320),
|
| 82 |
+
"observation/joint_position": curr_obs["joint_position"].astype(np.float64),
|
| 83 |
+
"observation/cartesian_position": np.zeros((6,), dtype=np.float64), # dummy cartesian position
|
| 84 |
+
"observation/gripper_position": curr_obs["gripper_position"].astype(np.float64),
|
| 85 |
+
"prompt": instruction,
|
| 86 |
+
"session_id": self.session_id,
|
| 87 |
+
}
|
| 88 |
+
for k, v in request_data.items():
|
| 89 |
+
print(f"{k}: {v.shape if not isinstance(v, str) else v}")
|
| 90 |
+
|
| 91 |
+
result = self.client.infer(request_data)
|
| 92 |
+
actions = result["actions"] if isinstance(result, dict) else result
|
| 93 |
+
assert len(actions.shape) == 2, f"Expected 2D array, got shape {actions.shape}"
|
| 94 |
+
assert actions.shape[-1] == 8, f"Expected 8 action dimensions (7 joints + 1 gripper), got {actions.shape[-1]}"
|
| 95 |
+
self.pred_action_chunk = actions
|
| 96 |
+
|
| 97 |
+
|
| 98 |
+
action = self.pred_action_chunk[self.actions_from_chunk_completed]
|
| 99 |
+
self.actions_from_chunk_completed += 1
|
| 100 |
+
|
| 101 |
+
# binarize gripper action
|
| 102 |
+
if action[-1].item() > 0.5:
|
| 103 |
+
action = np.concatenate([action[:-1], np.ones((1,))])
|
| 104 |
+
else:
|
| 105 |
+
action = np.concatenate([action[:-1], np.zeros((1,))])
|
| 106 |
+
|
| 107 |
+
img1 = image_tools.resize_with_pad(curr_obs["right_image"], 224, 224)
|
| 108 |
+
img2 = image_tools.resize_with_pad(curr_obs["wrist_image"], 224, 224)
|
| 109 |
+
img3 = image_tools.resize_with_pad(curr_obs["left_image"], 224, 224)
|
| 110 |
+
both = np.concatenate([img1, img2, img3], axis=1)
|
| 111 |
+
|
| 112 |
+
return {"action": action, "viz": both}
|
| 113 |
+
|
| 114 |
+
def _extract_observation(self, obs_dict, *, save_to_disk=False):
|
| 115 |
+
# Assign images
|
| 116 |
+
right_image = obs_dict["policy"]["external_cam"][0].clone().detach().cpu().numpy()
|
| 117 |
+
left_image = obs_dict["policy"]["external_cam_2"][0].clone().detach().cpu().numpy()
|
| 118 |
+
wrist_image = obs_dict["policy"]["wrist_cam"][0].clone().detach().cpu().numpy()
|
| 119 |
+
|
| 120 |
+
# Capture proprioceptive state
|
| 121 |
+
robot_state = obs_dict["policy"]
|
| 122 |
+
joint_position = robot_state["arm_joint_pos"].clone().detach().cpu().numpy()
|
| 123 |
+
gripper_position = robot_state["gripper_pos"].clone().detach().cpu().numpy()
|
| 124 |
+
|
| 125 |
+
if save_to_disk:
|
| 126 |
+
combined_image = np.concatenate([right_image, wrist_image], axis=1)
|
| 127 |
+
combined_image = Image.fromarray(combined_image)
|
| 128 |
+
combined_image.save("robot_camera_views.png")
|
| 129 |
+
|
| 130 |
+
return {
|
| 131 |
+
"right_image": right_image,
|
| 132 |
+
"left_image": left_image,
|
| 133 |
+
"wrist_image": wrist_image,
|
| 134 |
+
"joint_position": joint_position,
|
| 135 |
+
"gripper_position": gripper_position,
|
| 136 |
+
}
|
| 137 |
+
|
| 138 |
+
|
| 139 |
+
|
| 140 |
+
|
| 141 |
+
def main(
|
| 142 |
+
episodes: int = 10,
|
| 143 |
+
scene: int = 1,
|
| 144 |
+
headless: bool = True,
|
| 145 |
+
host: str = "localhost",
|
| 146 |
+
port: int = 6000,
|
| 147 |
+
):
|
| 148 |
+
# launch omniverse app with arguments (inside function to prevent overriding tyro)
|
| 149 |
+
from isaaclab.app import AppLauncher
|
| 150 |
+
parser = argparse.ArgumentParser(description="Tutorial on creating an empty stage.")
|
| 151 |
+
AppLauncher.add_app_launcher_args(parser)
|
| 152 |
+
args_cli, _ = parser.parse_known_args()
|
| 153 |
+
args_cli.enable_cameras = True
|
| 154 |
+
args_cli.headless = headless
|
| 155 |
+
app_launcher = AppLauncher(args_cli)
|
| 156 |
+
simulation_app = app_launcher.app
|
| 157 |
+
|
| 158 |
+
# All IsaacLab dependent modules should be imported after the app is launched
|
| 159 |
+
import sim_evals.environments # noqa: F401
|
| 160 |
+
from isaaclab_tasks.utils import parse_env_cfg
|
| 161 |
+
|
| 162 |
+
|
| 163 |
+
# Initialize the env
|
| 164 |
+
env_cfg = parse_env_cfg(
|
| 165 |
+
"DROID",
|
| 166 |
+
device=args_cli.device,
|
| 167 |
+
num_envs=1,
|
| 168 |
+
use_fabric=True,
|
| 169 |
+
)
|
| 170 |
+
instruction = None
|
| 171 |
+
match scene:
|
| 172 |
+
case 1:
|
| 173 |
+
instruction = "put the cube in the bowl"
|
| 174 |
+
case 2:
|
| 175 |
+
instruction = "pick up the can and put it in the mug"
|
| 176 |
+
case 3:
|
| 177 |
+
instruction = "put the banana in the bin"
|
| 178 |
+
case _:
|
| 179 |
+
raise ValueError(f"Scene {scene} not supported")
|
| 180 |
+
|
| 181 |
+
env_cfg.set_scene(scene)
|
| 182 |
+
env = gym.make("DROID", cfg=env_cfg)
|
| 183 |
+
|
| 184 |
+
obs, _ = env.reset()
|
| 185 |
+
obs, _ = env.reset() # need second render cycle to get correctly loaded materials
|
| 186 |
+
client = DreamZeroJointPosClient(remote_host=host, remote_port=port)
|
| 187 |
+
|
| 188 |
+
|
| 189 |
+
video_dir = Path("runs") / datetime.now().strftime("%Y-%m-%d") / datetime.now().strftime("%H-%M-%S")
|
| 190 |
+
video_dir.mkdir(parents=True, exist_ok=True)
|
| 191 |
+
video = []
|
| 192 |
+
ep = 0
|
| 193 |
+
max_steps = env.env.max_episode_length
|
| 194 |
+
with torch.no_grad():
|
| 195 |
+
for ep in range(episodes):
|
| 196 |
+
for _ in tqdm(range(max_steps), desc=f"Episode {ep+1}/{episodes}"):
|
| 197 |
+
ret = client.infer(obs, instruction)
|
| 198 |
+
if not headless:
|
| 199 |
+
cv2.imshow("Right Camera", cv2.cvtColor(ret["viz"], cv2.COLOR_RGB2BGR))
|
| 200 |
+
cv2.waitKey(1)
|
| 201 |
+
video.append(ret["viz"])
|
| 202 |
+
action = torch.tensor(ret["action"])[None]
|
| 203 |
+
obs, _, term, trunc, _ = env.step(action)
|
| 204 |
+
if term or trunc:
|
| 205 |
+
break
|
| 206 |
+
|
| 207 |
+
client.reset()
|
| 208 |
+
mediapy.write_video(
|
| 209 |
+
video_dir / f"episode_{ep}.mp4",
|
| 210 |
+
video,
|
| 211 |
+
fps=15,
|
| 212 |
+
)
|
| 213 |
+
video = []
|
| 214 |
+
|
| 215 |
+
env.close()
|
| 216 |
+
simulation_app.close()
|
| 217 |
+
|
| 218 |
+
if __name__ == "__main__":
|
| 219 |
+
args = tyro.cli(main)
|
eval_utils/serve_dreamzero_wan22.py
ADDED
|
@@ -0,0 +1,391 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""
|
| 2 |
+
Serve the DreamZero 5B implementation (Wan2.2-TI2V-5B) over the websocket policy server.
|
| 3 |
+
|
| 4 |
+
This is the 5B model: Wan2.2 diffusion backbone, 48-channel VAE38, frame_seqlen=50 (160×320
|
| 5 |
+
latent 10×20). Inference is causal with KV caching: first request in a session uses 1 frame
|
| 6 |
+
and warms the cache; subsequent requests use FRAMES_PER_CHUNK=4 frames and append to the cache.
|
| 7 |
+
On session_id change (or explicit reset), buffers and action_head.current_start_frame are cleared.
|
| 8 |
+
|
| 9 |
+
The checkpoint at model_path should be DreamZero with Wan22 5B (model/dreamzero/action_head=
|
| 10 |
+
wan_flow_matching_action_tf_wan22, data droid_relative_wan22 → 160×320). GrootSimPolicy loads
|
| 11 |
+
that checkpoint and runs inference; it is the correct policy class for DreamZero.
|
| 12 |
+
|
| 13 |
+
Usage (single GPU):
|
| 14 |
+
|
| 15 |
+
torchrun --nproc_per_node=1 eval_utils/serve_dreamzero_wan22.py --model_path ./checkpoints/dreamzero_droid_wan22_smoke --port 8000
|
| 16 |
+
|
| 17 |
+
# Or single process:
|
| 18 |
+
python eval_utils/serve_dreamzero_wan22.py --model_path ./checkpoints/dreamzero_droid_wan22_smoke --port 8000
|
| 19 |
+
|
| 20 |
+
Client: send observations per PolicyServerConfig (policy_server.py). Video is resized to the
|
| 21 |
+
checkpoint's expected resolution (e.g. 180×320) so the eval transform accepts it; the 5B action
|
| 22 |
+
head resizes to 160×320 internally. Override with --image_height/--image_width if needed.
|
| 23 |
+
Response is an action chunk (N, 8). Use session_id for episode boundaries.
|
| 24 |
+
"""
|
| 25 |
+
|
| 26 |
+
import datetime
|
| 27 |
+
import logging
|
| 28 |
+
import os
|
| 29 |
+
import sys
|
| 30 |
+
|
| 31 |
+
import imageio
|
| 32 |
+
|
| 33 |
+
logger = logging.getLogger(__name__)
|
| 34 |
+
|
| 35 |
+
import cv2
|
| 36 |
+
import numpy as np
|
| 37 |
+
import torch
|
| 38 |
+
import torch.distributed as dist
|
| 39 |
+
from torch.distributed.device_mesh import init_device_mesh
|
| 40 |
+
import tyro
|
| 41 |
+
|
| 42 |
+
# Avoid FailOnRecompileLimitHit when serving: the flow scheduler's torch.compile'd
|
| 43 |
+
# multistep_uni_p_bh_update recompiles under varying shapes/inputs (e.g. batch size,
|
| 44 |
+
# step_index, order). Increase limits so the server doesn't hit the default cap.
|
| 45 |
+
_dynamo = torch._dynamo.config
|
| 46 |
+
if hasattr(_dynamo, "cache_size_limit"):
|
| 47 |
+
_dynamo.cache_size_limit = 1000
|
| 48 |
+
if hasattr(_dynamo, "recompile_limit"):
|
| 49 |
+
_dynamo.recompile_limit = 800
|
| 50 |
+
if hasattr(_dynamo, "accumulated_cache_size_limit"):
|
| 51 |
+
_dynamo.accumulated_cache_size_limit = 1000
|
| 52 |
+
if hasattr(_dynamo, "accumulated_recompile_limit"):
|
| 53 |
+
_dynamo.accumulated_recompile_limit = 2000
|
| 54 |
+
from pathlib import Path
|
| 55 |
+
from tianshou.data import Batch
|
| 56 |
+
|
| 57 |
+
# Add repo root for imports
|
| 58 |
+
REPO_ROOT = Path(__file__).resolve().parents[1]
|
| 59 |
+
if str(REPO_ROOT) not in sys.path:
|
| 60 |
+
sys.path.insert(0, str(REPO_ROOT))
|
| 61 |
+
|
| 62 |
+
from openpi_client.base_policy import BasePolicy
|
| 63 |
+
|
| 64 |
+
from eval_utils.policy_server import WebsocketPolicyServer, PolicyServerConfig
|
| 65 |
+
from groot.vla.model.n1_5.sim_policy import GrootSimPolicy
|
| 66 |
+
from groot.vla.data.schema import EmbodimentTag
|
| 67 |
+
from groot.vla.data.transform import ComposedModalityTransform
|
| 68 |
+
|
| 69 |
+
|
| 70 |
+
# DreamZero Wan 5B is trained with 160×320 (droid_relative_wan22). Fallback if we cannot read from policy.
|
| 71 |
+
DEFAULT_IMAGE_HEIGHT = 160
|
| 72 |
+
DEFAULT_IMAGE_WIDTH = 320
|
| 73 |
+
FRAMES_PER_CHUNK = 4 # matches 5B num_frame_per_block for causal chunked inference
|
| 74 |
+
|
| 75 |
+
|
| 76 |
+
def _get_expected_video_resolution(policy: GrootSimPolicy) -> tuple[int, int]:
|
| 77 |
+
"""Get (height, width) the policy's eval_transform expects for video (from checkpoint
|
| 78 |
+
metadata). Resolution in metadata is (width, height); we return (height, width) for resize.
|
| 79 |
+
DreamZero Wan 5B (droid_relative_wan22) uses 160×320; other configs may use e.g. 180×320.
|
| 80 |
+
"""
|
| 81 |
+
eval_transform = getattr(policy, "eval_transform", None)
|
| 82 |
+
if eval_transform is None:
|
| 83 |
+
return (DEFAULT_IMAGE_HEIGHT, DEFAULT_IMAGE_WIDTH)
|
| 84 |
+
if not isinstance(eval_transform, ComposedModalityTransform):
|
| 85 |
+
return (DEFAULT_IMAGE_HEIGHT, DEFAULT_IMAGE_WIDTH)
|
| 86 |
+
for t in eval_transform.transforms:
|
| 87 |
+
if hasattr(t, "original_resolutions") and getattr(t, "original_resolutions", None):
|
| 88 |
+
res = t.original_resolutions
|
| 89 |
+
if res:
|
| 90 |
+
# original_resolutions values are (width, height)
|
| 91 |
+
w, h = next(iter(res.values()))
|
| 92 |
+
return (int(h), int(w))
|
| 93 |
+
return (DEFAULT_IMAGE_HEIGHT, DEFAULT_IMAGE_WIDTH)
|
| 94 |
+
|
| 95 |
+
|
| 96 |
+
def _resize_frames_to_resolution(frames: np.ndarray, target_h: int, target_w: int) -> np.ndarray:
|
| 97 |
+
"""Resize video frames to (target_h, target_w). Accepts (H,W,C) or (T,H,W,C)."""
|
| 98 |
+
if frames.ndim == 3:
|
| 99 |
+
if (frames.shape[0], frames.shape[1]) != (target_h, target_w):
|
| 100 |
+
frames = cv2.resize(frames, (target_w, target_h), interpolation=cv2.INTER_LINEAR)
|
| 101 |
+
return frames
|
| 102 |
+
out = np.stack(
|
| 103 |
+
[cv2.resize(f, (target_w, target_h), interpolation=cv2.INTER_LINEAR) for f in frames],
|
| 104 |
+
axis=0,
|
| 105 |
+
)
|
| 106 |
+
return out
|
| 107 |
+
|
| 108 |
+
|
| 109 |
+
def _maybe_init_distributed():
|
| 110 |
+
"""Initialize process group for single-GPU or multi-GPU. Required by GrootSimPolicy."""
|
| 111 |
+
if dist.is_initialized():
|
| 112 |
+
return
|
| 113 |
+
os.environ.setdefault("MASTER_ADDR", "localhost")
|
| 114 |
+
os.environ.setdefault("MASTER_PORT", "29500")
|
| 115 |
+
dist.init_process_group(backend="nccl", rank=0, world_size=1)
|
| 116 |
+
torch.cuda.set_device(0)
|
| 117 |
+
|
| 118 |
+
|
| 119 |
+
# Modality key mappings: client observation keys -> model input keys per embodiment.
|
| 120 |
+
# Client sends: observation/exterior_image_0_left, exterior_image_1_left, wrist_image_left.
|
| 121 |
+
VIDEO_KEY_MAPPING = {
|
| 122 |
+
"oxe_droid": {
|
| 123 |
+
"observation/exterior_image_0_left": "video.exterior_image_1_left",
|
| 124 |
+
"observation/exterior_image_1_left": "video.exterior_image_2_left",
|
| 125 |
+
"observation/wrist_image_left": "video.wrist_image_left",
|
| 126 |
+
},
|
| 127 |
+
}
|
| 128 |
+
STATE_KEY_MAPPING = {
|
| 129 |
+
"oxe_droid": ("state.joint_position", "state.gripper_position"),
|
| 130 |
+
}
|
| 131 |
+
LANGUAGE_KEY_MAPPING = {
|
| 132 |
+
"oxe_droid": "annotation.language.action_text",
|
| 133 |
+
}
|
| 134 |
+
|
| 135 |
+
|
| 136 |
+
class DreamZeroWan225BPolicy(BasePolicy):
|
| 137 |
+
"""
|
| 138 |
+
Wraps GrootSimPolicy for the DreamZero 5B implementation (Wan2.2-TI2V-5B).
|
| 139 |
+
|
| 140 |
+
Converts roboarena observation/action format to DROID/Batch. Video is resized to the
|
| 141 |
+
resolution expected by the policy's eval_transform (from checkpoint metadata) so
|
| 142 |
+
VideoToTensor validation passes. The 5B action head then resizes to 160×320 internally.
|
| 143 |
+
First call in a session uses 1 frame; later calls use 4 frames (FRAMES_PER_CHUNK).
|
| 144 |
+
Session reset clears frame buffers and action_head.current_start_frame.
|
| 145 |
+
"""
|
| 146 |
+
|
| 147 |
+
def __init__(
|
| 148 |
+
self,
|
| 149 |
+
groot_policy: GrootSimPolicy,
|
| 150 |
+
image_height: int,
|
| 151 |
+
image_width: int,
|
| 152 |
+
embodiment_tag: str = "oxe_droid",
|
| 153 |
+
save_video_pred: bool = False,
|
| 154 |
+
video_output_dir: str = "./video_pred_output",
|
| 155 |
+
):
|
| 156 |
+
super().__init__()
|
| 157 |
+
self._policy = groot_policy
|
| 158 |
+
self._image_height = image_height
|
| 159 |
+
self._image_width = image_width
|
| 160 |
+
self._embodiment_tag = (
|
| 161 |
+
embodiment_tag if embodiment_tag in VIDEO_KEY_MAPPING else "oxe_droid"
|
| 162 |
+
)
|
| 163 |
+
video_keys = list(VIDEO_KEY_MAPPING[self._embodiment_tag].values())
|
| 164 |
+
self._frame_buffers = {k: [] for k in video_keys}
|
| 165 |
+
self._is_first_call = True
|
| 166 |
+
self._current_session_id = None
|
| 167 |
+
self._save_video_pred = save_video_pred
|
| 168 |
+
self._video_output_dir = video_output_dir
|
| 169 |
+
self._video_pred_latents: list[torch.Tensor] = []
|
| 170 |
+
self._current_prompt: str = ""
|
| 171 |
+
|
| 172 |
+
def _convert_observation(self, obs: dict) -> dict:
|
| 173 |
+
"""Convert roboarena observation format to model Batch format.
|
| 174 |
+
Incoming frames are resized to the policy's expected (height, width) so
|
| 175 |
+
eval_transform's VideoToTensor check passes.
|
| 176 |
+
"""
|
| 177 |
+
image_key_mapping = VIDEO_KEY_MAPPING[self._embodiment_tag]
|
| 178 |
+
for roboarena_key, model_key in image_key_mapping.items():
|
| 179 |
+
if roboarena_key in obs:
|
| 180 |
+
data = obs[roboarena_key]
|
| 181 |
+
if isinstance(data, np.ndarray):
|
| 182 |
+
data = _resize_frames_to_resolution(
|
| 183 |
+
data, self._image_height, self._image_width
|
| 184 |
+
)
|
| 185 |
+
if data.ndim == 4:
|
| 186 |
+
self._frame_buffers[model_key].extend(list(data))
|
| 187 |
+
else:
|
| 188 |
+
self._frame_buffers[model_key].append(data)
|
| 189 |
+
|
| 190 |
+
num_frames = 1 if self._is_first_call else FRAMES_PER_CHUNK
|
| 191 |
+
converted = {}
|
| 192 |
+
for model_key, buffer in self._frame_buffers.items():
|
| 193 |
+
if len(buffer) > 0:
|
| 194 |
+
if len(buffer) >= num_frames:
|
| 195 |
+
frames_to_use = buffer[-num_frames:]
|
| 196 |
+
else:
|
| 197 |
+
frames_to_use = buffer.copy()
|
| 198 |
+
while len(frames_to_use) < num_frames:
|
| 199 |
+
frames_to_use.insert(0, buffer[0])
|
| 200 |
+
video = np.stack(frames_to_use, axis=0)
|
| 201 |
+
converted[model_key] = video
|
| 202 |
+
|
| 203 |
+
state_joint_key, state_gripper_key = STATE_KEY_MAPPING[self._embodiment_tag]
|
| 204 |
+
if "observation/joint_position" in obs:
|
| 205 |
+
joint_pos = np.asarray(obs["observation/joint_position"])
|
| 206 |
+
if joint_pos.ndim == 1:
|
| 207 |
+
joint_pos = joint_pos.reshape(1, -1)
|
| 208 |
+
converted[state_joint_key] = joint_pos.astype(np.float64)
|
| 209 |
+
else:
|
| 210 |
+
converted[state_joint_key] = np.zeros((1, 7), dtype=np.float64)
|
| 211 |
+
|
| 212 |
+
if "observation/gripper_position" in obs:
|
| 213 |
+
gripper_pos = np.asarray(obs["observation/gripper_position"])
|
| 214 |
+
if gripper_pos.ndim == 1:
|
| 215 |
+
gripper_pos = gripper_pos.reshape(1, -1)
|
| 216 |
+
converted[state_gripper_key] = gripper_pos.astype(np.float64)
|
| 217 |
+
else:
|
| 218 |
+
converted[state_gripper_key] = np.zeros((1,1), dtype=np.float64)
|
| 219 |
+
|
| 220 |
+
text_prompt = obs.get("prompt", "")
|
| 221 |
+
logger.info("Text prompt: %s", text_prompt)
|
| 222 |
+
if text_prompt:
|
| 223 |
+
self._current_prompt = text_prompt
|
| 224 |
+
lang_key = LANGUAGE_KEY_MAPPING[self._embodiment_tag]
|
| 225 |
+
converted[lang_key] = text_prompt
|
| 226 |
+
return converted
|
| 227 |
+
|
| 228 |
+
def _convert_action(self, action_dict: dict) -> np.ndarray:
|
| 229 |
+
"""Convert model action dict to (N, 8) array (7 joint + 1 gripper)."""
|
| 230 |
+
joint_action = None
|
| 231 |
+
gripper_action = None
|
| 232 |
+
for key, value in action_dict.items():
|
| 233 |
+
if ("joint_position" in key or "joint_pos" in key) and "gripper" not in key:
|
| 234 |
+
joint_action = value
|
| 235 |
+
elif "gripper_position" in key or "gripper" in key:
|
| 236 |
+
gripper_action = value
|
| 237 |
+
if joint_action is None:
|
| 238 |
+
return np.zeros((1, 8), dtype=np.float32)
|
| 239 |
+
if isinstance(joint_action, torch.Tensor):
|
| 240 |
+
joint_action = joint_action.cpu().numpy()
|
| 241 |
+
if joint_action.ndim == 1:
|
| 242 |
+
joint_action = joint_action.reshape(1, -1)
|
| 243 |
+
N = joint_action.shape[0]
|
| 244 |
+
if gripper_action is not None:
|
| 245 |
+
if isinstance(gripper_action, torch.Tensor):
|
| 246 |
+
gripper_action = gripper_action.cpu().numpy()
|
| 247 |
+
if gripper_action.ndim == 1:
|
| 248 |
+
gripper_action = gripper_action.reshape(-1, 1)
|
| 249 |
+
if gripper_action.shape[-1] > 1:
|
| 250 |
+
gripper_action = gripper_action[..., :1]
|
| 251 |
+
else:
|
| 252 |
+
gripper_action = np.zeros((N, 1), dtype=np.float32)
|
| 253 |
+
return np.concatenate([joint_action, gripper_action], axis=-1).astype(np.float32)
|
| 254 |
+
|
| 255 |
+
def infer(self, obs: dict) -> np.ndarray:
|
| 256 |
+
session_id = obs.get("session_id")
|
| 257 |
+
if session_id is not None and session_id != self._current_session_id:
|
| 258 |
+
if self._current_session_id is not None:
|
| 259 |
+
self.reset({})
|
| 260 |
+
self._current_session_id = session_id
|
| 261 |
+
|
| 262 |
+
converted_obs = self._convert_observation(obs)
|
| 263 |
+
batch = Batch(obs=converted_obs)
|
| 264 |
+
with torch.no_grad():
|
| 265 |
+
result_batch, video_pred = self._policy.lazy_joint_forward_causal(batch)
|
| 266 |
+
if self._save_video_pred and video_pred is not None:
|
| 267 |
+
self._video_pred_latents.append(video_pred.detach())
|
| 268 |
+
action_dict = {}
|
| 269 |
+
action_chunk_dict = result_batch.act
|
| 270 |
+
for k in dir(action_chunk_dict):
|
| 271 |
+
if k.startswith("action."):
|
| 272 |
+
action_dict[k] = getattr(action_chunk_dict, k)
|
| 273 |
+
action = self._convert_action(action_dict)
|
| 274 |
+
if self._is_first_call:
|
| 275 |
+
self._is_first_call = False
|
| 276 |
+
return action
|
| 277 |
+
|
| 278 |
+
def _save_predicted_video(self) -> None:
|
| 279 |
+
"""Decode accumulated video prediction latents through the VAE and save as mp4."""
|
| 280 |
+
if not self._video_pred_latents:
|
| 281 |
+
return
|
| 282 |
+
try:
|
| 283 |
+
from einops import rearrange
|
| 284 |
+
|
| 285 |
+
action_head = self._policy.trained_model.action_head
|
| 286 |
+
latents = torch.cat(self._video_pred_latents, dim=2)
|
| 287 |
+
with torch.no_grad():
|
| 288 |
+
frames = action_head.vae.decode(
|
| 289 |
+
latents,
|
| 290 |
+
tiled=action_head.tiled,
|
| 291 |
+
tile_size=(action_head.tile_size_height, action_head.tile_size_width),
|
| 292 |
+
tile_stride=(action_head.tile_stride_height, action_head.tile_stride_width),
|
| 293 |
+
)
|
| 294 |
+
frames = rearrange(frames, "B C T H W -> B T H W C")[0]
|
| 295 |
+
frames = ((frames.float() + 1) * 127.5).clip(0, 255).cpu().numpy().astype(np.uint8)
|
| 296 |
+
|
| 297 |
+
os.makedirs(self._video_output_dir, exist_ok=True)
|
| 298 |
+
timestamp = datetime.datetime.now().strftime("%m_%d_%H_%M_%S")
|
| 299 |
+
n_latent_frames = latents.shape[2]
|
| 300 |
+
existing = [f for f in os.listdir(self._video_output_dir) if f.endswith(".mp4")]
|
| 301 |
+
safe_prompt = self._current_prompt.replace(" ", "_")
|
| 302 |
+
safe_prompt = "".join(c for c in safe_prompt if c.isalnum() or c in "_-.")
|
| 303 |
+
if len(safe_prompt) > 80:
|
| 304 |
+
safe_prompt = safe_prompt[:80]
|
| 305 |
+
if not safe_prompt:
|
| 306 |
+
safe_prompt = "no_prompt"
|
| 307 |
+
output_path = os.path.join(
|
| 308 |
+
self._video_output_dir,
|
| 309 |
+
f"{len(existing):06}_{safe_prompt}_{timestamp}.mp4",
|
| 310 |
+
)
|
| 311 |
+
imageio.mimsave(output_path, list(frames), fps=5, codec="libx264")
|
| 312 |
+
logger.info("Saved video prediction (%d frames) to %s", len(frames), output_path)
|
| 313 |
+
except Exception as e:
|
| 314 |
+
logger.warning("Failed to save video prediction: %s", e)
|
| 315 |
+
|
| 316 |
+
def reset(self, reset_info: dict) -> None:
|
| 317 |
+
if self._save_video_pred:
|
| 318 |
+
self._save_predicted_video()
|
| 319 |
+
self._video_pred_latents.clear()
|
| 320 |
+
self._current_prompt = ""
|
| 321 |
+
for key in self._frame_buffers:
|
| 322 |
+
self._frame_buffers[key] = []
|
| 323 |
+
self._is_first_call = True
|
| 324 |
+
self._current_session_id = None
|
| 325 |
+
if hasattr(self._policy.trained_model, "action_head") and hasattr(
|
| 326 |
+
self._policy.trained_model.action_head, "current_start_frame"
|
| 327 |
+
):
|
| 328 |
+
self._policy.trained_model.action_head.current_start_frame = 0
|
| 329 |
+
|
| 330 |
+
|
| 331 |
+
def main(
|
| 332 |
+
model_path: str = "./checkpoints/dreamzero_droid_wan22_smoke",
|
| 333 |
+
embodiment_tag: str = "oxe_droid",
|
| 334 |
+
tokenizer_path: str | None = None,
|
| 335 |
+
port: int = 8000,
|
| 336 |
+
host: str = "0.0.0.0",
|
| 337 |
+
image_height: int | None = None,
|
| 338 |
+
image_width: int | None = None,
|
| 339 |
+
save_video_pred: bool = False,
|
| 340 |
+
video_output_dir: str = "./video_pred_output",
|
| 341 |
+
) -> None:
|
| 342 |
+
logging.basicConfig(level=logging.INFO, force=True)
|
| 343 |
+
|
| 344 |
+
_maybe_init_distributed()
|
| 345 |
+
device_mesh = init_device_mesh("cuda", mesh_shape=(1,), mesh_dim_names=("ip",))
|
| 346 |
+
|
| 347 |
+
logger.info("Loading DreamZero Wan22 policy from %s (embodiment=%s)", model_path, embodiment_tag)
|
| 348 |
+
checkpoint_name = os.path.basename(model_path.rstrip("/"))
|
| 349 |
+
video_output_dir = os.path.join(video_output_dir, checkpoint_name)
|
| 350 |
+
policy = GrootSimPolicy(
|
| 351 |
+
embodiment_tag=EmbodimentTag(embodiment_tag),
|
| 352 |
+
model_path=model_path,
|
| 353 |
+
tokenizer_path_override=tokenizer_path,
|
| 354 |
+
device="cuda" if torch.cuda.is_available() else "cpu",
|
| 355 |
+
device_mesh=device_mesh,
|
| 356 |
+
)
|
| 357 |
+
if image_height is not None and image_width is not None:
|
| 358 |
+
h, w = image_height, image_width
|
| 359 |
+
logger.info("Using CLI video resolution: %dx%d", h, w)
|
| 360 |
+
else:
|
| 361 |
+
h, w = _get_expected_video_resolution(policy)
|
| 362 |
+
logger.info("Using checkpoint video resolution: %dx%d (HxW)", h, w)
|
| 363 |
+
wrapper = DreamZeroWan225BPolicy(
|
| 364 |
+
groot_policy=policy,
|
| 365 |
+
image_height=h,
|
| 366 |
+
image_width=w,
|
| 367 |
+
embodiment_tag=embodiment_tag,
|
| 368 |
+
save_video_pred=save_video_pred,
|
| 369 |
+
video_output_dir=video_output_dir,
|
| 370 |
+
)
|
| 371 |
+
|
| 372 |
+
server_config = PolicyServerConfig(
|
| 373 |
+
image_resolution=(h, w),
|
| 374 |
+
needs_wrist_camera=True,
|
| 375 |
+
n_external_cameras=2,
|
| 376 |
+
needs_stereo_camera=False,
|
| 377 |
+
needs_session_id=True,
|
| 378 |
+
action_space="joint_position",
|
| 379 |
+
)
|
| 380 |
+
logger.info("Starting WebsocketPolicyServer on %s:%d (DreamZero 5B, %dx%d)", host, port, h, w)
|
| 381 |
+
server = WebsocketPolicyServer(
|
| 382 |
+
policy=wrapper,
|
| 383 |
+
server_config=server_config,
|
| 384 |
+
host=host,
|
| 385 |
+
port=port,
|
| 386 |
+
)
|
| 387 |
+
server.serve_forever()
|
| 388 |
+
|
| 389 |
+
|
| 390 |
+
if __name__ == "__main__":
|
| 391 |
+
tyro.cli(main)
|
groot/__init__.py
ADDED
|
File without changes
|
groot/control/__init__.py
ADDED
|
File without changes
|
groot/control/tensorrt_utils.py
ADDED
|
@@ -0,0 +1,852 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
import torch
|
| 2 |
+
import os
|
| 3 |
+
import subprocess
|
| 4 |
+
import tensorrt as trt
|
| 5 |
+
import sys
|
| 6 |
+
import atexit
|
| 7 |
+
import ctypes
|
| 8 |
+
import modelopt.torch.quantization as mtq
|
| 9 |
+
from typing import Dict, List, Tuple
|
| 10 |
+
import shutil
|
| 11 |
+
|
| 12 |
+
import numpy as np
|
| 13 |
+
import torch
|
| 14 |
+
|
| 15 |
+
|
| 16 |
+
FP8_DEFAULT_CONFIG = {
|
| 17 |
+
"quant_cfg": {
|
| 18 |
+
"*weight_quantizer": {"num_bits": (4, 3), "axis": None},
|
| 19 |
+
"*input_quantizer": {"num_bits": (4, 3), "axis": None},
|
| 20 |
+
"*output_quantizer": {"enable": False},
|
| 21 |
+
"*[qkv]_bmm_quantizer": {"num_bits": (4, 3), "axis": None},
|
| 22 |
+
"*softmax_quantizer": {
|
| 23 |
+
"num_bits": (4, 3),
|
| 24 |
+
"axis": None,
|
| 25 |
+
},
|
| 26 |
+
"default": {"enable": False},
|
| 27 |
+
},
|
| 28 |
+
"algorithm": "max",
|
| 29 |
+
}
|
| 30 |
+
|
| 31 |
+
NVFP4_DEFAULT_CONFIG = {
|
| 32 |
+
"quant_cfg": {
|
| 33 |
+
"*weight_quantizer": {
|
| 34 |
+
"num_bits": (2, 1),
|
| 35 |
+
"block_sizes": {-1: 16, "type": "dynamic", "scale_bits": (4, 3)},
|
| 36 |
+
"axis": None,
|
| 37 |
+
"enable": True,
|
| 38 |
+
},
|
| 39 |
+
"*input_quantizer": {
|
| 40 |
+
"num_bits": (2, 1),
|
| 41 |
+
"block_sizes": {-1: 16, "type": "dynamic", "scale_bits": (4, 3)},
|
| 42 |
+
"axis": None,
|
| 43 |
+
"enable": True,
|
| 44 |
+
},
|
| 45 |
+
"*output_quantizer": {"enable": False},
|
| 46 |
+
"*[qkv]_bmm_quantizer": {"num_bits": (4, 3), "axis": None},
|
| 47 |
+
"*softmax_quantizer": {
|
| 48 |
+
"num_bits": (4, 3),
|
| 49 |
+
"axis": None,
|
| 50 |
+
},
|
| 51 |
+
"default": {"enable": False},
|
| 52 |
+
},
|
| 53 |
+
"algorithm": "max",
|
| 54 |
+
}
|
| 55 |
+
|
| 56 |
+
|
| 57 |
+
|
| 58 |
+
def wan_quantize(
|
| 59 |
+
policy,
|
| 60 |
+
quantization_config,
|
| 61 |
+
model_type,
|
| 62 |
+
forward_loop,
|
| 63 |
+
):
|
| 64 |
+
"""Quantize the VLA model using ModelOpt - simplified to use calc_mse_for_single_trajectory."""
|
| 65 |
+
|
| 66 |
+
# Configure quantization - disable problematic layers
|
| 67 |
+
if "quant_cfg" in quantization_config:
|
| 68 |
+
quantization_config["quant_cfg"]["*patch_embedding*"] = {"enable": False}
|
| 69 |
+
# if model_type == "14B" or model_type == "ar_14B":
|
| 70 |
+
# # Workaround: until we understand the issue https://nvbugspro.nvidia.com/bug/5612316
|
| 71 |
+
# quantization_config["quant_cfg"]["*.self_attn.o.*"] = {"enable": False}
|
| 72 |
+
# quantization_config["quant_cfg"]["*.cross_attn.o.*"] = {"enable": False}
|
| 73 |
+
|
| 74 |
+
policy.trained_model.action_head.model = mtq.quantize(
|
| 75 |
+
policy.trained_model.action_head.model, quantization_config, forward_loop=forward_loop
|
| 76 |
+
)
|
| 77 |
+
mtq.print_quant_summary(policy.trained_model.action_head.model)
|
| 78 |
+
|
| 79 |
+
return
|
| 80 |
+
|
| 81 |
+
|
| 82 |
+
def wan_trt_quantize_and_load_engine(
|
| 83 |
+
policy,
|
| 84 |
+
cfg,
|
| 85 |
+
onnx_path,
|
| 86 |
+
engine_path,
|
| 87 |
+
model_type,
|
| 88 |
+
forward_loop,
|
| 89 |
+
):
|
| 90 |
+
if (
|
| 91 |
+
os.path.exists(os.path.dirname(engine_path))
|
| 92 |
+
and cfg.inference_mode == "trt_build"
|
| 93 |
+
):
|
| 94 |
+
shutil.rmtree(os.path.dirname(engine_path))
|
| 95 |
+
|
| 96 |
+
quantization_config = None
|
| 97 |
+
if cfg.quantize_dtype == "fp8":
|
| 98 |
+
quantization_config = FP8_DEFAULT_CONFIG.copy()
|
| 99 |
+
elif cfg.quantize_dtype == "nvfp4":
|
| 100 |
+
quantization_config = NVFP4_DEFAULT_CONFIG.copy()
|
| 101 |
+
else:
|
| 102 |
+
print(f"Quantization type {cfg.quantize_dtype} not supported. Skipping quantization.")
|
| 103 |
+
|
| 104 |
+
if quantization_config is not None and cfg.inference_mode == "trt_build":
|
| 105 |
+
#policy.trained_model.action_head.model.to(torch.float16)
|
| 106 |
+
wan_quantize(
|
| 107 |
+
policy,
|
| 108 |
+
quantization_config,
|
| 109 |
+
model_type=model_type,
|
| 110 |
+
forward_loop=forward_loop,
|
| 111 |
+
)
|
| 112 |
+
|
| 113 |
+
if cfg.inference_mode == "trt_build":
|
| 114 |
+
policy.trained_model.action_head.model.to(torch.float16)
|
| 115 |
+
|
| 116 |
+
print("Export model:", policy.trained_model.action_head.model)
|
| 117 |
+
|
| 118 |
+
test_inputs = create_wan_test_inputs(policy, device="cuda", model_type=model_type)
|
| 119 |
+
min_shape = None
|
| 120 |
+
max_shape = None
|
| 121 |
+
opt_shape = None
|
| 122 |
+
|
| 123 |
+
if model_type == "ar_14B":
|
| 124 |
+
|
| 125 |
+
policy.trained_model.action_head.model.forward = policy.trained_model.action_head.model._forward_inference_trt
|
| 126 |
+
dynamic_axes = {
|
| 127 |
+
"kv_cache_packed": {3: "kv_cache_len"},
|
| 128 |
+
}
|
| 129 |
+
min_shape = "kv_cache_packed:40x2x1x880x40x128"
|
| 130 |
+
max_shape = "kv_cache_packed:40x2x1x8800x40x128"
|
| 131 |
+
opt_shape = "kv_cache_packed:40x2x1x7920x40x128"
|
| 132 |
+
elif model_type == "ar_14B_droid":
|
| 133 |
+
policy.trained_model.action_head.model.forward = policy.trained_model.action_head.model._forward_inference_trt
|
| 134 |
+
dynamic_axes = {
|
| 135 |
+
"kv_cache_packed": {3: "kv_cache_len"},
|
| 136 |
+
}
|
| 137 |
+
min_shape = "kv_cache_packed:40x2x1x880x40x128"
|
| 138 |
+
max_shape = "kv_cache_packed:40x2x1x8800x40x128"
|
| 139 |
+
opt_shape = "kv_cache_packed:40x2x1x7920x40x128"
|
| 140 |
+
elif model_type == "ar_5B_n6":
|
| 141 |
+
policy.trained_model.action_head.model.forward = policy.trained_model.action_head.model._forward_inference_trt
|
| 142 |
+
dynamic_axes = {
|
| 143 |
+
"kv_cache_packed": {3: "kv_cache_len"},
|
| 144 |
+
}
|
| 145 |
+
min_shape = "kv_cache_packed:30x2x1x220x24x128"
|
| 146 |
+
max_shape = "kv_cache_packed:30x2x1x3080x24x128"
|
| 147 |
+
opt_shape = "kv_cache_packed:30x2x1x2860x24x128"
|
| 148 |
+
else:
|
| 149 |
+
dynamic_axes = None
|
| 150 |
+
|
| 151 |
+
if cfg.quantize_dtype == "nvfp4":
|
| 152 |
+
export_to_onnx_fp4(policy.trained_model.action_head.model, test_inputs, onnx_path, dynamic_axes=dynamic_axes)
|
| 153 |
+
else:
|
| 154 |
+
export_to_onnx(
|
| 155 |
+
policy.trained_model.action_head.model,
|
| 156 |
+
test_inputs,
|
| 157 |
+
onnx_path,
|
| 158 |
+
model_type=model_type,
|
| 159 |
+
quantization_mode=cfg.quantize_dtype,
|
| 160 |
+
dynamic_axes=dynamic_axes,
|
| 161 |
+
)
|
| 162 |
+
|
| 163 |
+
build_tensorrt_engine(onnx_path, engine_path, min_shape, max_shape, opt_shape)
|
| 164 |
+
|
| 165 |
+
trt_wan_model = load_tensorrt_engine(engine_path, model_type=model_type)
|
| 166 |
+
policy.trained_model.action_head.model = trt_wan_model
|
| 167 |
+
|
| 168 |
+
def export_to_onnx_fp4(model, test_inputs, onnx_save_path, dynamic_axes=None):
|
| 169 |
+
from modelopt.torch._deploy.utils.torch_onnx import OnnxBytes
|
| 170 |
+
from modelopt.torch._deploy.utils.torch_onnx import get_onnx_bytes_and_metadata
|
| 171 |
+
|
| 172 |
+
print("exporting to onnx fp4")
|
| 173 |
+
try:
|
| 174 |
+
onnx_bytes, _ = get_onnx_bytes_and_metadata(model=model, dummy_input=test_inputs, dynamic_axes=dynamic_axes)
|
| 175 |
+
onnx_model = OnnxBytes.from_bytes(onnx_bytes)
|
| 176 |
+
except Exception as e:
|
| 177 |
+
print(f"Error exporting model to ONNX: {e}")
|
| 178 |
+
return
|
| 179 |
+
save_dir = os.path.dirname(os.path.abspath(onnx_save_path))
|
| 180 |
+
os.makedirs(save_dir, exist_ok=True)
|
| 181 |
+
for filename, file_bytes in onnx_model.onnx_model.items():
|
| 182 |
+
file_path = os.path.join(save_dir, filename)
|
| 183 |
+
with open(file_path, "wb") as f:
|
| 184 |
+
f.write(file_bytes)
|
| 185 |
+
print(f"exported onnx to {file_path}")
|
| 186 |
+
|
| 187 |
+
|
| 188 |
+
def export_to_onnx(
|
| 189 |
+
pytorch_model,
|
| 190 |
+
test_inputs,
|
| 191 |
+
onnx_path="tensorrt/wan_model.onnx",
|
| 192 |
+
model_type="5B",
|
| 193 |
+
quantization_mode="fp8",
|
| 194 |
+
dynamic_axes=None,
|
| 195 |
+
):
|
| 196 |
+
#
|
| 197 |
+
if model_type == "5B":
|
| 198 |
+
return export_to_onnx_5B(pytorch_model, test_inputs, onnx_path, dynamic_axes)
|
| 199 |
+
elif model_type == "14B":
|
| 200 |
+
return export_to_onnx_14B(pytorch_model, test_inputs, onnx_path, dynamic_axes)
|
| 201 |
+
elif model_type == "ar_14B" or model_type == "ar_14B_droid":
|
| 202 |
+
return export_to_onnx_ar_14B(pytorch_model, test_inputs, onnx_path, dynamic_axes)
|
| 203 |
+
else:
|
| 204 |
+
raise ValueError(f"Model type {model_type} not supported")
|
| 205 |
+
|
| 206 |
+
|
| 207 |
+
def export_to_onnx_ar_14B(pytorch_model, test_inputs, onnx_path="tensorrt/wan_model.onnx", dynamic_axes=None):
|
| 208 |
+
"""Export PyTorch model to ONNX"""
|
| 209 |
+
print("Exporting AR 14B model to ONNX...", onnx_path)
|
| 210 |
+
|
| 211 |
+
# Create directory if it doesn't exist
|
| 212 |
+
os.makedirs(os.path.dirname(onnx_path), exist_ok=True)
|
| 213 |
+
pytorch_model.eval()
|
| 214 |
+
pytorch_model.to(torch.float16)
|
| 215 |
+
|
| 216 |
+
input_names = [
|
| 217 |
+
"x",
|
| 218 |
+
"timestep",
|
| 219 |
+
"context",
|
| 220 |
+
"kv_cache_packed",
|
| 221 |
+
"y",
|
| 222 |
+
"clip_feature",
|
| 223 |
+
"action",
|
| 224 |
+
"timestep_action",
|
| 225 |
+
"state",
|
| 226 |
+
]
|
| 227 |
+
output_names = ["video_noise_pred", "action_noise_pred"]
|
| 228 |
+
|
| 229 |
+
try:
|
| 230 |
+
with torch.no_grad():
|
| 231 |
+
torch.onnx.export(
|
| 232 |
+
pytorch_model,
|
| 233 |
+
test_inputs,
|
| 234 |
+
onnx_path,
|
| 235 |
+
export_params=True,
|
| 236 |
+
opset_version=20,
|
| 237 |
+
do_constant_folding=True,
|
| 238 |
+
input_names=input_names,
|
| 239 |
+
output_names=output_names,
|
| 240 |
+
dynamic_axes=dynamic_axes,
|
| 241 |
+
)
|
| 242 |
+
print(f" ONNX model exported to: {onnx_path}")
|
| 243 |
+
return onnx_path
|
| 244 |
+
|
| 245 |
+
except Exception as e:
|
| 246 |
+
import traceback
|
| 247 |
+
print(f" ERROR: ONNX export failed. Exception type: {type(e)}")
|
| 248 |
+
print("Traceback:")
|
| 249 |
+
traceback.print_exc()
|
| 250 |
+
return None
|
| 251 |
+
|
| 252 |
+
|
| 253 |
+
|
| 254 |
+
def export_to_onnx_5B(pytorch_model, test_inputs, onnx_path="tensorrt/wan_model.onnx"):
|
| 255 |
+
"""Export PyTorch model to ONNX"""
|
| 256 |
+
print("Exporting model to ONNX...")
|
| 257 |
+
|
| 258 |
+
# Create directory if it doesn't exist
|
| 259 |
+
os.makedirs(os.path.dirname(onnx_path), exist_ok=True)
|
| 260 |
+
pytorch_model.eval()
|
| 261 |
+
pytorch_model.to(torch.float16)
|
| 262 |
+
|
| 263 |
+
x, action, timestep, context, state, embodiment_id = test_inputs
|
| 264 |
+
|
| 265 |
+
# Define input names for better ONNX graph
|
| 266 |
+
input_names = ["x", "action", "timestep", "context", "state", "embodiment_id"]
|
| 267 |
+
output_names = ["video_noise_pred", "action_noise_pred"]
|
| 268 |
+
|
| 269 |
+
try:
|
| 270 |
+
with torch.no_grad():
|
| 271 |
+
torch.onnx.export(
|
| 272 |
+
pytorch_model,
|
| 273 |
+
(x, action, timestep, context, state, embodiment_id),
|
| 274 |
+
onnx_path,
|
| 275 |
+
export_params=True,
|
| 276 |
+
opset_version=20,
|
| 277 |
+
do_constant_folding=True,
|
| 278 |
+
input_names=input_names,
|
| 279 |
+
output_names=output_names,
|
| 280 |
+
)
|
| 281 |
+
print(f" ONNX model exported to: {onnx_path}")
|
| 282 |
+
return onnx_path
|
| 283 |
+
|
| 284 |
+
except Exception as e:
|
| 285 |
+
import traceback
|
| 286 |
+
print(f" ERROR: ONNX export failed. Exception type: {type(e)}")
|
| 287 |
+
print("Traceback:")
|
| 288 |
+
traceback.print_exc()
|
| 289 |
+
return None
|
| 290 |
+
|
| 291 |
+
|
| 292 |
+
def export_to_onnx_14B(pytorch_model, test_inputs, onnx_path="tensorrt/wan_model.onnx"):
|
| 293 |
+
"""Export PyTorch model to ONNX"""
|
| 294 |
+
print("Exporting model to ONNX...")
|
| 295 |
+
|
| 296 |
+
# Create directory if it doesn't exist
|
| 297 |
+
os.makedirs(os.path.dirname(onnx_path), exist_ok=True)
|
| 298 |
+
pytorch_model.eval()
|
| 299 |
+
pytorch_model.to(torch.float16)
|
| 300 |
+
|
| 301 |
+
x, action, timestep, context, state, embodiment_id, clip_feature, y = test_inputs
|
| 302 |
+
|
| 303 |
+
# Define input names for better ONNX graph
|
| 304 |
+
input_names = [
|
| 305 |
+
"x",
|
| 306 |
+
"action",
|
| 307 |
+
"timestep",
|
| 308 |
+
"context",
|
| 309 |
+
"state",
|
| 310 |
+
"embodiment_id",
|
| 311 |
+
"clip_feature",
|
| 312 |
+
"y",
|
| 313 |
+
]
|
| 314 |
+
output_names = ["video_noise_pred", "action_noise_pred"]
|
| 315 |
+
|
| 316 |
+
try:
|
| 317 |
+
with torch.no_grad():
|
| 318 |
+
torch.onnx.export(
|
| 319 |
+
pytorch_model,
|
| 320 |
+
(x, action, timestep, context, state, embodiment_id, clip_feature, y),
|
| 321 |
+
onnx_path,
|
| 322 |
+
export_params=True,
|
| 323 |
+
opset_version=20,
|
| 324 |
+
do_constant_folding=True,
|
| 325 |
+
input_names=input_names,
|
| 326 |
+
output_names=output_names,
|
| 327 |
+
)
|
| 328 |
+
print(f" ONNX model exported to: {onnx_path}")
|
| 329 |
+
return onnx_path
|
| 330 |
+
|
| 331 |
+
except Exception as e:
|
| 332 |
+
import traceback
|
| 333 |
+
print(f" ERROR: ONNX export failed. Exception type: {type(e)}")
|
| 334 |
+
print("Traceback:")
|
| 335 |
+
traceback.print_exc()
|
| 336 |
+
return None
|
| 337 |
+
|
| 338 |
+
|
| 339 |
+
def build_tensorrt_engine(onnx_path, engine_path="tensorrt/wan_model.trt", min_shape=None, max_shape=None, opt_shape=None):
|
| 340 |
+
"""Build TensorRT engine from ONNX using trtexec"""
|
| 341 |
+
print("Building TensorRT engine with trtexec...")
|
| 342 |
+
|
| 343 |
+
if not os.path.exists(onnx_path):
|
| 344 |
+
print(f" ERROR: ONNX file not found: {onnx_path}")
|
| 345 |
+
return None
|
| 346 |
+
|
| 347 |
+
# Create directory if it doesn't exist
|
| 348 |
+
os.makedirs(os.path.dirname(engine_path), exist_ok=True)
|
| 349 |
+
|
| 350 |
+
# Build engine using trtexec (much faster than torch_tensorrt)
|
| 351 |
+
trtexec_bin = shutil.which("trtexec") or "/opt/tensorrt/bin/trtexec"
|
| 352 |
+
cmd = [
|
| 353 |
+
trtexec_bin,
|
| 354 |
+
f"--onnx={onnx_path}",
|
| 355 |
+
f"--saveEngine={engine_path}",
|
| 356 |
+
"--fp8",
|
| 357 |
+
"--fp16",
|
| 358 |
+
"--bf16",
|
| 359 |
+
"--separateProfileRun",
|
| 360 |
+
"--profilingVerbosity=detailed",
|
| 361 |
+
"--memPoolSize=workspace:65536",
|
| 362 |
+
"--dumpProfile",
|
| 363 |
+
"--dumpLayerInfo",
|
| 364 |
+
"--useCudaGraph",
|
| 365 |
+
"--verbose",
|
| 366 |
+
]
|
| 367 |
+
|
| 368 |
+
if min_shape is not None:
|
| 369 |
+
cmd.append(f"--minShapes={min_shape}")
|
| 370 |
+
if max_shape is not None:
|
| 371 |
+
cmd.append(f"--maxShapes={max_shape}")
|
| 372 |
+
if opt_shape is not None:
|
| 373 |
+
cmd.append(f"--optShapes={opt_shape}")
|
| 374 |
+
|
| 375 |
+
# Create log file for trtexec output
|
| 376 |
+
log_file = engine_path.replace(".trt", "_build.log")
|
| 377 |
+
|
| 378 |
+
try:
|
| 379 |
+
print(f" Running: {' '.join(cmd)}")
|
| 380 |
+
print(f" Logging output to: {log_file}")
|
| 381 |
+
|
| 382 |
+
with open(log_file, "w") as f:
|
| 383 |
+
result = subprocess.run(cmd, stdout=f, stderr=subprocess.STDOUT, text=True, timeout=600)
|
| 384 |
+
|
| 385 |
+
if result.returncode == 0:
|
| 386 |
+
print(f" TensorRT engine built successfully: {engine_path}")
|
| 387 |
+
print(f" Build log saved to: {log_file}")
|
| 388 |
+
return engine_path
|
| 389 |
+
else:
|
| 390 |
+
print(f" ERROR: trtexec failed with return code {result.returncode}")
|
| 391 |
+
print(f" Check build log for details: {log_file}")
|
| 392 |
+
# Print last few lines of log file for immediate feedback
|
| 393 |
+
try:
|
| 394 |
+
with open(log_file, "r") as f:
|
| 395 |
+
lines = f.readlines()
|
| 396 |
+
if lines:
|
| 397 |
+
print(" Last few lines from build log:")
|
| 398 |
+
for line in lines[-10:]: # Show last 10 lines
|
| 399 |
+
print(f" {line.rstrip()}")
|
| 400 |
+
except:
|
| 401 |
+
pass
|
| 402 |
+
return None
|
| 403 |
+
|
| 404 |
+
except subprocess.TimeoutExpired:
|
| 405 |
+
print(" ERROR: trtexec timed out after 5 minutes")
|
| 406 |
+
print(f" Partial build log saved to: {log_file}")
|
| 407 |
+
return None
|
| 408 |
+
except Exception as e:
|
| 409 |
+
print(f" ERROR: Failed to run trtexec: {e}")
|
| 410 |
+
return None
|
| 411 |
+
|
| 412 |
+
|
| 413 |
+
def torch_type(trt_type):
|
| 414 |
+
mapping = {
|
| 415 |
+
trt.float32: torch.float32, # Added missing FLOAT mapping
|
| 416 |
+
trt.float16: torch.float16,
|
| 417 |
+
trt.bfloat16: torch.bfloat16,
|
| 418 |
+
trt.int8: torch.int8,
|
| 419 |
+
trt.int32: torch.int32,
|
| 420 |
+
trt.bool: torch.bool,
|
| 421 |
+
trt.uint8: torch.uint8,
|
| 422 |
+
trt.int64: torch.int64,
|
| 423 |
+
}
|
| 424 |
+
if trt_type in mapping:
|
| 425 |
+
return mapping[trt_type]
|
| 426 |
+
|
| 427 |
+
raise TypeError(
|
| 428 |
+
f"Could not resolve TensorRT datatype to an equivalent torch datatype. {trt_type}"
|
| 429 |
+
)
|
| 430 |
+
|
| 431 |
+
|
| 432 |
+
class Engine(object):
|
| 433 |
+
def __init__(self, file, plugins=[]):
|
| 434 |
+
super().__init__()
|
| 435 |
+
|
| 436 |
+
self.logger = trt.Logger(trt.Logger.ERROR)
|
| 437 |
+
trt.init_libnvinfer_plugins(self.logger, "")
|
| 438 |
+
|
| 439 |
+
self.plugins = [ctypes.CDLL(plugin, ctypes.RTLD_GLOBAL) for plugin in plugins]
|
| 440 |
+
self.file = file
|
| 441 |
+
self.load(file)
|
| 442 |
+
|
| 443 |
+
def destroy(self):
|
| 444 |
+
del self.execution_context
|
| 445 |
+
del self.handle
|
| 446 |
+
|
| 447 |
+
atexit.register(destroy, self)
|
| 448 |
+
self.print()
|
| 449 |
+
|
| 450 |
+
def print(self):
|
| 451 |
+
|
| 452 |
+
print("============= TRT Engine Detail =============")
|
| 453 |
+
print(f"Engine file: {self.file}")
|
| 454 |
+
print(f"Inputs: {len(self.in_meta)}")
|
| 455 |
+
for ib, item in enumerate(self.in_meta):
|
| 456 |
+
tensor_name, shape, dtype = item[:3]
|
| 457 |
+
print(f" {ib}. {tensor_name}: {'x'.join(map(str, shape))} [{dtype}]")
|
| 458 |
+
|
| 459 |
+
print(f"Outputs: {len(self.out_meta)}")
|
| 460 |
+
for ib, item in enumerate(self.out_meta):
|
| 461 |
+
tensor_name, shape, dtype = item[:3]
|
| 462 |
+
print(f" {ib}. {tensor_name}: {'x'.join(map(str, shape))} [{dtype}]")
|
| 463 |
+
print("=============================================")
|
| 464 |
+
|
| 465 |
+
def load(self, file):
|
| 466 |
+
runtime = trt.Runtime(self.logger)
|
| 467 |
+
|
| 468 |
+
with open(file, "rb") as f:
|
| 469 |
+
self.handle = runtime.deserialize_cuda_engine(f.read())
|
| 470 |
+
assert (
|
| 471 |
+
self.handle is not None
|
| 472 |
+
), f"Failed to deserialize the cuda engine from file: {file}"
|
| 473 |
+
|
| 474 |
+
self.execution_context = self.handle.create_execution_context()
|
| 475 |
+
self.meta, self.in_meta, self.out_meta = [], [], []
|
| 476 |
+
for tensor_name in self.handle:
|
| 477 |
+
shape = self.handle.get_tensor_shape(tensor_name)
|
| 478 |
+
print(f"Tensor name: {tensor_name}, shape: {shape}")
|
| 479 |
+
dtype = torch_type(self.handle.get_tensor_dtype(tensor_name))
|
| 480 |
+
if self.handle.get_tensor_mode(tensor_name) == trt.TensorIOMode.INPUT:
|
| 481 |
+
self.in_meta.append([tensor_name, shape, dtype])
|
| 482 |
+
else:
|
| 483 |
+
self.out_meta.append([tensor_name, shape, dtype])
|
| 484 |
+
|
| 485 |
+
def __call__(self, *args, **inputs):
|
| 486 |
+
return self.forward(*args, **inputs)
|
| 487 |
+
|
| 488 |
+
def set_runtime_tensor_shape(self, name, shape):
|
| 489 |
+
self.execution_context.set_input_shape(name, shape)
|
| 490 |
+
|
| 491 |
+
def forward(self, *args, **kwargs):
|
| 492 |
+
return_list = kwargs.pop("return_list", False)
|
| 493 |
+
reference_tensors = []
|
| 494 |
+
stream = torch.cuda.current_stream()
|
| 495 |
+
for iarg, x in enumerate(args):
|
| 496 |
+
name, shape, dtype = self.in_meta[iarg]
|
| 497 |
+
runtime_shape = self.execution_context.get_tensor_shape(name)
|
| 498 |
+
assert isinstance(x, torch.Tensor), f"Unsupported tensor type: {type(x)}"
|
| 499 |
+
assert runtime_shape == x.shape, f"Invalid input shape: {runtime_shape} != {x.shape}"
|
| 500 |
+
assert (
|
| 501 |
+
dtype == x.dtype
|
| 502 |
+
), f"Invalid tensor dtype, excepted dtype is {dtype}, but got {x.dtype}"
|
| 503 |
+
assert x.is_cuda, f"Invalid tensor device, excepted device is cuda, but got {x.device}"
|
| 504 |
+
x = x.cuda().contiguous()
|
| 505 |
+
self.execution_context.set_tensor_address(name, x.data_ptr())
|
| 506 |
+
reference_tensors.append(x)
|
| 507 |
+
|
| 508 |
+
for name, shape, dtype in self.in_meta:
|
| 509 |
+
if name not in kwargs:
|
| 510 |
+
continue
|
| 511 |
+
|
| 512 |
+
runtime_shape = self.execution_context.get_tensor_shape(name)
|
| 513 |
+
x = kwargs[name]
|
| 514 |
+
assert isinstance(x, torch.Tensor), f"Unsupported tensor[{name}] type: {type(x)}"
|
| 515 |
+
assert (
|
| 516 |
+
runtime_shape == x.shape
|
| 517 |
+
), f"Invalid input[{name}] shape: {x.shape}, but the expected shape is: {runtime_shape}"
|
| 518 |
+
assert (
|
| 519 |
+
dtype == x.dtype
|
| 520 |
+
), f"Invalid tensor[{name}] dtype, expected dtype is {dtype}, but got {x.dtype}"
|
| 521 |
+
assert (
|
| 522 |
+
x.is_cuda
|
| 523 |
+
), f"Invalid tensor[{name}] device, expected device is cuda, but got {x.device}"
|
| 524 |
+
x = x.cuda().contiguous()
|
| 525 |
+
self.execution_context.set_tensor_address(name, x.data_ptr())
|
| 526 |
+
reference_tensors.append(x)
|
| 527 |
+
|
| 528 |
+
for item in self.out_meta:
|
| 529 |
+
name = item[0]
|
| 530 |
+
runtime_shape = self.execution_context.get_tensor_shape(name)
|
| 531 |
+
output_tensor = torch.zeros(
|
| 532 |
+
*runtime_shape, dtype=item[2], device=reference_tensors[0].device
|
| 533 |
+
)
|
| 534 |
+
self.execution_context.set_tensor_address(name, output_tensor.data_ptr())
|
| 535 |
+
reference_tensors.append(output_tensor)
|
| 536 |
+
|
| 537 |
+
self.execution_context.execute_async_v3(stream.cuda_stream)
|
| 538 |
+
stream.synchronize()
|
| 539 |
+
assert len(reference_tensors) == len(self.in_meta) + len(
|
| 540 |
+
self.out_meta
|
| 541 |
+
), f"Invalid input tensors. The expected I/O tensors are {len(self.in_meta) + len(self.out_meta)}, but got {len(reference_tensors)}"
|
| 542 |
+
|
| 543 |
+
if return_list:
|
| 544 |
+
return [
|
| 545 |
+
reference_tensors[len(self.in_meta) + i] for i, item in enumerate(self.out_meta)
|
| 546 |
+
]
|
| 547 |
+
else:
|
| 548 |
+
return {
|
| 549 |
+
item[0]: reference_tensors[len(self.in_meta) + i]
|
| 550 |
+
for i, item in enumerate(self.out_meta)
|
| 551 |
+
}
|
| 552 |
+
|
| 553 |
+
|
| 554 |
+
class WanTrtModel5B(torch.nn.Module):
|
| 555 |
+
def __init__(self, eng_path: str):
|
| 556 |
+
super().__init__()
|
| 557 |
+
self.engine = Engine(eng_path)
|
| 558 |
+
|
| 559 |
+
def forward(
|
| 560 |
+
self,
|
| 561 |
+
x: torch.Tensor,
|
| 562 |
+
action: torch.Tensor,
|
| 563 |
+
timestep: torch.Tensor,
|
| 564 |
+
context: torch.Tensor,
|
| 565 |
+
state: torch.Tensor,
|
| 566 |
+
embodiment_id: torch.Tensor,
|
| 567 |
+
):
|
| 568 |
+
|
| 569 |
+
self.engine.set_runtime_tensor_shape("x", x.shape)
|
| 570 |
+
self.engine.set_runtime_tensor_shape("action", action.shape)
|
| 571 |
+
self.engine.set_runtime_tensor_shape("context", context.shape)
|
| 572 |
+
self.engine.set_runtime_tensor_shape("state", state.shape)
|
| 573 |
+
|
| 574 |
+
output = self.engine(
|
| 575 |
+
x=x.to(torch.float16),
|
| 576 |
+
action=action.to(torch.float16),
|
| 577 |
+
timestep=timestep.to(torch.float16),
|
| 578 |
+
context=context.to(torch.float16),
|
| 579 |
+
state=state.to(torch.float16),
|
| 580 |
+
embodiment_id=embodiment_id.to(torch.int32),
|
| 581 |
+
)
|
| 582 |
+
if "out.0" in output: # for nvfp4 model export through modelopt
|
| 583 |
+
return output["out.0"].to(torch.bfloat16).contiguous(), output["out.1"].to(torch.bfloat16).contiguous()
|
| 584 |
+
else:
|
| 585 |
+
return output["video_noise_pred"].to(torch.bfloat16).contiguous(), output["action_noise_pred"].to(torch.bfloat16).contiguous()
|
| 586 |
+
|
| 587 |
+
|
| 588 |
+
class WanTrtModel14B(torch.nn.Module):
|
| 589 |
+
def __init__(self, eng_path: str):
|
| 590 |
+
super().__init__()
|
| 591 |
+
self.engine = Engine(eng_path)
|
| 592 |
+
|
| 593 |
+
def forward(
|
| 594 |
+
self,
|
| 595 |
+
x: torch.Tensor,
|
| 596 |
+
action: torch.Tensor,
|
| 597 |
+
timestep: torch.Tensor,
|
| 598 |
+
context: torch.Tensor,
|
| 599 |
+
state: torch.Tensor,
|
| 600 |
+
embodiment_id: torch.Tensor,
|
| 601 |
+
clip_feature: torch.Tensor,
|
| 602 |
+
y: torch.Tensor,
|
| 603 |
+
):
|
| 604 |
+
|
| 605 |
+
self.engine.set_runtime_tensor_shape("x", x.shape)
|
| 606 |
+
self.engine.set_runtime_tensor_shape("action", action.shape)
|
| 607 |
+
self.engine.set_runtime_tensor_shape("context", context.shape)
|
| 608 |
+
self.engine.set_runtime_tensor_shape("state", state.shape)
|
| 609 |
+
self.engine.set_runtime_tensor_shape("clip_feature", clip_feature.shape)
|
| 610 |
+
self.engine.set_runtime_tensor_shape("y", y.shape)
|
| 611 |
+
|
| 612 |
+
output = self.engine(
|
| 613 |
+
x=x.to(torch.float16),
|
| 614 |
+
action=action.to(torch.float16),
|
| 615 |
+
timestep=timestep.to(torch.float16),
|
| 616 |
+
context=context.to(torch.float16),
|
| 617 |
+
state=state.to(torch.float16),
|
| 618 |
+
embodiment_id=embodiment_id.to(torch.int32),
|
| 619 |
+
clip_feature=clip_feature.to(torch.float16),
|
| 620 |
+
y=y.to(torch.float16),
|
| 621 |
+
)
|
| 622 |
+
if "out.0" in output: # for nvfp4 model export through modelopt
|
| 623 |
+
return output["out.0"].to(torch.bfloat16).contiguous(), output["out.1"].to(torch.bfloat16).contiguous()
|
| 624 |
+
else:
|
| 625 |
+
return output["video_noise_pred"].to(torch.bfloat16).contiguous(), output["action_noise_pred"].to(torch.bfloat16).contiguous()
|
| 626 |
+
|
| 627 |
+
|
| 628 |
+
class WanTrtModelAr5B(torch.nn.Module):
|
| 629 |
+
"""TRT wrapper for ar_5B_n6 model type - uses kv_cache but no clip_feature."""
|
| 630 |
+
def __init__(self, eng_path: str):
|
| 631 |
+
super().__init__()
|
| 632 |
+
self.engine = Engine(eng_path)
|
| 633 |
+
|
| 634 |
+
def forward(
|
| 635 |
+
self,
|
| 636 |
+
x,
|
| 637 |
+
timestep,
|
| 638 |
+
context,
|
| 639 |
+
kv_cache: list[torch.Tensor],
|
| 640 |
+
y=None,
|
| 641 |
+
action=None,
|
| 642 |
+
timestep_action=None,
|
| 643 |
+
state=None,
|
| 644 |
+
):
|
| 645 |
+
|
| 646 |
+
kv_cache_packed = torch.stack(kv_cache, dim=0)
|
| 647 |
+
|
| 648 |
+
self.engine.set_runtime_tensor_shape("x", x.shape)
|
| 649 |
+
self.engine.set_runtime_tensor_shape("timestep", timestep.shape)
|
| 650 |
+
self.engine.set_runtime_tensor_shape("context", context.shape)
|
| 651 |
+
self.engine.set_runtime_tensor_shape("kv_cache_packed", kv_cache_packed.shape)
|
| 652 |
+
# self.engine.set_runtime_tensor_shape("y", y.shape)
|
| 653 |
+
self.engine.set_runtime_tensor_shape("action", action.shape)
|
| 654 |
+
self.engine.set_runtime_tensor_shape("timestep_action", timestep_action.shape)
|
| 655 |
+
self.engine.set_runtime_tensor_shape("state", state.shape)
|
| 656 |
+
|
| 657 |
+
|
| 658 |
+
output = self.engine(
|
| 659 |
+
x.to(torch.float16),
|
| 660 |
+
timestep.to(torch.float16),
|
| 661 |
+
context.to(torch.float16),
|
| 662 |
+
kv_cache_packed.to(torch.float16),
|
| 663 |
+
# y.to(torch.float16),
|
| 664 |
+
action.to(torch.float16),
|
| 665 |
+
timestep_action.to(torch.float16),
|
| 666 |
+
state.to(torch.float16),
|
| 667 |
+
)
|
| 668 |
+
|
| 669 |
+
if "out.0" in output: # for nvfp4 model export through modelopt
|
| 670 |
+
return output["out.0"].to(torch.bfloat16).contiguous(), output["out.1"].to(torch.bfloat16).contiguous()
|
| 671 |
+
else:
|
| 672 |
+
return output["video_noise_pred"].to(torch.bfloat16).contiguous(), output["action_noise_pred"].to(torch.bfloat16).contiguous()
|
| 673 |
+
|
| 674 |
+
|
| 675 |
+
class WanTrtModelAr14B(torch.nn.Module):
|
| 676 |
+
def __init__(self, eng_path: str):
|
| 677 |
+
super().__init__()
|
| 678 |
+
self.engine = Engine(eng_path)
|
| 679 |
+
|
| 680 |
+
def forward(
|
| 681 |
+
self,
|
| 682 |
+
x,
|
| 683 |
+
timestep,
|
| 684 |
+
context,
|
| 685 |
+
kv_cache: list[torch.Tensor],
|
| 686 |
+
y=None,
|
| 687 |
+
clip_feature=None,
|
| 688 |
+
action=None,
|
| 689 |
+
timestep_action=None,
|
| 690 |
+
state=None,
|
| 691 |
+
):
|
| 692 |
+
|
| 693 |
+
kv_cache_packed = torch.stack(kv_cache, dim=0)
|
| 694 |
+
|
| 695 |
+
self.engine.set_runtime_tensor_shape("x", x.shape)
|
| 696 |
+
self.engine.set_runtime_tensor_shape("timestep", timestep.shape)
|
| 697 |
+
self.engine.set_runtime_tensor_shape("context", context.shape)
|
| 698 |
+
self.engine.set_runtime_tensor_shape("kv_cache_packed", kv_cache_packed.shape)
|
| 699 |
+
self.engine.set_runtime_tensor_shape("y", y.shape)
|
| 700 |
+
self.engine.set_runtime_tensor_shape("clip_feature", clip_feature.shape)
|
| 701 |
+
self.engine.set_runtime_tensor_shape("action", action.shape)
|
| 702 |
+
self.engine.set_runtime_tensor_shape("timestep_action", timestep_action.shape)
|
| 703 |
+
self.engine.set_runtime_tensor_shape("state", state.shape)
|
| 704 |
+
|
| 705 |
+
|
| 706 |
+
output = self.engine(
|
| 707 |
+
x.to(torch.float16),
|
| 708 |
+
timestep.to(torch.float16),
|
| 709 |
+
context.to(torch.float16),
|
| 710 |
+
kv_cache_packed.to(torch.float16),
|
| 711 |
+
y.to(torch.float16),
|
| 712 |
+
clip_feature.to(torch.float16),
|
| 713 |
+
action.to(torch.float16),
|
| 714 |
+
timestep_action.to(torch.float16),
|
| 715 |
+
state.to(torch.float16),
|
| 716 |
+
)
|
| 717 |
+
|
| 718 |
+
if "out.0" in output: # for nvfp4 model export through modelopt
|
| 719 |
+
return output["out.0"].to(torch.bfloat16).contiguous(), output["out.1"].to(torch.bfloat16).contiguous()
|
| 720 |
+
else:
|
| 721 |
+
return output["video_noise_pred"].to(torch.bfloat16).contiguous(), output["action_noise_pred"].to(torch.bfloat16).contiguous()
|
| 722 |
+
|
| 723 |
+
def load_tensorrt_engine(engine_path="tensorrt/wan_model.trt", model_type="5B"):
|
| 724 |
+
"""Load TensorRT engine"""
|
| 725 |
+
if model_type == "5B":
|
| 726 |
+
trt_inference = WanTrtModel5B(engine_path)
|
| 727 |
+
elif model_type == "ar_5B_n6" or model_type == "ar_5B":
|
| 728 |
+
trt_inference = WanTrtModelAr5B(engine_path)
|
| 729 |
+
elif model_type == "14B":
|
| 730 |
+
trt_inference = WanTrtModel14B(engine_path)
|
| 731 |
+
elif model_type == "ar_14B" or model_type == "ar_14B_droid":
|
| 732 |
+
trt_inference = WanTrtModelAr14B(engine_path)
|
| 733 |
+
else:
|
| 734 |
+
raise ValueError(f"Model type {model_type} not supported")
|
| 735 |
+
return trt_inference
|
| 736 |
+
|
| 737 |
+
|
| 738 |
+
def create_wan_test_inputs(policy, device="cuda", model_type="5B"):
|
| 739 |
+
# Get dtype from model parameters
|
| 740 |
+
dtype = torch.float16
|
| 741 |
+
|
| 742 |
+
# Use hardcoded dimensions from the original working version of the script
|
| 743 |
+
if model_type == "5B":
|
| 744 |
+
x = torch.randn(1, 48, 13, 22, 40, dtype=dtype, device=device)
|
| 745 |
+
action = torch.randn(1, 48, 32, dtype=dtype, device=device)
|
| 746 |
+
timestep = torch.randn(1, dtype=dtype, device=device)
|
| 747 |
+
context = torch.randn(1, 512, 4096, dtype=dtype, device=device)
|
| 748 |
+
state = torch.randn(1, 1, 64, dtype=dtype, device=device)
|
| 749 |
+
embodiment_id = torch.zeros(1, dtype=torch.int32, device=device)
|
| 750 |
+
timestep_action = torch.randn(1, 48, dtype=dtype, device=device)
|
| 751 |
+
seq_len = torch.tensor(440, dtype=torch.int32, device=device)
|
| 752 |
+
return x, action, timestep, context, state, embodiment_id, timestep_action, seq_len
|
| 753 |
+
elif model_type == "ar_5B_n6":
|
| 754 |
+
# ar_5B_n6 uses _forward_inference_trt which requires kv_cache_packed
|
| 755 |
+
# Shape from dynamic_axes: kv_cache_packed:30x2x1x220x24x128
|
| 756 |
+
# Note: 5B model doesn't use clip_feature (unlike 14B), but still needs y
|
| 757 |
+
x = torch.randn(1, 48, 2, 22, 40, dtype=dtype, device=device)
|
| 758 |
+
timestep = torch.randn(1, 2, dtype=dtype, device=device)
|
| 759 |
+
context = torch.randn(1, 512, 4096, dtype=dtype, device=device)
|
| 760 |
+
# y = torch.randn(1, 52, 2, 22, 40, dtype=dtype, device=device) # y is required by _forward_inference_trt
|
| 761 |
+
action = torch.randn(1, 48, 32, dtype=dtype, device=device)
|
| 762 |
+
timestep_action = torch.randn(1, 48, dtype=dtype, device=device)
|
| 763 |
+
state = torch.randn(1, 1, 64, dtype=dtype, device=device)
|
| 764 |
+
|
| 765 |
+
num_heads = 24
|
| 766 |
+
head_dim = 128
|
| 767 |
+
num_layers = 30
|
| 768 |
+
B = 1
|
| 769 |
+
|
| 770 |
+
kv_cache = []
|
| 771 |
+
for _ in range(num_layers):
|
| 772 |
+
kv_cache.append(
|
| 773 |
+
torch.zeros([2, B, 13*220, num_heads, head_dim], dtype=dtype, device=device)
|
| 774 |
+
)
|
| 775 |
+
|
| 776 |
+
kv_cache_packed = torch.stack(kv_cache, dim=0)
|
| 777 |
+
# Return order matches _forward_inference_trt signature: x, timestep, context, kv_cache_packed, y, action, timestep_action, state
|
| 778 |
+
return (x, timestep, context, kv_cache_packed, action, timestep_action, state)
|
| 779 |
+
elif model_type == "14B":
|
| 780 |
+
x = torch.randn(1, 16, 13, 44, 80, dtype=dtype, device=device)
|
| 781 |
+
action = torch.randn(1, 48, 32, dtype=dtype, device=device)
|
| 782 |
+
timestep = torch.randn(1, dtype=dtype, device=device)
|
| 783 |
+
context = torch.randn(1, 512, 4096, dtype=dtype, device=device)
|
| 784 |
+
state = torch.randn(1, 1, 64, dtype=dtype, device=device)
|
| 785 |
+
embodiment_id = torch.zeros(1, dtype=torch.int32, device=device)
|
| 786 |
+
clip_feature = torch.randn(1, 257, 1280, dtype=dtype, device=device)
|
| 787 |
+
y = torch.randn(1, 20, 13, 44, 80, dtype=dtype, device=device)
|
| 788 |
+
return x, action, timestep, context, state, embodiment_id, clip_feature, y
|
| 789 |
+
elif model_type == "ar_14B":
|
| 790 |
+
clip_feature = torch.randn(1, 257, 1280, dtype=dtype, device=device)
|
| 791 |
+
y = torch.randn(1, 20, 2, 44, 80, dtype=dtype, device=device)
|
| 792 |
+
timestep_action = torch.randn(1, 48, dtype=dtype, device=device)
|
| 793 |
+
x = torch.randn(1, 16, 2, 44, 80, dtype=dtype, device=device)
|
| 794 |
+
timestep = torch.randn(1, 2, dtype=dtype, device=device)
|
| 795 |
+
context = torch.randn(1, 512, 4096, dtype=dtype, device=device)
|
| 796 |
+
seq_len = torch.tensor(1760, dtype=torch.int32, device=device)
|
| 797 |
+
action = torch.randn(1, 48, 32, dtype=dtype, device=device)
|
| 798 |
+
state = torch.randn(1, 1, 64, dtype=dtype, device=device)
|
| 799 |
+
embodiment_id = torch.zeros(1, dtype=torch.int32, device=device)
|
| 800 |
+
|
| 801 |
+
num_heads = 40
|
| 802 |
+
head_dim = 5120 // num_heads
|
| 803 |
+
num_layers = 40
|
| 804 |
+
B = 1
|
| 805 |
+
|
| 806 |
+
kv_cache = []
|
| 807 |
+
for _ in range(num_layers):
|
| 808 |
+
kv_cache.append(
|
| 809 |
+
torch.zeros([2, B, 9*880, num_heads, head_dim], dtype=dtype, device=device)
|
| 810 |
+
)
|
| 811 |
+
|
| 812 |
+
crossattn_k_cache = []
|
| 813 |
+
for _ in range(num_layers):
|
| 814 |
+
crossattn_k_cache.append(
|
| 815 |
+
torch.zeros([2, B, 9*880, num_heads, head_dim], dtype=dtype, device=device)
|
| 816 |
+
)
|
| 817 |
+
kv_cache_packed = torch.stack(kv_cache, dim=0)
|
| 818 |
+
crossattn_packed = torch.stack(crossattn_k_cache, dim=0)
|
| 819 |
+
return (x, timestep, context, kv_cache_packed, y, clip_feature, action, timestep_action, state)
|
| 820 |
+
elif model_type == "ar_14B_droid":
|
| 821 |
+
clip_feature = torch.randn(1, 257, 1280, dtype=dtype, device=device)
|
| 822 |
+
y = torch.randn(1, 20, 2, 44, 80, dtype=dtype, device=device)
|
| 823 |
+
timestep_action = torch.randn(1, 24, dtype=dtype, device=device)
|
| 824 |
+
x = torch.randn(1, 16, 2, 44, 80, dtype=dtype, device=device)
|
| 825 |
+
timestep = torch.randn(1, 2, dtype=dtype, device=device)
|
| 826 |
+
context = torch.randn(1, 512, 4096, dtype=dtype, device=device)
|
| 827 |
+
seq_len = torch.tensor(1760, dtype=torch.int32, device=device)
|
| 828 |
+
action = torch.randn(1, 24, 32, dtype=dtype, device=device)
|
| 829 |
+
state = torch.randn(1, 1, 64, dtype=dtype, device=device)
|
| 830 |
+
embodiment_id = torch.zeros(1, dtype=torch.int32, device=device)
|
| 831 |
+
|
| 832 |
+
num_heads = 40
|
| 833 |
+
head_dim = 5120 // num_heads
|
| 834 |
+
num_layers = 40
|
| 835 |
+
B = 1
|
| 836 |
+
|
| 837 |
+
kv_cache = []
|
| 838 |
+
for _ in range(num_layers):
|
| 839 |
+
kv_cache.append(
|
| 840 |
+
torch.zeros([2, B, 9*880, num_heads, head_dim], dtype=dtype, device=device)
|
| 841 |
+
)
|
| 842 |
+
|
| 843 |
+
crossattn_k_cache = []
|
| 844 |
+
for _ in range(num_layers):
|
| 845 |
+
crossattn_k_cache.append(
|
| 846 |
+
torch.zeros([2, B, 9*880, num_heads, head_dim], dtype=dtype, device=device)
|
| 847 |
+
)
|
| 848 |
+
kv_cache_packed = torch.stack(kv_cache, dim=0)
|
| 849 |
+
crossattn_packed = torch.stack(crossattn_k_cache, dim=0)
|
| 850 |
+
return (x, timestep, context, kv_cache_packed, y, clip_feature, action, timestep_action, state)
|
| 851 |
+
|
| 852 |
+
|
groot/vla/__init__.py
ADDED
|
File without changes
|
groot/vla/common/__init__.py
ADDED
|
@@ -0,0 +1 @@
|
|
|
|
|
|
|
| 1 |
+
|
groot/vla/common/utils/__init__.py
ADDED
|
@@ -0,0 +1,3 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
from .data_structure import * # noqa: F403
|
| 2 |
+
from .io import * # noqa: F403
|
| 3 |
+
from .misc import * # noqa: F403
|
groot/vla/common/utils/data_structure/__init__.py
ADDED
|
@@ -0,0 +1,2 @@
|
|
|
|
|
|
|
|
|
|
| 1 |
+
from .shape_utils import * # noqa: F403
|
| 2 |
+
from .tree_utils import * # noqa: F403
|
groot/vla/common/utils/data_structure/shape_utils.py
ADDED
|
@@ -0,0 +1,283 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""
|
| 2 |
+
Shape inference methods
|
| 3 |
+
"""
|
| 4 |
+
|
| 5 |
+
from functools import partial
|
| 6 |
+
import math
|
| 7 |
+
from typing import List, Tuple, Union
|
| 8 |
+
import warnings
|
| 9 |
+
|
| 10 |
+
import numpy as np
|
| 11 |
+
import torch
|
| 12 |
+
|
| 13 |
+
# fmt: off
|
| 14 |
+
__all__ = [
|
| 15 |
+
"shape_convnd",
|
| 16 |
+
"shape_conv1d", "shape_conv2d", "shape_conv3d",
|
| 17 |
+
"shape_transpose_convnd",
|
| 18 |
+
"shape_transpose_conv1d", "shape_transpose_conv2d", "shape_transpose_conv3d",
|
| 19 |
+
"shape_poolnd",
|
| 20 |
+
"shape_maxpool1d", "shape_maxpool2d", "shape_maxpool3d",
|
| 21 |
+
"shape_avgpool1d", "shape_avgpool2d", "shape_avgpool3d",
|
| 22 |
+
"shape_slice",
|
| 23 |
+
"check_shape"
|
| 24 |
+
]
|
| 25 |
+
# fmt: on
|
| 26 |
+
|
| 27 |
+
|
| 28 |
+
def _get_shape(x):
|
| 29 |
+
"single object"
|
| 30 |
+
if isinstance(x, np.ndarray):
|
| 31 |
+
return tuple(x.shape)
|
| 32 |
+
else:
|
| 33 |
+
return tuple(x.size())
|
| 34 |
+
|
| 35 |
+
|
| 36 |
+
def _expands(dim, *xs):
|
| 37 |
+
"repeat vars like kernel and stride to match dim"
|
| 38 |
+
|
| 39 |
+
def _expand(x):
|
| 40 |
+
if isinstance(x, int):
|
| 41 |
+
return (x,) * dim
|
| 42 |
+
else:
|
| 43 |
+
assert len(x) == dim
|
| 44 |
+
return x
|
| 45 |
+
|
| 46 |
+
return map(lambda x: _expand(x), xs)
|
| 47 |
+
|
| 48 |
+
|
| 49 |
+
_HELPER_TENSOR = torch.zeros((1,))
|
| 50 |
+
|
| 51 |
+
|
| 52 |
+
def shape_slice(input_shape, slice):
|
| 53 |
+
"""
|
| 54 |
+
Credit to Adam Paszke for the trick. Shape inference without instantiating
|
| 55 |
+
an actual tensor.
|
| 56 |
+
The key is that `.expand()` does not actually allocate memory
|
| 57 |
+
Still needs to allocate a one-element HELPER_TENSOR.
|
| 58 |
+
"""
|
| 59 |
+
shape = _HELPER_TENSOR.expand(*input_shape)[slice]
|
| 60 |
+
if hasattr(shape, "size"):
|
| 61 |
+
return tuple(shape.size())
|
| 62 |
+
return (1,)
|
| 63 |
+
|
| 64 |
+
|
| 65 |
+
class ShapeSlice:
|
| 66 |
+
"""
|
| 67 |
+
shape_slice inference with easy []-operator
|
| 68 |
+
"""
|
| 69 |
+
|
| 70 |
+
def __init__(self, input_shape):
|
| 71 |
+
self.input_shape = input_shape
|
| 72 |
+
|
| 73 |
+
def __getitem__(self, slice):
|
| 74 |
+
return shape_slice(self.input_shape, slice)
|
| 75 |
+
|
| 76 |
+
|
| 77 |
+
def check_shape(
|
| 78 |
+
value: Union[Tuple, List, torch.Tensor, np.ndarray],
|
| 79 |
+
expected: Union[Tuple, List, torch.Tensor, np.ndarray],
|
| 80 |
+
err_msg="",
|
| 81 |
+
mode="raise",
|
| 82 |
+
):
|
| 83 |
+
"""
|
| 84 |
+
Args:
|
| 85 |
+
value: np array or torch Tensor
|
| 86 |
+
expected:
|
| 87 |
+
- list[int], tuple[int]: if any value is None, will match any dim
|
| 88 |
+
- np array or torch Tensor: must have the same dimensions
|
| 89 |
+
mode:
|
| 90 |
+
- "raise": raise ValueError, shape mismatch
|
| 91 |
+
- "return": returns True if shape matches, otherwise False
|
| 92 |
+
- "warning": warnings.warn
|
| 93 |
+
"""
|
| 94 |
+
assert mode in ["raise", "return", "warning"]
|
| 95 |
+
if torch.is_tensor(value):
|
| 96 |
+
actual_shape = value.size()
|
| 97 |
+
elif hasattr(value, "shape"):
|
| 98 |
+
actual_shape = value.shape
|
| 99 |
+
else:
|
| 100 |
+
assert isinstance(value, (list, tuple))
|
| 101 |
+
actual_shape = value
|
| 102 |
+
assert all(
|
| 103 |
+
isinstance(s, int) for s in actual_shape
|
| 104 |
+
), f"actual shape: {actual_shape} is not a list of ints"
|
| 105 |
+
|
| 106 |
+
if torch.is_tensor(expected):
|
| 107 |
+
expected_shape = expected.size()
|
| 108 |
+
elif hasattr(expected, "shape"):
|
| 109 |
+
expected_shape = expected.shape
|
| 110 |
+
else:
|
| 111 |
+
assert isinstance(expected, (list, tuple))
|
| 112 |
+
expected_shape = expected
|
| 113 |
+
|
| 114 |
+
err_msg = f" for {err_msg}" if err_msg else ""
|
| 115 |
+
|
| 116 |
+
if len(actual_shape) != len(expected_shape):
|
| 117 |
+
err_msg = (
|
| 118 |
+
f"Dimension mismatch{err_msg}: actual shape {actual_shape} "
|
| 119 |
+
f"!= expected shape {expected_shape}."
|
| 120 |
+
)
|
| 121 |
+
if mode == "raise":
|
| 122 |
+
raise ValueError(err_msg)
|
| 123 |
+
elif mode == "warning":
|
| 124 |
+
warnings.warn(err_msg)
|
| 125 |
+
return False
|
| 126 |
+
|
| 127 |
+
for s_a, s_e in zip(actual_shape, expected_shape):
|
| 128 |
+
if s_e is not None and s_a != s_e:
|
| 129 |
+
err_msg = (
|
| 130 |
+
f"Shape mismatch{err_msg}: actual shape {actual_shape} "
|
| 131 |
+
f"!= expected shape {expected_shape}."
|
| 132 |
+
)
|
| 133 |
+
if mode == "raise":
|
| 134 |
+
raise ValueError(err_msg)
|
| 135 |
+
elif mode == "warning":
|
| 136 |
+
warnings.warn(err_msg)
|
| 137 |
+
return False
|
| 138 |
+
return True
|
| 139 |
+
|
| 140 |
+
|
| 141 |
+
def shape_convnd(
|
| 142 |
+
dim,
|
| 143 |
+
input_shape,
|
| 144 |
+
out_channels,
|
| 145 |
+
kernel_size,
|
| 146 |
+
stride=1,
|
| 147 |
+
padding=0,
|
| 148 |
+
dilation=1,
|
| 149 |
+
has_batch=False,
|
| 150 |
+
):
|
| 151 |
+
"""
|
| 152 |
+
http://pytorch.org/docs/nn.html#conv1d
|
| 153 |
+
http://pytorch.org/docs/nn.html#conv2d
|
| 154 |
+
http://pytorch.org/docs/nn.html#conv3d
|
| 155 |
+
|
| 156 |
+
Args:
|
| 157 |
+
dim: supports 1D to 3D
|
| 158 |
+
input_shape:
|
| 159 |
+
- 1D: [channel, length]
|
| 160 |
+
- 2D: [channel, height, width]
|
| 161 |
+
- 3D: [channel, depth, height, width]
|
| 162 |
+
has_batch: whether the first dim is batch size or not
|
| 163 |
+
"""
|
| 164 |
+
if has_batch:
|
| 165 |
+
assert (
|
| 166 |
+
len(input_shape) == dim + 2
|
| 167 |
+
), "input shape with batch should be {}-dimensional".format(dim + 2)
|
| 168 |
+
else:
|
| 169 |
+
assert (
|
| 170 |
+
len(input_shape) == dim + 1
|
| 171 |
+
), "input shape without batch should be {}-dimensional".format(dim + 1)
|
| 172 |
+
if stride is None:
|
| 173 |
+
# for pooling convention in PyTorch
|
| 174 |
+
stride = kernel_size
|
| 175 |
+
kernel_size, stride, padding, dilation = _expands(dim, kernel_size, stride, padding, dilation)
|
| 176 |
+
if has_batch:
|
| 177 |
+
batch = input_shape[0]
|
| 178 |
+
input_shape = input_shape[1:]
|
| 179 |
+
else:
|
| 180 |
+
batch = None
|
| 181 |
+
_, *img = input_shape
|
| 182 |
+
new_img_shape = [
|
| 183 |
+
math.floor(
|
| 184 |
+
(img[i] + 2 * padding[i] - dilation[i] * (kernel_size[i] - 1) - 1) // stride[i] + 1
|
| 185 |
+
)
|
| 186 |
+
for i in range(dim)
|
| 187 |
+
]
|
| 188 |
+
return ((batch,) if has_batch else ()) + (out_channels, *new_img_shape)
|
| 189 |
+
|
| 190 |
+
|
| 191 |
+
def shape_poolnd(
|
| 192 |
+
dim, input_shape, kernel_size, stride=None, padding=0, dilation=1, has_batch=False
|
| 193 |
+
):
|
| 194 |
+
"""
|
| 195 |
+
The only difference from infer_shape_convnd is that `stride` default is None
|
| 196 |
+
"""
|
| 197 |
+
if has_batch:
|
| 198 |
+
out_channels = input_shape[1]
|
| 199 |
+
else:
|
| 200 |
+
out_channels = input_shape[0]
|
| 201 |
+
return shape_convnd(
|
| 202 |
+
dim,
|
| 203 |
+
input_shape,
|
| 204 |
+
out_channels,
|
| 205 |
+
kernel_size,
|
| 206 |
+
stride,
|
| 207 |
+
padding,
|
| 208 |
+
dilation,
|
| 209 |
+
has_batch,
|
| 210 |
+
)
|
| 211 |
+
|
| 212 |
+
|
| 213 |
+
def shape_transpose_convnd(
|
| 214 |
+
dim,
|
| 215 |
+
input_shape,
|
| 216 |
+
out_channels,
|
| 217 |
+
kernel_size,
|
| 218 |
+
stride=1,
|
| 219 |
+
padding=0,
|
| 220 |
+
output_padding=0,
|
| 221 |
+
dilation=1,
|
| 222 |
+
has_batch=False,
|
| 223 |
+
):
|
| 224 |
+
"""
|
| 225 |
+
http://pytorch.org/docs/nn.html#convtranspose1d
|
| 226 |
+
http://pytorch.org/docs/nn.html#convtranspose2d
|
| 227 |
+
http://pytorch.org/docs/nn.html#convtranspose3d
|
| 228 |
+
|
| 229 |
+
Args:
|
| 230 |
+
dim: supports 1D to 3D
|
| 231 |
+
input_shape:
|
| 232 |
+
- 1D: [channel, length]
|
| 233 |
+
- 2D: [channel, height, width]
|
| 234 |
+
- 3D: [channel, depth, height, width]
|
| 235 |
+
has_batch: whether the first dim is batch size or not
|
| 236 |
+
"""
|
| 237 |
+
if has_batch:
|
| 238 |
+
assert (
|
| 239 |
+
len(input_shape) == dim + 2
|
| 240 |
+
), "input shape with batch should be {}-dimensional".format(dim + 2)
|
| 241 |
+
else:
|
| 242 |
+
assert (
|
| 243 |
+
len(input_shape) == dim + 1
|
| 244 |
+
), "input shape without batch should be {}-dimensional".format(dim + 1)
|
| 245 |
+
kernel_size, stride, padding, output_padding, dilation = _expands(
|
| 246 |
+
dim, kernel_size, stride, padding, output_padding, dilation
|
| 247 |
+
)
|
| 248 |
+
if has_batch:
|
| 249 |
+
batch = input_shape[0]
|
| 250 |
+
input_shape = input_shape[1:]
|
| 251 |
+
else:
|
| 252 |
+
batch = None
|
| 253 |
+
_, *img = input_shape
|
| 254 |
+
new_img_shape = [
|
| 255 |
+
(img[i] - 1) * stride[i] - 2 * padding[i] + kernel_size[i] + output_padding[i]
|
| 256 |
+
for i in range(dim)
|
| 257 |
+
]
|
| 258 |
+
return ((batch,) if has_batch else ()) + (out_channels, *new_img_shape)
|
| 259 |
+
|
| 260 |
+
|
| 261 |
+
shape_conv1d = partial(shape_convnd, 1)
|
| 262 |
+
shape_conv2d = partial(shape_convnd, 2)
|
| 263 |
+
shape_conv3d = partial(shape_convnd, 3)
|
| 264 |
+
|
| 265 |
+
|
| 266 |
+
shape_transpose_conv1d = partial(shape_transpose_convnd, 1)
|
| 267 |
+
shape_transpose_conv2d = partial(shape_transpose_convnd, 2)
|
| 268 |
+
shape_transpose_conv3d = partial(shape_transpose_convnd, 3)
|
| 269 |
+
|
| 270 |
+
|
| 271 |
+
shape_maxpool1d = partial(shape_poolnd, 1)
|
| 272 |
+
shape_maxpool2d = partial(shape_poolnd, 2)
|
| 273 |
+
shape_maxpool3d = partial(shape_poolnd, 3)
|
| 274 |
+
|
| 275 |
+
|
| 276 |
+
"""
|
| 277 |
+
http://pytorch.org/docs/nn.html#avgpool1d
|
| 278 |
+
http://pytorch.org/docs/nn.html#avgpool2d
|
| 279 |
+
http://pytorch.org/docs/nn.html#avgpool3d
|
| 280 |
+
"""
|
| 281 |
+
shape_avgpool1d = partial(shape_maxpool1d, dilation=1)
|
| 282 |
+
shape_avgpool2d = partial(shape_maxpool2d, dilation=1)
|
| 283 |
+
shape_avgpool3d = partial(shape_maxpool3d, dilation=1)
|
groot/vla/common/utils/data_structure/tree_utils.py
ADDED
|
@@ -0,0 +1,219 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""
|
| 2 |
+
Utils to handle nested data structures
|
| 3 |
+
|
| 4 |
+
Install dm_tree first:
|
| 5 |
+
https://tree.readthedocs.io/en/latest/api.html
|
| 6 |
+
"""
|
| 7 |
+
|
| 8 |
+
import collections
|
| 9 |
+
from typing import Any, Iterable, List, Tuple, TypeVar
|
| 10 |
+
|
| 11 |
+
import numpy as np
|
| 12 |
+
|
| 13 |
+
try:
|
| 14 |
+
import tree
|
| 15 |
+
|
| 16 |
+
except ImportError:
|
| 17 |
+
raise ImportError("Please install dm_tree first: `pip install dm_tree`")
|
| 18 |
+
|
| 19 |
+
|
| 20 |
+
def is_sequence(obj):
|
| 21 |
+
"""
|
| 22 |
+
Returns:
|
| 23 |
+
True if the sequence is a collections.Sequence and not a string.
|
| 24 |
+
"""
|
| 25 |
+
return isinstance(obj, collections.abc.Sequence) and not isinstance(obj, str)
|
| 26 |
+
|
| 27 |
+
|
| 28 |
+
def is_mapping(obj):
|
| 29 |
+
"""
|
| 30 |
+
Returns:
|
| 31 |
+
True if the sequence is a collections.Mapping
|
| 32 |
+
"""
|
| 33 |
+
return isinstance(obj, collections.abc.Mapping)
|
| 34 |
+
|
| 35 |
+
|
| 36 |
+
def tree_value_at_path(obj, paths: Tuple):
|
| 37 |
+
try:
|
| 38 |
+
for p in paths:
|
| 39 |
+
obj = obj[p]
|
| 40 |
+
return obj
|
| 41 |
+
except Exception as e:
|
| 42 |
+
raise ValueError(f"{e}\n\n-- Incorrect nested path {paths} for object: {obj}.")
|
| 43 |
+
|
| 44 |
+
|
| 45 |
+
def tree_assign_at_path(obj, paths: Tuple, value):
|
| 46 |
+
try:
|
| 47 |
+
for p in paths[:-1]:
|
| 48 |
+
obj = obj[p]
|
| 49 |
+
if len(paths) > 0:
|
| 50 |
+
obj[paths[-1]] = value
|
| 51 |
+
except Exception as e:
|
| 52 |
+
raise ValueError(f"{e}\n\n-- Incorrect nested path {paths} for object: {obj}.")
|
| 53 |
+
|
| 54 |
+
|
| 55 |
+
def copy_non_leaf(obj):
|
| 56 |
+
"""
|
| 57 |
+
Deepcopy the nested structure, but does NOT copy the leaf values like Tensors
|
| 58 |
+
"""
|
| 59 |
+
return tree.map_structure(lambda x: x, obj)
|
| 60 |
+
|
| 61 |
+
|
| 62 |
+
# =======================================================================
|
| 63 |
+
# Copyright 2018 DeepMind Technologies Limited. All rights reserved.
|
| 64 |
+
#
|
| 65 |
+
# Licensed under the Apache License, Version 2.0 (the "License");
|
| 66 |
+
# you may not use this file except in compliance with the License.
|
| 67 |
+
# You may obtain a copy of the License at
|
| 68 |
+
#
|
| 69 |
+
# http://www.apache.org/licenses/LICENSE-2.0
|
| 70 |
+
#
|
| 71 |
+
# Unless required by applicable law or agreed to in writing, software
|
| 72 |
+
# distributed under the License is distributed on an "AS IS" BASIS,
|
| 73 |
+
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
| 74 |
+
# See the License for the specific language governing permissions and
|
| 75 |
+
# limitations under the License.
|
| 76 |
+
|
| 77 |
+
# Tensor framework-agnostic utilities for manipulating nested structures.
|
| 78 |
+
|
| 79 |
+
ElementType = TypeVar("ElementType")
|
| 80 |
+
|
| 81 |
+
|
| 82 |
+
def fast_map_structure(func, *structure):
|
| 83 |
+
"""Faster map_structure implementation which skips some error checking."""
|
| 84 |
+
flat_structure = (tree.flatten(s) for s in structure)
|
| 85 |
+
entries = zip(*flat_structure)
|
| 86 |
+
# Arbitrarily choose one of the structures of the original sequence (the last)
|
| 87 |
+
# to match the structure for the flattened sequence.
|
| 88 |
+
return tree.unflatten_as(structure[-1], [func(*x) for x in entries])
|
| 89 |
+
|
| 90 |
+
|
| 91 |
+
def stack_sequence_fields(sequence: Iterable[ElementType]) -> ElementType:
|
| 92 |
+
"""Stacks a list of identically nested objects.
|
| 93 |
+
|
| 94 |
+
This takes a sequence of identically nested objects and returns a single
|
| 95 |
+
nested object whose ith leaf is a stacked numpy array of the corresponding
|
| 96 |
+
ith leaf from each element of the sequence.
|
| 97 |
+
|
| 98 |
+
For example, if `sequence` is:
|
| 99 |
+
|
| 100 |
+
```python
|
| 101 |
+
[{
|
| 102 |
+
'action': np.array([1.0]),
|
| 103 |
+
'observation': (np.array([0.0, 1.0, 2.0]),),
|
| 104 |
+
'reward': 1.0
|
| 105 |
+
}, {
|
| 106 |
+
'action': np.array([0.5]),
|
| 107 |
+
'observation': (np.array([1.0, 2.0, 3.0]),),
|
| 108 |
+
'reward': 0.0
|
| 109 |
+
}, {
|
| 110 |
+
'action': np.array([0.3]),1
|
| 111 |
+
'observation': (np.array([2.0, 3.0, 4.0]),),
|
| 112 |
+
'reward': 0.5
|
| 113 |
+
}]
|
| 114 |
+
```
|
| 115 |
+
|
| 116 |
+
Then this function will return:
|
| 117 |
+
|
| 118 |
+
```python
|
| 119 |
+
{
|
| 120 |
+
'action': np.array([....]) # array shape = [3 x 1]
|
| 121 |
+
'observation': (np.array([...]),) # array shape = [3 x 3]
|
| 122 |
+
'reward': np.array([...]) # array shape = [3]
|
| 123 |
+
}
|
| 124 |
+
```
|
| 125 |
+
|
| 126 |
+
Note that the 'observation' entry in the above example has two levels of
|
| 127 |
+
nesting, i.e it is a tuple of arrays.
|
| 128 |
+
|
| 129 |
+
Args:
|
| 130 |
+
sequence: a list of identically nested objects.
|
| 131 |
+
|
| 132 |
+
Returns:
|
| 133 |
+
A nested object with numpy.
|
| 134 |
+
|
| 135 |
+
Raises:
|
| 136 |
+
ValueError: If `sequence` is an empty sequence.
|
| 137 |
+
"""
|
| 138 |
+
# Handle empty input sequences.
|
| 139 |
+
if not sequence:
|
| 140 |
+
raise ValueError("Input sequence must not be empty")
|
| 141 |
+
|
| 142 |
+
# Default to asarray when arrays don't have the same shape to be compatible
|
| 143 |
+
# with old behaviour.
|
| 144 |
+
try:
|
| 145 |
+
return fast_map_structure(lambda *values: np.stack(values), *sequence)
|
| 146 |
+
except ValueError:
|
| 147 |
+
return fast_map_structure(lambda *values: np.asarray(values), *sequence)
|
| 148 |
+
|
| 149 |
+
|
| 150 |
+
def unstack_sequence_fields(struct: ElementType, batch_size: int) -> List[ElementType]:
|
| 151 |
+
"""Converts a struct of batched arrays to a list of structs.
|
| 152 |
+
|
| 153 |
+
This is effectively the inverse of `stack_sequence_fields`.
|
| 154 |
+
|
| 155 |
+
Args:
|
| 156 |
+
struct: An (arbitrarily nested) structure of arrays.
|
| 157 |
+
batch_size: The length of the leading dimension of each array in the struct.
|
| 158 |
+
This is assumed to be static and known.
|
| 159 |
+
|
| 160 |
+
Returns:
|
| 161 |
+
A list of structs with the same structure as `struct`, where each leaf node
|
| 162 |
+
is an unbatched element of the original leaf node.
|
| 163 |
+
"""
|
| 164 |
+
|
| 165 |
+
return [tree.map_structure(lambda s, i=i: s[i], struct) for i in range(batch_size)]
|
| 166 |
+
|
| 167 |
+
|
| 168 |
+
def broadcast_structures(*args: Any) -> Any:
|
| 169 |
+
"""Returns versions of the arguments that give them the same nested structure.
|
| 170 |
+
|
| 171 |
+
Any nested items in *args must have the same structure.
|
| 172 |
+
|
| 173 |
+
Any non-nested item will be replaced with a nested version that shares that
|
| 174 |
+
structure. The leaves will all be references to the same original non-nested
|
| 175 |
+
item.
|
| 176 |
+
|
| 177 |
+
If all *args are nested, or all *args are non-nested, this function will
|
| 178 |
+
return *args unchanged.
|
| 179 |
+
|
| 180 |
+
Example:
|
| 181 |
+
```
|
| 182 |
+
a = ('a', 'b')
|
| 183 |
+
b = 'c'
|
| 184 |
+
tree_a, tree_b = broadcast_structure(a, b)
|
| 185 |
+
tree_a
|
| 186 |
+
> ('a', 'b')
|
| 187 |
+
tree_b
|
| 188 |
+
> ('c', 'c')
|
| 189 |
+
```
|
| 190 |
+
|
| 191 |
+
Args:
|
| 192 |
+
*args: A Sequence of nested or non-nested items.
|
| 193 |
+
|
| 194 |
+
Returns:
|
| 195 |
+
`*args`, except with all items sharing the same nest structure.
|
| 196 |
+
"""
|
| 197 |
+
if not args:
|
| 198 |
+
return
|
| 199 |
+
|
| 200 |
+
reference_tree = None
|
| 201 |
+
for arg in args:
|
| 202 |
+
if tree.is_nested(arg):
|
| 203 |
+
reference_tree = arg
|
| 204 |
+
break
|
| 205 |
+
|
| 206 |
+
if reference_tree is None:
|
| 207 |
+
reference_tree = args[0]
|
| 208 |
+
|
| 209 |
+
def mirror_structure(value, reference_tree):
|
| 210 |
+
if tree.is_nested(value):
|
| 211 |
+
# Use check_types=True so that the types of the trees we construct aren't
|
| 212 |
+
# dependent on our arbitrary choice of which nested arg to use as the
|
| 213 |
+
# reference_tree.
|
| 214 |
+
tree.assert_same_structure(value, reference_tree, check_types=True)
|
| 215 |
+
return value
|
| 216 |
+
else:
|
| 217 |
+
return tree.map_structure(lambda _: value, reference_tree)
|
| 218 |
+
|
| 219 |
+
return tuple(mirror_structure(arg, reference_tree) for arg in args)
|
groot/vla/common/utils/io/__init__.py
ADDED
|
@@ -0,0 +1,6 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
from .config_utils import * # noqa: F403
|
| 2 |
+
from .file_utils import * # noqa: F403
|
| 3 |
+
from .hdf5_utils import * # noqa: F403
|
| 4 |
+
from .json_utils import * # noqa: F403
|
| 5 |
+
from .print_utils import * # noqa: F403
|
| 6 |
+
from .termcolor import * # noqa: F403
|
groot/vla/common/utils/io/config_utils.py
ADDED
|
@@ -0,0 +1,260 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
from copy import deepcopy
|
| 2 |
+
import importlib.resources
|
| 3 |
+
import os
|
| 4 |
+
import sys
|
| 5 |
+
|
| 6 |
+
import hydra
|
| 7 |
+
from omegaconf import DictConfig, OmegaConf
|
| 8 |
+
import tree
|
| 9 |
+
|
| 10 |
+
from ..misc.functional_utils import call_once, is_mapping, is_sequence, meta_decorator
|
| 11 |
+
from .print_utils import to_scientific_str
|
| 12 |
+
|
| 13 |
+
_CLASS_REGISTRY = {} # for instantiation
|
| 14 |
+
|
| 15 |
+
|
| 16 |
+
def resource_file_path(pkg_name, fname) -> str:
|
| 17 |
+
with importlib.resources.path(pkg_name, fname) as p:
|
| 18 |
+
return str(p)
|
| 19 |
+
|
| 20 |
+
|
| 21 |
+
def print_config(cfg: DictConfig):
|
| 22 |
+
print(cfg.pretty(resolve=True))
|
| 23 |
+
|
| 24 |
+
|
| 25 |
+
def is_hydra_initialized():
|
| 26 |
+
return hydra.utils.HydraConfig.initialized()
|
| 27 |
+
|
| 28 |
+
|
| 29 |
+
def hydra_config():
|
| 30 |
+
# https://github.com/facebookresearch/hydra/issues/377
|
| 31 |
+
# HydraConfig() is a singleton
|
| 32 |
+
if is_hydra_initialized():
|
| 33 |
+
return hydra.utils.HydraConfig().cfg.hydra
|
| 34 |
+
else:
|
| 35 |
+
return None
|
| 36 |
+
|
| 37 |
+
|
| 38 |
+
def hydra_override_arg_list() -> list[str]:
|
| 39 |
+
"""
|
| 40 |
+
Returns:
|
| 41 |
+
list ["lr=0.2", "batch=64", ...]
|
| 42 |
+
"""
|
| 43 |
+
if is_hydra_initialized():
|
| 44 |
+
return hydra_config().overrides.task
|
| 45 |
+
else:
|
| 46 |
+
return []
|
| 47 |
+
|
| 48 |
+
|
| 49 |
+
def hydra_override_name():
|
| 50 |
+
if is_hydra_initialized():
|
| 51 |
+
return hydra_config().job.override_dirname
|
| 52 |
+
else:
|
| 53 |
+
return ""
|
| 54 |
+
|
| 55 |
+
|
| 56 |
+
def hydra_original_dir(*subpaths):
|
| 57 |
+
return os.path.join(hydra.utils.get_original_cwd(), *subpaths)
|
| 58 |
+
|
| 59 |
+
|
| 60 |
+
@call_once(on_second_call="noop")
|
| 61 |
+
def register_omegaconf_resolvers():
|
| 62 |
+
import numpy as np
|
| 63 |
+
|
| 64 |
+
OmegaConf.register_new_resolver("scientific", lambda v, i=0: to_scientific_str(v, i))
|
| 65 |
+
OmegaConf.register_new_resolver("_optional", lambda v: f"_{v}" if v else "")
|
| 66 |
+
OmegaConf.register_new_resolver("optional_", lambda v: f"{v}_" if v else "")
|
| 67 |
+
OmegaConf.register_new_resolver("_optional_", lambda v: f"_{v}_" if v else "")
|
| 68 |
+
OmegaConf.register_new_resolver("__optional", lambda v: f"__{v}" if v else "")
|
| 69 |
+
OmegaConf.register_new_resolver("optional__", lambda v: f"{v}__" if v else "")
|
| 70 |
+
OmegaConf.register_new_resolver("__optional__", lambda v: f"__{v}__" if v else "")
|
| 71 |
+
OmegaConf.register_new_resolver("iftrue", lambda cond, v_default: cond if cond else v_default)
|
| 72 |
+
OmegaConf.register_new_resolver("ifelse", lambda cond, v1, v2="": v1 if cond else v2)
|
| 73 |
+
OmegaConf.register_new_resolver(
|
| 74 |
+
"ifequal", lambda query, key, v1, v2: v1 if query == key else v2
|
| 75 |
+
)
|
| 76 |
+
OmegaConf.register_new_resolver("intbool", lambda cond: 1 if cond else 0)
|
| 77 |
+
OmegaConf.register_new_resolver("mult", lambda *x: np.prod(x).tolist())
|
| 78 |
+
OmegaConf.register_new_resolver("add", lambda *x: sum(x))
|
| 79 |
+
OmegaConf.register_new_resolver("div", lambda x, y: x / y)
|
| 80 |
+
OmegaConf.register_new_resolver("intdiv", lambda x, y: x // y)
|
| 81 |
+
|
| 82 |
+
# try each key until the key exists. Useful for multiple classes that have different
|
| 83 |
+
# names for the same key
|
| 84 |
+
def _try_key(cfg, *keys):
|
| 85 |
+
for k in keys:
|
| 86 |
+
if k in cfg:
|
| 87 |
+
return cfg[k]
|
| 88 |
+
raise KeyError(f"no key in {keys} is valid")
|
| 89 |
+
|
| 90 |
+
OmegaConf.register_new_resolver("trykey", _try_key)
|
| 91 |
+
# replace `resnet.gn.ws` -> `resnet_gn_ws`, because omegaconf doesn't support
|
| 92 |
+
# keys with dots. Useful for generating run name with dots
|
| 93 |
+
OmegaConf.register_new_resolver("underscore_to_dots", lambda s: s.replace("_", "."))
|
| 94 |
+
|
| 95 |
+
def _no_instantiate(cfg):
|
| 96 |
+
cfg = deepcopy(cfg)
|
| 97 |
+
cfg[_NO_INSTANTIATE] = True
|
| 98 |
+
return cfg
|
| 99 |
+
|
| 100 |
+
OmegaConf.register_new_resolver("no_instantiate", _no_instantiate)
|
| 101 |
+
|
| 102 |
+
|
| 103 |
+
# ========================================================
|
| 104 |
+
# ================== Instantiation tools ================
|
| 105 |
+
# ========================================================
|
| 106 |
+
|
| 107 |
+
|
| 108 |
+
def register_callable(name, class_type):
|
| 109 |
+
if isinstance(class_type, str):
|
| 110 |
+
class_type, name = name, class_type
|
| 111 |
+
assert callable(class_type)
|
| 112 |
+
_CLASS_REGISTRY[name] = class_type
|
| 113 |
+
|
| 114 |
+
|
| 115 |
+
@meta_decorator
|
| 116 |
+
def register_class(cls, alias=None):
|
| 117 |
+
"""
|
| 118 |
+
Decorator
|
| 119 |
+
"""
|
| 120 |
+
assert callable(cls)
|
| 121 |
+
_CLASS_REGISTRY[cls.__name__] = cls
|
| 122 |
+
if alias:
|
| 123 |
+
assert is_sequence(alias)
|
| 124 |
+
for a in alias:
|
| 125 |
+
_CLASS_REGISTRY[str(a)] = cls
|
| 126 |
+
return cls
|
| 127 |
+
|
| 128 |
+
|
| 129 |
+
def omegaconf_to_dict(cfg, resolve: bool = True, enum_to_str: bool = False):
|
| 130 |
+
"""
|
| 131 |
+
Convert arbitrary nested omegaconf objects to primitive containers
|
| 132 |
+
|
| 133 |
+
WARNING: cannot use tree lib because it gets confused on DictConfig and ListConfig
|
| 134 |
+
"""
|
| 135 |
+
kw = dict(resolve=resolve, enum_to_str=enum_to_str)
|
| 136 |
+
if OmegaConf.is_config(cfg):
|
| 137 |
+
return OmegaConf.to_container(cfg, **kw)
|
| 138 |
+
elif is_sequence(cfg):
|
| 139 |
+
return type(cfg)(omegaconf_to_dict(c, **kw) for c in cfg)
|
| 140 |
+
elif is_mapping(cfg):
|
| 141 |
+
return {k: omegaconf_to_dict(c, **kw) for k, c in cfg.items()}
|
| 142 |
+
else:
|
| 143 |
+
return cfg
|
| 144 |
+
|
| 145 |
+
|
| 146 |
+
def omegaconf_save(cfg, *paths: str, resolve: bool = True):
|
| 147 |
+
"""
|
| 148 |
+
Save omegaconf to yaml
|
| 149 |
+
"""
|
| 150 |
+
from .file_utils import f_join
|
| 151 |
+
|
| 152 |
+
OmegaConf.save(cfg, f_join(*paths), resolve=resolve)
|
| 153 |
+
|
| 154 |
+
|
| 155 |
+
def get_class(path):
|
| 156 |
+
"""
|
| 157 |
+
First try to find the class in the registry first,
|
| 158 |
+
if it doesn't exist, use importlib to locate it
|
| 159 |
+
"""
|
| 160 |
+
if path in _CLASS_REGISTRY:
|
| 161 |
+
return _CLASS_REGISTRY[path]
|
| 162 |
+
else:
|
| 163 |
+
assert "." in path, (
|
| 164 |
+
f"Because {path} is not found in class registry, " f"it must be a full module path"
|
| 165 |
+
)
|
| 166 |
+
try:
|
| 167 |
+
from importlib import import_module
|
| 168 |
+
|
| 169 |
+
module_path, _, class_name = path.rpartition(".")
|
| 170 |
+
mod = import_module(module_path)
|
| 171 |
+
try:
|
| 172 |
+
class_type = getattr(mod, class_name)
|
| 173 |
+
except AttributeError:
|
| 174 |
+
raise ImportError("Class {} is not in module {}".format(class_name, module_path))
|
| 175 |
+
return class_type
|
| 176 |
+
except ValueError as e:
|
| 177 |
+
print("Error initializing class " + path, file=sys.stderr)
|
| 178 |
+
raise e
|
| 179 |
+
|
| 180 |
+
|
| 181 |
+
_DELETE_ARG = "__delete__"
|
| 182 |
+
_NO_INSTANTIATE = "__no_instantiate__" # return config as-is
|
| 183 |
+
_OMEGA_MISSING = "???"
|
| 184 |
+
|
| 185 |
+
|
| 186 |
+
def _get_instantiate_params(cfg, kwargs=None):
|
| 187 |
+
params = cfg
|
| 188 |
+
f_args, f_kwargs = (), {}
|
| 189 |
+
for k, value in params.items():
|
| 190 |
+
if k in ["cls", "class"]:
|
| 191 |
+
continue
|
| 192 |
+
elif k == "*args":
|
| 193 |
+
assert is_sequence(value), '"*args" value must be a sequence'
|
| 194 |
+
f_args = list(value)
|
| 195 |
+
continue
|
| 196 |
+
if value == _OMEGA_MISSING:
|
| 197 |
+
if kwargs and k in kwargs:
|
| 198 |
+
value = kwargs[k]
|
| 199 |
+
else:
|
| 200 |
+
raise ValueError(f'Missing required keyword arg "{k}" in cfg: {cfg}')
|
| 201 |
+
if value == _DELETE_ARG:
|
| 202 |
+
continue
|
| 203 |
+
else:
|
| 204 |
+
f_kwargs[k] = value
|
| 205 |
+
return f_args, f_kwargs
|
| 206 |
+
|
| 207 |
+
|
| 208 |
+
def _instantiate_single(cfg):
|
| 209 |
+
if is_mapping(cfg) and ("cls" in cfg or "class" in cfg):
|
| 210 |
+
assert bool("cls" in cfg) != bool("class" in cfg), (
|
| 211 |
+
"to instantiate from config, "
|
| 212 |
+
'one and only one of "cls" or "class" key should be provided'
|
| 213 |
+
)
|
| 214 |
+
if _NO_INSTANTIATE in cfg:
|
| 215 |
+
no_instantiate = cfg.pop(_NO_INSTANTIATE)
|
| 216 |
+
if no_instantiate:
|
| 217 |
+
cfg = deepcopy(cfg)
|
| 218 |
+
return cfg
|
| 219 |
+
else:
|
| 220 |
+
return _instantiate_single(cfg)
|
| 221 |
+
|
| 222 |
+
cls = cfg.get("class", cfg.get("cls"))
|
| 223 |
+
args, kwargs = _get_instantiate_params(cfg)
|
| 224 |
+
try:
|
| 225 |
+
class_type = get_class(cls)
|
| 226 |
+
return class_type(*args, **kwargs)
|
| 227 |
+
except Exception as e:
|
| 228 |
+
raise RuntimeError(f"Error instantiating {cls}: {e}")
|
| 229 |
+
else:
|
| 230 |
+
return None
|
| 231 |
+
|
| 232 |
+
|
| 233 |
+
def instantiate(_cfg_, **kwargs):
|
| 234 |
+
"""
|
| 235 |
+
Any dict with "cls" or "class" key is considered instantiable.
|
| 236 |
+
|
| 237 |
+
Any key that has the special value "__delete__"
|
| 238 |
+
will not be passed to the constructor
|
| 239 |
+
|
| 240 |
+
**kwargs only apply to the top level object if it's a dict, otherwise raise error
|
| 241 |
+
"""
|
| 242 |
+
assert OmegaConf.is_config(_cfg_) or isinstance(_cfg_, (list, tuple)) or is_mapping(_cfg_), (
|
| 243 |
+
'"cfg" must be a dict, list, tuple, or OmegaConf config to be instantiated. '
|
| 244 |
+
f"Current its type is {type(_cfg_)}"
|
| 245 |
+
)
|
| 246 |
+
|
| 247 |
+
_cfg_ = omegaconf_to_dict(_cfg_, resolve=True)
|
| 248 |
+
|
| 249 |
+
if kwargs:
|
| 250 |
+
if is_mapping(_cfg_):
|
| 251 |
+
_cfg_ = _cfg_.copy()
|
| 252 |
+
_cfg_.update(kwargs)
|
| 253 |
+
_cfg_ = {k: v for k, v in _cfg_.items() if v != _DELETE_ARG}
|
| 254 |
+
else:
|
| 255 |
+
raise RuntimeError(
|
| 256 |
+
f"**kwargs specified, but the top-level cfg is not a dict. "
|
| 257 |
+
f"It has type {type(_cfg_)}"
|
| 258 |
+
)
|
| 259 |
+
|
| 260 |
+
return tree.traverse(_instantiate_single, _cfg_, top_down=False)
|
groot/vla/common/utils/io/file_utils.py
ADDED
|
@@ -0,0 +1,707 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""
|
| 2 |
+
File system utils.
|
| 3 |
+
"""
|
| 4 |
+
|
| 5 |
+
import glob
|
| 6 |
+
import os
|
| 7 |
+
import pickle
|
| 8 |
+
import shutil
|
| 9 |
+
import sys
|
| 10 |
+
from typing import Callable, Union
|
| 11 |
+
|
| 12 |
+
from ..data_structure.tree_utils import is_sequence
|
| 13 |
+
|
| 14 |
+
__all__ = [
|
| 15 |
+
"create_tar",
|
| 16 |
+
"dump_pickle",
|
| 17 |
+
"dump_text",
|
| 18 |
+
"dump_text_lines",
|
| 19 |
+
"extract_tar",
|
| 20 |
+
"f_add_ext",
|
| 21 |
+
"f_append_before_ext",
|
| 22 |
+
"f_copy",
|
| 23 |
+
"f_copytree",
|
| 24 |
+
"f_exists",
|
| 25 |
+
"f_expand",
|
| 26 |
+
"f_ext",
|
| 27 |
+
"f_glob",
|
| 28 |
+
"f_has_ext",
|
| 29 |
+
"f_join",
|
| 30 |
+
"f_listdir",
|
| 31 |
+
"f_mkdir",
|
| 32 |
+
"f_mkdir_in_path",
|
| 33 |
+
"f_move",
|
| 34 |
+
"f_not_empty",
|
| 35 |
+
"f_remove",
|
| 36 |
+
"f_size",
|
| 37 |
+
"f_split_path",
|
| 38 |
+
"f_time",
|
| 39 |
+
"get_dir",
|
| 40 |
+
"get_file_lock",
|
| 41 |
+
"get_package_root",
|
| 42 |
+
"get_parent_dir",
|
| 43 |
+
"get_script_dir",
|
| 44 |
+
"get_script_file_name",
|
| 45 |
+
"get_script_self_path",
|
| 46 |
+
"host_id",
|
| 47 |
+
"host_name",
|
| 48 |
+
"insert_before_ext",
|
| 49 |
+
"is_abs_path",
|
| 50 |
+
"is_dir",
|
| 51 |
+
"is_file",
|
| 52 |
+
"is_relative_path",
|
| 53 |
+
"last_part_in_path",
|
| 54 |
+
"load_pickle",
|
| 55 |
+
"load_text",
|
| 56 |
+
"load_text_lines",
|
| 57 |
+
"md5_checksum",
|
| 58 |
+
"move_with_backup",
|
| 59 |
+
"next_available_file_name",
|
| 60 |
+
"owner_name",
|
| 61 |
+
"pickle_dump",
|
| 62 |
+
"pickle_load",
|
| 63 |
+
"read_text",
|
| 64 |
+
"read_text_lines",
|
| 65 |
+
"text_dump",
|
| 66 |
+
"text_load",
|
| 67 |
+
"timestamp_file_name",
|
| 68 |
+
"utf_open",
|
| 69 |
+
"write_text",
|
| 70 |
+
"write_text_lines",
|
| 71 |
+
]
|
| 72 |
+
|
| 73 |
+
f_ext = os.path.splitext
|
| 74 |
+
|
| 75 |
+
f_size = os.path.getsize
|
| 76 |
+
|
| 77 |
+
is_file = os.path.isfile
|
| 78 |
+
|
| 79 |
+
is_dir = os.path.isdir
|
| 80 |
+
|
| 81 |
+
get_dir = os.path.dirname
|
| 82 |
+
|
| 83 |
+
|
| 84 |
+
def owner_name(filepath):
|
| 85 |
+
"""
|
| 86 |
+
Returns: file owner name, unix only
|
| 87 |
+
"""
|
| 88 |
+
import pwd
|
| 89 |
+
|
| 90 |
+
return pwd.getpwuid(os.stat(filepath).st_uid).pw_name
|
| 91 |
+
|
| 92 |
+
|
| 93 |
+
def host_name():
|
| 94 |
+
"Get host name, alias with ``socket.gethostname()``"
|
| 95 |
+
from socket import gethostname
|
| 96 |
+
|
| 97 |
+
return gethostname()
|
| 98 |
+
|
| 99 |
+
|
| 100 |
+
def host_id():
|
| 101 |
+
"""
|
| 102 |
+
Returns: first part of hostname up to '.'
|
| 103 |
+
"""
|
| 104 |
+
return host_name().split(".")[0]
|
| 105 |
+
|
| 106 |
+
|
| 107 |
+
def utf_open(fname, mode):
|
| 108 |
+
"""
|
| 109 |
+
Wrapper for codecs.open
|
| 110 |
+
"""
|
| 111 |
+
import codecs
|
| 112 |
+
|
| 113 |
+
return codecs.open(fname, mode=mode, encoding="utf-8")
|
| 114 |
+
|
| 115 |
+
|
| 116 |
+
def f_not_empty(*fpaths):
|
| 117 |
+
"""
|
| 118 |
+
Returns:
|
| 119 |
+
True if and only if the file exists and file size > 0
|
| 120 |
+
if fpath is a dir, if and only if dir exists and has at least 1 file
|
| 121 |
+
"""
|
| 122 |
+
fpath = f_join(*fpaths)
|
| 123 |
+
if not os.path.exists(fpath):
|
| 124 |
+
return False
|
| 125 |
+
|
| 126 |
+
if os.path.isdir(fpath):
|
| 127 |
+
return len(os.listdir(fpath)) > 0
|
| 128 |
+
else:
|
| 129 |
+
return os.path.getsize(fpath) > 0
|
| 130 |
+
|
| 131 |
+
|
| 132 |
+
def f_expand(fpath):
|
| 133 |
+
return os.path.expandvars(os.path.expanduser(fpath))
|
| 134 |
+
|
| 135 |
+
|
| 136 |
+
def f_exists(*fpaths):
|
| 137 |
+
return os.path.exists(f_join(*fpaths))
|
| 138 |
+
|
| 139 |
+
|
| 140 |
+
def f_join(*fpaths):
|
| 141 |
+
"""
|
| 142 |
+
Join file paths and expand special symbols like `~` for home dir
|
| 143 |
+
"""
|
| 144 |
+
|
| 145 |
+
def pack_varargs(args):
|
| 146 |
+
"""
|
| 147 |
+
Pack *args or a single list arg as list
|
| 148 |
+
|
| 149 |
+
def f(*args):
|
| 150 |
+
arg_list = pack_varargs(args)
|
| 151 |
+
# arg_list is now packed as a list
|
| 152 |
+
"""
|
| 153 |
+
assert isinstance(args, tuple), "please input the tuple `args` as in *args"
|
| 154 |
+
if len(args) == 1 and is_sequence(args[0]):
|
| 155 |
+
return args[0]
|
| 156 |
+
else:
|
| 157 |
+
return args
|
| 158 |
+
|
| 159 |
+
fpaths = pack_varargs(fpaths)
|
| 160 |
+
fpath = f_expand(os.path.join(*fpaths))
|
| 161 |
+
if isinstance(fpath, str):
|
| 162 |
+
fpath = fpath.strip()
|
| 163 |
+
return fpath
|
| 164 |
+
|
| 165 |
+
|
| 166 |
+
def f_listdir(
|
| 167 |
+
*fpaths,
|
| 168 |
+
filter_ext=None,
|
| 169 |
+
filter=None,
|
| 170 |
+
sort=True,
|
| 171 |
+
full_path=False,
|
| 172 |
+
nonexist_ok=True,
|
| 173 |
+
recursive=False,
|
| 174 |
+
):
|
| 175 |
+
"""
|
| 176 |
+
Args:
|
| 177 |
+
full_path: True to return full paths to the dir contents
|
| 178 |
+
filter: function that takes in file name and returns True to include
|
| 179 |
+
nonexist_ok: True to return [] if the dir is non-existent, False to raise
|
| 180 |
+
sort: sort the file names by alphabetical
|
| 181 |
+
recursive: True to use os.walk to recursively list files. Note that `filter`
|
| 182 |
+
will be applied to the relative path string to the root dir.
|
| 183 |
+
e.g. filter will take "a/data1.txt" and "a/b/data3.txt" as input, instead of
|
| 184 |
+
just the base file names "data1.txt" and "data3.txt".
|
| 185 |
+
if False, will simply call os.listdir()
|
| 186 |
+
"""
|
| 187 |
+
assert not (filter_ext and filter), "filter_ext and filter are mutually exclusive"
|
| 188 |
+
dir_path = f_join(*fpaths)
|
| 189 |
+
if not os.path.exists(dir_path) and nonexist_ok:
|
| 190 |
+
return []
|
| 191 |
+
if recursive:
|
| 192 |
+
files = [
|
| 193 |
+
os.path.join(os.path.relpath(root, dir_path), file)
|
| 194 |
+
for root, _, files in os.walk(dir_path)
|
| 195 |
+
for file in files
|
| 196 |
+
]
|
| 197 |
+
else:
|
| 198 |
+
files = os.listdir(dir_path)
|
| 199 |
+
if filter is not None:
|
| 200 |
+
files = [f for f in files if filter(f)]
|
| 201 |
+
elif filter_ext is not None:
|
| 202 |
+
files = [f for f in files if f.endswith(filter_ext)]
|
| 203 |
+
if sort:
|
| 204 |
+
files.sort()
|
| 205 |
+
if full_path:
|
| 206 |
+
return [os.path.join(dir_path, f) for f in files]
|
| 207 |
+
else:
|
| 208 |
+
return files
|
| 209 |
+
|
| 210 |
+
|
| 211 |
+
def f_mkdir(*fpaths):
|
| 212 |
+
"""
|
| 213 |
+
Recursively creates all the subdirs
|
| 214 |
+
If exist, do nothing.
|
| 215 |
+
"""
|
| 216 |
+
fpath = f_join(*fpaths)
|
| 217 |
+
os.makedirs(fpath, exist_ok=True)
|
| 218 |
+
return fpath
|
| 219 |
+
|
| 220 |
+
|
| 221 |
+
def f_mkdir_in_path(*fpaths):
|
| 222 |
+
"""
|
| 223 |
+
fpath is a file,
|
| 224 |
+
recursively creates all the parent dirs that lead to the file
|
| 225 |
+
If exist, do nothing.
|
| 226 |
+
"""
|
| 227 |
+
os.makedirs(get_dir(f_join(*fpaths)), exist_ok=True)
|
| 228 |
+
|
| 229 |
+
|
| 230 |
+
def last_part_in_path(fpath):
|
| 231 |
+
"""
|
| 232 |
+
https://stackoverflow.com/questions/3925096/how-to-get-only-the-last-part-of-a-path-in-python
|
| 233 |
+
"""
|
| 234 |
+
return os.path.basename(os.path.normpath(f_expand(fpath)))
|
| 235 |
+
|
| 236 |
+
|
| 237 |
+
def is_abs_path(*fpath):
|
| 238 |
+
return os.path.isabs(f_join(*fpath))
|
| 239 |
+
|
| 240 |
+
|
| 241 |
+
def is_relative_path(*fpath):
|
| 242 |
+
return not is_abs_path(f_join(*fpath))
|
| 243 |
+
|
| 244 |
+
|
| 245 |
+
def f_time(*fpath):
|
| 246 |
+
"File modification time"
|
| 247 |
+
return str(os.path.getctime(f_join(*fpath)))
|
| 248 |
+
|
| 249 |
+
|
| 250 |
+
def f_append_before_ext(fpath, suffix):
|
| 251 |
+
"""
|
| 252 |
+
Append a suffix to file name and retain its extension
|
| 253 |
+
"""
|
| 254 |
+
name, ext = f_ext(fpath)
|
| 255 |
+
return name + suffix + ext
|
| 256 |
+
|
| 257 |
+
|
| 258 |
+
def f_add_ext(fpath, ext):
|
| 259 |
+
"""
|
| 260 |
+
Append an extension if not already there
|
| 261 |
+
Args:
|
| 262 |
+
ext: will add a preceding `.` if doesn't exist
|
| 263 |
+
"""
|
| 264 |
+
if not ext.startswith("."):
|
| 265 |
+
ext = "." + ext
|
| 266 |
+
if fpath.endswith(ext):
|
| 267 |
+
return fpath
|
| 268 |
+
else:
|
| 269 |
+
return fpath + ext
|
| 270 |
+
|
| 271 |
+
|
| 272 |
+
def f_has_ext(fpath, ext):
|
| 273 |
+
"Test if file path is a text file"
|
| 274 |
+
_, actual_ext = f_ext(fpath)
|
| 275 |
+
return actual_ext == "." + ext.lstrip(".")
|
| 276 |
+
|
| 277 |
+
|
| 278 |
+
def f_glob(*fpath):
|
| 279 |
+
return glob.glob(f_join(*fpath), recursive=True)
|
| 280 |
+
|
| 281 |
+
|
| 282 |
+
def f_remove(*fpath, verbose=False, dry_run=False):
|
| 283 |
+
"""
|
| 284 |
+
If exist, remove. Supports both dir and file. Supports glob wildcard.
|
| 285 |
+
"""
|
| 286 |
+
import errno
|
| 287 |
+
|
| 288 |
+
assert isinstance(verbose, bool)
|
| 289 |
+
fpath = f_join(fpath)
|
| 290 |
+
if dry_run:
|
| 291 |
+
print("Dry run, delete:", fpath)
|
| 292 |
+
return
|
| 293 |
+
for f in glob.glob(fpath):
|
| 294 |
+
try:
|
| 295 |
+
shutil.rmtree(f)
|
| 296 |
+
except OSError as e:
|
| 297 |
+
if e.errno == errno.ENOTDIR:
|
| 298 |
+
try:
|
| 299 |
+
os.remove(f)
|
| 300 |
+
except Exception as e: # final resort safeguard
|
| 301 |
+
pass
|
| 302 |
+
if verbose:
|
| 303 |
+
print(f'Deleted "{fpath}"')
|
| 304 |
+
|
| 305 |
+
|
| 306 |
+
def f_copy(fsrc, fdst, ignore=None, include=None, exists_ok=True, verbose=False):
|
| 307 |
+
"""
|
| 308 |
+
Supports both dir and file. Supports glob wildcard.
|
| 309 |
+
"""
|
| 310 |
+
import errno
|
| 311 |
+
|
| 312 |
+
fsrc, fdst = f_expand(fsrc), f_expand(fdst)
|
| 313 |
+
for f in glob.glob(fsrc):
|
| 314 |
+
try:
|
| 315 |
+
f_copytree(f, fdst, ignore=ignore, include=include, exist_ok=exists_ok)
|
| 316 |
+
except OSError as e:
|
| 317 |
+
if e.errno == errno.ENOTDIR:
|
| 318 |
+
shutil.copy(f, fdst)
|
| 319 |
+
else:
|
| 320 |
+
raise
|
| 321 |
+
if verbose:
|
| 322 |
+
print(f'Copied "{fsrc}" to "{fdst}"')
|
| 323 |
+
|
| 324 |
+
|
| 325 |
+
def _f_copytree(
|
| 326 |
+
src,
|
| 327 |
+
dst,
|
| 328 |
+
symlinks=False,
|
| 329 |
+
ignore=None,
|
| 330 |
+
exist_ok=True,
|
| 331 |
+
copy_function=shutil.copy2,
|
| 332 |
+
ignore_dangling_symlinks=False,
|
| 333 |
+
):
|
| 334 |
+
"""Copied from python standard lib shutil.copytree
|
| 335 |
+
except that we allow exist_ok
|
| 336 |
+
Use f_copytree as entry
|
| 337 |
+
"""
|
| 338 |
+
names = os.listdir(src)
|
| 339 |
+
if ignore is not None:
|
| 340 |
+
ignored_names = ignore(src, names)
|
| 341 |
+
else:
|
| 342 |
+
ignored_names = set()
|
| 343 |
+
|
| 344 |
+
os.makedirs(dst, exist_ok=exist_ok)
|
| 345 |
+
errors = []
|
| 346 |
+
for name in names:
|
| 347 |
+
if name in ignored_names:
|
| 348 |
+
continue
|
| 349 |
+
srcname = os.path.join(src, name)
|
| 350 |
+
dstname = os.path.join(dst, name)
|
| 351 |
+
try:
|
| 352 |
+
if os.path.islink(srcname):
|
| 353 |
+
linkto = os.readlink(srcname)
|
| 354 |
+
if symlinks:
|
| 355 |
+
# We can't just leave it to `copy_function` because legacy
|
| 356 |
+
# code with a custom `copy_function` may rely on copytree
|
| 357 |
+
# doing the right thing.
|
| 358 |
+
os.symlink(linkto, dstname)
|
| 359 |
+
shutil.copystat(srcname, dstname, follow_symlinks=not symlinks)
|
| 360 |
+
else:
|
| 361 |
+
# ignore dangling symlink if the flag is on
|
| 362 |
+
if not os.path.exists(linkto) and ignore_dangling_symlinks:
|
| 363 |
+
continue
|
| 364 |
+
# otherwise let the copy occurs. copy2 will raise an error
|
| 365 |
+
if os.path.isdir(srcname):
|
| 366 |
+
_f_copytree(srcname, dstname, symlinks, ignore, exist_ok, copy_function)
|
| 367 |
+
else:
|
| 368 |
+
copy_function(srcname, dstname)
|
| 369 |
+
elif os.path.isdir(srcname):
|
| 370 |
+
_f_copytree(srcname, dstname, symlinks, ignore, exist_ok, copy_function)
|
| 371 |
+
else:
|
| 372 |
+
# Will raise a SpecialFileError for unsupported file types
|
| 373 |
+
copy_function(srcname, dstname)
|
| 374 |
+
# catch the Error from the recursive copytree so that we can
|
| 375 |
+
# continue with other files
|
| 376 |
+
except shutil.Error as err:
|
| 377 |
+
errors.extend(err.args[0])
|
| 378 |
+
except OSError as why:
|
| 379 |
+
errors.append((srcname, dstname, str(why)))
|
| 380 |
+
try:
|
| 381 |
+
shutil.copystat(src, dst)
|
| 382 |
+
except OSError as why:
|
| 383 |
+
# Copying file access times may fail on Windows
|
| 384 |
+
if getattr(why, "winerror", None) is None:
|
| 385 |
+
errors.append((src, dst, str(why)))
|
| 386 |
+
if errors:
|
| 387 |
+
raise shutil.Error(errors)
|
| 388 |
+
return dst
|
| 389 |
+
|
| 390 |
+
|
| 391 |
+
def _include_patterns(*patterns):
|
| 392 |
+
"""Factory function that can be used with copytree() ignore parameter.
|
| 393 |
+
|
| 394 |
+
Arguments define a sequence of glob-style patterns
|
| 395 |
+
that are used to specify what files to NOT ignore.
|
| 396 |
+
Creates and returns a function that determines this for each directory
|
| 397 |
+
in the file hierarchy rooted at the source directory when used with
|
| 398 |
+
shutil.copytree().
|
| 399 |
+
"""
|
| 400 |
+
|
| 401 |
+
def _ignore_patterns(path, names):
|
| 402 |
+
import fnmatch
|
| 403 |
+
|
| 404 |
+
keep = set(name for pattern in patterns for name in fnmatch.filter(names, pattern))
|
| 405 |
+
ignore = set(
|
| 406 |
+
name
|
| 407 |
+
for name in names
|
| 408 |
+
if name not in keep and not os.path.isdir(os.path.join(path, name))
|
| 409 |
+
)
|
| 410 |
+
return ignore
|
| 411 |
+
|
| 412 |
+
return _ignore_patterns
|
| 413 |
+
|
| 414 |
+
|
| 415 |
+
def f_copytree(fsrc, fdst, symlinks=False, ignore=None, include=None, exist_ok=True):
|
| 416 |
+
fsrc, fdst = f_expand(fsrc), f_expand(fdst)
|
| 417 |
+
assert (ignore is None) or (include is None), "ignore= and include= are mutually exclusive"
|
| 418 |
+
if ignore:
|
| 419 |
+
ignore = shutil.ignore_patterns(*ignore)
|
| 420 |
+
elif include:
|
| 421 |
+
ignore = _include_patterns(*include)
|
| 422 |
+
_f_copytree(fsrc, fdst, ignore=ignore, symlinks=symlinks, exist_ok=exist_ok)
|
| 423 |
+
|
| 424 |
+
|
| 425 |
+
def f_move(fsrc, fdst):
|
| 426 |
+
fsrc, fdst = f_expand(fsrc), f_expand(fdst)
|
| 427 |
+
for f in glob.glob(fsrc):
|
| 428 |
+
shutil.move(f, fdst)
|
| 429 |
+
|
| 430 |
+
|
| 431 |
+
def f_split_path(fpath, normpath=True):
|
| 432 |
+
"""
|
| 433 |
+
Splits path into a list of its component folders
|
| 434 |
+
|
| 435 |
+
Args:
|
| 436 |
+
normpath: call os.path.normpath to remove redundant '/' and
|
| 437 |
+
up-level references like ".."
|
| 438 |
+
"""
|
| 439 |
+
if normpath:
|
| 440 |
+
fpath = os.path.normpath(fpath)
|
| 441 |
+
allparts = []
|
| 442 |
+
while 1:
|
| 443 |
+
parts = os.path.split(fpath)
|
| 444 |
+
if parts[0] == fpath: # sentinel for absolute paths
|
| 445 |
+
allparts.insert(0, parts[0])
|
| 446 |
+
break
|
| 447 |
+
elif parts[1] == fpath: # sentinel for relative paths
|
| 448 |
+
allparts.insert(0, parts[1])
|
| 449 |
+
break
|
| 450 |
+
else:
|
| 451 |
+
fpath = parts[0]
|
| 452 |
+
allparts.insert(0, parts[1])
|
| 453 |
+
return allparts
|
| 454 |
+
|
| 455 |
+
|
| 456 |
+
def get_script_dir():
|
| 457 |
+
"""
|
| 458 |
+
Returns: the dir of current script
|
| 459 |
+
"""
|
| 460 |
+
return os.path.dirname(os.path.realpath(sys.argv[0]))
|
| 461 |
+
|
| 462 |
+
|
| 463 |
+
def get_script_file_name():
|
| 464 |
+
"""
|
| 465 |
+
Returns: the dir of current script
|
| 466 |
+
"""
|
| 467 |
+
return os.path.basename(sys.argv[0])
|
| 468 |
+
|
| 469 |
+
|
| 470 |
+
def get_script_self_path():
|
| 471 |
+
"""
|
| 472 |
+
Returns: the dir of current script
|
| 473 |
+
"""
|
| 474 |
+
return os.path.realpath(sys.argv[0])
|
| 475 |
+
|
| 476 |
+
|
| 477 |
+
def get_parent_dir(location, abspath=False):
|
| 478 |
+
"""
|
| 479 |
+
Args:
|
| 480 |
+
location: current directory or file
|
| 481 |
+
|
| 482 |
+
Returns:
|
| 483 |
+
parent directory absolute or relative path
|
| 484 |
+
"""
|
| 485 |
+
_path = os.path.abspath if abspath else os.path.relpath
|
| 486 |
+
return _path(f_join(location, os.pardir))
|
| 487 |
+
|
| 488 |
+
|
| 489 |
+
def md5_checksum(*fpath):
|
| 490 |
+
"""
|
| 491 |
+
File md5 signature
|
| 492 |
+
"""
|
| 493 |
+
import hashlib
|
| 494 |
+
|
| 495 |
+
hash_md5 = hashlib.md5()
|
| 496 |
+
with open(f_join(*fpath), "rb") as f:
|
| 497 |
+
for chunk in iter(lambda: f.read(65536), b""):
|
| 498 |
+
hash_md5.update(chunk)
|
| 499 |
+
return hash_md5.hexdigest()
|
| 500 |
+
|
| 501 |
+
|
| 502 |
+
def create_tar(fsrc, output_tarball, include=None, ignore=None, compress_mode="gz"):
|
| 503 |
+
"""
|
| 504 |
+
Args:
|
| 505 |
+
fsrc: source file or folder
|
| 506 |
+
output_tarball: output tar file name
|
| 507 |
+
compress_mode: ``gz``, ``bz2``, ``xz`` or ``''`` (empty for uncompressed write)
|
| 508 |
+
include: include pattern, will trigger copy to temp directory
|
| 509 |
+
ignore: ignore pattern, will trigger copy to temp directory
|
| 510 |
+
"""
|
| 511 |
+
import tarfile
|
| 512 |
+
import tempfile
|
| 513 |
+
|
| 514 |
+
fsrc, output_tarball = f_expand(fsrc), f_expand(output_tarball)
|
| 515 |
+
assert compress_mode in ["gz", "bz2", "xz", ""]
|
| 516 |
+
src_base = os.path.basename(fsrc)
|
| 517 |
+
|
| 518 |
+
tempdir = None
|
| 519 |
+
if include or ignore:
|
| 520 |
+
tempdir = tempfile.mkdtemp()
|
| 521 |
+
tempdest = f_join(tempdir, src_base)
|
| 522 |
+
f_copy(fsrc, tempdest, include=include, ignore=ignore)
|
| 523 |
+
fsrc = tempdest
|
| 524 |
+
|
| 525 |
+
with tarfile.open(output_tarball, "w:" + compress_mode) as tar:
|
| 526 |
+
tar.add(fsrc, arcname=src_base)
|
| 527 |
+
|
| 528 |
+
if tempdir:
|
| 529 |
+
f_remove(tempdir)
|
| 530 |
+
|
| 531 |
+
|
| 532 |
+
def extract_tar(source_tarball, output_dir=".", members=None):
|
| 533 |
+
"""
|
| 534 |
+
Args:
|
| 535 |
+
source_tarball: extract members from archive
|
| 536 |
+
output_dir: default to current working dir
|
| 537 |
+
members: must be a subset of the list returned by getmembers()
|
| 538 |
+
"""
|
| 539 |
+
import tarfile
|
| 540 |
+
|
| 541 |
+
source_tarball, output_dir = f_expand(source_tarball), f_expand(output_dir)
|
| 542 |
+
with tarfile.open(source_tarball, "r:*") as tar:
|
| 543 |
+
tar.extractall(output_dir, members=members)
|
| 544 |
+
|
| 545 |
+
|
| 546 |
+
def move_with_backup(*fpath, suffix=".bak"):
|
| 547 |
+
"""
|
| 548 |
+
Ensures that a path is not occupied. If there is a file, rename it by
|
| 549 |
+
adding @suffix. Resursively backs up everything.
|
| 550 |
+
|
| 551 |
+
Args:
|
| 552 |
+
fpath: file path to clear
|
| 553 |
+
suffix: Add to backed up files (default: {'.bak'})
|
| 554 |
+
"""
|
| 555 |
+
fpath = str(f_join(*fpath))
|
| 556 |
+
if os.path.exists(fpath):
|
| 557 |
+
move_with_backup(fpath + suffix)
|
| 558 |
+
shutil.move(fpath, fpath + suffix)
|
| 559 |
+
|
| 560 |
+
|
| 561 |
+
def insert_before_ext(name, insert):
|
| 562 |
+
"""
|
| 563 |
+
log.txt -> log.ep50.txt
|
| 564 |
+
"""
|
| 565 |
+
name, ext = os.path.splitext(name)
|
| 566 |
+
return name + insert + ext
|
| 567 |
+
|
| 568 |
+
|
| 569 |
+
def timestamp_file_name(fname):
|
| 570 |
+
from datetime import datetime
|
| 571 |
+
|
| 572 |
+
timestr = datetime.now().strftime("_%H-%M-%S_%m-%d-%y")
|
| 573 |
+
return insert_before_ext(fname, timestr)
|
| 574 |
+
|
| 575 |
+
|
| 576 |
+
def next_available_file_name(
|
| 577 |
+
*fpath,
|
| 578 |
+
suffix_template: Union[str, Callable[[int], str]] = "_v{i+1}",
|
| 579 |
+
before_ext: bool = True,
|
| 580 |
+
):
|
| 581 |
+
"""
|
| 582 |
+
Args:
|
| 583 |
+
suffix_template: a format string using "i" variable or
|
| 584 |
+
lambda int -> str
|
| 585 |
+
before_ext: True to insert suffix before the extension
|
| 586 |
+
"""
|
| 587 |
+
|
| 588 |
+
def fstring(fmt_str, **kwargs):
|
| 589 |
+
"""
|
| 590 |
+
Simulate python f-string but without `f`
|
| 591 |
+
"""
|
| 592 |
+
import shlex
|
| 593 |
+
|
| 594 |
+
locals().update(kwargs)
|
| 595 |
+
return eval("f" + shlex.quote(fmt_str))
|
| 596 |
+
|
| 597 |
+
orig_file_path = f_join(*fpath)
|
| 598 |
+
i = 0
|
| 599 |
+
fpath = orig_file_path
|
| 600 |
+
while os.path.exists(fpath):
|
| 601 |
+
if isinstance(suffix_template, str):
|
| 602 |
+
suffix = fstring(suffix_template, i=i)
|
| 603 |
+
elif callable(suffix_template):
|
| 604 |
+
suffix = suffix_template(i)
|
| 605 |
+
assert isinstance(suffix, str)
|
| 606 |
+
else:
|
| 607 |
+
raise NotImplementedError(f"Unsupported suffix template {suffix_template}")
|
| 608 |
+
if before_ext:
|
| 609 |
+
fpath = insert_before_ext(orig_file_path, suffix)
|
| 610 |
+
else:
|
| 611 |
+
fpath = orig_file_path + suffix
|
| 612 |
+
i += 1
|
| 613 |
+
return fpath
|
| 614 |
+
|
| 615 |
+
|
| 616 |
+
def get_file_lock(*fpath, timeout: int = 15, logging_level="critical"):
|
| 617 |
+
"""
|
| 618 |
+
NFS-safe filesystem-backed lock. `pip install flufl.lock`
|
| 619 |
+
https://flufllock.readthedocs.io/en/stable/apiref.html
|
| 620 |
+
|
| 621 |
+
Args:
|
| 622 |
+
fpath: should be a path on NFS so that every process can see it
|
| 623 |
+
timeout: seconds
|
| 624 |
+
"""
|
| 625 |
+
import logging
|
| 626 |
+
|
| 627 |
+
from flufl.lock import Lock
|
| 628 |
+
|
| 629 |
+
logging.getLogger("flufl.lock").setLevel(logging_level.upper())
|
| 630 |
+
return Lock(f_join(*fpath), lifetime=timeout)
|
| 631 |
+
|
| 632 |
+
|
| 633 |
+
def load_pickle(*fpaths):
|
| 634 |
+
with open(f_join(*fpaths), "rb") as fp:
|
| 635 |
+
return pickle.load(fp)
|
| 636 |
+
|
| 637 |
+
|
| 638 |
+
def dump_pickle(data, *fpaths):
|
| 639 |
+
with open(f_join(*fpaths), "wb") as fp:
|
| 640 |
+
pickle.dump(data, fp)
|
| 641 |
+
|
| 642 |
+
|
| 643 |
+
def load_text(*fpaths, by_lines=False):
|
| 644 |
+
with open(f_join(*fpaths), "r") as fp:
|
| 645 |
+
if by_lines:
|
| 646 |
+
return fp.readlines()
|
| 647 |
+
else:
|
| 648 |
+
return fp.read()
|
| 649 |
+
|
| 650 |
+
|
| 651 |
+
def load_text_lines(*fpaths):
|
| 652 |
+
return load_text(*fpaths, by_lines=True)
|
| 653 |
+
|
| 654 |
+
|
| 655 |
+
def dump_text(s, *fpaths):
|
| 656 |
+
with open(f_join(*fpaths), "w") as fp:
|
| 657 |
+
fp.write(s)
|
| 658 |
+
|
| 659 |
+
|
| 660 |
+
def dump_text_lines(lines: list[str], *fpaths, add_newline=True):
|
| 661 |
+
with open(f_join(*fpaths), "w") as fp:
|
| 662 |
+
for line in lines:
|
| 663 |
+
print(line, file=fp, end="\n" if add_newline else "")
|
| 664 |
+
|
| 665 |
+
|
| 666 |
+
def get_package_root() -> str:
|
| 667 |
+
import importlib.util
|
| 668 |
+
import inspect
|
| 669 |
+
|
| 670 |
+
# Get the current frame
|
| 671 |
+
current_frame = inspect.currentframe()
|
| 672 |
+
if current_frame is None:
|
| 673 |
+
raise ImportError("Cannot determine the package name from __package__")
|
| 674 |
+
|
| 675 |
+
# Get the caller module
|
| 676 |
+
caller_module = inspect.getmodule(current_frame.f_back)
|
| 677 |
+
if caller_module is None:
|
| 678 |
+
raise ImportError("Cannot determine the package name from __package__")
|
| 679 |
+
|
| 680 |
+
# Get the package name
|
| 681 |
+
package_name = caller_module.__package__
|
| 682 |
+
if not package_name:
|
| 683 |
+
raise ImportError("Cannot determine the package name from __package__")
|
| 684 |
+
|
| 685 |
+
# Get the top-level package name
|
| 686 |
+
top_package_name = package_name.split(".")[0]
|
| 687 |
+
|
| 688 |
+
# Find the package specification
|
| 689 |
+
spec = importlib.util.find_spec(top_package_name)
|
| 690 |
+
|
| 691 |
+
if spec and spec.origin:
|
| 692 |
+
# Get the directory containing the package's __init__.py file
|
| 693 |
+
package_dir = os.path.dirname(spec.origin)
|
| 694 |
+
return package_dir
|
| 695 |
+
else:
|
| 696 |
+
raise ImportError(f"Cannot find the package {top_package_name}")
|
| 697 |
+
|
| 698 |
+
|
| 699 |
+
# aliases to be consistent with other load_* and dump_*
|
| 700 |
+
pickle_load = load_pickle
|
| 701 |
+
pickle_dump = dump_pickle
|
| 702 |
+
text_load = load_text
|
| 703 |
+
read_text = load_text
|
| 704 |
+
read_text_lines = load_text_lines
|
| 705 |
+
write_text = dump_text
|
| 706 |
+
write_text_lines = dump_text_lines
|
| 707 |
+
text_dump = dump_text
|
groot/vla/common/utils/io/hdf5_utils.py
ADDED
|
@@ -0,0 +1,84 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
import h5py
|
| 2 |
+
import numpy as np
|
| 3 |
+
from pydantic import BaseModel
|
| 4 |
+
|
| 5 |
+
|
| 6 |
+
def hdf5_save(data: BaseModel | dict, group: h5py.Group) -> None:
|
| 7 |
+
"""Recursively save Pydantic model or dict to HDF5 group."""
|
| 8 |
+
if isinstance(data, BaseModel):
|
| 9 |
+
# Convert to dict and exclude None values
|
| 10 |
+
data_dict = data.model_dump(mode="python", exclude_none=True)
|
| 11 |
+
else:
|
| 12 |
+
data_dict = data
|
| 13 |
+
|
| 14 |
+
for key, value in data_dict.items():
|
| 15 |
+
if isinstance(value, np.ndarray):
|
| 16 |
+
group.create_dataset(key, data=value)
|
| 17 |
+
elif isinstance(value, (BaseModel, dict)):
|
| 18 |
+
subgroup = group.create_group(key)
|
| 19 |
+
hdf5_save(value, subgroup)
|
| 20 |
+
else:
|
| 21 |
+
# For primitive types, convert to numpy array
|
| 22 |
+
try:
|
| 23 |
+
group.create_dataset(key, data=np.array(value))
|
| 24 |
+
except TypeError:
|
| 25 |
+
raise ValueError(f"Unsupported type: {type(value)} for key: {key}")
|
| 26 |
+
|
| 27 |
+
|
| 28 |
+
def hdf5_load(group: h5py.Group) -> dict:
|
| 29 |
+
"""Recursively load HDF5 group to Pydantic model or dict."""
|
| 30 |
+
data_dict = {}
|
| 31 |
+
for key, value in group.items():
|
| 32 |
+
if isinstance(value, h5py.Dataset):
|
| 33 |
+
data_dict[key] = value[()]
|
| 34 |
+
elif isinstance(value, h5py.Group):
|
| 35 |
+
data_dict[key] = hdf5_load(value)
|
| 36 |
+
return data_dict
|
| 37 |
+
|
| 38 |
+
|
| 39 |
+
def hdf5_is_subset(this: h5py.Group, other: h5py.Group, verbose: bool = False) -> bool:
|
| 40 |
+
"""Check if this HDF5 group is a subset of another HDF5 group."""
|
| 41 |
+
for key, value in this.items():
|
| 42 |
+
if key not in other:
|
| 43 |
+
if verbose:
|
| 44 |
+
print(f"Key {key} not in other")
|
| 45 |
+
return False
|
| 46 |
+
elif isinstance(value, h5py.Group):
|
| 47 |
+
if not isinstance(other[key], h5py.Group):
|
| 48 |
+
if verbose:
|
| 49 |
+
print(f"Key {key} is not a group in other")
|
| 50 |
+
return False
|
| 51 |
+
if not hdf5_is_subset(value, other[key], verbose):
|
| 52 |
+
if verbose:
|
| 53 |
+
print(f"Key {key} is not a subset of other")
|
| 54 |
+
return False
|
| 55 |
+
elif isinstance(value, h5py.Dataset):
|
| 56 |
+
if not isinstance(other[key], h5py.Dataset):
|
| 57 |
+
if verbose:
|
| 58 |
+
print(f"Key {key} is not a dataset in other")
|
| 59 |
+
return False
|
| 60 |
+
if not np.array_equal(value, other[key]):
|
| 61 |
+
if verbose:
|
| 62 |
+
print(f"Key {key} is not equal in other")
|
| 63 |
+
return False
|
| 64 |
+
elif isinstance(value, h5py.Datatype):
|
| 65 |
+
if not isinstance(other[key], h5py.Datatype):
|
| 66 |
+
if verbose:
|
| 67 |
+
print(f"Key {key} is not a datatype in other")
|
| 68 |
+
return False
|
| 69 |
+
if value != other[key]:
|
| 70 |
+
if verbose:
|
| 71 |
+
print(f"Key {key} is not equal in other")
|
| 72 |
+
return False
|
| 73 |
+
else:
|
| 74 |
+
# try to compare
|
| 75 |
+
if value != other[key]:
|
| 76 |
+
if verbose:
|
| 77 |
+
print(f"Key {key} is not equal in other")
|
| 78 |
+
return False
|
| 79 |
+
return True
|
| 80 |
+
|
| 81 |
+
|
| 82 |
+
def hdf5_is_equal(this: h5py.Group, other: h5py.Group, verbose: bool = False) -> bool:
|
| 83 |
+
"""Check if this HDF5 group is equal to another HDF5 group."""
|
| 84 |
+
return hdf5_is_subset(this, other, verbose) and hdf5_is_subset(other, this, verbose)
|
groot/vla/common/utils/io/json_utils.py
ADDED
|
@@ -0,0 +1,270 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""
|
| 2 |
+
JSON, YAML, and python config file utilities
|
| 3 |
+
"""
|
| 4 |
+
|
| 5 |
+
from io import StringIO
|
| 6 |
+
import json
|
| 7 |
+
import os.path as path
|
| 8 |
+
|
| 9 |
+
import yaml
|
| 10 |
+
|
| 11 |
+
from ..misc.functional_utils import make_recursive_func
|
| 12 |
+
from .file_utils import f_join
|
| 13 |
+
|
| 14 |
+
__all__ = [
|
| 15 |
+
"json_load",
|
| 16 |
+
"json_loads",
|
| 17 |
+
"jsonl_load",
|
| 18 |
+
"yaml_load",
|
| 19 |
+
"yaml_loads",
|
| 20 |
+
"json_dump",
|
| 21 |
+
"json_dumps",
|
| 22 |
+
"jsonl_dump",
|
| 23 |
+
"yaml_dump",
|
| 24 |
+
"yaml_dumps",
|
| 25 |
+
"json_or_yaml_load",
|
| 26 |
+
"json_or_yaml_dump",
|
| 27 |
+
"Jsonl",
|
| 28 |
+
# ---------------- Aliases -----------------
|
| 29 |
+
"load_json",
|
| 30 |
+
"loads_json",
|
| 31 |
+
"load_jsonl",
|
| 32 |
+
"load_yaml",
|
| 33 |
+
"loads_yaml",
|
| 34 |
+
"dump_json",
|
| 35 |
+
"dumps_json",
|
| 36 |
+
"dump_jsonl",
|
| 37 |
+
"dump_yaml",
|
| 38 |
+
"dumps_yaml",
|
| 39 |
+
"load_json_or_yaml",
|
| 40 |
+
"dump_json_or_yaml",
|
| 41 |
+
]
|
| 42 |
+
|
| 43 |
+
from typing import Dict, List
|
| 44 |
+
|
| 45 |
+
from typing_extensions import Literal
|
| 46 |
+
|
| 47 |
+
|
| 48 |
+
def json_load(*file_path, **kwargs):
|
| 49 |
+
file_path = f_join(file_path)
|
| 50 |
+
with open(file_path, "r") as fp:
|
| 51 |
+
return json.load(fp, **kwargs)
|
| 52 |
+
|
| 53 |
+
|
| 54 |
+
def json_loads(string, **kwargs):
|
| 55 |
+
return json.loads(string, **kwargs)
|
| 56 |
+
|
| 57 |
+
|
| 58 |
+
def jsonl_load(*file_path, **kwargs):
|
| 59 |
+
file_path = f_join(file_path)
|
| 60 |
+
data = []
|
| 61 |
+
for line in open(file_path):
|
| 62 |
+
data.append(json.loads(line, **kwargs))
|
| 63 |
+
return data
|
| 64 |
+
|
| 65 |
+
|
| 66 |
+
@make_recursive_func
|
| 67 |
+
def any_to_primitive(x):
|
| 68 |
+
try:
|
| 69 |
+
import torch
|
| 70 |
+
except ImportError:
|
| 71 |
+
raise ImportError("torch is required for any_to_primitive")
|
| 72 |
+
import numpy as np
|
| 73 |
+
|
| 74 |
+
if isinstance(x, (np.ndarray, np.number, torch.Tensor)):
|
| 75 |
+
return x.tolist()
|
| 76 |
+
else:
|
| 77 |
+
return x
|
| 78 |
+
|
| 79 |
+
|
| 80 |
+
def json_dump(data, *file_path, convert_to_primitive=False, **kwargs):
|
| 81 |
+
if convert_to_primitive:
|
| 82 |
+
data = any_to_primitive(data)
|
| 83 |
+
file_path = f_join(file_path)
|
| 84 |
+
with open(file_path, "w") as fp:
|
| 85 |
+
json.dump(data, fp, **kwargs)
|
| 86 |
+
|
| 87 |
+
|
| 88 |
+
def json_dumps(data, convert_to_primitive=False, **kwargs):
|
| 89 |
+
"""
|
| 90 |
+
Returns: string
|
| 91 |
+
"""
|
| 92 |
+
if convert_to_primitive:
|
| 93 |
+
data = any_to_primitive(data)
|
| 94 |
+
return json.dumps(data, **kwargs)
|
| 95 |
+
|
| 96 |
+
|
| 97 |
+
def jsonl_dump(data, *file_path):
|
| 98 |
+
from .file_utils import is_sequence
|
| 99 |
+
|
| 100 |
+
assert is_sequence(data)
|
| 101 |
+
data = any_to_primitive(data)
|
| 102 |
+
file_path = f_join(file_path)
|
| 103 |
+
with open(file_path, "w") as fp:
|
| 104 |
+
for line in data:
|
| 105 |
+
print(json.dumps(line), file=fp, flush=True)
|
| 106 |
+
|
| 107 |
+
|
| 108 |
+
def yaml_load(*file_path, loader=yaml.safe_load, **kwargs):
|
| 109 |
+
file_path = f_join(file_path)
|
| 110 |
+
with open(file_path, "r") as fp:
|
| 111 |
+
return loader(fp, **kwargs)
|
| 112 |
+
|
| 113 |
+
|
| 114 |
+
def yaml_loads(string, *, loader=yaml.safe_load, **kwargs):
|
| 115 |
+
return loader(string, **kwargs)
|
| 116 |
+
|
| 117 |
+
|
| 118 |
+
def yaml_dump(data, *file_path, dumper=yaml.safe_dump, convert_to_primitive=False, **kwargs):
|
| 119 |
+
if convert_to_primitive:
|
| 120 |
+
data = any_to_primitive(data)
|
| 121 |
+
file_path = f_join(file_path)
|
| 122 |
+
indent = kwargs.pop("indent", 2)
|
| 123 |
+
default_flow_style = kwargs.pop("default_flow_style", False)
|
| 124 |
+
sort_keys = kwargs.pop("sort_keys", False) # preserves original dict order
|
| 125 |
+
with open(file_path, "w") as fp:
|
| 126 |
+
dumper(
|
| 127 |
+
data,
|
| 128 |
+
stream=fp,
|
| 129 |
+
indent=indent,
|
| 130 |
+
default_flow_style=default_flow_style,
|
| 131 |
+
sort_keys=sort_keys,
|
| 132 |
+
**kwargs,
|
| 133 |
+
)
|
| 134 |
+
|
| 135 |
+
|
| 136 |
+
def yaml_dumps(data, *, dumper=yaml.safe_dump, convert_to_primitive=False, **kwargs):
|
| 137 |
+
"Returns: string"
|
| 138 |
+
if convert_to_primitive:
|
| 139 |
+
data = any_to_primitive(data)
|
| 140 |
+
stream = StringIO()
|
| 141 |
+
indent = kwargs.pop("indent", 2)
|
| 142 |
+
default_flow_style = kwargs.pop("default_flow_style", False)
|
| 143 |
+
sort_keys = kwargs.pop("sort_keys", False) # preserves original dict order
|
| 144 |
+
dumper(
|
| 145 |
+
data,
|
| 146 |
+
stream,
|
| 147 |
+
indent=indent,
|
| 148 |
+
default_flow_style=default_flow_style,
|
| 149 |
+
sort_keys=sort_keys,
|
| 150 |
+
**kwargs,
|
| 151 |
+
)
|
| 152 |
+
return stream.getvalue()
|
| 153 |
+
|
| 154 |
+
|
| 155 |
+
# ==================== auto-recognize extension ====================
|
| 156 |
+
def json_or_yaml_load(*file_path, **loader_kwargs):
|
| 157 |
+
"""
|
| 158 |
+
Args:
|
| 159 |
+
file_path: JSON or YAML loader depends on the file extension
|
| 160 |
+
|
| 161 |
+
Raises:
|
| 162 |
+
IOError: if extension is not ".json", ".yml", or ".yaml"
|
| 163 |
+
"""
|
| 164 |
+
file_path = str(f_join(file_path))
|
| 165 |
+
if file_path.endswith(".json"):
|
| 166 |
+
return json_load(file_path, **loader_kwargs)
|
| 167 |
+
elif file_path.endswith(".yml") or file_path.endswith(".yaml"):
|
| 168 |
+
return yaml_load(file_path, **loader_kwargs)
|
| 169 |
+
else:
|
| 170 |
+
raise IOError(
|
| 171 |
+
f'unknown file extension: "{file_path}", '
|
| 172 |
+
f'loader supports only ".json", ".yml", ".yaml"'
|
| 173 |
+
)
|
| 174 |
+
|
| 175 |
+
|
| 176 |
+
def json_or_yaml_dump(data, *file_path, **dumper_kwargs):
|
| 177 |
+
"""
|
| 178 |
+
Args:
|
| 179 |
+
file_path: JSON or YAML loader depends on the file extension
|
| 180 |
+
|
| 181 |
+
Raises:
|
| 182 |
+
IOError: if extension is not ".json", ".yml", or ".yaml"
|
| 183 |
+
"""
|
| 184 |
+
file_path = str(f_join(file_path))
|
| 185 |
+
if file_path.endswith(".json"):
|
| 186 |
+
return json_dump(data, file_path, **dumper_kwargs)
|
| 187 |
+
elif file_path.endswith(".yml") or file_path.endswith(".yaml"):
|
| 188 |
+
return yaml_dump(data, file_path, **dumper_kwargs)
|
| 189 |
+
else:
|
| 190 |
+
raise IOError(
|
| 191 |
+
f'unknown file extension: "{file_path}", '
|
| 192 |
+
f'dumper supports only ".json", ".yml", ".yaml"'
|
| 193 |
+
)
|
| 194 |
+
|
| 195 |
+
|
| 196 |
+
# ---------------- Aliases -----------------
|
| 197 |
+
# add aliases where verb goes first, json_load -> load_json
|
| 198 |
+
load_json = json_load
|
| 199 |
+
load_yaml = yaml_load
|
| 200 |
+
load_jsonl = jsonl_load
|
| 201 |
+
loads_json = json_loads
|
| 202 |
+
loads_yaml = yaml_loads
|
| 203 |
+
dump_json = json_dump
|
| 204 |
+
dump_jsonl = jsonl_dump
|
| 205 |
+
dump_yaml = yaml_dump
|
| 206 |
+
dumps_json = json_dumps
|
| 207 |
+
dumps_yaml = yaml_dumps
|
| 208 |
+
load_json_or_yaml = json_or_yaml_load
|
| 209 |
+
dump_json_or_yaml = json_or_yaml_dump
|
| 210 |
+
|
| 211 |
+
# ==================== Jsonl ====================
|
| 212 |
+
|
| 213 |
+
|
| 214 |
+
class Jsonl:
|
| 215 |
+
"""
|
| 216 |
+
Both reader and writer, as if everything's in-memory
|
| 217 |
+
"""
|
| 218 |
+
|
| 219 |
+
def __init__(self, *file_path, mode: Literal["r", "w", "a"] = "a"):
|
| 220 |
+
"""
|
| 221 |
+
Args:
|
| 222 |
+
mode:
|
| 223 |
+
- 'r': file must already exists
|
| 224 |
+
- 'w': overwrite the file regardless of whether it exists or not
|
| 225 |
+
- 'a': create a new file if doesn't exist, or append to an existing file
|
| 226 |
+
"""
|
| 227 |
+
assert mode in "rwa"
|
| 228 |
+
self._file_path = str(f_join(file_path))
|
| 229 |
+
self._mode = mode
|
| 230 |
+
if mode == "r":
|
| 231 |
+
assert path.exists(self._file_path)
|
| 232 |
+
self._fp = None
|
| 233 |
+
else:
|
| 234 |
+
self._fp = open(self._file_path, mode)
|
| 235 |
+
if path.exists(self._file_path) and mode != "w":
|
| 236 |
+
self.data = jsonl_load(self._file_path)
|
| 237 |
+
else:
|
| 238 |
+
self.data = []
|
| 239 |
+
|
| 240 |
+
def append(self, data: Dict):
|
| 241 |
+
if self._mode == "r":
|
| 242 |
+
raise RuntimeError("Jsonl read mode cannot call append()")
|
| 243 |
+
self.data.append(data)
|
| 244 |
+
print(json_dumps(data), file=self._fp, flush=True)
|
| 245 |
+
|
| 246 |
+
def extend(self, data_list: List[Dict]):
|
| 247 |
+
for data in data_list:
|
| 248 |
+
self.append(data)
|
| 249 |
+
|
| 250 |
+
def close(self):
|
| 251 |
+
if self._fp is not None:
|
| 252 |
+
self._fp.close()
|
| 253 |
+
|
| 254 |
+
def __getitem__(self, idx):
|
| 255 |
+
return self.data[idx]
|
| 256 |
+
|
| 257 |
+
def __len__(self):
|
| 258 |
+
return len(self.data)
|
| 259 |
+
|
| 260 |
+
def __iter__(self):
|
| 261 |
+
return iter(self.data)
|
| 262 |
+
|
| 263 |
+
def __enter__(self):
|
| 264 |
+
return self
|
| 265 |
+
|
| 266 |
+
def __exit__(self, type, value, traceback):
|
| 267 |
+
self.close()
|
| 268 |
+
|
| 269 |
+
def __bool__(self):
|
| 270 |
+
return bool(self.data)
|
groot/vla/common/utils/io/print_utils.py
ADDED
|
@@ -0,0 +1,362 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
from datetime import datetime
|
| 2 |
+
import io
|
| 3 |
+
import logging
|
| 4 |
+
import os
|
| 5 |
+
import pprint
|
| 6 |
+
import shlex
|
| 7 |
+
import string
|
| 8 |
+
import sys
|
| 9 |
+
import textwrap
|
| 10 |
+
import time
|
| 11 |
+
import traceback
|
| 12 |
+
from typing import Callable, Union
|
| 13 |
+
|
| 14 |
+
import numpy as np
|
| 15 |
+
from typing_extensions import Literal
|
| 16 |
+
|
| 17 |
+
from ..misc.functional_utils import meta_decorator
|
| 18 |
+
from ..misc.misc_utils import match_patterns
|
| 19 |
+
|
| 20 |
+
|
| 21 |
+
def to_readable_count_str(value: int, precision: int = 2) -> str:
|
| 22 |
+
assert value >= 0
|
| 23 |
+
labels = [" ", "K", "M", "B", "T"]
|
| 24 |
+
num_digits = int(np.floor(np.log10(value)) + 1 if value > 0 else 1)
|
| 25 |
+
num_groups = int(np.ceil(num_digits / 3))
|
| 26 |
+
num_groups = min(num_groups, len(labels)) # don't abbreviate beyond trillions
|
| 27 |
+
shift = -3 * (num_groups - 1)
|
| 28 |
+
value = value * (10**shift)
|
| 29 |
+
index = num_groups - 1
|
| 30 |
+
rem = value - int(value)
|
| 31 |
+
if precision > 0 and rem > 0.01:
|
| 32 |
+
fmt = f"{{:.{precision}f}}"
|
| 33 |
+
rem_str = fmt.format(rem).lstrip("0")
|
| 34 |
+
else:
|
| 35 |
+
rem_str = ""
|
| 36 |
+
return f"{int(value):,d}{rem_str} {labels[index]}"
|
| 37 |
+
|
| 38 |
+
|
| 39 |
+
def to_scientific_str(value, precision: int = 1, capitalize: bool = False) -> str:
|
| 40 |
+
"""
|
| 41 |
+
0.0015 -> "1.5e-3"
|
| 42 |
+
"""
|
| 43 |
+
if value == 0:
|
| 44 |
+
return "0"
|
| 45 |
+
return f"{value:.{precision}e}".replace("e-0", "E-" if capitalize else "e-")
|
| 46 |
+
|
| 47 |
+
|
| 48 |
+
def print_str(*args, **kwargs):
|
| 49 |
+
"""
|
| 50 |
+
Same as print() signature but returns a string
|
| 51 |
+
"""
|
| 52 |
+
sstream = io.StringIO()
|
| 53 |
+
kwargs.pop("file", None)
|
| 54 |
+
print(*args, **kwargs, file=sstream)
|
| 55 |
+
return sstream.getvalue()
|
| 56 |
+
|
| 57 |
+
|
| 58 |
+
def fstring(fmt_str, **kwargs):
|
| 59 |
+
"""
|
| 60 |
+
Simulate python f-string but without `f`
|
| 61 |
+
"""
|
| 62 |
+
locals().update(kwargs)
|
| 63 |
+
return eval("f" + shlex.quote(fmt_str))
|
| 64 |
+
|
| 65 |
+
|
| 66 |
+
def get_format_keys(fmt_str):
|
| 67 |
+
keys = []
|
| 68 |
+
for literal, field_name, fmt_spec, conversion in string.Formatter().parse(fmt_str):
|
| 69 |
+
if field_name:
|
| 70 |
+
keys.append(field_name)
|
| 71 |
+
return keys
|
| 72 |
+
|
| 73 |
+
|
| 74 |
+
def get_timestamp(milli_precision: int = 3):
|
| 75 |
+
fmt = "%y-%m-%d %H:%M:%S"
|
| 76 |
+
if milli_precision > 0:
|
| 77 |
+
fmt += ".%f"
|
| 78 |
+
stamp = datetime.now().strftime(fmt)
|
| 79 |
+
if milli_precision > 0:
|
| 80 |
+
stamp = stamp[:-milli_precision]
|
| 81 |
+
return stamp
|
| 82 |
+
|
| 83 |
+
|
| 84 |
+
def pretty_repr_str(obj, **kwargs):
|
| 85 |
+
"""
|
| 86 |
+
Useful to produce __repr__()
|
| 87 |
+
"""
|
| 88 |
+
if isinstance(obj, str):
|
| 89 |
+
cls_name = obj
|
| 90 |
+
else:
|
| 91 |
+
cls_name = obj.__class__.__name__
|
| 92 |
+
kw_strs = [k + "=" + pprint.pformat(v, indent=2, compact=True) for k, v in kwargs.items()]
|
| 93 |
+
new_line = len(cls_name) + sum(len(kw) for kw in kw_strs) > 84
|
| 94 |
+
if new_line:
|
| 95 |
+
kw = ",\n".join(kw_strs)
|
| 96 |
+
return f"{cls_name}(\n{textwrap.indent(kw, ' ')}\n)"
|
| 97 |
+
else:
|
| 98 |
+
kw = ", ".join(kw_strs)
|
| 99 |
+
return f"{cls_name}({kw})"
|
| 100 |
+
|
| 101 |
+
|
| 102 |
+
def pprint_(*objs, **kwargs):
|
| 103 |
+
"""
|
| 104 |
+
Use pprint to format the objects
|
| 105 |
+
"""
|
| 106 |
+
print(
|
| 107 |
+
*[pprint.pformat(obj, indent=2) if not isinstance(obj, str) else obj for obj in objs],
|
| 108 |
+
**kwargs,
|
| 109 |
+
)
|
| 110 |
+
|
| 111 |
+
|
| 112 |
+
def get_exception_info(to_str: bool = False):
|
| 113 |
+
"""
|
| 114 |
+
Returns:
|
| 115 |
+
{'type': ExceptionType, 'value': ExceptionObject, 'trace': <traceback str>}
|
| 116 |
+
"""
|
| 117 |
+
typ_, value, trace = sys.exc_info()
|
| 118 |
+
return {
|
| 119 |
+
"type": typ_.__name__ if to_str else typ_,
|
| 120 |
+
"value": str(value) if to_str else value,
|
| 121 |
+
"trace": "".join(traceback.format_exception(typ_, value, trace)),
|
| 122 |
+
}
|
| 123 |
+
|
| 124 |
+
|
| 125 |
+
class DebugPrinter:
|
| 126 |
+
"""
|
| 127 |
+
Debug print, usage: dprint = DebugPrint(enabled=True)
|
| 128 |
+
dprint(...)
|
| 129 |
+
"""
|
| 130 |
+
|
| 131 |
+
def __init__(self, enabled, tensor_summary: Literal["shape", "shape+dtype", "none"] = "shape"):
|
| 132 |
+
"""
|
| 133 |
+
Args:
|
| 134 |
+
tensor_summary:
|
| 135 |
+
- shape: only prints shape
|
| 136 |
+
- shape+dtype: also prints dtype and device
|
| 137 |
+
- none: print full tensor
|
| 138 |
+
"""
|
| 139 |
+
self.enabled = enabled
|
| 140 |
+
assert tensor_summary in ["shape", "shape+dtype", "none"]
|
| 141 |
+
self.tensor_summary = tensor_summary
|
| 142 |
+
|
| 143 |
+
def __call__(self, *args, **kwargs):
|
| 144 |
+
if not self.enabled:
|
| 145 |
+
return
|
| 146 |
+
args = [self._process_arg(a) for a in args]
|
| 147 |
+
pprint_(*args, **kwargs)
|
| 148 |
+
|
| 149 |
+
def _process_arg(self, arg):
|
| 150 |
+
import numpy as np
|
| 151 |
+
import torch
|
| 152 |
+
|
| 153 |
+
if torch.is_tensor(arg):
|
| 154 |
+
if self.tensor_summary == "shape":
|
| 155 |
+
return str(list(arg.size()))
|
| 156 |
+
elif self.tensor_summary == "shape+dtype":
|
| 157 |
+
return f"{arg.dtype}{list(arg.size())}|{arg.device}"
|
| 158 |
+
elif isinstance(arg, np.ndarray):
|
| 159 |
+
if self.tensor_summary == "shape":
|
| 160 |
+
return str(list(arg.shape))
|
| 161 |
+
elif self.tensor_summary == "shape+dtype":
|
| 162 |
+
return f"{arg.dtype}{list(arg.shape)}"
|
| 163 |
+
return arg
|
| 164 |
+
|
| 165 |
+
|
| 166 |
+
@meta_decorator
|
| 167 |
+
def watch(func, seconds: int = 5, max_times: int = 0, keep_returns: bool = False):
|
| 168 |
+
"""
|
| 169 |
+
Decorator: executes a function repeated with the args and
|
| 170 |
+
emulate `watch -n` capability
|
| 171 |
+
|
| 172 |
+
See `gpustat` repo: https://github.com/wookayin/gpustat/pull/41/files
|
| 173 |
+
|
| 174 |
+
Args:
|
| 175 |
+
max_times: watch for `max_times` and then exit. If 0, never exits
|
| 176 |
+
keep_returns: if True, will keep the return value from the function
|
| 177 |
+
and return as a list at the end
|
| 178 |
+
"""
|
| 179 |
+
from blessings import Terminal
|
| 180 |
+
|
| 181 |
+
def _wrapped(*args, **kwargs):
|
| 182 |
+
term = Terminal()
|
| 183 |
+
N = 0
|
| 184 |
+
returns = []
|
| 185 |
+
with term.fullscreen():
|
| 186 |
+
while True:
|
| 187 |
+
try:
|
| 188 |
+
with term.location(0, 0):
|
| 189 |
+
ret = func(*args, **kwargs)
|
| 190 |
+
print(term.clear_eos, end="")
|
| 191 |
+
if keep_returns:
|
| 192 |
+
returns.append(ret)
|
| 193 |
+
N += 1
|
| 194 |
+
if max_times > 0 and N >= max_times:
|
| 195 |
+
break
|
| 196 |
+
time.sleep(seconds)
|
| 197 |
+
except KeyboardInterrupt:
|
| 198 |
+
break
|
| 199 |
+
return returns
|
| 200 |
+
|
| 201 |
+
return _wrapped
|
| 202 |
+
|
| 203 |
+
|
| 204 |
+
class PrintRedirection(object):
|
| 205 |
+
"""
|
| 206 |
+
Context manager: temporarily redirects stdout and stderr
|
| 207 |
+
"""
|
| 208 |
+
|
| 209 |
+
def __init__(self, stdout=None, stderr=None):
|
| 210 |
+
"""
|
| 211 |
+
Args:
|
| 212 |
+
stdout: if None, defaults to sys.stdout, unchanged
|
| 213 |
+
stderr: if None, defaults to sys.stderr, unchanged
|
| 214 |
+
"""
|
| 215 |
+
if stdout is None:
|
| 216 |
+
stdout = sys.stdout
|
| 217 |
+
if stderr is None:
|
| 218 |
+
stderr = sys.stderr
|
| 219 |
+
self._stdout, self._stderr = stdout, stderr
|
| 220 |
+
|
| 221 |
+
def __enter__(self):
|
| 222 |
+
self._old_out, self._old_err = sys.stdout, sys.stderr
|
| 223 |
+
self._old_out.flush()
|
| 224 |
+
self._old_err.flush()
|
| 225 |
+
sys.stdout, sys.stderr = self._stdout, self._stderr
|
| 226 |
+
return self
|
| 227 |
+
|
| 228 |
+
def __exit__(self, exc_type, exc_value, traceback):
|
| 229 |
+
self.flush()
|
| 230 |
+
# restore the normal stdout and stderr
|
| 231 |
+
sys.stdout, sys.stderr = self._old_out, self._old_err
|
| 232 |
+
|
| 233 |
+
def flush(self):
|
| 234 |
+
"Manually flush the replaced stdout/stderr buffers."
|
| 235 |
+
self._stdout.flush()
|
| 236 |
+
self._stderr.flush()
|
| 237 |
+
|
| 238 |
+
|
| 239 |
+
class PrintToFile(PrintRedirection):
|
| 240 |
+
"""
|
| 241 |
+
Print to file and save/close the handle at the end.
|
| 242 |
+
"""
|
| 243 |
+
|
| 244 |
+
def __init__(self, out_file=None, err_file=None):
|
| 245 |
+
"""
|
| 246 |
+
Args:
|
| 247 |
+
out_file: file path
|
| 248 |
+
err_file: file path. If the same as out_file, print both stdout
|
| 249 |
+
and stderr to one file in order.
|
| 250 |
+
"""
|
| 251 |
+
self.out_file, self.err_file = out_file, err_file
|
| 252 |
+
if out_file:
|
| 253 |
+
out_file = os.path.expanduser(out_file)
|
| 254 |
+
self.out_file = open(out_file, "w")
|
| 255 |
+
if err_file:
|
| 256 |
+
err_file = os.path.expanduser(out_file)
|
| 257 |
+
if err_file == out_file: # redirect both stdout/err to one file
|
| 258 |
+
self.err_file = self.out_file
|
| 259 |
+
else:
|
| 260 |
+
self.err_file = open(os.path.expanduser(out_file), "w")
|
| 261 |
+
super().__init__(stdout=self.out_file, stderr=self.err_file)
|
| 262 |
+
|
| 263 |
+
def __exit__(self, *args):
|
| 264 |
+
super().__exit__(*args)
|
| 265 |
+
if self.out_file:
|
| 266 |
+
self.out_file.close()
|
| 267 |
+
if self.err_file:
|
| 268 |
+
self.err_file.close()
|
| 269 |
+
|
| 270 |
+
|
| 271 |
+
def PrintSuppress(no_out=True, no_err=False):
|
| 272 |
+
"""
|
| 273 |
+
Args:
|
| 274 |
+
no_out: stdout writes to sys.devnull
|
| 275 |
+
no_err: stderr writes to sys.devnull
|
| 276 |
+
"""
|
| 277 |
+
out_file = os.devnull if no_out else None
|
| 278 |
+
err_file = os.devnull if no_err else None
|
| 279 |
+
return PrintToFile(out_file=out_file, err_file=err_file)
|
| 280 |
+
|
| 281 |
+
|
| 282 |
+
class PrintString(PrintRedirection):
|
| 283 |
+
"""
|
| 284 |
+
Redirect stdout and stderr to strings.
|
| 285 |
+
"""
|
| 286 |
+
|
| 287 |
+
def __init__(self):
|
| 288 |
+
self.out_stream = io.StringIO()
|
| 289 |
+
self.err_stream = io.StringIO()
|
| 290 |
+
super().__init__(stdout=self.out_stream, stderr=self.err_stream)
|
| 291 |
+
|
| 292 |
+
def stdout(self):
|
| 293 |
+
"Returns: stdout as one string."
|
| 294 |
+
return self.out_stream.getvalue()
|
| 295 |
+
|
| 296 |
+
def stderr(self):
|
| 297 |
+
"Returns: stderr as one string."
|
| 298 |
+
return self.err_stream.getvalue()
|
| 299 |
+
|
| 300 |
+
def stdout_by_line(self):
|
| 301 |
+
"Returns: a list of stdout line by line, ignore trailing blanks"
|
| 302 |
+
return self.stdout().rstrip().split("\n")
|
| 303 |
+
|
| 304 |
+
def stderr_by_line(self):
|
| 305 |
+
"Returns: a list of stderr line by line, ignore trailing blanks"
|
| 306 |
+
return self.stderr().rstrip().split("\n")
|
| 307 |
+
|
| 308 |
+
|
| 309 |
+
# ==================== Logging filters ====================
|
| 310 |
+
class ExcludeLoggingFilter(logging.Filter):
|
| 311 |
+
"""
|
| 312 |
+
Usage: logging.getLogger('name').addFilter(
|
| 313 |
+
ExcludeLoggingFilter(['info mess*age', 'Warning: *'])
|
| 314 |
+
)
|
| 315 |
+
Supports wildcard.
|
| 316 |
+
https://relaxdiego.com/2014/07/logging-in-python.html
|
| 317 |
+
"""
|
| 318 |
+
|
| 319 |
+
def __init__(self, patterns):
|
| 320 |
+
super().__init__()
|
| 321 |
+
self._patterns = patterns
|
| 322 |
+
|
| 323 |
+
def filter(self, record):
|
| 324 |
+
if match_patterns(record.msg, include=self._patterns):
|
| 325 |
+
return False
|
| 326 |
+
else:
|
| 327 |
+
return True
|
| 328 |
+
|
| 329 |
+
|
| 330 |
+
class ReplaceStringLoggingFilter(logging.Filter):
|
| 331 |
+
def __init__(self, patterns, replacer: Callable):
|
| 332 |
+
super().__init__()
|
| 333 |
+
self._patterns = patterns
|
| 334 |
+
assert callable(replacer)
|
| 335 |
+
self._replacer = replacer
|
| 336 |
+
|
| 337 |
+
def filter(self, record):
|
| 338 |
+
if match_patterns(record.msg, include=self._patterns):
|
| 339 |
+
record.msg = self._replacer(record.msg)
|
| 340 |
+
|
| 341 |
+
|
| 342 |
+
def logging_exclude_pattern(
|
| 343 |
+
logger_name,
|
| 344 |
+
patterns: Union[str, list[str], Callable, list[Callable], None],
|
| 345 |
+
):
|
| 346 |
+
"""
|
| 347 |
+
Args:
|
| 348 |
+
patterns: see groot.vla.common.utils.misc_utils.match_patterns
|
| 349 |
+
"""
|
| 350 |
+
logging.getLogger(logger_name).addFilter(ExcludeLoggingFilter(patterns))
|
| 351 |
+
|
| 352 |
+
|
| 353 |
+
def logging_replace_string(
|
| 354 |
+
logger_name,
|
| 355 |
+
patterns: Union[str, list[str], Callable, list[Callable], None],
|
| 356 |
+
replacer: Callable,
|
| 357 |
+
):
|
| 358 |
+
"""
|
| 359 |
+
Args:
|
| 360 |
+
patterns: see groot.vla.common.utils.misc_utils.match_patterns
|
| 361 |
+
"""
|
| 362 |
+
logging.getLogger(logger_name).addFilter(ReplaceStringLoggingFilter(patterns, replacer))
|
groot/vla/common/utils/io/termcolor.py
ADDED
|
@@ -0,0 +1,186 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# coding: utf-8
|
| 2 |
+
# Copyright (c) 2008-2011 Volvox Development Team
|
| 3 |
+
#
|
| 4 |
+
# Permission is hereby granted, free of charge, to any person obtaining a copy
|
| 5 |
+
# of this software and associated documentation files (the "Software"), to deal
|
| 6 |
+
# in the Software without restriction, including without limitation the rights
|
| 7 |
+
# to use, copy, modify, merge, publish, distribute, sublicense, and/or sell
|
| 8 |
+
# copies of the Software, and to permit persons to whom the Software is
|
| 9 |
+
# furnished to do so, subject to the following conditions:
|
| 10 |
+
#
|
| 11 |
+
# The above copyright notice and this permission notice shall be included in
|
| 12 |
+
# all copies or substantial portions of the Software.
|
| 13 |
+
#
|
| 14 |
+
# THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
|
| 15 |
+
# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY,
|
| 16 |
+
# FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE
|
| 17 |
+
# AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER
|
| 18 |
+
# LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM,
|
| 19 |
+
# OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN
|
| 20 |
+
# THE SOFTWARE.
|
| 21 |
+
#
|
| 22 |
+
# Original Author: Konstantin Lepa <konstantin.lepa@gmail.com>
|
| 23 |
+
# Updated by Jim Fan
|
| 24 |
+
|
| 25 |
+
"""ANSII Color formatting for output in terminal."""
|
| 26 |
+
import io
|
| 27 |
+
import os
|
| 28 |
+
from typing import List, Optional, Union
|
| 29 |
+
|
| 30 |
+
__ALL__ = ["color_text", "cprint"]
|
| 31 |
+
|
| 32 |
+
STYLES = dict(
|
| 33 |
+
list(
|
| 34 |
+
zip(
|
| 35 |
+
["bold", "dark", "", "underline", "blink", "", "reverse", "concealed"],
|
| 36 |
+
list(range(1, 9)),
|
| 37 |
+
)
|
| 38 |
+
)
|
| 39 |
+
)
|
| 40 |
+
del STYLES[""]
|
| 41 |
+
|
| 42 |
+
|
| 43 |
+
HIGHLIGHTS = dict(
|
| 44 |
+
list(
|
| 45 |
+
zip(
|
| 46 |
+
["grey", "red", "green", "yellow", "blue", "magenta", "cyan", "white"],
|
| 47 |
+
list(range(40, 48)),
|
| 48 |
+
)
|
| 49 |
+
)
|
| 50 |
+
)
|
| 51 |
+
|
| 52 |
+
|
| 53 |
+
COLORS = dict(
|
| 54 |
+
list(
|
| 55 |
+
zip(
|
| 56 |
+
["grey", "red", "green", "yellow", "blue", "magenta", "cyan", "white"],
|
| 57 |
+
list(range(30, 38)),
|
| 58 |
+
)
|
| 59 |
+
)
|
| 60 |
+
)
|
| 61 |
+
|
| 62 |
+
|
| 63 |
+
def _strip_bg_prefix(color):
|
| 64 |
+
"on_red -> red"
|
| 65 |
+
if color.startswith("on_"):
|
| 66 |
+
return color[len("on_") :]
|
| 67 |
+
else:
|
| 68 |
+
return color
|
| 69 |
+
|
| 70 |
+
|
| 71 |
+
RESET = "\033[0m"
|
| 72 |
+
|
| 73 |
+
|
| 74 |
+
def color_text(
|
| 75 |
+
text,
|
| 76 |
+
color: Optional[str] = None,
|
| 77 |
+
bg_color: Optional[str] = None,
|
| 78 |
+
styles: Optional[Union[str, List[str]]] = None,
|
| 79 |
+
):
|
| 80 |
+
"""Colorize text.
|
| 81 |
+
|
| 82 |
+
Available text colors:
|
| 83 |
+
red, green, yellow, blue, magenta, cyan, white.
|
| 84 |
+
|
| 85 |
+
Available text highlights:
|
| 86 |
+
on_red, on_green, on_yellow, on_blue, on_magenta, on_cyan, on_white.
|
| 87 |
+
|
| 88 |
+
Available attributes:
|
| 89 |
+
bold, dark, underline, blink, reverse, concealed.
|
| 90 |
+
|
| 91 |
+
Example:
|
| 92 |
+
colored('Hello, World!', 'red', 'on_grey', ['blue', 'blink'])
|
| 93 |
+
colored('Hello, World!', 'green')
|
| 94 |
+
"""
|
| 95 |
+
if os.getenv("ANSI_COLORS_DISABLED") is None:
|
| 96 |
+
fmt_str = "\033[%dm%s"
|
| 97 |
+
if color is not None:
|
| 98 |
+
text = fmt_str % (COLORS[color], text)
|
| 99 |
+
|
| 100 |
+
if bg_color is not None:
|
| 101 |
+
bg_color = _strip_bg_prefix(bg_color)
|
| 102 |
+
text = fmt_str % (HIGHLIGHTS[bg_color], text)
|
| 103 |
+
|
| 104 |
+
if styles is not None:
|
| 105 |
+
if isinstance(styles, str):
|
| 106 |
+
styles = [styles]
|
| 107 |
+
for style in styles:
|
| 108 |
+
text = fmt_str % (STYLES[style], text)
|
| 109 |
+
|
| 110 |
+
text += RESET
|
| 111 |
+
return text
|
| 112 |
+
|
| 113 |
+
|
| 114 |
+
def cprint(
|
| 115 |
+
*args,
|
| 116 |
+
color: Optional[str] = None,
|
| 117 |
+
bg_color: Optional[str] = None,
|
| 118 |
+
styles: Optional[Union[str, List[str]]] = None,
|
| 119 |
+
**kwargs,
|
| 120 |
+
):
|
| 121 |
+
"""Print colorize text.
|
| 122 |
+
|
| 123 |
+
It accepts arguments of print function.
|
| 124 |
+
"""
|
| 125 |
+
sstream = io.StringIO()
|
| 126 |
+
print(*args, sep=kwargs.pop("sep", None), end="", file=sstream)
|
| 127 |
+
text = sstream.getvalue()
|
| 128 |
+
print((color_text(text, color, bg_color, styles)), **kwargs)
|
| 129 |
+
|
| 130 |
+
|
| 131 |
+
if __name__ == "__main__":
|
| 132 |
+
print("Current terminal type: %s" % os.getenv("TERM"))
|
| 133 |
+
print("Test basic colors:")
|
| 134 |
+
cprint("Grey color", color="grey")
|
| 135 |
+
cprint("Red color", color="red")
|
| 136 |
+
cprint("Green color", color="green")
|
| 137 |
+
cprint("Yellow color", color="yellow")
|
| 138 |
+
cprint("Blue color", color="blue")
|
| 139 |
+
cprint("Magenta color", color="magenta")
|
| 140 |
+
cprint("Cyan color", color="cyan")
|
| 141 |
+
cprint("White color", color="white")
|
| 142 |
+
print(("-" * 78))
|
| 143 |
+
|
| 144 |
+
print("Test highlights:")
|
| 145 |
+
cprint("On grey color", bg_color="on_grey")
|
| 146 |
+
cprint("On red color", bg_color="on_red")
|
| 147 |
+
cprint("On green color", bg_color="on_green")
|
| 148 |
+
cprint("On yellow color", bg_color="on_yellow")
|
| 149 |
+
cprint("On blue color", bg_color="on_blue")
|
| 150 |
+
cprint("On magenta color", bg_color="on_magenta")
|
| 151 |
+
cprint("On cyan color", bg_color="on_cyan")
|
| 152 |
+
cprint("On white color", color="grey", bg_color="on_white")
|
| 153 |
+
print("-" * 78)
|
| 154 |
+
|
| 155 |
+
print("Test attributes:")
|
| 156 |
+
cprint("Bold grey color", color="grey", styles="bold")
|
| 157 |
+
cprint("Dark red color", color="red", styles=["dark"])
|
| 158 |
+
cprint("Underline green color", color="green", styles=["underline"])
|
| 159 |
+
cprint("Blink yellow color", color="yellow", styles=["blink"])
|
| 160 |
+
cprint("Reversed blue color", color="blue", styles=["reverse"])
|
| 161 |
+
cprint("Concealed Magenta color", color="magenta", styles=["concealed"])
|
| 162 |
+
cprint(
|
| 163 |
+
"Bold underline reverse cyan color",
|
| 164 |
+
color="cyan",
|
| 165 |
+
styles=["bold", "underline", "reverse"],
|
| 166 |
+
)
|
| 167 |
+
cprint(
|
| 168 |
+
"Dark blink concealed white color",
|
| 169 |
+
color="white",
|
| 170 |
+
styles=["dark", "blink", "concealed"],
|
| 171 |
+
)
|
| 172 |
+
print(("-" * 78))
|
| 173 |
+
|
| 174 |
+
print("Test mixing:")
|
| 175 |
+
cprint(
|
| 176 |
+
"Underline red on grey color",
|
| 177 |
+
color="red",
|
| 178 |
+
bg_color="on_grey",
|
| 179 |
+
styles="underline",
|
| 180 |
+
)
|
| 181 |
+
cprint(
|
| 182 |
+
"Reversed green on red color",
|
| 183 |
+
color="green",
|
| 184 |
+
bg_color="on_red",
|
| 185 |
+
styles="reverse",
|
| 186 |
+
)
|
groot/vla/common/utils/misc/__init__.py
ADDED
|
@@ -0,0 +1,5 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
from .functional_utils import * # noqa: F403
|
| 2 |
+
from .image_utils import * # noqa: F403
|
| 3 |
+
from .misc_utils import * # noqa: F403
|
| 4 |
+
from .torch_utils import * # noqa: F403
|
| 5 |
+
from .video_utils import * # noqa: F403
|
groot/vla/common/utils/misc/array_tensor_utils.py
ADDED
|
@@ -0,0 +1,372 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""
|
| 2 |
+
Functions that work on nested structures of torch.Tensor or numpy array
|
| 3 |
+
"""
|
| 4 |
+
|
| 5 |
+
from typing import Any, Dict, List, Optional, Union
|
| 6 |
+
|
| 7 |
+
import numpy as np
|
| 8 |
+
import torch
|
| 9 |
+
import tree
|
| 10 |
+
|
| 11 |
+
from ..data_structure.tree_utils import (
|
| 12 |
+
copy_non_leaf,
|
| 13 |
+
is_sequence,
|
| 14 |
+
tree_assign_at_path,
|
| 15 |
+
tree_value_at_path,
|
| 16 |
+
)
|
| 17 |
+
from .functional_utils import make_recursive_func
|
| 18 |
+
|
| 19 |
+
|
| 20 |
+
def is_array_tensor(obj):
|
| 21 |
+
return isinstance(obj, (np.ndarray, torch.Tensor))
|
| 22 |
+
|
| 23 |
+
|
| 24 |
+
def is_numpy(obj):
|
| 25 |
+
return isinstance(obj, np.ndarray)
|
| 26 |
+
|
| 27 |
+
|
| 28 |
+
def is_tensor(obj):
|
| 29 |
+
return torch.is_tensor(obj)
|
| 30 |
+
|
| 31 |
+
|
| 32 |
+
def any_stack(xs: List, *, dim: int = 0):
|
| 33 |
+
"""
|
| 34 |
+
Works for both torch Tensor and numpy array
|
| 35 |
+
"""
|
| 36 |
+
|
| 37 |
+
def _any_stack_helper(*xs):
|
| 38 |
+
x = xs[0]
|
| 39 |
+
if isinstance(x, np.ndarray):
|
| 40 |
+
return np.stack(xs, axis=dim)
|
| 41 |
+
elif torch.is_tensor(x):
|
| 42 |
+
return torch.stack(xs, dim=dim)
|
| 43 |
+
elif isinstance(x, float):
|
| 44 |
+
# special treatment for float, defaults to float32
|
| 45 |
+
return np.array(xs, dtype=np.float32)
|
| 46 |
+
else:
|
| 47 |
+
return np.array(xs)
|
| 48 |
+
|
| 49 |
+
return tree.map_structure(_any_stack_helper, *xs)
|
| 50 |
+
|
| 51 |
+
|
| 52 |
+
def any_concat(xs: List, *, dim: int = 0):
|
| 53 |
+
"""
|
| 54 |
+
Works for both torch Tensor and numpy array
|
| 55 |
+
"""
|
| 56 |
+
|
| 57 |
+
def _any_concat_helper(*xs):
|
| 58 |
+
x = xs[0]
|
| 59 |
+
if isinstance(x, np.ndarray):
|
| 60 |
+
return np.concatenate(xs, axis=dim)
|
| 61 |
+
elif torch.is_tensor(x):
|
| 62 |
+
return torch.cat(xs, dim=dim)
|
| 63 |
+
elif isinstance(x, float):
|
| 64 |
+
# special treatment for float, defaults to float32
|
| 65 |
+
return np.array(xs, dtype=np.float32)
|
| 66 |
+
else:
|
| 67 |
+
return np.array(xs)
|
| 68 |
+
|
| 69 |
+
return tree.map_structure(_any_concat_helper, *xs)
|
| 70 |
+
|
| 71 |
+
|
| 72 |
+
def any_chunk(x, chunks: int, *, dim: int = 0, strict: bool = True) -> List[Any]:
|
| 73 |
+
"""
|
| 74 |
+
Works for both torch Tensor and numpy array
|
| 75 |
+
|
| 76 |
+
Returns:
|
| 77 |
+
list of chunked nested structures
|
| 78 |
+
"""
|
| 79 |
+
assert chunks >= 1
|
| 80 |
+
|
| 81 |
+
x_copies = [copy_non_leaf(x) for _ in range(chunks)]
|
| 82 |
+
|
| 83 |
+
def _any_chunk_helper(path, x):
|
| 84 |
+
if is_array_tensor(x):
|
| 85 |
+
if isinstance(x, np.ndarray):
|
| 86 |
+
chunked_values = np.split(x, chunks, axis=dim)
|
| 87 |
+
else:
|
| 88 |
+
chunked_values = torch.chunk(x, chunks, dim=dim)
|
| 89 |
+
|
| 90 |
+
if path:
|
| 91 |
+
for xc, chunked in zip(x_copies, chunked_values):
|
| 92 |
+
tree_assign_at_path(xc, path, chunked)
|
| 93 |
+
else: # top-level, no nested path
|
| 94 |
+
for i, chunked in enumerate(chunked_values):
|
| 95 |
+
x_copies[i] = chunked
|
| 96 |
+
else:
|
| 97 |
+
if strict:
|
| 98 |
+
raise NotImplementedError(f"Cannot chunk type {type(x)}")
|
| 99 |
+
else:
|
| 100 |
+
return
|
| 101 |
+
|
| 102 |
+
tree.map_structure_with_path(_any_chunk_helper, x)
|
| 103 |
+
return x_copies
|
| 104 |
+
|
| 105 |
+
|
| 106 |
+
def chunk_seq(arr, chunks: int, check_divide=True):
|
| 107 |
+
"""
|
| 108 |
+
Args:
|
| 109 |
+
check_divide: True to force arr must divide n
|
| 110 |
+
"""
|
| 111 |
+
k, m = divmod(len(arr), chunks)
|
| 112 |
+
if check_divide and m != 0:
|
| 113 |
+
raise ValueError(f"Array len {len(arr)} does not divide chunks {chunks}")
|
| 114 |
+
return (arr[i * k + min(i, m) : (i + 1) * k + min(i + 1, m)] for i in range(chunks))
|
| 115 |
+
|
| 116 |
+
|
| 117 |
+
@make_recursive_func
|
| 118 |
+
def any_zeros_like(x: Union[Dict, np.ndarray, torch.Tensor, int, float, np.number]):
|
| 119 |
+
"""Returns a zero-filled object of the same (d)type and shape as the input.
|
| 120 |
+
|
| 121 |
+
The difference between this and `np.zeros_like()` is that this works well
|
| 122 |
+
with `np.number`, `int`, `float`, and `jax.numpy.DeviceArray` objects without
|
| 123 |
+
converting them to `np.ndarray`s.
|
| 124 |
+
|
| 125 |
+
Args:
|
| 126 |
+
x: The object to replace with 0s.
|
| 127 |
+
|
| 128 |
+
Returns:
|
| 129 |
+
A zero-filed object of the same (d)type and shape as the input.
|
| 130 |
+
"""
|
| 131 |
+
if isinstance(x, (int, float, np.number)):
|
| 132 |
+
return type(x)(0)
|
| 133 |
+
elif is_tensor(x):
|
| 134 |
+
return torch.zeros_like(x)
|
| 135 |
+
elif is_numpy(x):
|
| 136 |
+
return np.zeros_like(x)
|
| 137 |
+
else:
|
| 138 |
+
raise ValueError(
|
| 139 |
+
f"Input ({type(x)}) must be either a numpy array, a tensor, an int, or a float."
|
| 140 |
+
)
|
| 141 |
+
|
| 142 |
+
|
| 143 |
+
@make_recursive_func
|
| 144 |
+
def any_ones_like(x: Union[Dict, np.ndarray, torch.Tensor, int, float, np.number]):
|
| 145 |
+
"""Returns a one-filled object of the same (d)type and shape as the input.
|
| 146 |
+
The difference between this and `np.ones_like()` is that this works well
|
| 147 |
+
with `np.number`, `int`, `float`, and `jax.numpy.DeviceArray` objects without
|
| 148 |
+
converting them to `np.ndarray`s.
|
| 149 |
+
Args:
|
| 150 |
+
x: The object to replace with 1s.
|
| 151 |
+
Returns:
|
| 152 |
+
A one-filed object of the same (d)type and shape as the input.
|
| 153 |
+
"""
|
| 154 |
+
if isinstance(x, (int, float, np.number)):
|
| 155 |
+
return type(x)(1)
|
| 156 |
+
elif is_tensor(x):
|
| 157 |
+
return torch.ones_like(x)
|
| 158 |
+
elif is_numpy(x):
|
| 159 |
+
return np.ones_like(x)
|
| 160 |
+
else:
|
| 161 |
+
raise ValueError(
|
| 162 |
+
f"Input ({type(x)}) must be either a numpy array, a tensor, an int, or a float."
|
| 163 |
+
)
|
| 164 |
+
|
| 165 |
+
|
| 166 |
+
@make_recursive_func
|
| 167 |
+
def any_zero_(x: Union[Dict, np.ndarray, torch.Tensor]):
|
| 168 |
+
"""
|
| 169 |
+
Apply in-place zero-out to a tensor, i.e. x.zero_()
|
| 170 |
+
"""
|
| 171 |
+
if is_tensor(x):
|
| 172 |
+
x.zero_()
|
| 173 |
+
elif is_numpy(x):
|
| 174 |
+
x.fill(0)
|
| 175 |
+
else:
|
| 176 |
+
raise ValueError(f"Input ({type(x)}) must be either a numpy array or a tensor")
|
| 177 |
+
|
| 178 |
+
|
| 179 |
+
@make_recursive_func
|
| 180 |
+
def any_fill_(x: Union[Dict, np.ndarray, torch.Tensor], value):
|
| 181 |
+
"""
|
| 182 |
+
Apply in-place zero-out to a tensor, i.e. x.zero_()
|
| 183 |
+
"""
|
| 184 |
+
if is_tensor(x):
|
| 185 |
+
x.fill_(value)
|
| 186 |
+
elif is_numpy(x):
|
| 187 |
+
x.fill(value)
|
| 188 |
+
else:
|
| 189 |
+
raise ValueError(f"Input ({type(x)}) must be either a numpy array or a tensor")
|
| 190 |
+
|
| 191 |
+
|
| 192 |
+
def get_batch_size(x, strict: bool = False) -> int:
|
| 193 |
+
"""
|
| 194 |
+
Args:
|
| 195 |
+
x: can be any arbitrary nested structure of np array and torch tensor
|
| 196 |
+
strict: True to check all batch sizes are the same
|
| 197 |
+
"""
|
| 198 |
+
|
| 199 |
+
def _get_batch_size(x):
|
| 200 |
+
if isinstance(x, np.ndarray):
|
| 201 |
+
return x.shape[0]
|
| 202 |
+
elif torch.is_tensor(x):
|
| 203 |
+
return x.size(0)
|
| 204 |
+
else:
|
| 205 |
+
return len(x)
|
| 206 |
+
|
| 207 |
+
xs = tree.flatten(x)
|
| 208 |
+
|
| 209 |
+
if strict:
|
| 210 |
+
batch_sizes = [_get_batch_size(x) for x in xs]
|
| 211 |
+
assert all(
|
| 212 |
+
b == batch_sizes[0] for b in batch_sizes
|
| 213 |
+
), f"batch sizes must all be the same in nested structure: {batch_sizes}"
|
| 214 |
+
return batch_sizes[0]
|
| 215 |
+
else:
|
| 216 |
+
return _get_batch_size(xs[0])
|
| 217 |
+
|
| 218 |
+
|
| 219 |
+
@make_recursive_func
|
| 220 |
+
def add_batch_dim(x):
|
| 221 |
+
if is_numpy(x):
|
| 222 |
+
return np.expand_dims(x, axis=0)
|
| 223 |
+
elif is_tensor(x):
|
| 224 |
+
return x.unsqueeze(0)
|
| 225 |
+
else:
|
| 226 |
+
raise NotImplementedError(f"Unsupported data structure: {type(x)}")
|
| 227 |
+
|
| 228 |
+
|
| 229 |
+
@make_recursive_func
|
| 230 |
+
def remove_batch_dim(x):
|
| 231 |
+
if is_numpy(x):
|
| 232 |
+
return np.squeeze(x, axis=0)
|
| 233 |
+
elif is_tensor(x):
|
| 234 |
+
return x.squeeze(0)
|
| 235 |
+
else:
|
| 236 |
+
raise NotImplementedError(f"Unsupported data structure: {type(x)}")
|
| 237 |
+
|
| 238 |
+
|
| 239 |
+
@make_recursive_func
|
| 240 |
+
def any_to_primitive(x):
|
| 241 |
+
if isinstance(x, (np.ndarray, np.number, torch.Tensor)):
|
| 242 |
+
return x.tolist()
|
| 243 |
+
else:
|
| 244 |
+
return x
|
| 245 |
+
|
| 246 |
+
|
| 247 |
+
@make_recursive_func
|
| 248 |
+
def any_get_shape(x):
|
| 249 |
+
if is_numpy(x):
|
| 250 |
+
return tuple(x.shape)
|
| 251 |
+
elif is_tensor(x):
|
| 252 |
+
return tuple(x.size())
|
| 253 |
+
else:
|
| 254 |
+
raise NotImplementedError(f"Unsupported data structure: {type(x)}")
|
| 255 |
+
|
| 256 |
+
|
| 257 |
+
@make_recursive_func
|
| 258 |
+
def any_mean(x, dim: Optional[int] = None, keepdim: bool = False):
|
| 259 |
+
if is_numpy(x):
|
| 260 |
+
return np.mean(x, axis=dim, keepdims=keepdim)
|
| 261 |
+
elif is_tensor(x):
|
| 262 |
+
return torch.mean(x, dim=dim, keepdim=keepdim)
|
| 263 |
+
else:
|
| 264 |
+
raise NotImplementedError(f"Unsupported data structure: {type(x)}")
|
| 265 |
+
|
| 266 |
+
|
| 267 |
+
@make_recursive_func
|
| 268 |
+
def any_variance(x, dim: Optional[int] = None, keepdim: bool = False, unbiased: bool = False):
|
| 269 |
+
if is_numpy(x):
|
| 270 |
+
return np.var(x, axis=dim, keepdims=keepdim, ddof=1 if unbiased else 0)
|
| 271 |
+
elif is_tensor(x):
|
| 272 |
+
return torch.var(x, dim=dim, keepdim=keepdim, unbiased=unbiased)
|
| 273 |
+
else:
|
| 274 |
+
raise NotImplementedError(f"Unsupported data structure: {type(x)}")
|
| 275 |
+
|
| 276 |
+
|
| 277 |
+
@make_recursive_func
|
| 278 |
+
def any_describe_str(x, shape_only=False):
|
| 279 |
+
"""
|
| 280 |
+
Describe type, shape, device, data type (of np array/tensor)
|
| 281 |
+
Very useful for debugging
|
| 282 |
+
"""
|
| 283 |
+
t = type(x)
|
| 284 |
+
tname = type(x).__name__
|
| 285 |
+
if is_numpy(x):
|
| 286 |
+
shape = list(x.shape)
|
| 287 |
+
if x.size == 1:
|
| 288 |
+
if shape_only:
|
| 289 |
+
return f"np scalar: {x.item()} {shape}"
|
| 290 |
+
else:
|
| 291 |
+
return f"np scalar: {x.item()} {shape} {x.dtype}"
|
| 292 |
+
else:
|
| 293 |
+
if shape_only:
|
| 294 |
+
return f"np: {shape}"
|
| 295 |
+
else:
|
| 296 |
+
return f"np: {shape} {x.dtype}"
|
| 297 |
+
elif is_tensor(x):
|
| 298 |
+
shape = list(x.size())
|
| 299 |
+
if x.numel() == 1:
|
| 300 |
+
if shape_only:
|
| 301 |
+
return f"torch scalar: {x.item()} {shape}"
|
| 302 |
+
else:
|
| 303 |
+
return f"torch scalar: {x.item()} {shape} {x.dtype} {x.device}"
|
| 304 |
+
else:
|
| 305 |
+
if shape_only:
|
| 306 |
+
return f"torch: {shape}"
|
| 307 |
+
else:
|
| 308 |
+
return f"torch: {shape} {x.dtype} {x.device}"
|
| 309 |
+
elif is_sequence(x):
|
| 310 |
+
return f"{tname}[{len(x)}]"
|
| 311 |
+
elif isinstance(x, str):
|
| 312 |
+
return x
|
| 313 |
+
elif x is None:
|
| 314 |
+
return "None"
|
| 315 |
+
elif np.issubdtype(t, np.number) or np.issubdtype(t, np.bool_):
|
| 316 |
+
return f"{tname}: {x}"
|
| 317 |
+
else:
|
| 318 |
+
return f"{tname}"
|
| 319 |
+
|
| 320 |
+
|
| 321 |
+
def any_describe(x, msg="", *, shape_only=False):
|
| 322 |
+
# from omlet.utils import yaml_dumps
|
| 323 |
+
from pprint import pprint
|
| 324 |
+
|
| 325 |
+
if isinstance(x, str) and msg != "":
|
| 326 |
+
x, msg = msg, x
|
| 327 |
+
|
| 328 |
+
if msg:
|
| 329 |
+
msg += ": "
|
| 330 |
+
print(msg, end="")
|
| 331 |
+
pprint(any_describe_str(x, shape_only=shape_only))
|
| 332 |
+
|
| 333 |
+
|
| 334 |
+
@make_recursive_func
|
| 335 |
+
def any_slice(x, slice):
|
| 336 |
+
"""
|
| 337 |
+
Args:
|
| 338 |
+
slice: you can use np.s_[...] to return the slice object
|
| 339 |
+
"""
|
| 340 |
+
if is_array_tensor(x):
|
| 341 |
+
return x[slice]
|
| 342 |
+
else:
|
| 343 |
+
return x
|
| 344 |
+
|
| 345 |
+
|
| 346 |
+
def any_assign(x, assign_value, slice):
|
| 347 |
+
"""
|
| 348 |
+
Recursive version of x[slice] = assign_value
|
| 349 |
+
If structures of x and assign_value do not match, we will respect `assign_value`
|
| 350 |
+
E.g. x = {'a': ..., 'b': ...}, assign_value = {'a': ...}, then 'b' will not change
|
| 351 |
+
|
| 352 |
+
Use np.s_[...] to get advanced slicing
|
| 353 |
+
"""
|
| 354 |
+
|
| 355 |
+
def _any_assign_helper(path, v):
|
| 356 |
+
y = tree_value_at_path(x, path)
|
| 357 |
+
y[slice] = v
|
| 358 |
+
|
| 359 |
+
tree.map_structure_with_path(_any_assign_helper, assign_value)
|
| 360 |
+
|
| 361 |
+
|
| 362 |
+
@make_recursive_func
|
| 363 |
+
def any_transpose_first_two_axes(x):
|
| 364 |
+
"""
|
| 365 |
+
util to convert between (L, B, ...) and (B, L, ...)
|
| 366 |
+
"""
|
| 367 |
+
if is_numpy(x):
|
| 368 |
+
return np.swapaxes(x, 0, 1)
|
| 369 |
+
elif is_tensor(x):
|
| 370 |
+
return torch.swapaxes(x, 0, 1)
|
| 371 |
+
else:
|
| 372 |
+
raise ValueError(f"Input ({type(x)}) must be either a numpy array or a tensor.")
|
groot/vla/common/utils/misc/functional_utils.py
ADDED
|
@@ -0,0 +1,635 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""
|
| 2 |
+
Inspect, meta, etc.
|
| 3 |
+
"""
|
| 4 |
+
|
| 5 |
+
from __future__ import annotations
|
| 6 |
+
|
| 7 |
+
import functools
|
| 8 |
+
import inspect
|
| 9 |
+
import pprint
|
| 10 |
+
import sys
|
| 11 |
+
import types
|
| 12 |
+
from typing import Any, Dict, Literal
|
| 13 |
+
import warnings
|
| 14 |
+
|
| 15 |
+
from ..data_structure.tree_utils import is_mapping, is_sequence
|
| 16 |
+
|
| 17 |
+
|
| 18 |
+
def state_dict_class(keys: list[str]):
|
| 19 |
+
"""
|
| 20 |
+
Just like pytorch nn.Module
|
| 21 |
+
Add the following methods to the class:
|
| 22 |
+
state_dict() -> dict of attribute keys
|
| 23 |
+
load_state_dict(sdict) restore states
|
| 24 |
+
"""
|
| 25 |
+
|
| 26 |
+
def _wrap_class(cls):
|
| 27 |
+
assert inspect.isclass(cls)
|
| 28 |
+
|
| 29 |
+
def state_dict(self):
|
| 30 |
+
return {k: getattr(self, k) for k in keys}
|
| 31 |
+
|
| 32 |
+
def load_state_dict(self, states: Dict[str, Any]):
|
| 33 |
+
if not set(keys).issubset(set(states.keys())):
|
| 34 |
+
raise ValueError(f"states does not have all the required keys: {keys}")
|
| 35 |
+
for k in keys:
|
| 36 |
+
setattr(self, k, states[k])
|
| 37 |
+
|
| 38 |
+
@property
|
| 39 |
+
def state_keys(self):
|
| 40 |
+
return keys
|
| 41 |
+
|
| 42 |
+
cls.state_dict = state_dict
|
| 43 |
+
cls.load_state_dict = load_state_dict
|
| 44 |
+
cls.state_keys = state_keys
|
| 45 |
+
return cls
|
| 46 |
+
|
| 47 |
+
return _wrap_class
|
| 48 |
+
|
| 49 |
+
|
| 50 |
+
def implements_method(object, method: str):
|
| 51 |
+
"""
|
| 52 |
+
Returns:
|
| 53 |
+
True if object implements a method
|
| 54 |
+
"""
|
| 55 |
+
return hasattr(object, method) and callable(getattr(object, method))
|
| 56 |
+
|
| 57 |
+
|
| 58 |
+
def assert_implements_method(object, method: str | list[str]):
|
| 59 |
+
if isinstance(method, str):
|
| 60 |
+
method = [method]
|
| 61 |
+
for m in method:
|
| 62 |
+
assert implements_method(object, m), (
|
| 63 |
+
f"object {object.__class__} does not " f"implement method {m}()"
|
| 64 |
+
)
|
| 65 |
+
|
| 66 |
+
|
| 67 |
+
def meta_decorator(decor):
|
| 68 |
+
"""
|
| 69 |
+
a decorator, allowing the wrapped decorator to be used as:
|
| 70 |
+
@decorator(*args, **kwargs)
|
| 71 |
+
def callable()
|
| 72 |
+
-- or --
|
| 73 |
+
@decorator # without parenthesis, args and kwargs will use default
|
| 74 |
+
def callable()
|
| 75 |
+
|
| 76 |
+
Args:
|
| 77 |
+
decor: a decorator whose first argument is a callable (function or class
|
| 78 |
+
to be decorated), and the rest of the arguments can be omitted as default.
|
| 79 |
+
decor(f, ... the other arguments must have default values)
|
| 80 |
+
|
| 81 |
+
Warning:
|
| 82 |
+
decor can NOT be a function that receives a single, callable argument.
|
| 83 |
+
See stackoverflow: http://goo.gl/UEYbDB
|
| 84 |
+
"""
|
| 85 |
+
import functools
|
| 86 |
+
|
| 87 |
+
def single_callable(args, kwargs):
|
| 88 |
+
return len(args) == 1 and len(kwargs) == 0 and callable(args[0])
|
| 89 |
+
|
| 90 |
+
@functools.wraps(decor)
|
| 91 |
+
def new_decor(*args, **kwargs):
|
| 92 |
+
if single_callable(args, kwargs):
|
| 93 |
+
# this is the double-decorated f.
|
| 94 |
+
# It should not run on a single callable.
|
| 95 |
+
return decor(args[0])
|
| 96 |
+
else:
|
| 97 |
+
# decorator arguments
|
| 98 |
+
return lambda real_f: decor(real_f, *args, **kwargs)
|
| 99 |
+
|
| 100 |
+
return new_decor
|
| 101 |
+
|
| 102 |
+
|
| 103 |
+
@meta_decorator
|
| 104 |
+
def make_recursive_func(fn, *, with_path=False):
|
| 105 |
+
"""
|
| 106 |
+
Decorator that turns a function that works on a single array/tensor to working on
|
| 107 |
+
arbitrary nested structures.
|
| 108 |
+
"""
|
| 109 |
+
import functools
|
| 110 |
+
|
| 111 |
+
import tree
|
| 112 |
+
|
| 113 |
+
@functools.wraps(fn)
|
| 114 |
+
def _wrapper(tensor_struct, *args, **kwargs):
|
| 115 |
+
if with_path:
|
| 116 |
+
return tree.map_structure_with_path(
|
| 117 |
+
lambda paths, x: fn(paths, x, *args, **kwargs), tensor_struct
|
| 118 |
+
)
|
| 119 |
+
else:
|
| 120 |
+
return tree.map_structure(lambda x: fn(x, *args, **kwargs), tensor_struct)
|
| 121 |
+
|
| 122 |
+
return _wrapper
|
| 123 |
+
|
| 124 |
+
|
| 125 |
+
@meta_decorator
|
| 126 |
+
def deprecated(func, msg="", action="warning", type=""):
|
| 127 |
+
"""
|
| 128 |
+
Function/class decorator: designate deprecation.
|
| 129 |
+
|
| 130 |
+
Args:
|
| 131 |
+
msg: string message.
|
| 132 |
+
action: string mode
|
| 133 |
+
- 'warning': (default) prints `msg` to stderr
|
| 134 |
+
- 'noop': do nothing, just for source code annotation purposes
|
| 135 |
+
- 'raise': raise DeprecatedError(`msg`)
|
| 136 |
+
"""
|
| 137 |
+
action = action.lower()
|
| 138 |
+
type = type.lower()
|
| 139 |
+
ALL_ACTIONS = ["warn", "warning", "noop", "raise"]
|
| 140 |
+
if action not in ALL_ACTIONS:
|
| 141 |
+
raise ValueError(f"Unknown action {action}. Choose from {ALL_ACTIONS}.")
|
| 142 |
+
ALL_TYPES = {
|
| 143 |
+
"": DeprecationWarning,
|
| 144 |
+
"pending": PendingDeprecationWarning,
|
| 145 |
+
"future": FutureWarning,
|
| 146 |
+
}
|
| 147 |
+
if type not in ALL_TYPES:
|
| 148 |
+
raise ValueError(f"Unknown type {type}. Choose from {ALL_TYPES.keys()}.")
|
| 149 |
+
if not msg:
|
| 150 |
+
msg = "This is a deprecated feature."
|
| 151 |
+
|
| 152 |
+
WarningExceptionCls = ALL_TYPES[type]
|
| 153 |
+
|
| 154 |
+
# only does the deprecation when being called
|
| 155 |
+
@functools.wraps(func)
|
| 156 |
+
def _deprecated(*args, **kwargs):
|
| 157 |
+
if action in ["warning", "warn"]:
|
| 158 |
+
warnings.warn(msg, WarningExceptionCls)
|
| 159 |
+
elif action == "raise":
|
| 160 |
+
raise WarningExceptionCls(msg)
|
| 161 |
+
return func(*args, **kwargs)
|
| 162 |
+
|
| 163 |
+
return _deprecated
|
| 164 |
+
|
| 165 |
+
|
| 166 |
+
@meta_decorator
|
| 167 |
+
def call_once(func, on_second_call: Literal["noop", "raise", "warn"] = "noop"):
|
| 168 |
+
"""
|
| 169 |
+
Decorator to ensure that a function is only called once.
|
| 170 |
+
|
| 171 |
+
Args:
|
| 172 |
+
on_second_call (str): what happens when the function is called a second time.
|
| 173 |
+
"""
|
| 174 |
+
assert on_second_call in [
|
| 175 |
+
"noop",
|
| 176 |
+
"raise",
|
| 177 |
+
"warn",
|
| 178 |
+
], "mode must be one of 'noop', 'raise', 'warn'"
|
| 179 |
+
|
| 180 |
+
@functools.wraps(func)
|
| 181 |
+
def wrapper(*args, **kwargs):
|
| 182 |
+
if wrapper._called:
|
| 183 |
+
if on_second_call == "raise":
|
| 184 |
+
raise RuntimeError(f"{func.__name__} has already been called. Can only call once.")
|
| 185 |
+
elif on_second_call == "warn":
|
| 186 |
+
warnings.warn(f"{func.__name__} has already been called. Should only call once.")
|
| 187 |
+
else:
|
| 188 |
+
wrapper._called = True
|
| 189 |
+
return func(*args, **kwargs)
|
| 190 |
+
|
| 191 |
+
wrapper._called = False
|
| 192 |
+
return wrapper
|
| 193 |
+
|
| 194 |
+
|
| 195 |
+
class NoopObject:
|
| 196 |
+
"""
|
| 197 |
+
Object that does nothing when called any method
|
| 198 |
+
"""
|
| 199 |
+
|
| 200 |
+
def __init__(self, *args, **kwargs):
|
| 201 |
+
self.init_args = args
|
| 202 |
+
self.init_kwargs = kwargs
|
| 203 |
+
|
| 204 |
+
def __getattr__(self, name):
|
| 205 |
+
def _func(*args, **kwargs):
|
| 206 |
+
pass
|
| 207 |
+
|
| 208 |
+
return _func
|
| 209 |
+
|
| 210 |
+
|
| 211 |
+
class NoopContext:
|
| 212 |
+
"""
|
| 213 |
+
Placeholder context manager that does nothing.
|
| 214 |
+
We could have written simply as:
|
| 215 |
+
|
| 216 |
+
@contextmanager
|
| 217 |
+
def noop_context(*args, **kwargs):
|
| 218 |
+
yield
|
| 219 |
+
|
| 220 |
+
but the returned context manager cannot be called twice, i.e.
|
| 221 |
+
my_noop = NoopContext()
|
| 222 |
+
with my_noop:
|
| 223 |
+
do1()
|
| 224 |
+
with my_noop: # trigger generator error
|
| 225 |
+
do2()
|
| 226 |
+
"""
|
| 227 |
+
|
| 228 |
+
def __init__(self, *args, **kwargs):
|
| 229 |
+
self.args = args
|
| 230 |
+
self.kwargs = kwargs
|
| 231 |
+
|
| 232 |
+
def __enter__(self):
|
| 233 |
+
return self
|
| 234 |
+
|
| 235 |
+
def __exit__(self, exc_type, exc_val, exc_tb):
|
| 236 |
+
pass
|
| 237 |
+
|
| 238 |
+
|
| 239 |
+
def make_registry_metaclass(class_name):
|
| 240 |
+
"""
|
| 241 |
+
Usage:
|
| 242 |
+
|
| 243 |
+
TrainerRegistry = make_registry_metaclass('TrainerRegistry')
|
| 244 |
+
|
| 245 |
+
class BaseTrainer(metaclass=TrainerRegistry):
|
| 246 |
+
pass
|
| 247 |
+
|
| 248 |
+
class MyTrainer(BaseTrainer):
|
| 249 |
+
pass
|
| 250 |
+
|
| 251 |
+
TrainerRegistry['MyTrainer'] -> MyTrainer class # syntax enabled by metaclass
|
| 252 |
+
TrainerRegistry.get_class('MyTrainer') # same as above
|
| 253 |
+
TrainerRegistry.registry -> full dict of {name: trainer_class}
|
| 254 |
+
|
| 255 |
+
Templated definition:
|
| 256 |
+
class TrainerRegistry(type):
|
| 257 |
+
registry = {}
|
| 258 |
+
|
| 259 |
+
def __new__(cls, name, bases, attr):
|
| 260 |
+
new_cls = super().__new__(cls, name, bases, attr)
|
| 261 |
+
TrainerRegistry.registry[name] = new_cls
|
| 262 |
+
return new_cls
|
| 263 |
+
|
| 264 |
+
def get_class(cls, name):
|
| 265 |
+
if name not in cls.registry:
|
| 266 |
+
raise KeyError(
|
| 267 |
+
f"Trainer class {name} not found in registry: "
|
| 268 |
+
f"{pprint.pformat(cls.registry)}"
|
| 269 |
+
)
|
| 270 |
+
return cls.registry[name]"""
|
| 271 |
+
|
| 272 |
+
def new__(cls, name, bases, attr):
|
| 273 |
+
"""
|
| 274 |
+
Change the attr dict to dynamically add methods and attributes
|
| 275 |
+
"""
|
| 276 |
+
new_cls = type.__new__(cls, name, bases, attr)
|
| 277 |
+
cls.registry[name] = new_cls
|
| 278 |
+
return new_cls
|
| 279 |
+
|
| 280 |
+
def get_class(cls, name):
|
| 281 |
+
if name not in cls.registry:
|
| 282 |
+
existing_cls = list(cls.registry.keys())
|
| 283 |
+
raise KeyError(f"{class_name} class '{name}' not found in registry: {existing_cls}")
|
| 284 |
+
return cls.registry[name]
|
| 285 |
+
|
| 286 |
+
def instantiate(cls_, cls, **kwargs):
|
| 287 |
+
Cls = cls_.get_class(cls)
|
| 288 |
+
return Cls(**kwargs)
|
| 289 |
+
|
| 290 |
+
class _BracketOperator(type):
|
| 291 |
+
def __getitem__(cls, name):
|
| 292 |
+
return get_class(cls, name)
|
| 293 |
+
|
| 294 |
+
return types.new_class(
|
| 295 |
+
class_name,
|
| 296 |
+
bases=(type,),
|
| 297 |
+
kwds={"metaclass": _BracketOperator},
|
| 298 |
+
exec_body=lambda ns: ns.update(
|
| 299 |
+
{
|
| 300 |
+
"registry": {},
|
| 301 |
+
"__new__": new__,
|
| 302 |
+
"get_class": classmethod(get_class),
|
| 303 |
+
"instantiate": classmethod(instantiate),
|
| 304 |
+
}
|
| 305 |
+
),
|
| 306 |
+
)
|
| 307 |
+
|
| 308 |
+
|
| 309 |
+
class ClassRegistry:
|
| 310 |
+
"""
|
| 311 |
+
May be a preferred way over make_registry_metaclass if your code does not support
|
| 312 |
+
metaclass well, e.g. pickle or Ray
|
| 313 |
+
|
| 314 |
+
Use in conjunction with `__init_subclass__` hook in your base class
|
| 315 |
+
|
| 316 |
+
class BaseClass:
|
| 317 |
+
registry = ClassRegistry()
|
| 318 |
+
|
| 319 |
+
def __init_subclass__(cls, **kwargs):
|
| 320 |
+
cls.registry.add(cls)
|
| 321 |
+
super().__init_subclass__(**kwargs)
|
| 322 |
+
|
| 323 |
+
print(registry)
|
| 324 |
+
"""
|
| 325 |
+
|
| 326 |
+
def __init__(self, base_class_name: str = None):
|
| 327 |
+
self.registry = {}
|
| 328 |
+
self._base_class_name = base_class_name
|
| 329 |
+
|
| 330 |
+
def add(self, cls):
|
| 331 |
+
self.registry[cls.__name__] = cls
|
| 332 |
+
|
| 333 |
+
def get(self, name):
|
| 334 |
+
if name not in self.registry:
|
| 335 |
+
existing_cls = list(self.registry.keys())
|
| 336 |
+
base_name = self._base_class_name + " " if self._base_class_name else ""
|
| 337 |
+
raise KeyError(f"{base_name} subclass '{name}' not found in registry: {existing_cls}")
|
| 338 |
+
return self.registry[name]
|
| 339 |
+
|
| 340 |
+
def __str__(self):
|
| 341 |
+
return pprint.pformat(self.registry)
|
| 342 |
+
|
| 343 |
+
def __getitem__(self, name):
|
| 344 |
+
return self.get(name)
|
| 345 |
+
|
| 346 |
+
def instantiate(self, cls, **kwargs):
|
| 347 |
+
return self.get(cls)(**kwargs)
|
| 348 |
+
|
| 349 |
+
|
| 350 |
+
# ========================================================
|
| 351 |
+
# =================== Inspect utils ====================
|
| 352 |
+
# ========================================================
|
| 353 |
+
|
| 354 |
+
|
| 355 |
+
def func_parameters(func):
|
| 356 |
+
return inspect.signature(func).parameters
|
| 357 |
+
|
| 358 |
+
|
| 359 |
+
def func_has_arg(func, arg_name):
|
| 360 |
+
return arg_name in func_parameters(func)
|
| 361 |
+
|
| 362 |
+
|
| 363 |
+
def pack_varargs(args):
|
| 364 |
+
"""
|
| 365 |
+
Pack *args or a single list arg as list
|
| 366 |
+
|
| 367 |
+
def f(*args):
|
| 368 |
+
arg_list = pack_varargs(args)
|
| 369 |
+
# arg_list is now packed as a list
|
| 370 |
+
"""
|
| 371 |
+
assert isinstance(args, tuple), "please input the tuple `args` as in *args"
|
| 372 |
+
if len(args) == 1 and is_sequence(args[0]):
|
| 373 |
+
return args[0]
|
| 374 |
+
else:
|
| 375 |
+
return args
|
| 376 |
+
|
| 377 |
+
|
| 378 |
+
def enable_list_arg(func):
|
| 379 |
+
"""
|
| 380 |
+
Function decorator.
|
| 381 |
+
If a function only accepts varargs (*args),
|
| 382 |
+
make it support a single list arg as well
|
| 383 |
+
"""
|
| 384 |
+
|
| 385 |
+
@functools.wraps(func)
|
| 386 |
+
def wrapper(*args, **kwargs):
|
| 387 |
+
args = pack_varargs(args)
|
| 388 |
+
return func(*args, **kwargs)
|
| 389 |
+
|
| 390 |
+
return wrapper
|
| 391 |
+
|
| 392 |
+
|
| 393 |
+
def enable_varargs(func):
|
| 394 |
+
"""
|
| 395 |
+
Function decorator.
|
| 396 |
+
If a function only accepts a list arg, make it support varargs as well
|
| 397 |
+
"""
|
| 398 |
+
|
| 399 |
+
@functools.wraps(func)
|
| 400 |
+
def wrapper(*args, **kwargs):
|
| 401 |
+
args = pack_varargs(args)
|
| 402 |
+
return func(args, **kwargs)
|
| 403 |
+
|
| 404 |
+
return wrapper
|
| 405 |
+
|
| 406 |
+
|
| 407 |
+
def pack_kwargs(args, kwargs):
|
| 408 |
+
"""
|
| 409 |
+
Pack **kwargs or a single dict arg as dict
|
| 410 |
+
|
| 411 |
+
def f(*args, **kwargs):
|
| 412 |
+
kwdict = pack_kwargs(args, kwargs)
|
| 413 |
+
# kwdict is now packed as a dict
|
| 414 |
+
"""
|
| 415 |
+
if len(args) == 1 and is_mapping(args[0]):
|
| 416 |
+
assert not kwargs, "cannot have both **kwargs and a dict arg"
|
| 417 |
+
return args[0] # single-dict
|
| 418 |
+
else:
|
| 419 |
+
assert not args, "cannot have positional args if **kwargs exist"
|
| 420 |
+
return kwargs
|
| 421 |
+
|
| 422 |
+
|
| 423 |
+
def merge_kwargs(args, kwargs) -> Dict:
|
| 424 |
+
"""
|
| 425 |
+
Merge all dicts in `args` and keywords in kwargs.
|
| 426 |
+
|
| 427 |
+
E.g. merge_kwargs({"a.b": 1, "a.c": 2}, foo=6, bar=8)
|
| 428 |
+
-> {"a.b": 1, "a.c": 2, "foo": 6, "bar": 8}
|
| 429 |
+
"""
|
| 430 |
+
kw_all = {}
|
| 431 |
+
for arg in args:
|
| 432 |
+
assert is_mapping(arg), f"{arg} is not a dict."
|
| 433 |
+
kw_all.update(arg)
|
| 434 |
+
kw_all.update(kwargs)
|
| 435 |
+
return kw_all
|
| 436 |
+
|
| 437 |
+
|
| 438 |
+
def enable_dict_arg(func):
|
| 439 |
+
"""
|
| 440 |
+
Function decorator.
|
| 441 |
+
If a function only accepts varargs (*args),
|
| 442 |
+
make it support a single list arg as well
|
| 443 |
+
"""
|
| 444 |
+
|
| 445 |
+
@functools.wraps(func)
|
| 446 |
+
def wrapper(*args, **kwargs):
|
| 447 |
+
kwargs = pack_kwargs(args, kwargs)
|
| 448 |
+
return func(**kwargs)
|
| 449 |
+
|
| 450 |
+
return wrapper
|
| 451 |
+
|
| 452 |
+
|
| 453 |
+
def enable_kwargs(func):
|
| 454 |
+
"""
|
| 455 |
+
Function decorator.
|
| 456 |
+
If a function only accepts a dict arg, make it support kwargs as well
|
| 457 |
+
"""
|
| 458 |
+
|
| 459 |
+
@functools.wraps(func)
|
| 460 |
+
def wrapper(*args, **kwargs):
|
| 461 |
+
kwargs = pack_kwargs(args, kwargs)
|
| 462 |
+
return func(kwargs)
|
| 463 |
+
|
| 464 |
+
return wrapper
|
| 465 |
+
|
| 466 |
+
|
| 467 |
+
def has_keys(D, keys: list):
|
| 468 |
+
assert is_mapping(D)
|
| 469 |
+
return all(key in D for key in keys)
|
| 470 |
+
|
| 471 |
+
|
| 472 |
+
def assert_has_keys(D, keys: list):
|
| 473 |
+
assert is_mapping(D), "Input is not a dict"
|
| 474 |
+
for key in keys:
|
| 475 |
+
if key not in D:
|
| 476 |
+
raise KeyError(f'Required key "{key}" is missing in dict {D}')
|
| 477 |
+
return True
|
| 478 |
+
|
| 479 |
+
|
| 480 |
+
def method_decorator(decorator):
|
| 481 |
+
"""
|
| 482 |
+
Decorator of decorator: transform a decorator that only works on normal
|
| 483 |
+
functions to a decorator that works on class methods
|
| 484 |
+
From Django form: https://goo.gl/XLjxKK
|
| 485 |
+
"""
|
| 486 |
+
|
| 487 |
+
@functools.wraps(decorator)
|
| 488 |
+
def wrapped_decorator(method):
|
| 489 |
+
@functools.wraps(method)
|
| 490 |
+
def wrapper(self, *args, **kwargs):
|
| 491 |
+
def bound_func(*args2, **kwargs2):
|
| 492 |
+
return method(self, *args2, **kwargs2)
|
| 493 |
+
|
| 494 |
+
return decorator(bound_func)(*args, **kwargs)
|
| 495 |
+
|
| 496 |
+
return wrapper
|
| 497 |
+
|
| 498 |
+
return wrapped_decorator
|
| 499 |
+
|
| 500 |
+
|
| 501 |
+
def accepts_varargs(func):
|
| 502 |
+
"""
|
| 503 |
+
If a function accepts *args
|
| 504 |
+
"""
|
| 505 |
+
params = inspect.signature(func).parameters
|
| 506 |
+
return any(param.kind == inspect.Parameter.VAR_POSITIONAL for param in params.values())
|
| 507 |
+
|
| 508 |
+
|
| 509 |
+
def accepts_kwargs(func):
|
| 510 |
+
"""
|
| 511 |
+
If a function accepts **kwargs
|
| 512 |
+
"""
|
| 513 |
+
params = inspect.signature(func).parameters
|
| 514 |
+
return any(param.kind == inspect.Parameter.VAR_KEYWORD for param in params.values())
|
| 515 |
+
|
| 516 |
+
|
| 517 |
+
def is_signature_compatible(func, *args, **kwargs):
|
| 518 |
+
sig = inspect.signature(func)
|
| 519 |
+
try:
|
| 520 |
+
sig.bind(*args, **kwargs)
|
| 521 |
+
return True
|
| 522 |
+
except TypeError:
|
| 523 |
+
return False
|
| 524 |
+
|
| 525 |
+
|
| 526 |
+
def make_list(x):
|
| 527 |
+
"""
|
| 528 |
+
Turns a singleton object to a list. If already a list, no change.
|
| 529 |
+
"""
|
| 530 |
+
if is_sequence(x):
|
| 531 |
+
return x
|
| 532 |
+
else:
|
| 533 |
+
return [x]
|
| 534 |
+
|
| 535 |
+
|
| 536 |
+
def make_tuple(elem, repeats):
|
| 537 |
+
"""
|
| 538 |
+
E.g. expand a singleton x into (x, x, x)
|
| 539 |
+
useful for things like image_size or kernal, which can be a single int/float
|
| 540 |
+
or a tuple of fixed size
|
| 541 |
+
"""
|
| 542 |
+
if is_sequence(elem):
|
| 543 |
+
assert len(elem) == repeats, f"length of input must be {repeats}: {elem}"
|
| 544 |
+
return elem
|
| 545 |
+
else:
|
| 546 |
+
return (elem,) * repeats
|
| 547 |
+
|
| 548 |
+
|
| 549 |
+
def accumulate(iterable, fn=lambda x, y: x + y):
|
| 550 |
+
"""
|
| 551 |
+
Return running totals
|
| 552 |
+
# _accumulate([1,2,3,4,5]) --> 1 3 6 10 15
|
| 553 |
+
# _accumulate([1,2,3,4,5], operator.mul) --> 1 2 6 24 120
|
| 554 |
+
"""
|
| 555 |
+
it = iter(iterable)
|
| 556 |
+
try:
|
| 557 |
+
total = next(it)
|
| 558 |
+
except StopIteration:
|
| 559 |
+
return
|
| 560 |
+
yield total
|
| 561 |
+
for element in it:
|
| 562 |
+
total = fn(total, element)
|
| 563 |
+
yield total
|
| 564 |
+
|
| 565 |
+
|
| 566 |
+
class DecoratorContextManager:
|
| 567 |
+
"""
|
| 568 |
+
Allow a context manager to be used as a decorator
|
| 569 |
+
From torch.auto_grad.grad_mode
|
| 570 |
+
"""
|
| 571 |
+
|
| 572 |
+
def __call__(self, func):
|
| 573 |
+
if inspect.isgeneratorfunction(func):
|
| 574 |
+
return self._wrap_generator(func)
|
| 575 |
+
|
| 576 |
+
@functools.wraps(func)
|
| 577 |
+
def decorate_context(*args, **kwargs):
|
| 578 |
+
with self.__class__():
|
| 579 |
+
return func(*args, **kwargs)
|
| 580 |
+
|
| 581 |
+
return decorate_context
|
| 582 |
+
|
| 583 |
+
def _wrap_generator(self, func):
|
| 584 |
+
"""Wrap each generator invocation with the context manager"""
|
| 585 |
+
|
| 586 |
+
@functools.wraps(func)
|
| 587 |
+
def generator_context(*args, **kwargs):
|
| 588 |
+
gen = func(*args, **kwargs)
|
| 589 |
+
|
| 590 |
+
# Generators are suspended and unsuspended at `yield`, hence we
|
| 591 |
+
# make sure the grad mode is properly set every time the execution
|
| 592 |
+
# flow returns into the wrapped generator and restored when it
|
| 593 |
+
# returns through our `yield` to our caller (see PR #49017).
|
| 594 |
+
cls = type(self)
|
| 595 |
+
try:
|
| 596 |
+
# Issuing `None` to a generator fires it up
|
| 597 |
+
with cls():
|
| 598 |
+
response = gen.send(None)
|
| 599 |
+
|
| 600 |
+
while True:
|
| 601 |
+
try:
|
| 602 |
+
# Forward the response to our caller and get its next request
|
| 603 |
+
request = yield response
|
| 604 |
+
|
| 605 |
+
except GeneratorExit:
|
| 606 |
+
# Inform the still active generator about its imminent closure
|
| 607 |
+
with cls():
|
| 608 |
+
gen.close()
|
| 609 |
+
raise
|
| 610 |
+
|
| 611 |
+
except BaseException:
|
| 612 |
+
# Propagate the exception thrown at us by the caller
|
| 613 |
+
with cls():
|
| 614 |
+
response = gen.throw(*sys.exc_info())
|
| 615 |
+
|
| 616 |
+
else:
|
| 617 |
+
# Pass the last request to the generator and get its response
|
| 618 |
+
with cls():
|
| 619 |
+
response = gen.send(request)
|
| 620 |
+
|
| 621 |
+
# We let the exceptions raised above by the generator's `.throw` or
|
| 622 |
+
# `.send` methods bubble up to our caller, except for StopIteration
|
| 623 |
+
except StopIteration as e:
|
| 624 |
+
# The generator informed us that it is done: take whatever its
|
| 625 |
+
# returned value (if any) was and indicate that we're done too
|
| 626 |
+
# by returning it (see docs for python's return-statement).
|
| 627 |
+
return e.value
|
| 628 |
+
|
| 629 |
+
return generator_context
|
| 630 |
+
|
| 631 |
+
def __enter__(self) -> None:
|
| 632 |
+
raise NotImplementedError
|
| 633 |
+
|
| 634 |
+
def __exit__(self, exc_type: Any, exc_value: Any, traceback: Any) -> None:
|
| 635 |
+
raise NotImplementedError
|
groot/vla/common/utils/misc/image_utils.py
ADDED
|
@@ -0,0 +1,225 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""
|
| 2 |
+
Visualizations
|
| 3 |
+
"""
|
| 4 |
+
|
| 5 |
+
from __future__ import annotations
|
| 6 |
+
|
| 7 |
+
import os
|
| 8 |
+
import time
|
| 9 |
+
from typing import Literal
|
| 10 |
+
import warnings
|
| 11 |
+
|
| 12 |
+
import cv2
|
| 13 |
+
import imageio
|
| 14 |
+
import matplotlib.pyplot as plt
|
| 15 |
+
import numpy as np
|
| 16 |
+
import torch
|
| 17 |
+
|
| 18 |
+
from .array_tensor_utils import any_describe
|
| 19 |
+
from .misc_utils import global_once
|
| 20 |
+
from .torch_utils import torch_normalize
|
| 21 |
+
|
| 22 |
+
|
| 23 |
+
def to_image(img, channel_order="auto"):
|
| 24 |
+
"""
|
| 25 |
+
Returns:
|
| 26 |
+
numpy image of shape [H, W, C]
|
| 27 |
+
in "auto" mode, we assume C == 3
|
| 28 |
+
"""
|
| 29 |
+
assert channel_order in ["hwc", "chw", "auto"]
|
| 30 |
+
if torch.is_tensor(img):
|
| 31 |
+
img = img.cpu().numpy()
|
| 32 |
+
assert isinstance(img, np.ndarray)
|
| 33 |
+
if img.ndim == 4:
|
| 34 |
+
assert img.shape[0] == 1
|
| 35 |
+
img = img[0]
|
| 36 |
+
assert img.ndim == 3
|
| 37 |
+
if channel_order == "auto":
|
| 38 |
+
# use C==3 to detect order
|
| 39 |
+
if img.shape[0] == 3:
|
| 40 |
+
channel_order = "chw"
|
| 41 |
+
else:
|
| 42 |
+
assert img.shape[-1] == 3, "image should either have [3,H,W] or [H,W,3]"
|
| 43 |
+
channel_order = "hwc"
|
| 44 |
+
img = img.astype(np.uint8)
|
| 45 |
+
if channel_order == "chw":
|
| 46 |
+
return np.transpose(img, (1, 2, 0))
|
| 47 |
+
else:
|
| 48 |
+
return img
|
| 49 |
+
|
| 50 |
+
|
| 51 |
+
def imshow(img):
|
| 52 |
+
plt.imshow(to_image(img))
|
| 53 |
+
|
| 54 |
+
|
| 55 |
+
def imsave(img, path):
|
| 56 |
+
imageio.imsave(os.path.expanduser(path), to_image(img))
|
| 57 |
+
|
| 58 |
+
|
| 59 |
+
def imread(path, channel_order="chw", format="torch"):
|
| 60 |
+
assert channel_order in ["hwc", "chw"]
|
| 61 |
+
assert format in ["numpy", "torch"]
|
| 62 |
+
img = imageio.imread(path)
|
| 63 |
+
if channel_order == "chw":
|
| 64 |
+
img = np.transpose(img, (2, 0, 1)) # hwc -> chw
|
| 65 |
+
if format == "torch":
|
| 66 |
+
return torch.from_numpy(img)
|
| 67 |
+
else:
|
| 68 |
+
return img
|
| 69 |
+
|
| 70 |
+
|
| 71 |
+
class Cv2Display:
|
| 72 |
+
def __init__(
|
| 73 |
+
self,
|
| 74 |
+
window_name="display",
|
| 75 |
+
image_size=None,
|
| 76 |
+
channel_order="auto",
|
| 77 |
+
bgr2rgb=True,
|
| 78 |
+
step_sleep=0,
|
| 79 |
+
enabled=True,
|
| 80 |
+
):
|
| 81 |
+
"""
|
| 82 |
+
Use cv2.imshow() to pop a window, requires virtual desktop GUI
|
| 83 |
+
|
| 84 |
+
Args:
|
| 85 |
+
channel_order: auto, hwc, or chw
|
| 86 |
+
image_size: None to use the original image size, otherwise resize
|
| 87 |
+
step_sleep: sleep for a few seconds
|
| 88 |
+
"""
|
| 89 |
+
self._window_name = window_name
|
| 90 |
+
if isinstance(image_size, int):
|
| 91 |
+
image_size = (image_size, image_size)
|
| 92 |
+
else:
|
| 93 |
+
assert image_size is None or len(image_size) == 2
|
| 94 |
+
self._image_size = image_size
|
| 95 |
+
assert channel_order in ["auto", "chw", "hwc"]
|
| 96 |
+
self._channel_order = channel_order
|
| 97 |
+
self._bgr2rgb = bgr2rgb
|
| 98 |
+
self._step_sleep = step_sleep
|
| 99 |
+
self._enabled = enabled
|
| 100 |
+
|
| 101 |
+
def _resize(self, img):
|
| 102 |
+
if self._image_size is None:
|
| 103 |
+
return img
|
| 104 |
+
H, W = img.shape[:2]
|
| 105 |
+
Ht, Wt = self._image_size # target
|
| 106 |
+
return cv2.resize(
|
| 107 |
+
img,
|
| 108 |
+
self._image_size,
|
| 109 |
+
interpolation=cv2.INTER_AREA if Ht < H else cv2.INTER_LINEAR,
|
| 110 |
+
)
|
| 111 |
+
|
| 112 |
+
def _reorder(self, img):
|
| 113 |
+
if self._channel_order == "chw":
|
| 114 |
+
return np.transpose(img, (1, 2, 0))
|
| 115 |
+
elif self._channel_order == "hwc":
|
| 116 |
+
return img
|
| 117 |
+
else:
|
| 118 |
+
if img.shape[0] in [1, 3]: # chw
|
| 119 |
+
return np.transpose(img, (1, 2, 0))
|
| 120 |
+
else:
|
| 121 |
+
return img
|
| 122 |
+
|
| 123 |
+
def __call__(self, img):
|
| 124 |
+
if not self._enabled:
|
| 125 |
+
return
|
| 126 |
+
import torch
|
| 127 |
+
|
| 128 |
+
# prevent segfault in IsaacGym
|
| 129 |
+
display_var = os.environ.get("DISPLAY", None)
|
| 130 |
+
if not display_var:
|
| 131 |
+
os.environ["DISPLAY"] = ":0.0"
|
| 132 |
+
|
| 133 |
+
if torch.is_tensor(img):
|
| 134 |
+
img = img.detach().cpu().numpy()
|
| 135 |
+
|
| 136 |
+
img = self._resize(self._reorder(img))
|
| 137 |
+
if self._bgr2rgb:
|
| 138 |
+
img = cv2.cvtColor(img, cv2.COLOR_BGR2RGB)
|
| 139 |
+
time.sleep(self._step_sleep)
|
| 140 |
+
cv2.imshow(self._window_name, img)
|
| 141 |
+
cv2.waitKey(1)
|
| 142 |
+
|
| 143 |
+
if display_var is not None:
|
| 144 |
+
os.environ["DISPLAY"] = display_var
|
| 145 |
+
|
| 146 |
+
def close(self):
|
| 147 |
+
if not self._enabled:
|
| 148 |
+
return
|
| 149 |
+
cv2.destroyWindow(self._window_name)
|
| 150 |
+
|
| 151 |
+
|
| 152 |
+
# ---------------- Image tensor handling -----------------
|
| 153 |
+
def sanity_check_image_tensor(
|
| 154 |
+
img: torch.Tensor, on_error: Literal["raise", "warn", "ignore"] = "raise"
|
| 155 |
+
):
|
| 156 |
+
"""
|
| 157 |
+
Check if the input image tensor is all integers, which is wrong for any NN input.
|
| 158 |
+
This is a common case if the user forgets to normalize the image first
|
| 159 |
+
"""
|
| 160 |
+
assert on_error in [
|
| 161 |
+
"raise",
|
| 162 |
+
"warn",
|
| 163 |
+
"ignore",
|
| 164 |
+
], 'on_error must be "raise", "warn", or "ignore"'
|
| 165 |
+
if not img.dtype.is_floating_point:
|
| 166 |
+
msg = f"Image tensor is not floating point format, but {img.dtype}!"
|
| 167 |
+
if on_error == "raise":
|
| 168 |
+
raise ValueError(msg)
|
| 169 |
+
elif on_error == "warn":
|
| 170 |
+
warnings.warn(msg)
|
| 171 |
+
else:
|
| 172 |
+
return False
|
| 173 |
+
# check if all values in the image are close to an integer
|
| 174 |
+
if (img - torch.round(img)).abs().max() < 1e-5:
|
| 175 |
+
msg = (
|
| 176 |
+
"Input image is all close to integers, "
|
| 177 |
+
"are you sure you have normalized it before passing it to a NN?"
|
| 178 |
+
)
|
| 179 |
+
if on_error == "raise":
|
| 180 |
+
raise ValueError(msg)
|
| 181 |
+
elif on_error == "warn":
|
| 182 |
+
warnings.warn(msg)
|
| 183 |
+
else:
|
| 184 |
+
return False
|
| 185 |
+
return True
|
| 186 |
+
|
| 187 |
+
|
| 188 |
+
@torch.no_grad()
|
| 189 |
+
def basic_image_tensor_preprocess(
|
| 190 |
+
img,
|
| 191 |
+
mean: tuple[float, float, float] = (0.5, 0.5, 0.5),
|
| 192 |
+
std: tuple[float, float, float] = (0.5, 0.5, 0.5),
|
| 193 |
+
shape: tuple[int, int] | None = None,
|
| 194 |
+
):
|
| 195 |
+
"""
|
| 196 |
+
Check for resize, and divide by 255
|
| 197 |
+
"""
|
| 198 |
+
import kornia
|
| 199 |
+
|
| 200 |
+
assert torch.is_tensor(img)
|
| 201 |
+
assert img.dim() >= 4, any_describe(img)
|
| 202 |
+
original_shape = list(img.size())
|
| 203 |
+
img = img.float()
|
| 204 |
+
img = img.flatten(0, img.dim() - 4)
|
| 205 |
+
assert img.dim() == 4
|
| 206 |
+
|
| 207 |
+
input_size = img.size()[-2:]
|
| 208 |
+
if global_once("groot.vla.common.utils.image_utils.basic_image_preprocess:input_size"):
|
| 209 |
+
assert img.max() > 2, "img should be between [0, 255] before normalize"
|
| 210 |
+
|
| 211 |
+
if shape and input_size != shape:
|
| 212 |
+
if global_once("groot.vla.common.utils.image_utils.basic_image_preprocess:transform"):
|
| 213 |
+
warnings.warn(
|
| 214 |
+
f'{"Down" if shape < input_size else "Up"}sampling image'
|
| 215 |
+
f" from original resolution {input_size}x{input_size}"
|
| 216 |
+
f" to {shape}x{shape}"
|
| 217 |
+
)
|
| 218 |
+
img = kornia.geometry.transform.resize(img, shape).clamp(0.0, 255.0)
|
| 219 |
+
|
| 220 |
+
B, C, H, W = img.size()
|
| 221 |
+
assert C % 3 == 0, "channel must divide 3"
|
| 222 |
+
img = img.view(B * C // 3, 3, H, W)
|
| 223 |
+
img = torch_normalize(img / 255.0, mean=mean, std=std)
|
| 224 |
+
original_shape[-2:] = H, W
|
| 225 |
+
return img.view(original_shape)
|
groot/vla/common/utils/misc/misc_utils.py
ADDED
|
@@ -0,0 +1,261 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
import codecs
|
| 2 |
+
from collections import Counter
|
| 3 |
+
import fnmatch
|
| 4 |
+
import hashlib
|
| 5 |
+
import os
|
| 6 |
+
import pickle
|
| 7 |
+
from typing import Any, Callable, Dict, List, Optional, Union
|
| 8 |
+
|
| 9 |
+
from typing_extensions import Literal
|
| 10 |
+
|
| 11 |
+
|
| 12 |
+
def set_os_envs(envs: Optional[Dict[str, Any]] = None):
|
| 13 |
+
"""
|
| 14 |
+
Special value __delete__ or None indicates that the ENV_VAR should be removed
|
| 15 |
+
"""
|
| 16 |
+
if envs is None:
|
| 17 |
+
envs = {}
|
| 18 |
+
DEL = {None, "__delete__"}
|
| 19 |
+
for k, v in envs.items():
|
| 20 |
+
if v in DEL:
|
| 21 |
+
os.environ.pop(k, None)
|
| 22 |
+
os.environ.update({k: str(v) for k, v in envs.items() if v not in DEL})
|
| 23 |
+
|
| 24 |
+
|
| 25 |
+
def argmax(L):
|
| 26 |
+
return max(zip(L, range(len(L))))[1]
|
| 27 |
+
|
| 28 |
+
|
| 29 |
+
def _match_patterns_helper(element, patterns):
|
| 30 |
+
for p in patterns:
|
| 31 |
+
if callable(p) and p(element):
|
| 32 |
+
return True
|
| 33 |
+
if fnmatch.fnmatch(element, p):
|
| 34 |
+
return True
|
| 35 |
+
return False
|
| 36 |
+
|
| 37 |
+
|
| 38 |
+
def match_patterns(
|
| 39 |
+
item: str,
|
| 40 |
+
include: Union[str, List[str], Callable, List[Callable], None] = None,
|
| 41 |
+
exclude: Union[str, List[str], Callable, List[Callable], None] = None,
|
| 42 |
+
*,
|
| 43 |
+
precedence: Literal["include", "exclude"] = "exclude",
|
| 44 |
+
):
|
| 45 |
+
"""
|
| 46 |
+
Args:
|
| 47 |
+
include: None to disable `include` filter and delegate to exclude
|
| 48 |
+
precedence: "include" or "exclude"
|
| 49 |
+
"""
|
| 50 |
+
assert precedence in ["include", "exclude"]
|
| 51 |
+
if exclude is None:
|
| 52 |
+
exclude = []
|
| 53 |
+
if isinstance(exclude, (str, Callable)):
|
| 54 |
+
exclude = [exclude]
|
| 55 |
+
if isinstance(include, (str, Callable)):
|
| 56 |
+
include = [include]
|
| 57 |
+
if include is None:
|
| 58 |
+
# exclude is the sole veto vote
|
| 59 |
+
return not _match_patterns_helper(item, exclude)
|
| 60 |
+
|
| 61 |
+
if precedence == "include":
|
| 62 |
+
return _match_patterns_helper(item, include)
|
| 63 |
+
else:
|
| 64 |
+
if _match_patterns_helper(item, exclude):
|
| 65 |
+
return False
|
| 66 |
+
else:
|
| 67 |
+
return _match_patterns_helper(item, include)
|
| 68 |
+
|
| 69 |
+
|
| 70 |
+
def filter_patterns(
|
| 71 |
+
items: List[str],
|
| 72 |
+
include: Union[str, List[str], Callable, List[Callable], None] = None,
|
| 73 |
+
exclude: Union[str, List[str], Callable, List[Callable], None] = None,
|
| 74 |
+
*,
|
| 75 |
+
precedence: Literal["include", "exclude"] = "exclude",
|
| 76 |
+
ordering: Literal["original", "include"] = "original",
|
| 77 |
+
):
|
| 78 |
+
"""
|
| 79 |
+
Args:
|
| 80 |
+
ordering: affects the order of items in the returned list. Does not affect the
|
| 81 |
+
content of the returned list.
|
| 82 |
+
- "original": keep the ordering of items in the input list
|
| 83 |
+
- "include": order items by the order of include patterns
|
| 84 |
+
"""
|
| 85 |
+
assert ordering in ["original", "include"]
|
| 86 |
+
if include is None or isinstance(include, str) or ordering == "original":
|
| 87 |
+
return [
|
| 88 |
+
x
|
| 89 |
+
for x in items
|
| 90 |
+
if match_patterns(x, include=include, exclude=exclude, precedence=precedence)
|
| 91 |
+
]
|
| 92 |
+
else:
|
| 93 |
+
items = items.copy()
|
| 94 |
+
ret = []
|
| 95 |
+
for inc in include:
|
| 96 |
+
for i, item in enumerate(items):
|
| 97 |
+
if item is None:
|
| 98 |
+
continue
|
| 99 |
+
if match_patterns(item, include=inc, exclude=exclude, precedence=precedence):
|
| 100 |
+
ret.append(item)
|
| 101 |
+
items[i] = None
|
| 102 |
+
return ret
|
| 103 |
+
|
| 104 |
+
|
| 105 |
+
def getitem_nested(cfg, key: str):
|
| 106 |
+
"""
|
| 107 |
+
Recursively get key, if key has '.' in it
|
| 108 |
+
"""
|
| 109 |
+
keys = key.split(".")
|
| 110 |
+
for k in keys:
|
| 111 |
+
assert k in cfg, f'{k} in key "{key}" does not exist in config'
|
| 112 |
+
cfg = cfg[k]
|
| 113 |
+
return cfg
|
| 114 |
+
|
| 115 |
+
|
| 116 |
+
def setitem_nested(cfg, key: str, value):
|
| 117 |
+
"""
|
| 118 |
+
Recursively get key, if key has '.' in it
|
| 119 |
+
"""
|
| 120 |
+
keys = key.split(".")
|
| 121 |
+
for k in keys[:-1]:
|
| 122 |
+
assert k in cfg, f'{k} in key "{key}" does not exist in config'
|
| 123 |
+
cfg = cfg[k]
|
| 124 |
+
cfg[keys[-1]] = value
|
| 125 |
+
|
| 126 |
+
|
| 127 |
+
def getattr_nested(obj, key: str):
|
| 128 |
+
"""
|
| 129 |
+
Recursively get attribute
|
| 130 |
+
"""
|
| 131 |
+
keys = key.split(".")
|
| 132 |
+
for k in keys:
|
| 133 |
+
assert hasattr(obj, k), f'{k} in attribute "{key}" does not exist'
|
| 134 |
+
obj = getattr(obj, k)
|
| 135 |
+
return obj
|
| 136 |
+
|
| 137 |
+
|
| 138 |
+
def setattr_nested(obj, key: str, value):
|
| 139 |
+
"""
|
| 140 |
+
Recursively set attribute
|
| 141 |
+
"""
|
| 142 |
+
keys = key.split(".")
|
| 143 |
+
for k in keys[:-1]:
|
| 144 |
+
assert hasattr(obj, k), f'{k} in attribute "{key}" does not exist'
|
| 145 |
+
obj = getattr(obj, k)
|
| 146 |
+
setattr(obj, keys[-1], value)
|
| 147 |
+
|
| 148 |
+
|
| 149 |
+
class PeriodicEvent:
|
| 150 |
+
"""
|
| 151 |
+
triggers every period
|
| 152 |
+
"""
|
| 153 |
+
|
| 154 |
+
def __init__(self, period: int, initial_value=0):
|
| 155 |
+
self._period = period
|
| 156 |
+
assert self._period >= 1
|
| 157 |
+
self._last_threshold = initial_value
|
| 158 |
+
self._last_value = initial_value
|
| 159 |
+
self._trigger_counts = 0
|
| 160 |
+
|
| 161 |
+
def __call__(self, new_value=None, increment=None):
|
| 162 |
+
assert bool(new_value is None) != bool(increment is None), (
|
| 163 |
+
"you must specify one and only one of new_value or increment, " "but not both"
|
| 164 |
+
)
|
| 165 |
+
d = self._period
|
| 166 |
+
if new_value is None:
|
| 167 |
+
new_value = self._last_value + increment
|
| 168 |
+
assert new_value >= self._last_value, (
|
| 169 |
+
f"value must be monotonically increasing. "
|
| 170 |
+
f"Current value {new_value} < last value {self._last_value}"
|
| 171 |
+
)
|
| 172 |
+
self._last_value = new_value
|
| 173 |
+
if new_value - self._last_threshold >= d:
|
| 174 |
+
self._last_threshold += (new_value - self._last_threshold) // d * d
|
| 175 |
+
self._trigger_counts += 1
|
| 176 |
+
return True
|
| 177 |
+
else:
|
| 178 |
+
return False
|
| 179 |
+
|
| 180 |
+
@property
|
| 181 |
+
def trigger_counts(self):
|
| 182 |
+
return self._trigger_counts
|
| 183 |
+
|
| 184 |
+
@property
|
| 185 |
+
def current_value(self):
|
| 186 |
+
return self._last_value
|
| 187 |
+
|
| 188 |
+
|
| 189 |
+
class Once:
|
| 190 |
+
def __init__(self):
|
| 191 |
+
self._triggered = False
|
| 192 |
+
|
| 193 |
+
def __call__(self):
|
| 194 |
+
if not self._triggered:
|
| 195 |
+
self._triggered = True
|
| 196 |
+
return True
|
| 197 |
+
else:
|
| 198 |
+
return False
|
| 199 |
+
|
| 200 |
+
def __bool__(self):
|
| 201 |
+
raise RuntimeError("`Once` objects should be used by calling ()")
|
| 202 |
+
|
| 203 |
+
|
| 204 |
+
_GLOBAL_ONCE_SET = set()
|
| 205 |
+
_GLOBAL_NTIMES_COUNTER = Counter()
|
| 206 |
+
|
| 207 |
+
|
| 208 |
+
def global_once(name):
|
| 209 |
+
"""
|
| 210 |
+
Try this to automate the name:
|
| 211 |
+
https://gist.github.com/techtonik/2151727#gistcomment-2333747
|
| 212 |
+
"""
|
| 213 |
+
if name in _GLOBAL_ONCE_SET:
|
| 214 |
+
return False
|
| 215 |
+
else:
|
| 216 |
+
_GLOBAL_ONCE_SET.add(name)
|
| 217 |
+
return True
|
| 218 |
+
|
| 219 |
+
|
| 220 |
+
def global_n_times(name, n: int):
|
| 221 |
+
"""
|
| 222 |
+
Triggers N times
|
| 223 |
+
"""
|
| 224 |
+
assert n >= 1
|
| 225 |
+
if _GLOBAL_NTIMES_COUNTER[name] < n:
|
| 226 |
+
_GLOBAL_NTIMES_COUNTER[name] += 1
|
| 227 |
+
return True
|
| 228 |
+
else:
|
| 229 |
+
return False
|
| 230 |
+
|
| 231 |
+
|
| 232 |
+
class Every:
|
| 233 |
+
def __init__(self, n: int, on_first: bool = False):
|
| 234 |
+
assert n > 0
|
| 235 |
+
self._i = 0 if on_first else 1
|
| 236 |
+
self._n = n
|
| 237 |
+
|
| 238 |
+
def __call__(self):
|
| 239 |
+
return self._i % self._n == 0
|
| 240 |
+
|
| 241 |
+
def __bool__(self):
|
| 242 |
+
raise RuntimeError("`Every` objects should be used by calling ()")
|
| 243 |
+
|
| 244 |
+
|
| 245 |
+
def encode_base64(obj) -> str:
|
| 246 |
+
return codecs.encode(pickle.dumps(obj), "base64").decode()
|
| 247 |
+
|
| 248 |
+
|
| 249 |
+
def decode_base64(s: str):
|
| 250 |
+
return pickle.loads(codecs.decode(s.encode(), "base64"))
|
| 251 |
+
|
| 252 |
+
|
| 253 |
+
def safe_hash(input_tuple):
|
| 254 |
+
# keep 128 bits of the hash
|
| 255 |
+
tuple_string = repr(input_tuple).encode("utf-8")
|
| 256 |
+
sha256 = hashlib.sha256()
|
| 257 |
+
sha256.update(tuple_string)
|
| 258 |
+
|
| 259 |
+
seed = int(sha256.hexdigest(), 16)
|
| 260 |
+
|
| 261 |
+
return seed & 0xFFFFFFFFFFFFFFFFFFFFFFFFFFFFFFFF
|
groot/vla/common/utils/misc/torch_utils.py
ADDED
|
@@ -0,0 +1,748 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
from __future__ import annotations
|
| 2 |
+
|
| 3 |
+
from copy import deepcopy
|
| 4 |
+
import os
|
| 5 |
+
import random
|
| 6 |
+
import time
|
| 7 |
+
from typing import List, Optional, Tuple, Union
|
| 8 |
+
|
| 9 |
+
import numpy as np
|
| 10 |
+
import torch
|
| 11 |
+
import torch.nn as nn
|
| 12 |
+
import tree
|
| 13 |
+
from typing_extensions import Literal
|
| 14 |
+
|
| 15 |
+
from ..data_structure.tree_utils import tree_value_at_path
|
| 16 |
+
from ..io.file_utils import f_join
|
| 17 |
+
from ..io.print_utils import to_readable_count_str
|
| 18 |
+
from .functional_utils import assert_implements_method, implements_method
|
| 19 |
+
|
| 20 |
+
|
| 21 |
+
def weight_init(m):
|
| 22 |
+
"""Custom weight init for Conv2D and Linear layers."""
|
| 23 |
+
if isinstance(m, nn.Linear):
|
| 24 |
+
nn.init.orthogonal_(m.weight.data)
|
| 25 |
+
if hasattr(m.bias, "data"):
|
| 26 |
+
m.bias.data.fill_(0.0)
|
| 27 |
+
elif isinstance(m, nn.Conv2d) or isinstance(m, nn.ConvTranspose2d):
|
| 28 |
+
gain = nn.init.calculate_gain("relu")
|
| 29 |
+
nn.init.orthogonal_(m.weight.data, gain)
|
| 30 |
+
if hasattr(m.bias, "data"):
|
| 31 |
+
m.bias.data.fill_(0.0)
|
| 32 |
+
|
| 33 |
+
|
| 34 |
+
def get_seed(
|
| 35 |
+
seed: Union[int, str, None],
|
| 36 |
+
handle_invalid_seed: Literal["none", "system", "raise"] = "none",
|
| 37 |
+
) -> Optional[int]:
|
| 38 |
+
"""
|
| 39 |
+
Args:
|
| 40 |
+
seed:
|
| 41 |
+
"system": use scrambled int based on system time
|
| 42 |
+
None or int < 0: invalid seed values, see `handle_invalid_seed`
|
| 43 |
+
int >= 0: returns seed
|
| 44 |
+
handle_invalid_seed: None or int < 0
|
| 45 |
+
- "none": returns None
|
| 46 |
+
- "system": returns scrambled int based on system time
|
| 47 |
+
- "raise": raise Exception
|
| 48 |
+
"""
|
| 49 |
+
handle_invalid_seed = handle_invalid_seed.lower()
|
| 50 |
+
assert handle_invalid_seed in ["none", "system", "raise"]
|
| 51 |
+
if isinstance(seed, str):
|
| 52 |
+
assert seed in ["system"]
|
| 53 |
+
invalid = False
|
| 54 |
+
else:
|
| 55 |
+
assert seed is None or isinstance(seed, int)
|
| 56 |
+
invalid = seed is None or seed < 0
|
| 57 |
+
|
| 58 |
+
if seed == "system" or invalid and handle_invalid_seed == "system":
|
| 59 |
+
# https://stackoverflow.com/questions/27276135/python-random-system-time-seed
|
| 60 |
+
t = int(time.time() * 100000)
|
| 61 |
+
return (
|
| 62 |
+
((t & 0xFF000000) >> 24)
|
| 63 |
+
+ ((t & 0x00FF0000) >> 8)
|
| 64 |
+
+ ((t & 0x0000FF00) << 8)
|
| 65 |
+
+ ((t & 0x000000FF) << 24)
|
| 66 |
+
)
|
| 67 |
+
elif invalid:
|
| 68 |
+
if handle_invalid_seed == "none":
|
| 69 |
+
return None
|
| 70 |
+
elif handle_invalid_seed == "raise":
|
| 71 |
+
raise ValueError(
|
| 72 |
+
f"Invalid random seed: {seed}, " f'must be a non-negative integer or "system"'
|
| 73 |
+
)
|
| 74 |
+
else:
|
| 75 |
+
raise NotImplementedError
|
| 76 |
+
else:
|
| 77 |
+
return seed
|
| 78 |
+
|
| 79 |
+
|
| 80 |
+
def set_deterministic(flag: bool = True):
|
| 81 |
+
if not flag:
|
| 82 |
+
return
|
| 83 |
+
|
| 84 |
+
os.environ["CUBLAS_WORKSPACE_CONFIG"] = ":4096:8"
|
| 85 |
+
os.environ["HOROVOD_FUSION_THRESHOLD"] = "0"
|
| 86 |
+
import torch.backends.cudnn as cudnn
|
| 87 |
+
|
| 88 |
+
cudnn.deterministic = True
|
| 89 |
+
cudnn.benchmark = False
|
| 90 |
+
if hasattr(torch, "use_deterministic_algorithms"):
|
| 91 |
+
# only available in PyTorch >= 1.9
|
| 92 |
+
torch.use_deterministic_algorithms(True)
|
| 93 |
+
elif hasattr(torch, "set_deterministic"):
|
| 94 |
+
# only available in PyTorch >= 1.7
|
| 95 |
+
torch.set_deterministic(True)
|
| 96 |
+
|
| 97 |
+
|
| 98 |
+
def set_seed_everywhere(
|
| 99 |
+
seed: Optional[Union[int, str]],
|
| 100 |
+
deterministic=False,
|
| 101 |
+
set_tensorflow=False,
|
| 102 |
+
handle_invalid_seed: Literal["none", "system", "raise"] = "none",
|
| 103 |
+
) -> Optional[int]:
|
| 104 |
+
"""
|
| 105 |
+
References:
|
| 106 |
+
- https://github.com/NVIDIA/framework-determinism/blob/master/pytorch.md
|
| 107 |
+
- https://pytorch.org/docs/stable/notes/randomness.html
|
| 108 |
+
- CUBLAS env var:
|
| 109 |
+
https://docs.nvidia.com/cuda/cublas/index.html#cublasApi_reproducibility
|
| 110 |
+
|
| 111 |
+
Args:
|
| 112 |
+
seed: see `get_seed()`
|
| 113 |
+
handle_invalid_seed: see `get_seed()`
|
| 114 |
+
"""
|
| 115 |
+
set_deterministic(deterministic)
|
| 116 |
+
|
| 117 |
+
seed = get_seed(seed, handle_invalid_seed=handle_invalid_seed)
|
| 118 |
+
if seed is None:
|
| 119 |
+
return None
|
| 120 |
+
|
| 121 |
+
os.environ["PYTHONHASHSEED"] = str(seed)
|
| 122 |
+
random.seed(seed)
|
| 123 |
+
np.random.seed(seed)
|
| 124 |
+
torch.manual_seed(seed)
|
| 125 |
+
if torch.cuda.is_available():
|
| 126 |
+
torch.cuda.manual_seed_all(seed)
|
| 127 |
+
if set_tensorflow:
|
| 128 |
+
try:
|
| 129 |
+
import tensorflow as tf
|
| 130 |
+
|
| 131 |
+
tf.random.set_seed(seed)
|
| 132 |
+
except ImportError:
|
| 133 |
+
pass
|
| 134 |
+
return seed
|
| 135 |
+
|
| 136 |
+
|
| 137 |
+
class eval_mode(object):
|
| 138 |
+
def __init__(self, *models):
|
| 139 |
+
self.models = models
|
| 140 |
+
|
| 141 |
+
def __enter__(self):
|
| 142 |
+
self.prev_states = []
|
| 143 |
+
for model in self.models:
|
| 144 |
+
self.prev_states.append(model.training)
|
| 145 |
+
model.train(False)
|
| 146 |
+
|
| 147 |
+
def __exit__(self, *args):
|
| 148 |
+
for model, state in zip(self.models, self.prev_states):
|
| 149 |
+
model.train(state)
|
| 150 |
+
return False
|
| 151 |
+
|
| 152 |
+
|
| 153 |
+
def get_device(x, strict: bool = False) -> int:
|
| 154 |
+
"""
|
| 155 |
+
Args:
|
| 156 |
+
x: can be any arbitrary nested structure of np array and torch tensor
|
| 157 |
+
strict: True to check all batch sizes are the same
|
| 158 |
+
"""
|
| 159 |
+
xs = tree.flatten(x)
|
| 160 |
+
|
| 161 |
+
def _get_device(x):
|
| 162 |
+
if torch.is_tensor(x):
|
| 163 |
+
return x.device
|
| 164 |
+
elif isinstance(x, nn.Module):
|
| 165 |
+
return get_module_device(x)
|
| 166 |
+
else:
|
| 167 |
+
return None
|
| 168 |
+
|
| 169 |
+
if strict:
|
| 170 |
+
devices = [_get_device(x) for x in xs]
|
| 171 |
+
assert all(
|
| 172 |
+
b == devices[0] for b in devices
|
| 173 |
+
), f"devices must all be the same in nested structure: {devices}"
|
| 174 |
+
return devices[0]
|
| 175 |
+
else:
|
| 176 |
+
return _get_device(xs[0])
|
| 177 |
+
|
| 178 |
+
|
| 179 |
+
def load_torch(*fpath: str, map_location="cpu") -> dict:
|
| 180 |
+
"""
|
| 181 |
+
Default maps to "cpu"
|
| 182 |
+
"""
|
| 183 |
+
fpath = str(f_join(fpath))
|
| 184 |
+
try:
|
| 185 |
+
return torch.load(fpath, map_location=map_location)
|
| 186 |
+
except RuntimeError as e:
|
| 187 |
+
raise RuntimeError(f"{e}\n\n --- Error loading {fpath}")
|
| 188 |
+
|
| 189 |
+
|
| 190 |
+
def save_torch(D, *fpath):
|
| 191 |
+
"""
|
| 192 |
+
Supports both (D, fpath) and (fpath, D) arg order, as long as one of them is a str
|
| 193 |
+
"""
|
| 194 |
+
if isinstance(D, str):
|
| 195 |
+
assert not isinstance(fpath, str), "Either torch_save(D, fpath) " "or torch_save(fpath, D)"
|
| 196 |
+
fpath, D = D, fpath
|
| 197 |
+
torch.save(D, str(f_join(fpath)))
|
| 198 |
+
|
| 199 |
+
|
| 200 |
+
# Aliases for consistency with load_pickle, load_text, load_json/yaml, etc.
|
| 201 |
+
torch_load = load_torch
|
| 202 |
+
torch_save = save_torch
|
| 203 |
+
dump_torch = save_torch
|
| 204 |
+
|
| 205 |
+
|
| 206 |
+
def torch_compute_stats(x, precision: int = 2):
|
| 207 |
+
x = x.to(dtype=torch.float32)
|
| 208 |
+
return (
|
| 209 |
+
f"mean|std: {torch.mean(x):.{precision}f} +/- {torch.std(x):.{precision}f}, "
|
| 210 |
+
f"median: {torch.median(x):.{precision}f}, "
|
| 211 |
+
f"max: {torch.max(x):.{precision}f}, min: {torch.min(x):.{precision}f}"
|
| 212 |
+
)
|
| 213 |
+
|
| 214 |
+
|
| 215 |
+
def tensor_hash(x: torch.Tensor, mode: str = "mean"):
|
| 216 |
+
if isinstance(x, np.ndarray):
|
| 217 |
+
x = torch.from_numpy(x)
|
| 218 |
+
x = x.float().abs()
|
| 219 |
+
if mode == "sum":
|
| 220 |
+
x = x.sum()
|
| 221 |
+
elif mode == "mean":
|
| 222 |
+
x = x.mean()
|
| 223 |
+
else:
|
| 224 |
+
raise NotImplementedError
|
| 225 |
+
return float(x)
|
| 226 |
+
|
| 227 |
+
|
| 228 |
+
def torch_flatten_indices(indices: torch.Tensor, shape: Tuple[int]):
|
| 229 |
+
"""
|
| 230 |
+
Convert M dim indices to 1D indices with the given shape
|
| 231 |
+
|
| 232 |
+
Args:
|
| 233 |
+
indices: BxM, batch_size x M-dimensional
|
| 234 |
+
"""
|
| 235 |
+
offsets = np.array(shape) # e.g. [3, 4, 5, 6]
|
| 236 |
+
offsets = np.append(offsets[1:], 1) # [4, 5, 6, 1]
|
| 237 |
+
offsets = np.cumprod(offsets[::-1])[::-1] # [4*5*6, 5*6, 6, 1]
|
| 238 |
+
offsets = torch.tensor(offsets.copy(), dtype=torch.long)
|
| 239 |
+
assert offsets.size() == (len(shape),)
|
| 240 |
+
return (indices * offsets.to(device=indices.device)).sum(dim=1)
|
| 241 |
+
|
| 242 |
+
|
| 243 |
+
def torch_multi_index_select(x: torch.Tensor, indices: torch.Tensor):
|
| 244 |
+
"""
|
| 245 |
+
Args:
|
| 246 |
+
x: N dim
|
| 247 |
+
indices: [B x M], M <= N, will select the first M-D from N-D
|
| 248 |
+
|
| 249 |
+
Returns:
|
| 250 |
+
(N - M + 1) dim
|
| 251 |
+
"""
|
| 252 |
+
assert indices.ndim == 2
|
| 253 |
+
B, idx_dim = indices.size()
|
| 254 |
+
x_shape = x.size()
|
| 255 |
+
assert len(x_shape) >= idx_dim
|
| 256 |
+
remainder_dim = len(x_shape) - idx_dim
|
| 257 |
+
if remainder_dim == 0:
|
| 258 |
+
x = torch.flatten(x)
|
| 259 |
+
else:
|
| 260 |
+
x = x.view(-1, *x_shape[-remainder_dim:]) # flatten the first M dims
|
| 261 |
+
# convert indices to a 1D flattened array
|
| 262 |
+
indices = torch_flatten_indices(indices, x_shape[:idx_dim])
|
| 263 |
+
selected = x[indices]
|
| 264 |
+
return selected
|
| 265 |
+
|
| 266 |
+
|
| 267 |
+
# ========== module operations =========
|
| 268 |
+
def set_requires_grad(model, requires_grad):
|
| 269 |
+
if torch.is_tensor(model):
|
| 270 |
+
model.requires_grad = requires_grad
|
| 271 |
+
else:
|
| 272 |
+
for param in model.parameters():
|
| 273 |
+
param.requires_grad = requires_grad
|
| 274 |
+
|
| 275 |
+
|
| 276 |
+
def freeze_params(model):
|
| 277 |
+
set_requires_grad(model, False)
|
| 278 |
+
if not torch.is_tensor(model):
|
| 279 |
+
model.eval()
|
| 280 |
+
|
| 281 |
+
|
| 282 |
+
def unfreeze_params(model):
|
| 283 |
+
set_requires_grad(model, True)
|
| 284 |
+
if not torch.is_tensor(model):
|
| 285 |
+
model.train()
|
| 286 |
+
|
| 287 |
+
|
| 288 |
+
def clip_grad_value(model, max_value):
|
| 289 |
+
with torch.no_grad():
|
| 290 |
+
nn.utils.clip_grad_value_(model.parameters(), max_value)
|
| 291 |
+
|
| 292 |
+
|
| 293 |
+
def clip_grad_norm(model, max_norm, norm_type=2):
|
| 294 |
+
"""
|
| 295 |
+
Returns:
|
| 296 |
+
Total norm of the parameters (viewed as a single vector).
|
| 297 |
+
"""
|
| 298 |
+
with torch.no_grad():
|
| 299 |
+
return nn.utils.clip_grad_norm_(model.parameters(), max_norm=max_norm, norm_type=norm_type)
|
| 300 |
+
|
| 301 |
+
|
| 302 |
+
def implements_state_dict(object, requires_load_method: bool = False):
|
| 303 |
+
cond = implements_method(object, "state_dict")
|
| 304 |
+
if requires_load_method:
|
| 305 |
+
return cond and implements_method(object, "load_state_dict")
|
| 306 |
+
else:
|
| 307 |
+
return cond
|
| 308 |
+
|
| 309 |
+
|
| 310 |
+
def unwrap_ddp_model(model):
|
| 311 |
+
if hasattr(model, "module") and len(list(model.children())) == 1:
|
| 312 |
+
model = model.module
|
| 313 |
+
return model
|
| 314 |
+
|
| 315 |
+
|
| 316 |
+
class DDPMethodWrapper(nn.Module):
|
| 317 |
+
"""
|
| 318 |
+
Wraps another module's method as forward(), because DDP only works on forward()
|
| 319 |
+
This module can be wrapped with DDP and directly called.
|
| 320 |
+
It will not save any extra parameters
|
| 321 |
+
"""
|
| 322 |
+
|
| 323 |
+
def __init__(self, net: nn.Module, method_name: str):
|
| 324 |
+
super().__init__()
|
| 325 |
+
self.net = net
|
| 326 |
+
assert_implements_method(net, method_name)
|
| 327 |
+
self._method_name = method_name
|
| 328 |
+
|
| 329 |
+
def forward(self, *args, **kwargs):
|
| 330 |
+
return getattr(self.net, self._method_name)(*args, **kwargs)
|
| 331 |
+
|
| 332 |
+
def state_dict(self):
|
| 333 |
+
return {}
|
| 334 |
+
|
| 335 |
+
|
| 336 |
+
def to_state_dict(objects, to_cpu: bool = False, copy: bool = False, unwrap_ddp: bool = False):
|
| 337 |
+
"""
|
| 338 |
+
Anything that has state_dict() method, e.g. nn.Module, Optimizer, LRScheduler, etc.
|
| 339 |
+
|
| 340 |
+
Args:
|
| 341 |
+
to_cpu: True to copy to CPU. The original tensors will still be on GPU.
|
| 342 |
+
copy: takes effect if and only if to_cpu is False
|
| 343 |
+
"""
|
| 344 |
+
|
| 345 |
+
def _transfer(x):
|
| 346 |
+
if torch.is_tensor(x):
|
| 347 |
+
x = x.detach()
|
| 348 |
+
if to_cpu:
|
| 349 |
+
return x.cpu()
|
| 350 |
+
elif copy:
|
| 351 |
+
return x.clone()
|
| 352 |
+
return x
|
| 353 |
+
|
| 354 |
+
def _to_state_dict(m):
|
| 355 |
+
if implements_state_dict(m):
|
| 356 |
+
if isinstance(m, nn.Module) and unwrap_ddp:
|
| 357 |
+
m = unwrap_ddp_model(m)
|
| 358 |
+
return tree.map_structure(_transfer, m.state_dict())
|
| 359 |
+
else:
|
| 360 |
+
return _transfer(m)
|
| 361 |
+
|
| 362 |
+
return tree.map_structure(_to_state_dict, objects)
|
| 363 |
+
|
| 364 |
+
|
| 365 |
+
def load_state_dict(objects, states, strip_prefix=None, strict=False):
|
| 366 |
+
"""
|
| 367 |
+
Args:
|
| 368 |
+
strict: objects and states must match exactly
|
| 369 |
+
strip_prefix: only match the keys that have the prefix, and strip it
|
| 370 |
+
"""
|
| 371 |
+
|
| 372 |
+
def _load(paths, obj):
|
| 373 |
+
if not implements_method(obj, "load_state_dict"):
|
| 374 |
+
raise ValueError(f"Object {type(obj)} does not support load_state_dict() method")
|
| 375 |
+
try:
|
| 376 |
+
state = tree_value_at_path(states, paths)
|
| 377 |
+
except ValueError: # paths do not exist in `states` structure
|
| 378 |
+
if strict:
|
| 379 |
+
raise
|
| 380 |
+
else:
|
| 381 |
+
return
|
| 382 |
+
if strip_prefix:
|
| 383 |
+
assert isinstance(strip_prefix, str)
|
| 384 |
+
state = {
|
| 385 |
+
k[len(strip_prefix) :]: v for k, v in state.items() if k.startswith(strip_prefix)
|
| 386 |
+
}
|
| 387 |
+
if isinstance(obj, nn.Module):
|
| 388 |
+
return obj.load_state_dict(state, strict=strict)
|
| 389 |
+
else:
|
| 390 |
+
return obj.load_state_dict(state)
|
| 391 |
+
|
| 392 |
+
return tree.map_structure_with_path(_load, objects)
|
| 393 |
+
|
| 394 |
+
|
| 395 |
+
def count_parameters(model):
|
| 396 |
+
return sum(x.numel() for x in model.parameters())
|
| 397 |
+
|
| 398 |
+
|
| 399 |
+
def readable_count_parameters(model, precision: int = 2):
|
| 400 |
+
return to_readable_count_str(count_parameters(model), precision=precision)
|
| 401 |
+
|
| 402 |
+
|
| 403 |
+
def get_module_device(model):
|
| 404 |
+
"""
|
| 405 |
+
Returns:
|
| 406 |
+
first model parameter's device
|
| 407 |
+
"""
|
| 408 |
+
return next(model.parameters()).device
|
| 409 |
+
|
| 410 |
+
|
| 411 |
+
def maybe_transfer_module(model, device):
|
| 412 |
+
"""
|
| 413 |
+
Transfer a module to another device if and only if they are on different devices.
|
| 414 |
+
Assumes that the module's first parameter determines the module device, i.e.
|
| 415 |
+
no model parallelism.
|
| 416 |
+
|
| 417 |
+
Returns:
|
| 418 |
+
True if module is transferred to a different device, False otherwise
|
| 419 |
+
"""
|
| 420 |
+
if device is None:
|
| 421 |
+
return False
|
| 422 |
+
device = torch.device(device)
|
| 423 |
+
if get_module_device(model) != device:
|
| 424 |
+
model.to(device=device)
|
| 425 |
+
return True
|
| 426 |
+
else:
|
| 427 |
+
return False
|
| 428 |
+
|
| 429 |
+
|
| 430 |
+
def clone_model(model):
|
| 431 |
+
with torch.no_grad():
|
| 432 |
+
new_model = deepcopy(model).to(get_module_device(model))
|
| 433 |
+
# new_model.load_state_dict(model.state_dict())
|
| 434 |
+
return new_model
|
| 435 |
+
|
| 436 |
+
|
| 437 |
+
def update_soft_params(net, target_net, tau):
|
| 438 |
+
for param, target_param in zip(net.parameters(), target_net.parameters()):
|
| 439 |
+
target_param.data.copy_(tau * param.data + (1 - tau) * target_param.data)
|
| 440 |
+
|
| 441 |
+
|
| 442 |
+
def tie_weights(src, trg):
|
| 443 |
+
# TODO deprecate this
|
| 444 |
+
assert type(src) is type(trg)
|
| 445 |
+
trg.weight = src.weight
|
| 446 |
+
trg.bias = src.bias
|
| 447 |
+
|
| 448 |
+
|
| 449 |
+
def torch_normalize(tensor: torch.Tensor, mean, std, inplace=False):
|
| 450 |
+
"""
|
| 451 |
+
Adapted from https://pytorch.org/docs/stable/_modules/torchvision/transforms/functional.html#normalize
|
| 452 |
+
|
| 453 |
+
Normalize a tensor image with mean and standard deviation.
|
| 454 |
+
|
| 455 |
+
.. note::
|
| 456 |
+
This transform acts out of place by default, i.e., it does not mutates the input tensor.
|
| 457 |
+
|
| 458 |
+
See :class:`~torchvision.transforms.Normalize` for more details.
|
| 459 |
+
|
| 460 |
+
Args:
|
| 461 |
+
tensor (Tensor): Tensor image of size (C, H, W) to be normalized.
|
| 462 |
+
mean (sequence): Sequence of means for each channel.
|
| 463 |
+
std (sequence): Sequence of standard deviations for each channel.
|
| 464 |
+
inplace(bool,optional): Bool to make this operation inplace.
|
| 465 |
+
|
| 466 |
+
Returns:
|
| 467 |
+
Tensor: Normalized Tensor image.
|
| 468 |
+
"""
|
| 469 |
+
if not torch.is_tensor(tensor):
|
| 470 |
+
raise TypeError("tensor should be a torch tensor. Got {}.".format(type(tensor)))
|
| 471 |
+
|
| 472 |
+
if not inplace:
|
| 473 |
+
tensor = tensor.clone()
|
| 474 |
+
|
| 475 |
+
dtype = tensor.dtype
|
| 476 |
+
mean = torch.as_tensor(mean, dtype=dtype, device=tensor.device)
|
| 477 |
+
std = torch.as_tensor(std, dtype=dtype, device=tensor.device)
|
| 478 |
+
if (std == 0).any():
|
| 479 |
+
raise ValueError(
|
| 480 |
+
f"std evaluated to zero after conversion to {dtype}, leading to division by zero."
|
| 481 |
+
)
|
| 482 |
+
if mean.ndim == 1:
|
| 483 |
+
mean = mean[:, None, None]
|
| 484 |
+
if std.ndim == 1:
|
| 485 |
+
std = std[:, None, None]
|
| 486 |
+
tensor.sub_(mean).div_(std)
|
| 487 |
+
return tensor
|
| 488 |
+
|
| 489 |
+
|
| 490 |
+
def contains_rnn(net: nn.Module) -> bool:
|
| 491 |
+
for m in net.modules():
|
| 492 |
+
if isinstance(m, nn.RNNBase):
|
| 493 |
+
return True
|
| 494 |
+
return False
|
| 495 |
+
|
| 496 |
+
|
| 497 |
+
def multi_one_hot(x, num_classes: List[int], to_float=True):
|
| 498 |
+
"""
|
| 499 |
+
Concatenates multiple one-hot matrices, useful for embedding MultiDiscrete action space
|
| 500 |
+
|
| 501 |
+
Args:
|
| 502 |
+
x: torch.long, [*N, D]
|
| 503 |
+
num_classes: list len == D, match the last dim of x
|
| 504 |
+
|
| 505 |
+
Returns:
|
| 506 |
+
[*N, sum(num_classes)]
|
| 507 |
+
"""
|
| 508 |
+
from torch.nn.functional import one_hot
|
| 509 |
+
|
| 510 |
+
assert x.dtype == torch.long
|
| 511 |
+
assert x.dim() >= 2, x.size()
|
| 512 |
+
assert len(num_classes) == x.size(-1), f"{len(num_classes)} != {x.size(1)}"
|
| 513 |
+
result = torch.cat(
|
| 514 |
+
[one_hot(t, c) for t, c in zip(torch.unbind(x, dim=-1), num_classes)], dim=-1
|
| 515 |
+
)
|
| 516 |
+
if to_float:
|
| 517 |
+
return result.float()
|
| 518 |
+
else:
|
| 519 |
+
return result
|
| 520 |
+
|
| 521 |
+
|
| 522 |
+
def _random_derangement(n):
|
| 523 |
+
while True:
|
| 524 |
+
v = [i for i in range(n)]
|
| 525 |
+
for j in range(n - 1, -1, -1):
|
| 526 |
+
p = random.randint(0, j)
|
| 527 |
+
if v[p] == j:
|
| 528 |
+
break
|
| 529 |
+
else:
|
| 530 |
+
v[j], v[p] = v[p], v[j]
|
| 531 |
+
else:
|
| 532 |
+
if v[0] != 0:
|
| 533 |
+
return tuple(v)
|
| 534 |
+
|
| 535 |
+
|
| 536 |
+
def random_derangement(n, format: Literal["list", "numpy", "torch"] = "torch"):
|
| 537 |
+
"""
|
| 538 |
+
Early refusal algorithm, described at
|
| 539 |
+
https://stackoverflow.com/questions/25200220/generate-a-random-derangement-of-a-list
|
| 540 |
+
Derangement is permuation without fixed point, useful for constructing negative
|
| 541 |
+
pairs in contrastive learning.
|
| 542 |
+
"""
|
| 543 |
+
assert format in ["list", "numpy", "torch"]
|
| 544 |
+
D = _random_derangement(n)
|
| 545 |
+
if format == "list":
|
| 546 |
+
return D
|
| 547 |
+
elif format == "numpy":
|
| 548 |
+
return np.array(D, dtype=np.long)
|
| 549 |
+
elif format == "torch":
|
| 550 |
+
return torch.tensor(D, dtype=torch.long)
|
| 551 |
+
else:
|
| 552 |
+
raise NotImplementedError(f"Unknown format {format}")
|
| 553 |
+
|
| 554 |
+
|
| 555 |
+
def classify_accuracy(
|
| 556 |
+
output,
|
| 557 |
+
target,
|
| 558 |
+
topk: Union[int, List[int], Tuple[int]] = 1,
|
| 559 |
+
mask=None,
|
| 560 |
+
reduction="mean",
|
| 561 |
+
scale_100=False,
|
| 562 |
+
):
|
| 563 |
+
"""
|
| 564 |
+
Computes the accuracy over the k top predictions for the specified values of k.
|
| 565 |
+
Accuracy is a float between 0.0 and 1.0
|
| 566 |
+
|
| 567 |
+
Args:
|
| 568 |
+
topk: if int, return a single acc. If tuple, return a tuple of accs
|
| 569 |
+
mask: shape [batch_size,], binary mask of whether to include this sample or not
|
| 570 |
+
"""
|
| 571 |
+
if isinstance(topk, int):
|
| 572 |
+
topk = [topk]
|
| 573 |
+
is_int = True
|
| 574 |
+
else:
|
| 575 |
+
is_int = False
|
| 576 |
+
|
| 577 |
+
batch_size = target.size(0)
|
| 578 |
+
assert output.size(0) == batch_size
|
| 579 |
+
if mask is not None:
|
| 580 |
+
assert mask.dim() == 1
|
| 581 |
+
assert mask.size(0) == batch_size
|
| 582 |
+
|
| 583 |
+
assert reduction in ["sum", "mean", "none"]
|
| 584 |
+
if reduction != "mean":
|
| 585 |
+
assert not scale_100, f"reduce={reduction} does not support scale_100=True"
|
| 586 |
+
|
| 587 |
+
with torch.no_grad():
|
| 588 |
+
maxk = max(topk)
|
| 589 |
+
|
| 590 |
+
_, pred = output.topk(maxk, 1, True, True)
|
| 591 |
+
pred = pred.t()
|
| 592 |
+
correct = pred.eq(target.view(1, -1).expand_as(pred))
|
| 593 |
+
if mask is not None:
|
| 594 |
+
correct = mask * correct
|
| 595 |
+
|
| 596 |
+
mult = 100.0 if scale_100 else 1.0
|
| 597 |
+
res = []
|
| 598 |
+
for k in topk:
|
| 599 |
+
correct_k = correct[:k].int().sum(dim=0)
|
| 600 |
+
if reduction == "mean":
|
| 601 |
+
if mask is not None:
|
| 602 |
+
# fmt: off
|
| 603 |
+
res.append(
|
| 604 |
+
float(correct_k.float().sum().mul_(mult / mask.sum().item()).item())
|
| 605 |
+
)
|
| 606 |
+
# fmt: on
|
| 607 |
+
else:
|
| 608 |
+
res.append(float(correct_k.float().sum().mul_(mult / batch_size).item()))
|
| 609 |
+
elif reduction == "sum":
|
| 610 |
+
res.append(int(correct_k.sum().item()))
|
| 611 |
+
elif reduction == "none":
|
| 612 |
+
res.append(correct_k)
|
| 613 |
+
else:
|
| 614 |
+
raise NotImplementedError(f"Unknown reduce={reduction}")
|
| 615 |
+
|
| 616 |
+
if is_int:
|
| 617 |
+
assert len(res) == 1, "INTERNAL"
|
| 618 |
+
return res[0]
|
| 619 |
+
else:
|
| 620 |
+
return res
|
| 621 |
+
|
| 622 |
+
|
| 623 |
+
def sequential_split_dataset(dataset: torch.utils.data.Dataset, split_portions: list[float]):
|
| 624 |
+
"""
|
| 625 |
+
Split a dataset into multiple datasets, each with a different portion of the
|
| 626 |
+
original dataset. Uses torch.utils.data.Subset.
|
| 627 |
+
"""
|
| 628 |
+
from .functional_utils import accumulate
|
| 629 |
+
|
| 630 |
+
assert len(split_portions) > 0, "split_portions must be a non-empty list"
|
| 631 |
+
assert all(0.0 <= p <= 1.0 for p in split_portions), f"{split_portions=}"
|
| 632 |
+
assert abs(sum(split_portions) - 1.0) < 1e-6, f"{sum(split_portions)=} != 1.0"
|
| 633 |
+
L = len(dataset)
|
| 634 |
+
assert L > 0, "dataset must be non-empty"
|
| 635 |
+
# split the list with proportions
|
| 636 |
+
lengths = [int(p * L) for p in split_portions]
|
| 637 |
+
# make sure the last split fills the full dataset
|
| 638 |
+
lengths[-1] += L - sum(lengths)
|
| 639 |
+
indices = list(range(L))
|
| 640 |
+
|
| 641 |
+
return [
|
| 642 |
+
torch.utils.data.Subset(dataset, indices[offset - length : offset])
|
| 643 |
+
for offset, length in zip(accumulate(lengths), lengths)
|
| 644 |
+
]
|
| 645 |
+
|
| 646 |
+
|
| 647 |
+
class RunningMeanStd:
|
| 648 |
+
def __init__(self):
|
| 649 |
+
"""
|
| 650 |
+
Calulates the running mean and std of a data stream
|
| 651 |
+
https://en.wikipedia.org/wiki/Algorithms_for_calculating_variance#Parallel_algorithm
|
| 652 |
+
"""
|
| 653 |
+
self._mean = None
|
| 654 |
+
self._var = None
|
| 655 |
+
self._count = 0
|
| 656 |
+
|
| 657 |
+
@property
|
| 658 |
+
def mean(self):
|
| 659 |
+
return self._mean
|
| 660 |
+
|
| 661 |
+
@property
|
| 662 |
+
def var(self):
|
| 663 |
+
return self._var
|
| 664 |
+
|
| 665 |
+
@property
|
| 666 |
+
def std(self):
|
| 667 |
+
if isinstance(self._var, np.ndarray):
|
| 668 |
+
return np.sqrt(self._var)
|
| 669 |
+
else:
|
| 670 |
+
return self._var.sqrt()
|
| 671 |
+
|
| 672 |
+
@property
|
| 673 |
+
def count(self):
|
| 674 |
+
return self._count
|
| 675 |
+
|
| 676 |
+
def update(self, values: np.ndarray | torch.Tensor) -> None:
|
| 677 |
+
from .array_tensor_utils import any_mean, any_variance, get_batch_size
|
| 678 |
+
|
| 679 |
+
batch_mean = any_mean(values, dim=0)
|
| 680 |
+
# our running var calculation currently only supports unbiased=False
|
| 681 |
+
batch_var = any_variance(values, dim=0, unbiased=False)
|
| 682 |
+
batch_count = get_batch_size(values)
|
| 683 |
+
self.update_from_moments(batch_mean, batch_var, batch_count)
|
| 684 |
+
|
| 685 |
+
def update_from_moments(
|
| 686 |
+
self,
|
| 687 |
+
batch_mean: np.ndarray | torch.Tensor,
|
| 688 |
+
batch_var: np.ndarray | torch.Tensor,
|
| 689 |
+
batch_count: int,
|
| 690 |
+
) -> None:
|
| 691 |
+
from .array_tensor_utils import any_get_shape
|
| 692 |
+
|
| 693 |
+
is_tensor = torch.is_tensor(batch_mean)
|
| 694 |
+
_zeros = batch_mean.new_zeros if is_tensor else np.zeros
|
| 695 |
+
if self._mean is None:
|
| 696 |
+
self._mean = _zeros(any_get_shape(batch_mean))
|
| 697 |
+
if self._var is None:
|
| 698 |
+
self._var = _zeros(any_get_shape(batch_var)) + 1.0
|
| 699 |
+
|
| 700 |
+
delta = batch_mean - self._mean
|
| 701 |
+
tot_count = self._count + batch_count
|
| 702 |
+
assert tot_count > 0, "count must be > 0"
|
| 703 |
+
|
| 704 |
+
new_mean = self._mean + delta * batch_count / tot_count
|
| 705 |
+
m_a = self._var * self._count
|
| 706 |
+
m_b = batch_var * batch_count
|
| 707 |
+
m_2 = m_a + m_b + delta * delta * self._count * batch_count / tot_count
|
| 708 |
+
new_var = m_2 / tot_count
|
| 709 |
+
|
| 710 |
+
self._mean = new_mean
|
| 711 |
+
self._var = new_var
|
| 712 |
+
self._count = tot_count
|
| 713 |
+
|
| 714 |
+
|
| 715 |
+
class AverageMeter:
|
| 716 |
+
"""Computes and stores the average and current value"""
|
| 717 |
+
|
| 718 |
+
def __init__(self, name="", fmt="f"):
|
| 719 |
+
self._name = name
|
| 720 |
+
self._fmt = fmt
|
| 721 |
+
self.reset()
|
| 722 |
+
|
| 723 |
+
def reset(self):
|
| 724 |
+
self._sum = 0.0
|
| 725 |
+
self._count = 0.0
|
| 726 |
+
|
| 727 |
+
@torch.no_grad()
|
| 728 |
+
def update(self, value, n=1):
|
| 729 |
+
if torch.is_tensor(value):
|
| 730 |
+
value = value.detach()
|
| 731 |
+
self._sum += value * n
|
| 732 |
+
self._count += n
|
| 733 |
+
|
| 734 |
+
@torch.no_grad()
|
| 735 |
+
def compute(self):
|
| 736 |
+
return float(self._sum / self._count)
|
| 737 |
+
|
| 738 |
+
def __float__(self):
|
| 739 |
+
return self.compute()
|
| 740 |
+
|
| 741 |
+
def __str__(self):
|
| 742 |
+
if self._fmt:
|
| 743 |
+
s = f"{float(self):{self._fmt}}"
|
| 744 |
+
else:
|
| 745 |
+
s = str(float(self))
|
| 746 |
+
if self._name:
|
| 747 |
+
return f"{self._name}: {s}"
|
| 748 |
+
return s
|
groot/vla/common/utils/misc/video_utils.py
ADDED
|
@@ -0,0 +1,487 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
import json
|
| 2 |
+
import subprocess
|
| 3 |
+
|
| 4 |
+
import av
|
| 5 |
+
import cv2
|
| 6 |
+
import numpy as np
|
| 7 |
+
import torchvision
|
| 8 |
+
|
| 9 |
+
# Import decord with graceful fallback
|
| 10 |
+
try:
|
| 11 |
+
import decord
|
| 12 |
+
|
| 13 |
+
DECORD_AVAILABLE = True
|
| 14 |
+
except ImportError:
|
| 15 |
+
DECORD_AVAILABLE = False
|
| 16 |
+
|
| 17 |
+
try:
|
| 18 |
+
import torchcodec
|
| 19 |
+
|
| 20 |
+
TORCHCODEC_AVAILABLE = True
|
| 21 |
+
except (ImportError, RuntimeError):
|
| 22 |
+
TORCHCODEC_AVAILABLE = False
|
| 23 |
+
|
| 24 |
+
|
| 25 |
+
def _get_video_info_ffmpeg(video_path: str) -> dict:
|
| 26 |
+
"""Get video metadata using ffprobe."""
|
| 27 |
+
cmd = [
|
| 28 |
+
"ffprobe",
|
| 29 |
+
"-v",
|
| 30 |
+
"error",
|
| 31 |
+
"-select_streams",
|
| 32 |
+
"v:0",
|
| 33 |
+
"-show_entries",
|
| 34 |
+
"stream=nb_frames,duration,r_frame_rate",
|
| 35 |
+
"-of",
|
| 36 |
+
"json",
|
| 37 |
+
video_path,
|
| 38 |
+
]
|
| 39 |
+
|
| 40 |
+
try:
|
| 41 |
+
output = subprocess.check_output(cmd, stderr=subprocess.STDOUT).decode("utf-8")
|
| 42 |
+
probe_data = json.loads(output)
|
| 43 |
+
stream = probe_data["streams"][0]
|
| 44 |
+
|
| 45 |
+
# Parse frame rate (comes as fraction like "15/1")
|
| 46 |
+
if "/" in stream["r_frame_rate"]:
|
| 47 |
+
num, den = map(int, stream["r_frame_rate"].split("/"))
|
| 48 |
+
fps = num / den
|
| 49 |
+
else:
|
| 50 |
+
fps = float(stream["r_frame_rate"])
|
| 51 |
+
|
| 52 |
+
# Get frame count and duration
|
| 53 |
+
nb_frames = int(stream.get("nb_frames", 0))
|
| 54 |
+
duration = float(stream.get("duration", 0))
|
| 55 |
+
|
| 56 |
+
# If nb_frames is not available, estimate from duration and fps
|
| 57 |
+
if nb_frames == 0 and duration > 0:
|
| 58 |
+
nb_frames = int(duration * fps)
|
| 59 |
+
|
| 60 |
+
return {
|
| 61 |
+
"nb_frames": nb_frames,
|
| 62 |
+
"fps": fps,
|
| 63 |
+
"duration": duration,
|
| 64 |
+
}
|
| 65 |
+
except (subprocess.CalledProcessError, json.JSONDecodeError, KeyError) as e:
|
| 66 |
+
raise ValueError(f"Failed to get video info for {video_path}: {e}")
|
| 67 |
+
|
| 68 |
+
|
| 69 |
+
def _extract_frames_ffmpeg(video_path: str, frame_indices: list[int]) -> np.ndarray:
|
| 70 |
+
"""Extract specific frames using ffmpeg."""
|
| 71 |
+
frames = []
|
| 72 |
+
|
| 73 |
+
for idx in frame_indices:
|
| 74 |
+
# Use ffmpeg to extract a specific frame
|
| 75 |
+
cmd = [
|
| 76 |
+
"ffmpeg",
|
| 77 |
+
"-i",
|
| 78 |
+
video_path,
|
| 79 |
+
"-vf",
|
| 80 |
+
f"select=eq(n\\,{idx})",
|
| 81 |
+
"-vframes",
|
| 82 |
+
"1",
|
| 83 |
+
"-f",
|
| 84 |
+
"image2pipe",
|
| 85 |
+
"-pix_fmt",
|
| 86 |
+
"rgb24",
|
| 87 |
+
"-vcodec",
|
| 88 |
+
"rawvideo",
|
| 89 |
+
"-",
|
| 90 |
+
]
|
| 91 |
+
|
| 92 |
+
try:
|
| 93 |
+
output = subprocess.check_output(cmd, stderr=subprocess.DEVNULL)
|
| 94 |
+
|
| 95 |
+
# Check if output is empty (frame doesn't exist)
|
| 96 |
+
if len(output) == 0:
|
| 97 |
+
raise subprocess.CalledProcessError(1, cmd)
|
| 98 |
+
|
| 99 |
+
# Get frame dimensions by probing first
|
| 100 |
+
if len(frames) == 0:
|
| 101 |
+
info_cmd = [
|
| 102 |
+
"ffprobe",
|
| 103 |
+
"-v",
|
| 104 |
+
"error",
|
| 105 |
+
"-select_streams",
|
| 106 |
+
"v:0",
|
| 107 |
+
"-show_entries",
|
| 108 |
+
"stream=width,height",
|
| 109 |
+
"-of",
|
| 110 |
+
"json",
|
| 111 |
+
video_path,
|
| 112 |
+
]
|
| 113 |
+
info_output = subprocess.check_output(info_cmd).decode("utf-8")
|
| 114 |
+
info_data = json.loads(info_output)
|
| 115 |
+
width = info_data["streams"][0]["width"]
|
| 116 |
+
height = info_data["streams"][0]["height"]
|
| 117 |
+
|
| 118 |
+
# Decode raw RGB data
|
| 119 |
+
frame_data = np.frombuffer(output, dtype=np.uint8)
|
| 120 |
+
frame = frame_data.reshape((height, width, 3))
|
| 121 |
+
frames.append(frame)
|
| 122 |
+
|
| 123 |
+
except subprocess.CalledProcessError:
|
| 124 |
+
# Frame might not exist, create a black frame
|
| 125 |
+
if len(frames) > 0:
|
| 126 |
+
frames.append(np.zeros_like(frames[0]))
|
| 127 |
+
else:
|
| 128 |
+
# Default fallback frame
|
| 129 |
+
frames.append(np.zeros((480, 640, 3), dtype=np.uint8))
|
| 130 |
+
|
| 131 |
+
return np.array(frames)
|
| 132 |
+
|
| 133 |
+
|
| 134 |
+
def _extract_frames_at_timestamps_ffmpeg(video_path: str, timestamps: list[float]) -> np.ndarray:
|
| 135 |
+
"""Extract frames at specific timestamps using ffmpeg."""
|
| 136 |
+
frames = []
|
| 137 |
+
|
| 138 |
+
for timestamp in timestamps:
|
| 139 |
+
cmd = [
|
| 140 |
+
"ffmpeg",
|
| 141 |
+
"-ss",
|
| 142 |
+
str(timestamp),
|
| 143 |
+
"-i",
|
| 144 |
+
video_path,
|
| 145 |
+
"-vframes",
|
| 146 |
+
"1",
|
| 147 |
+
"-f",
|
| 148 |
+
"image2pipe",
|
| 149 |
+
"-pix_fmt",
|
| 150 |
+
"rgb24",
|
| 151 |
+
"-vcodec",
|
| 152 |
+
"rawvideo",
|
| 153 |
+
"-",
|
| 154 |
+
]
|
| 155 |
+
|
| 156 |
+
try:
|
| 157 |
+
output = subprocess.check_output(cmd, stderr=subprocess.DEVNULL)
|
| 158 |
+
|
| 159 |
+
# Check if output is empty (timestamp doesn't exist)
|
| 160 |
+
if len(output) == 0:
|
| 161 |
+
raise subprocess.CalledProcessError(1, cmd)
|
| 162 |
+
|
| 163 |
+
# Get frame dimensions
|
| 164 |
+
if len(frames) == 0:
|
| 165 |
+
info_cmd = [
|
| 166 |
+
"ffprobe",
|
| 167 |
+
"-v",
|
| 168 |
+
"error",
|
| 169 |
+
"-select_streams",
|
| 170 |
+
"v:0",
|
| 171 |
+
"-show_entries",
|
| 172 |
+
"stream=width,height",
|
| 173 |
+
"-of",
|
| 174 |
+
"json",
|
| 175 |
+
video_path,
|
| 176 |
+
]
|
| 177 |
+
info_output = subprocess.check_output(info_cmd).decode("utf-8")
|
| 178 |
+
info_data = json.loads(info_output)
|
| 179 |
+
width = info_data["streams"][0]["width"]
|
| 180 |
+
height = info_data["streams"][0]["height"]
|
| 181 |
+
|
| 182 |
+
# Decode raw RGB data
|
| 183 |
+
frame_data = np.frombuffer(output, dtype=np.uint8)
|
| 184 |
+
frame = frame_data.reshape((height, width, 3))
|
| 185 |
+
frames.append(frame)
|
| 186 |
+
|
| 187 |
+
except subprocess.CalledProcessError:
|
| 188 |
+
# Timestamp might be out of bounds, use last frame or black frame
|
| 189 |
+
if len(frames) > 0:
|
| 190 |
+
frames.append(frames[-1])
|
| 191 |
+
else:
|
| 192 |
+
frames.append(np.zeros((480, 640, 3), dtype=np.uint8))
|
| 193 |
+
|
| 194 |
+
return np.array(frames)
|
| 195 |
+
|
| 196 |
+
|
| 197 |
+
def _extract_all_frames_ffmpeg(video_path: str) -> tuple[np.ndarray, np.ndarray]:
|
| 198 |
+
"""Extract all frames and their timestamps using ffmpeg."""
|
| 199 |
+
# Get video info
|
| 200 |
+
info = _get_video_info_ffmpeg(video_path)
|
| 201 |
+
fps = info["fps"]
|
| 202 |
+
|
| 203 |
+
# Extract all frames
|
| 204 |
+
cmd = [
|
| 205 |
+
"ffmpeg",
|
| 206 |
+
"-i",
|
| 207 |
+
video_path,
|
| 208 |
+
"-f",
|
| 209 |
+
"image2pipe",
|
| 210 |
+
"-pix_fmt",
|
| 211 |
+
"rgb24",
|
| 212 |
+
"-vcodec",
|
| 213 |
+
"rawvideo",
|
| 214 |
+
"-",
|
| 215 |
+
]
|
| 216 |
+
|
| 217 |
+
try:
|
| 218 |
+
output = subprocess.check_output(cmd, stderr=subprocess.DEVNULL)
|
| 219 |
+
|
| 220 |
+
# Get frame dimensions
|
| 221 |
+
info_cmd = [
|
| 222 |
+
"ffprobe",
|
| 223 |
+
"-v",
|
| 224 |
+
"error",
|
| 225 |
+
"-select_streams",
|
| 226 |
+
"v:0",
|
| 227 |
+
"-show_entries",
|
| 228 |
+
"stream=width,height",
|
| 229 |
+
"-of",
|
| 230 |
+
"json",
|
| 231 |
+
video_path,
|
| 232 |
+
]
|
| 233 |
+
info_output = subprocess.check_output(info_cmd).decode("utf-8")
|
| 234 |
+
info_data = json.loads(info_output)
|
| 235 |
+
width = info_data["streams"][0]["width"]
|
| 236 |
+
height = info_data["streams"][0]["height"]
|
| 237 |
+
|
| 238 |
+
# Decode all frames
|
| 239 |
+
frame_data = np.frombuffer(output, dtype=np.uint8)
|
| 240 |
+
total_pixels = len(frame_data) // 3
|
| 241 |
+
actual_frames = total_pixels // (width * height)
|
| 242 |
+
|
| 243 |
+
frames = frame_data[: actual_frames * width * height * 3].reshape(
|
| 244 |
+
(actual_frames, height, width, 3)
|
| 245 |
+
)
|
| 246 |
+
|
| 247 |
+
# Generate timestamps
|
| 248 |
+
timestamps = np.arange(actual_frames) / fps
|
| 249 |
+
|
| 250 |
+
return frames, timestamps
|
| 251 |
+
|
| 252 |
+
except subprocess.CalledProcessError as e:
|
| 253 |
+
raise ValueError(f"Failed to extract frames from {video_path}: {e}")
|
| 254 |
+
|
| 255 |
+
|
| 256 |
+
def get_frames_by_indices(
|
| 257 |
+
video_path: str,
|
| 258 |
+
indices: list[int] | np.ndarray,
|
| 259 |
+
video_backend: str = "ffmpeg",
|
| 260 |
+
video_backend_kwargs: dict = {},
|
| 261 |
+
) -> np.ndarray:
|
| 262 |
+
if video_backend == "decord":
|
| 263 |
+
if not DECORD_AVAILABLE:
|
| 264 |
+
raise ImportError("decord is not available. Install it with: pip install decord")
|
| 265 |
+
vr = decord.VideoReader(video_path, **video_backend_kwargs)
|
| 266 |
+
frames = vr.get_batch(indices)
|
| 267 |
+
return frames.asnumpy()
|
| 268 |
+
elif video_backend == "torchcodec":
|
| 269 |
+
if not TORCHCODEC_AVAILABLE:
|
| 270 |
+
raise ImportError("torchcodec is not available.")
|
| 271 |
+
decoder = torchcodec.decoders.VideoDecoder(
|
| 272 |
+
video_path, device="cpu", dimension_order="NHWC", num_ffmpeg_threads=0
|
| 273 |
+
)
|
| 274 |
+
return decoder.get_frames_at(indices=indices).data.numpy()
|
| 275 |
+
elif video_backend == "ffmpeg":
|
| 276 |
+
return _extract_frames_ffmpeg(video_path, list(indices))
|
| 277 |
+
elif video_backend == "opencv":
|
| 278 |
+
frames = []
|
| 279 |
+
cap = cv2.VideoCapture(video_path, **video_backend_kwargs)
|
| 280 |
+
for idx in indices:
|
| 281 |
+
cap.set(cv2.CAP_PROP_POS_FRAMES, idx)
|
| 282 |
+
ret, frame = cap.read()
|
| 283 |
+
if not ret:
|
| 284 |
+
raise ValueError(f"Unable to read frame at index {idx}")
|
| 285 |
+
frames.append(frame)
|
| 286 |
+
cap.release()
|
| 287 |
+
frames = np.array(frames)
|
| 288 |
+
return frames
|
| 289 |
+
else:
|
| 290 |
+
raise NotImplementedError
|
| 291 |
+
|
| 292 |
+
|
| 293 |
+
def get_frames_by_timestamps(
|
| 294 |
+
video_path: str,
|
| 295 |
+
timestamps: list[float] | np.ndarray,
|
| 296 |
+
video_backend: str = "ffmpeg",
|
| 297 |
+
video_backend_kwargs: dict = {},
|
| 298 |
+
fps: None | float = None,
|
| 299 |
+
) -> np.ndarray:
|
| 300 |
+
"""Get frames from a video at specified timestamps.
|
| 301 |
+
|
| 302 |
+
Args:
|
| 303 |
+
video_path (str): Path to the video file.
|
| 304 |
+
timestamps (list[int] | np.ndarray): Timestamps to retrieve frames for, in seconds.
|
| 305 |
+
video_backend (str, optional): Video backend to use. Defaults to "ffmpeg".
|
| 306 |
+
fps (float, optional): FPS of the video. Defaults to 30.
|
| 307 |
+
Returns:
|
| 308 |
+
np.ndarray: Frames at the specified timestamps.
|
| 309 |
+
"""
|
| 310 |
+
if video_backend == "decord":
|
| 311 |
+
if not DECORD_AVAILABLE:
|
| 312 |
+
raise ImportError("decord is not available. Install it with: pip install decord")
|
| 313 |
+
vr = decord.VideoReader(video_path, **video_backend_kwargs)
|
| 314 |
+
num_frames = len(vr)
|
| 315 |
+
# Retrieve the timestamps for each frame in the video
|
| 316 |
+
frame_ts: np.ndarray = vr.get_frame_timestamp(range(num_frames))
|
| 317 |
+
# Map each requested timestamp to the closest frame index
|
| 318 |
+
# Only take the first element of the frame_ts array which corresponds to start_seconds
|
| 319 |
+
indices = np.abs(frame_ts[:, :1] - timestamps).argmin(axis=0)
|
| 320 |
+
frames = vr.get_batch(indices)
|
| 321 |
+
return frames.asnumpy()
|
| 322 |
+
elif video_backend == "torchcodec":
|
| 323 |
+
if not TORCHCODEC_AVAILABLE:
|
| 324 |
+
raise ImportError("torchcodec is not available.")
|
| 325 |
+
decoder = torchcodec.decoders.VideoDecoder(
|
| 326 |
+
video_path, device="cpu", dimension_order="NHWC", num_ffmpeg_threads=0
|
| 327 |
+
)
|
| 328 |
+
|
| 329 |
+
# https://docs.pytorch.org/torchcodec/stable/generated/torchcodec.decoders.VideoStreamMetadata.html#torchcodec.decoders.VideoStreamMetadata
|
| 330 |
+
# Temporary fix: use 30 fps as the fps of the video (agibot)
|
| 331 |
+
# TODO: get fps as parameter
|
| 332 |
+
if fps is None:
|
| 333 |
+
fps = decoder.metadata.average_fps
|
| 334 |
+
interval = 1 / fps
|
| 335 |
+
timestamps = np.array(timestamps).astype(np.float64)
|
| 336 |
+
|
| 337 |
+
if np.all(timestamps == 0):
|
| 338 |
+
timestamps = np.arange(len(timestamps)) / fps
|
| 339 |
+
|
| 340 |
+
# Get video duration range from first and last frames
|
| 341 |
+
# This is a robust way to get valid timestamp range without depending on specific metadata attributes
|
| 342 |
+
first_frame = decoder.get_frames_at(indices=[0])
|
| 343 |
+
last_frame = decoder.get_frames_at(indices=[len(decoder) - 1])
|
| 344 |
+
min_pts = float(first_frame.pts_seconds[0])
|
| 345 |
+
max_pts = float(last_frame.pts_seconds[0])
|
| 346 |
+
|
| 347 |
+
# Clamp timestamps to valid range to avoid RuntimeError
|
| 348 |
+
timestamps = np.clip(timestamps, min_pts, max_pts)
|
| 349 |
+
|
| 350 |
+
# Correct float precision issues in timestamps
|
| 351 |
+
# E.g. for 5fps video: [1.0, 1.20000005, 1.39999998] -> [1.0, 1.2, 1.4]
|
| 352 |
+
# Without this, the torchcodec will read the delayed frame (e.g. 1.39999998 -> 1.2)
|
| 353 |
+
# Round to nearest frame interval to prevent torchcodec from reading wrong frames
|
| 354 |
+
# Allow max 1% error from expected interval
|
| 355 |
+
if fps is None:
|
| 356 |
+
closest_timestamps = np.round(timestamps / interval) * interval
|
| 357 |
+
# Re-clamp after rounding to ensure still in valid range
|
| 358 |
+
closest_timestamps = np.clip(closest_timestamps, min_pts, max_pts)
|
| 359 |
+
timestamp_errors = np.abs(closest_timestamps - timestamps) / interval
|
| 360 |
+
invalid_mask = timestamp_errors >= 0.01
|
| 361 |
+
if np.any(invalid_mask):
|
| 362 |
+
invalid_indices = np.where(invalid_mask)[0]
|
| 363 |
+
invalid_timestamps = timestamps[invalid_indices]
|
| 364 |
+
raise ValueError(
|
| 365 |
+
f"Try to read invalid timestamps {invalid_timestamps} from video {video_path} (FPS: {fps})"
|
| 366 |
+
)
|
| 367 |
+
|
| 368 |
+
timestamps = closest_timestamps
|
| 369 |
+
|
| 370 |
+
return decoder.get_frames_played_at(seconds=timestamps).data.numpy()
|
| 371 |
+
elif video_backend == "ffmpeg":
|
| 372 |
+
return _extract_frames_at_timestamps_ffmpeg(video_path, list(timestamps))
|
| 373 |
+
elif video_backend == "opencv":
|
| 374 |
+
# Open the video file
|
| 375 |
+
cap = cv2.VideoCapture(video_path, **video_backend_kwargs)
|
| 376 |
+
if not cap.isOpened():
|
| 377 |
+
raise ValueError(f"Unable to open video file: {video_path}")
|
| 378 |
+
# Retrieve the total number of frames
|
| 379 |
+
num_frames = int(cap.get(cv2.CAP_PROP_FRAME_COUNT))
|
| 380 |
+
# Calculate timestamps for each frame
|
| 381 |
+
fps = cap.get(cv2.CAP_PROP_FPS)
|
| 382 |
+
frame_ts = np.arange(num_frames) / fps
|
| 383 |
+
frame_ts = frame_ts[:, np.newaxis] # Reshape to (num_frames, 1) for broadcasting
|
| 384 |
+
# Map each requested timestamp to the closest frame index
|
| 385 |
+
indices = np.abs(frame_ts - timestamps).argmin(axis=0)
|
| 386 |
+
frames = []
|
| 387 |
+
for idx in indices:
|
| 388 |
+
cap.set(cv2.CAP_PROP_POS_FRAMES, idx)
|
| 389 |
+
ret, frame = cap.read()
|
| 390 |
+
if not ret:
|
| 391 |
+
raise ValueError(f"Unable to read frame at index {idx}")
|
| 392 |
+
frames.append(frame)
|
| 393 |
+
cap.release()
|
| 394 |
+
frames = np.array(frames)
|
| 395 |
+
return frames
|
| 396 |
+
|
| 397 |
+
elif video_backend == "torchvision_av":
|
| 398 |
+
# set backend
|
| 399 |
+
torchvision.set_video_backend("pyav")
|
| 400 |
+
|
| 401 |
+
# set a video stream reader
|
| 402 |
+
reader = torchvision.io.VideoReader(video_path, "video")
|
| 403 |
+
|
| 404 |
+
# set the first and last requested timestamps
|
| 405 |
+
# Note: previous timestamps are usually loaded, since we need to access the previous key frame
|
| 406 |
+
first_ts = timestamps[0]
|
| 407 |
+
last_ts = timestamps[-1]
|
| 408 |
+
|
| 409 |
+
# access closest key frame of the first requested frame
|
| 410 |
+
# Note: closest key frame timestamp is usally smaller than `first_ts` (e.g. key frame can be the first frame of the video)
|
| 411 |
+
# for details on what `seek` is doing see: https://pyav.basswood-io.com/docs/stable/api/container.html?highlight=inputcontainer#av.container.InputContainer.seek
|
| 412 |
+
reader.seek(first_ts, keyframes_only=True)
|
| 413 |
+
|
| 414 |
+
# Decode frames sequentially, storing the ones we need in a dictionary
|
| 415 |
+
# to map timestamps to frame data. This allows for easy re-ordering later.
|
| 416 |
+
found_frames_map = {}
|
| 417 |
+
tolerance = 0.001 # 1ms tolerance for timestamp matching
|
| 418 |
+
|
| 419 |
+
for frame in reader:
|
| 420 |
+
current_ts = frame["pts"]
|
| 421 |
+
|
| 422 |
+
# Use tolerance-based matching instead of exact match
|
| 423 |
+
for ts in timestamps:
|
| 424 |
+
if ts not in found_frames_map and abs(current_ts - ts) < tolerance:
|
| 425 |
+
found_frames_map[ts] = frame["data"]
|
| 426 |
+
break
|
| 427 |
+
|
| 428 |
+
if current_ts >= last_ts + tolerance or len(found_frames_map) == len(timestamps):
|
| 429 |
+
break
|
| 430 |
+
|
| 431 |
+
reader.container.close()
|
| 432 |
+
reader = None
|
| 433 |
+
|
| 434 |
+
# Debug: print timestamp matching results
|
| 435 |
+
print(f"[video_utils] Requested {len(timestamps)} timestamps: {timestamps[:4]}{'...' if len(timestamps) > 4 else ''}")
|
| 436 |
+
print(f"[video_utils] Found {len(found_frames_map)} frames with tolerance={tolerance}s")
|
| 437 |
+
if len(found_frames_map) < len(timestamps):
|
| 438 |
+
missing = [ts for ts in timestamps if ts not in found_frames_map]
|
| 439 |
+
print(f"[video_utils] WARNING: Missing timestamps: {missing[:4]}{'...' if len(missing) > 4 else ''}")
|
| 440 |
+
|
| 441 |
+
frames = np.array(list(found_frames_map.values()))
|
| 442 |
+
return frames.transpose(0, 2, 3, 1)
|
| 443 |
+
|
| 444 |
+
else:
|
| 445 |
+
raise NotImplementedError
|
| 446 |
+
|
| 447 |
+
|
| 448 |
+
def get_all_frames(
|
| 449 |
+
video_path: str,
|
| 450 |
+
video_backend: str = "ffmpeg",
|
| 451 |
+
video_backend_kwargs: dict = {},
|
| 452 |
+
) -> tuple[np.ndarray, np.ndarray]:
|
| 453 |
+
"""Get all frames from a video.
|
| 454 |
+
|
| 455 |
+
Returns:
|
| 456 |
+
tuple[np.ndarray, np.ndarray]: Frames and timestamps.
|
| 457 |
+
"""
|
| 458 |
+
if video_backend == "decord":
|
| 459 |
+
if not DECORD_AVAILABLE:
|
| 460 |
+
raise ImportError("decord is not available. Install it with: pip install decord")
|
| 461 |
+
vr = decord.VideoReader(video_path, **video_backend_kwargs)
|
| 462 |
+
frames = vr.get_batch(range(len(vr))).asnumpy()
|
| 463 |
+
return frames, vr.get_frame_timestamp(range(len(vr)))[:, 0]
|
| 464 |
+
elif video_backend == "torchcodec":
|
| 465 |
+
if not TORCHCODEC_AVAILABLE:
|
| 466 |
+
raise ImportError("torchcodec is not available.")
|
| 467 |
+
decoder = torchcodec.decoders.VideoDecoder(
|
| 468 |
+
video_path, device="cpu", dimension_order="NHWC", num_ffmpeg_threads=0
|
| 469 |
+
)
|
| 470 |
+
frames = decoder.get_frames_at(indices=range(len(decoder)))
|
| 471 |
+
return frames.data.numpy(), frames.pts_seconds.numpy()
|
| 472 |
+
elif video_backend == "ffmpeg":
|
| 473 |
+
return _extract_all_frames_ffmpeg(video_path)
|
| 474 |
+
elif video_backend == "pyav":
|
| 475 |
+
container = av.open(video_path)
|
| 476 |
+
stream = container.streams.video[0]
|
| 477 |
+
assert stream.time_base is not None
|
| 478 |
+
frames = []
|
| 479 |
+
timestamps = []
|
| 480 |
+
for frame in container.decode(video=0):
|
| 481 |
+
frames.append(frame.to_ndarray(format="rgb24"))
|
| 482 |
+
timestamps.append(frame.pts * stream.time_base)
|
| 483 |
+
container.close()
|
| 484 |
+
return np.stack(frames), np.array(timestamps)
|
| 485 |
+
|
| 486 |
+
else:
|
| 487 |
+
raise NotImplementedError
|
groot/vla/configs/conf.yaml
ADDED
|
@@ -0,0 +1,135 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
defaults:
|
| 2 |
+
- _self_ # all below configs will override this conf.yaml
|
| 3 |
+
- model: dreamzero/vla
|
| 4 |
+
- data: dreamzero/droid_horizon_relative
|
| 5 |
+
- override hydra/hydra_logging: disabled # disable hydra logging
|
| 6 |
+
- override hydra/job_logging: disabled # disable hydra job logging
|
| 7 |
+
|
| 8 |
+
|
| 9 |
+
# === Model Arguments ===
|
| 10 |
+
model: ???
|
| 11 |
+
|
| 12 |
+
# === Data Arguments ===
|
| 13 |
+
train_dataset: ???
|
| 14 |
+
|
| 15 |
+
# ======== Trainer ========
|
| 16 |
+
trainer:
|
| 17 |
+
_target_: groot.vla.experiment.VLATrainer
|
| 18 |
+
_partial_: true
|
| 19 |
+
_recursive_: false
|
| 20 |
+
callbacks:
|
| 21 |
+
model: ??? # model
|
| 22 |
+
train_dataset: ??? # train_dataset
|
| 23 |
+
compute_dtype: ??? # dtype_from_string(model.config.model_dtype)
|
| 24 |
+
benchmark_time: false # whether or not to benchmark time for training
|
| 25 |
+
# Legacy per-step profiling (profiles every N steps)
|
| 26 |
+
enable_profiling: false # (legacy) enable per-step profiling in training_step
|
| 27 |
+
profiling_steps: 5 # (legacy) profile every N steps
|
| 28 |
+
# ProfCallback: window-based profiling
|
| 29 |
+
enable_prof_callback: false # enable ProfCallback for window-based profiling
|
| 30 |
+
profile_start_step: 50 # session step to start profiling
|
| 31 |
+
profile_warmup_steps: 1 # warmup steps for profiler
|
| 32 |
+
profile_active_steps: 3 # active profiling steps
|
| 33 |
+
profile_record_shapes: false # record tensor shapes (adds overhead)
|
| 34 |
+
profile_with_stack: false # record Python stack traces
|
| 35 |
+
profile_memory: false # record memory allocation
|
| 36 |
+
|
| 37 |
+
# === Training Arguments ===
|
| 38 |
+
|
| 39 |
+
wandb_project: ??? # needs to be specified by user
|
| 40 |
+
output_dir: ??? # need to be specified by user
|
| 41 |
+
load_from_yaml: # need to be specified by user, will override the current config
|
| 42 |
+
gear_credentials: null
|
| 43 |
+
upload_checkpoints: false
|
| 44 |
+
upload_every: 1000
|
| 45 |
+
upload_last_n_checkpoints: 5
|
| 46 |
+
remove_unused_columns: false
|
| 47 |
+
bf16: false
|
| 48 |
+
tf32: false
|
| 49 |
+
global_batch_size: null
|
| 50 |
+
raise_error_if_global_batch_size_not_set: false
|
| 51 |
+
per_device_train_batch_size: 256
|
| 52 |
+
per_device_eval_batch_size: 64
|
| 53 |
+
gradient_accumulation_steps: 1
|
| 54 |
+
dataloader_num_workers: 10
|
| 55 |
+
dataloader_pin_memory: true
|
| 56 |
+
dataloader_persistent_workers: true
|
| 57 |
+
optim: adamw_torch
|
| 58 |
+
learning_rate: 1e-4
|
| 59 |
+
adam_beta1: 0.95
|
| 60 |
+
adam_beta2: 0.999
|
| 61 |
+
adam_epsilon: 1e-8
|
| 62 |
+
weight_decay: 1e-6
|
| 63 |
+
lr_scheduler_type: cosine
|
| 64 |
+
warmup_ratio: 0.05
|
| 65 |
+
logging_steps: 10.0
|
| 66 |
+
num_train_epochs: 1000
|
| 67 |
+
max_steps: -1
|
| 68 |
+
save_strategy: steps
|
| 69 |
+
save_steps: 500
|
| 70 |
+
eval_strategy: "no" # there has to be a double quote; otherwise a bare `no` will be interpreted as False
|
| 71 |
+
save_total_limit: 8
|
| 72 |
+
report_to: wandb
|
| 73 |
+
seed: 42
|
| 74 |
+
do_eval: false
|
| 75 |
+
gradient_checkpointing: false
|
| 76 |
+
ddp_find_unused_parameters: false
|
| 77 |
+
ddp_bucket_cap_mb: 100
|
| 78 |
+
ray_num_workers: ???
|
| 79 |
+
eval_bf16: true
|
| 80 |
+
torch_compile_mode: null
|
| 81 |
+
|
| 82 |
+
pretrained_model_path: null
|
| 83 |
+
only_tune_projectors: false
|
| 84 |
+
|
| 85 |
+
save_llm: false
|
| 86 |
+
save_lora_only: false
|
| 87 |
+
save_value_model: false
|
| 88 |
+
save_q_model: false
|
| 89 |
+
|
| 90 |
+
download_cache: false
|
| 91 |
+
|
| 92 |
+
training_args:
|
| 93 |
+
_target_: transformers.TrainingArguments
|
| 94 |
+
output_dir: ${output_dir}
|
| 95 |
+
run_name: ??? # training_args.output_dir.split("/")[-1]
|
| 96 |
+
remove_unused_columns: ${remove_unused_columns}
|
| 97 |
+
deepspeed: ""
|
| 98 |
+
gradient_checkpointing: ${gradient_checkpointing}
|
| 99 |
+
bf16: ${bf16}
|
| 100 |
+
tf32: ${tf32}
|
| 101 |
+
per_device_train_batch_size: ${per_device_train_batch_size}
|
| 102 |
+
per_device_eval_batch_size: ${per_device_eval_batch_size}
|
| 103 |
+
gradient_accumulation_steps: ${gradient_accumulation_steps}
|
| 104 |
+
dataloader_num_workers: ${dataloader_num_workers}
|
| 105 |
+
dataloader_pin_memory: ${dataloader_pin_memory}
|
| 106 |
+
dataloader_persistent_workers: ${dataloader_persistent_workers}
|
| 107 |
+
optim: ${optim}
|
| 108 |
+
adam_beta1: ${adam_beta1}
|
| 109 |
+
adam_beta2: ${adam_beta2}
|
| 110 |
+
adam_epsilon: ${adam_epsilon}
|
| 111 |
+
learning_rate: ${learning_rate}
|
| 112 |
+
weight_decay: ${weight_decay}
|
| 113 |
+
warmup_ratio: ${warmup_ratio}
|
| 114 |
+
lr_scheduler_type: ${lr_scheduler_type}
|
| 115 |
+
logging_steps: ${logging_steps}
|
| 116 |
+
num_train_epochs: ${num_train_epochs}
|
| 117 |
+
max_steps: ${max_steps}
|
| 118 |
+
save_strategy: ${save_strategy}
|
| 119 |
+
save_steps: ${save_steps}
|
| 120 |
+
save_total_limit: ${save_total_limit}
|
| 121 |
+
report_to: ${report_to}
|
| 122 |
+
seed: ${seed}
|
| 123 |
+
do_eval: ${do_eval}
|
| 124 |
+
ddp_find_unused_parameters: ${ddp_find_unused_parameters}
|
| 125 |
+
ddp_bucket_cap_mb: ${ddp_bucket_cap_mb}
|
| 126 |
+
torch_compile_mode: ${torch_compile_mode}
|
| 127 |
+
|
| 128 |
+
# === Profiling Arguments ===
|
| 129 |
+
profile_dir: null
|
| 130 |
+
|
| 131 |
+
# === Disable Hydra Config ===
|
| 132 |
+
hydra:
|
| 133 |
+
output_subdir: null
|
| 134 |
+
run:
|
| 135 |
+
dir: .
|
groot/vla/configs/data/dreamzero/agibot_relative.yaml
ADDED
|
@@ -0,0 +1,54 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# @package _global_
|
| 2 |
+
|
| 3 |
+
defaults:
|
| 4 |
+
- dreamzero/base_48_wan_fine_aug_relative
|
| 5 |
+
- _self_ # this file will override the base
|
| 6 |
+
|
| 7 |
+
max_state_dim: 64
|
| 8 |
+
use_global_metadata: false
|
| 9 |
+
relative_action: true
|
| 10 |
+
relative_action_per_horizon: false
|
| 11 |
+
relative_action_keys:
|
| 12 |
+
- left_arm_joint_position
|
| 13 |
+
- right_arm_joint_position
|
| 14 |
+
- left_effector_position
|
| 15 |
+
- right_effector_position
|
| 16 |
+
- head_position
|
| 17 |
+
- waist_position
|
| 18 |
+
max_chunk_size: 5
|
| 19 |
+
# Use 10% of data in shards before moving to next shard
|
| 20 |
+
dataset_shard_sampling_rate: 0.1
|
| 21 |
+
mixture_dataset_cls: groot.vla.data.dataset.lerobot_sharded.ShardedLeRobotMixtureDataset.from_mixture_spec
|
| 22 |
+
single_dataset_cls: groot.vla.data.dataset.lerobot_sharded.ShardedLeRobotSubLangSingleActionChunkDatasetDROID
|
| 23 |
+
|
| 24 |
+
# Set your AGIbot dataset path here or override via CLI:
|
| 25 |
+
# agibot_data_root=/path/to/your/agibot_dataset
|
| 26 |
+
agibot_data_root: ???
|
| 27 |
+
|
| 28 |
+
train_dataset:
|
| 29 |
+
_target_: ${mixture_dataset_cls}
|
| 30 |
+
_convert_: object
|
| 31 |
+
mixture_spec:
|
| 32 |
+
- dataset_path:
|
| 33 |
+
agibot:
|
| 34 |
+
- ${agibot_data_root}
|
| 35 |
+
dataset_weight: 1.0
|
| 36 |
+
distribute_weights: true
|
| 37 |
+
|
| 38 |
+
dataset_class: ${single_dataset_cls}
|
| 39 |
+
all_modality_configs: ${modality_configs}
|
| 40 |
+
all_transforms: ${transforms}
|
| 41 |
+
metadata_versions: ${metadata_versions}
|
| 42 |
+
fps: ${fps}
|
| 43 |
+
dataset_kwargs:
|
| 44 |
+
video_backend: decord
|
| 45 |
+
use_global_metadata: ${use_global_metadata}
|
| 46 |
+
max_chunk_size: ${max_chunk_size}
|
| 47 |
+
relative_action: ${relative_action}
|
| 48 |
+
relative_action_keys: ${relative_action_keys}
|
| 49 |
+
relative_action_per_horizon: ${relative_action_per_horizon}
|
| 50 |
+
mixture_kwargs:
|
| 51 |
+
training: true
|
| 52 |
+
balance_dataset_weights: false
|
| 53 |
+
seed: 42
|
| 54 |
+
shard_sampling_rate: ${dataset_shard_sampling_rate}
|
groot/vla/configs/data/dreamzero/base_48_wan_fine_aug_relative.yaml
ADDED
|
@@ -0,0 +1,439 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# @package _global_
|
| 2 |
+
|
| 3 |
+
# Assume `model_specific_transform` is defined in the model config file
|
| 4 |
+
|
| 5 |
+
################################################################################
|
| 6 |
+
# Normalization Statistics
|
| 7 |
+
# By default, we compute the normalization statistics for the datasets actually
|
| 8 |
+
# used in the mixture. If you want to use the global metadata, set this to true.
|
| 9 |
+
################################################################################
|
| 10 |
+
|
| 11 |
+
use_global_metadata: false
|
| 12 |
+
|
| 13 |
+
################################################################################
|
| 14 |
+
# Dimension Information
|
| 15 |
+
################################################################################
|
| 16 |
+
|
| 17 |
+
num_frames: 49
|
| 18 |
+
action_horizon: 48
|
| 19 |
+
state_horizon: 1
|
| 20 |
+
|
| 21 |
+
# image_resolution_width: 832
|
| 22 |
+
# image_resolution_height: 480
|
| 23 |
+
|
| 24 |
+
image_resolution_width: 480
|
| 25 |
+
image_resolution_height: 256
|
| 26 |
+
|
| 27 |
+
image_resolution_width_single_frame: 256
|
| 28 |
+
image_resolution_height_single_frame: 256
|
| 29 |
+
|
| 30 |
+
################################################################################
|
| 31 |
+
# Anchored Video Transforms
|
| 32 |
+
################################################################################
|
| 33 |
+
totensor_cfg: &totensor_cfg
|
| 34 |
+
_target_: groot.vla.data.transform.VideoToTensor
|
| 35 |
+
apply_to: ???
|
| 36 |
+
|
| 37 |
+
crop_cfg: &crop_cfg
|
| 38 |
+
_target_: groot.vla.data.transform.VideoCrop
|
| 39 |
+
apply_to: ???
|
| 40 |
+
scale: 0.95
|
| 41 |
+
mode: random
|
| 42 |
+
|
| 43 |
+
|
| 44 |
+
resize_cfg: &resize_cfg
|
| 45 |
+
_target_: groot.vla.data.transform.VideoResize
|
| 46 |
+
apply_to: ???
|
| 47 |
+
height: ${image_resolution_height}
|
| 48 |
+
width: ${image_resolution_width}
|
| 49 |
+
interpolation: linear
|
| 50 |
+
|
| 51 |
+
resize_cfg_single_frame: &resize_cfg_single_frame
|
| 52 |
+
_target_: groot.vla.data.transform.VideoResize
|
| 53 |
+
apply_to: ???
|
| 54 |
+
height: ${image_resolution_height_single_frame}
|
| 55 |
+
width: ${image_resolution_width_single_frame}
|
| 56 |
+
interpolation: linear
|
| 57 |
+
|
| 58 |
+
color_jitter_cfg: &color_jitter_cfg
|
| 59 |
+
_target_: groot.vla.data.transform.VideoColorJitter
|
| 60 |
+
apply_to: ???
|
| 61 |
+
brightness: 0.3
|
| 62 |
+
contrast: 0.4
|
| 63 |
+
saturation: 0.5
|
| 64 |
+
hue: 0.08
|
| 65 |
+
|
| 66 |
+
random_grayscale_cfg: &random_grayscale_cfg
|
| 67 |
+
_target_: groot.vla.data.transform.VideoRandomGrayscale
|
| 68 |
+
apply_to: ???
|
| 69 |
+
p: 0.1
|
| 70 |
+
|
| 71 |
+
random_posterize_cfg: &random_posterize_cfg
|
| 72 |
+
_target_: groot.vla.data.transform.VideoRandomPosterize
|
| 73 |
+
apply_to: ???
|
| 74 |
+
bits: 4
|
| 75 |
+
p: 0.1
|
| 76 |
+
|
| 77 |
+
normalize_cfg: &normalize_cfg
|
| 78 |
+
_target_: groot.vla.data.transform.VideoNormalize
|
| 79 |
+
apply_to: ???
|
| 80 |
+
mean: [0.5, 0.5, 0.5]
|
| 81 |
+
std: [0.5, 0.5, 0.5]
|
| 82 |
+
|
| 83 |
+
|
| 84 |
+
to_numpy_cfg: &to_numpy_cfg
|
| 85 |
+
_target_: groot.vla.data.transform.VideoToNumpy
|
| 86 |
+
apply_to: ???
|
| 87 |
+
|
| 88 |
+
|
| 89 |
+
################################################################################
|
| 90 |
+
# oxe_droid (OXE Droid)
|
| 91 |
+
################################################################################
|
| 92 |
+
|
| 93 |
+
# Modality Configs
|
| 94 |
+
modality_config_oxe_droid:
|
| 95 |
+
video:
|
| 96 |
+
_target_: groot.vla.data.dataset.ModalityConfig
|
| 97 |
+
delta_indices: [0, 1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 12, 13, 14, 15, 16, 17, 18, 19, 20, 21, 22, 23, 24]
|
| 98 |
+
eval_delta_indices: [0]
|
| 99 |
+
modality_keys:
|
| 100 |
+
- video.exterior_image_1_left
|
| 101 |
+
- video.exterior_image_2_left
|
| 102 |
+
- video.wrist_image_left
|
| 103 |
+
state:
|
| 104 |
+
_target_: groot.vla.data.dataset.ModalityConfig
|
| 105 |
+
delta_indices: [0]
|
| 106 |
+
modality_keys:
|
| 107 |
+
- state.joint_position
|
| 108 |
+
- state.gripper_position
|
| 109 |
+
action:
|
| 110 |
+
_target_: groot.vla.data.dataset.ModalityConfig
|
| 111 |
+
delta_indices: [0, 1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 12, 13, 14, 15, 16, 17, 18, 19, 20, 21, 22, 23]
|
| 112 |
+
modality_keys:
|
| 113 |
+
- action.joint_position
|
| 114 |
+
- action.gripper_position
|
| 115 |
+
language:
|
| 116 |
+
_target_: groot.vla.data.dataset.ModalityConfig
|
| 117 |
+
delta_indices: [0]
|
| 118 |
+
modality_keys:
|
| 119 |
+
- annotation.language.language_instruction
|
| 120 |
+
- annotation.language.language_instruction_2
|
| 121 |
+
- annotation.language.language_instruction_3
|
| 122 |
+
lapa_action:
|
| 123 |
+
_target_: groot.vla.data.dataset.ModalityConfig
|
| 124 |
+
delta_indices: [0]
|
| 125 |
+
modality_keys:
|
| 126 |
+
- lapa_action
|
| 127 |
+
|
| 128 |
+
# Transforms
|
| 129 |
+
transform_oxe_droid:
|
| 130 |
+
_target_: groot.vla.data.transform.ComposedModalityTransform
|
| 131 |
+
transforms:
|
| 132 |
+
# Video transforms
|
| 133 |
+
- <<: *totensor_cfg
|
| 134 |
+
apply_to: ${modality_config_oxe_droid.video.modality_keys}
|
| 135 |
+
- <<: *crop_cfg
|
| 136 |
+
apply_to: ${modality_config_oxe_droid.video.modality_keys}
|
| 137 |
+
- <<: *resize_cfg
|
| 138 |
+
apply_to: ${modality_config_oxe_droid.video.modality_keys}
|
| 139 |
+
- <<: *color_jitter_cfg
|
| 140 |
+
apply_to: ${modality_config_oxe_droid.video.modality_keys}
|
| 141 |
+
- <<: *to_numpy_cfg
|
| 142 |
+
apply_to: ${modality_config_oxe_droid.video.modality_keys}
|
| 143 |
+
|
| 144 |
+
# State transforms
|
| 145 |
+
- _target_: groot.vla.data.transform.StateActionToTensor
|
| 146 |
+
apply_to: ${modality_config_oxe_droid.state.modality_keys}
|
| 147 |
+
- _target_: groot.vla.data.transform.StateActionTransform
|
| 148 |
+
apply_to: ${modality_config_oxe_droid.state.modality_keys}
|
| 149 |
+
normalization_modes:
|
| 150 |
+
state.joint_position: q99
|
| 151 |
+
state.gripper_position: q99
|
| 152 |
+
|
| 153 |
+
# Action transforms
|
| 154 |
+
- _target_: groot.vla.data.transform.StateActionToTensor
|
| 155 |
+
apply_to: ${modality_config_oxe_droid.action.modality_keys}
|
| 156 |
+
- _target_: groot.vla.data.transform.StateActionTransform
|
| 157 |
+
apply_to: ${modality_config_oxe_droid.action.modality_keys}
|
| 158 |
+
normalization_modes:
|
| 159 |
+
action.joint_position: q99
|
| 160 |
+
action.gripper_position: q99
|
| 161 |
+
|
| 162 |
+
# ConcatTransform
|
| 163 |
+
- _target_: groot.vla.data.transform.ConcatTransform
|
| 164 |
+
video_concat_order: ${modality_config_oxe_droid.video.modality_keys}
|
| 165 |
+
state_concat_order: ${modality_config_oxe_droid.state.modality_keys}
|
| 166 |
+
action_concat_order: ${modality_config_oxe_droid.action.modality_keys}
|
| 167 |
+
|
| 168 |
+
# Model-specific transform
|
| 169 |
+
- ${model_specific_transform}
|
| 170 |
+
|
| 171 |
+
################################################################################
|
| 172 |
+
# agibot (AGIbot: state 32, action 22, 3 views)
|
| 173 |
+
################################################################################
|
| 174 |
+
|
| 175 |
+
modality_config_agibot:
|
| 176 |
+
video:
|
| 177 |
+
_target_: groot.vla.data.dataset.ModalityConfig
|
| 178 |
+
delta_indices: [0, 1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 12, 13, 14, 15, 16, 17, 18, 19, 20, 21, 22, 23, 24]
|
| 179 |
+
eval_delta_indices: [-3, -2, -1, 0]
|
| 180 |
+
modality_keys:
|
| 181 |
+
- video.top_head
|
| 182 |
+
- video.hand_left
|
| 183 |
+
- video.hand_right
|
| 184 |
+
state:
|
| 185 |
+
_target_: groot.vla.data.dataset.ModalityConfig
|
| 186 |
+
delta_indices: [0]
|
| 187 |
+
modality_keys:
|
| 188 |
+
- state.left_arm_joint_position
|
| 189 |
+
- state.right_arm_joint_position
|
| 190 |
+
- state.left_effector_position
|
| 191 |
+
- state.right_effector_position
|
| 192 |
+
- state.head_position
|
| 193 |
+
- state.waist_position
|
| 194 |
+
action:
|
| 195 |
+
_target_: groot.vla.data.dataset.ModalityConfig
|
| 196 |
+
delta_indices: [0, 1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 12, 13, 14, 15, 16, 17, 18, 19, 20, 21, 22, 23]
|
| 197 |
+
modality_keys:
|
| 198 |
+
- action.left_arm_joint_position
|
| 199 |
+
- action.right_arm_joint_position
|
| 200 |
+
- action.left_effector_position
|
| 201 |
+
- action.right_effector_position
|
| 202 |
+
- action.head_position
|
| 203 |
+
- action.waist_position
|
| 204 |
+
- action.robot_velocity
|
| 205 |
+
language:
|
| 206 |
+
_target_: groot.vla.data.dataset.ModalityConfig
|
| 207 |
+
delta_indices: [0]
|
| 208 |
+
modality_keys:
|
| 209 |
+
- annotation.language.action_text
|
| 210 |
+
|
| 211 |
+
transform_agibot:
|
| 212 |
+
_target_: groot.vla.data.transform.ComposedModalityTransform
|
| 213 |
+
transforms:
|
| 214 |
+
# Video transforms
|
| 215 |
+
- <<: *totensor_cfg
|
| 216 |
+
apply_to: ${modality_config_agibot.video.modality_keys}
|
| 217 |
+
- <<: *crop_cfg
|
| 218 |
+
apply_to: ${modality_config_agibot.video.modality_keys}
|
| 219 |
+
- <<: *resize_cfg
|
| 220 |
+
apply_to: ${modality_config_agibot.video.modality_keys}
|
| 221 |
+
- <<: *color_jitter_cfg
|
| 222 |
+
apply_to: ${modality_config_agibot.video.modality_keys}
|
| 223 |
+
- <<: *to_numpy_cfg
|
| 224 |
+
apply_to: ${modality_config_agibot.video.modality_keys}
|
| 225 |
+
|
| 226 |
+
# State transforms
|
| 227 |
+
- _target_: groot.vla.data.transform.StateActionToTensor
|
| 228 |
+
apply_to: ${modality_config_agibot.state.modality_keys}
|
| 229 |
+
- _target_: groot.vla.data.transform.StateActionTransform
|
| 230 |
+
apply_to: ${modality_config_agibot.state.modality_keys}
|
| 231 |
+
normalization_modes:
|
| 232 |
+
state.left_arm_joint_position: q99
|
| 233 |
+
state.right_arm_joint_position: q99
|
| 234 |
+
state.left_effector_position: q99
|
| 235 |
+
state.right_effector_position: q99
|
| 236 |
+
state.head_position: q99
|
| 237 |
+
state.waist_position: q99
|
| 238 |
+
|
| 239 |
+
# Action transforms
|
| 240 |
+
- _target_: groot.vla.data.transform.StateActionToTensor
|
| 241 |
+
apply_to: ${modality_config_agibot.action.modality_keys}
|
| 242 |
+
- _target_: groot.vla.data.transform.StateActionTransform
|
| 243 |
+
apply_to: ${modality_config_agibot.action.modality_keys}
|
| 244 |
+
normalization_modes:
|
| 245 |
+
action.left_arm_joint_position: q99
|
| 246 |
+
action.right_arm_joint_position: q99
|
| 247 |
+
action.left_effector_position: q99
|
| 248 |
+
action.right_effector_position: q99
|
| 249 |
+
action.head_position: q99
|
| 250 |
+
action.waist_position: q99
|
| 251 |
+
action.robot_velocity: q99
|
| 252 |
+
|
| 253 |
+
# ConcatTransform
|
| 254 |
+
- _target_: groot.vla.data.transform.ConcatTransform
|
| 255 |
+
video_concat_order: ${modality_config_agibot.video.modality_keys}
|
| 256 |
+
state_concat_order: ${modality_config_agibot.state.modality_keys}
|
| 257 |
+
action_concat_order: ${modality_config_agibot.action.modality_keys}
|
| 258 |
+
|
| 259 |
+
# Model-specific transform
|
| 260 |
+
- ${model_specific_transform}
|
| 261 |
+
|
| 262 |
+
################################################################################
|
| 263 |
+
# yam (YAM: joint+gripper only from Dataset/YAM_play_data/meta/modality.json;
|
| 264 |
+
# state 14 dims, action 14 dims, 3 views)
|
| 265 |
+
################################################################################
|
| 266 |
+
|
| 267 |
+
modality_config_yam:
|
| 268 |
+
video:
|
| 269 |
+
_target_: groot.vla.data.dataset.ModalityConfig
|
| 270 |
+
delta_indices: [0, 1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 12, 13, 14, 15, 16, 17, 18, 19, 20, 21, 22, 23, 24]
|
| 271 |
+
eval_delta_indices: [0]
|
| 272 |
+
modality_keys:
|
| 273 |
+
- video.top_camera-images-rgb
|
| 274 |
+
- video.left_camera-images-rgb
|
| 275 |
+
- video.right_camera-images-rgb
|
| 276 |
+
state:
|
| 277 |
+
_target_: groot.vla.data.dataset.ModalityConfig
|
| 278 |
+
delta_indices: [0]
|
| 279 |
+
modality_keys:
|
| 280 |
+
- state.left_joint_pos
|
| 281 |
+
- state.left_gripper_pos
|
| 282 |
+
- state.right_joint_pos
|
| 283 |
+
- state.right_gripper_pos
|
| 284 |
+
action:
|
| 285 |
+
_target_: groot.vla.data.dataset.ModalityConfig
|
| 286 |
+
delta_indices: [0, 1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 12, 13, 14, 15, 16, 17, 18, 19, 20, 21, 22, 23]
|
| 287 |
+
modality_keys:
|
| 288 |
+
- action.left_joint_pos
|
| 289 |
+
- action.left_gripper_pos
|
| 290 |
+
- action.right_joint_pos
|
| 291 |
+
- action.right_gripper_pos
|
| 292 |
+
language:
|
| 293 |
+
_target_: groot.vla.data.dataset.ModalityConfig
|
| 294 |
+
delta_indices: [0]
|
| 295 |
+
modality_keys:
|
| 296 |
+
- annotation.task
|
| 297 |
+
|
| 298 |
+
transform_yam:
|
| 299 |
+
_target_: groot.vla.data.transform.ComposedModalityTransform
|
| 300 |
+
transforms:
|
| 301 |
+
# Video transforms
|
| 302 |
+
- <<: *totensor_cfg
|
| 303 |
+
apply_to: ${modality_config_yam.video.modality_keys}
|
| 304 |
+
- <<: *crop_cfg
|
| 305 |
+
apply_to: ${modality_config_yam.video.modality_keys}
|
| 306 |
+
- <<: *resize_cfg
|
| 307 |
+
apply_to: ${modality_config_yam.video.modality_keys}
|
| 308 |
+
- <<: *color_jitter_cfg
|
| 309 |
+
apply_to: ${modality_config_yam.video.modality_keys}
|
| 310 |
+
- <<: *to_numpy_cfg
|
| 311 |
+
apply_to: ${modality_config_yam.video.modality_keys}
|
| 312 |
+
|
| 313 |
+
# State transforms
|
| 314 |
+
- _target_: groot.vla.data.transform.StateActionToTensor
|
| 315 |
+
apply_to: ${modality_config_yam.state.modality_keys}
|
| 316 |
+
- _target_: groot.vla.data.transform.StateActionTransform
|
| 317 |
+
apply_to: ${modality_config_yam.state.modality_keys}
|
| 318 |
+
normalization_modes:
|
| 319 |
+
state.left_joint_pos: q99
|
| 320 |
+
state.left_gripper_pos: q99
|
| 321 |
+
state.right_joint_pos: q99
|
| 322 |
+
state.right_gripper_pos: q99
|
| 323 |
+
|
| 324 |
+
# Action transforms
|
| 325 |
+
- _target_: groot.vla.data.transform.StateActionToTensor
|
| 326 |
+
apply_to: ${modality_config_yam.action.modality_keys}
|
| 327 |
+
- _target_: groot.vla.data.transform.StateActionTransform
|
| 328 |
+
apply_to: ${modality_config_yam.action.modality_keys}
|
| 329 |
+
normalization_modes:
|
| 330 |
+
action.left_joint_pos: q99
|
| 331 |
+
action.left_gripper_pos: q99
|
| 332 |
+
action.right_joint_pos: q99
|
| 333 |
+
action.right_gripper_pos: q99
|
| 334 |
+
|
| 335 |
+
# ConcatTransform
|
| 336 |
+
- _target_: groot.vla.data.transform.ConcatTransform
|
| 337 |
+
video_concat_order: ${modality_config_yam.video.modality_keys}
|
| 338 |
+
state_concat_order: ${modality_config_yam.state.modality_keys}
|
| 339 |
+
action_concat_order: ${modality_config_yam.action.modality_keys}
|
| 340 |
+
|
| 341 |
+
# Model-specific transform
|
| 342 |
+
- ${model_specific_transform}
|
| 343 |
+
|
| 344 |
+
|
| 345 |
+
# libero_sim (LIBERO: 1 camera, 8-dim state, 7-dim action)
|
| 346 |
+
modality_config_libero_sim:
|
| 347 |
+
video:
|
| 348 |
+
_target_: groot.vla.data.dataset.ModalityConfig
|
| 349 |
+
delta_indices: [0, 1, 2, 3, 4, 5, 6, 7]
|
| 350 |
+
modality_keys:
|
| 351 |
+
- video.image
|
| 352 |
+
state:
|
| 353 |
+
_target_: groot.vla.data.dataset.ModalityConfig
|
| 354 |
+
delta_indices: [0]
|
| 355 |
+
modality_keys:
|
| 356 |
+
- state.joint_pos
|
| 357 |
+
- state.gripper_pos
|
| 358 |
+
action:
|
| 359 |
+
_target_: groot.vla.data.dataset.ModalityConfig
|
| 360 |
+
delta_indices: [0, 1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11]
|
| 361 |
+
modality_keys:
|
| 362 |
+
- action.joint_pos
|
| 363 |
+
language:
|
| 364 |
+
_target_: groot.vla.data.dataset.ModalityConfig
|
| 365 |
+
delta_indices: [0]
|
| 366 |
+
modality_keys:
|
| 367 |
+
- annotation.task
|
| 368 |
+
|
| 369 |
+
transform_libero_sim:
|
| 370 |
+
_target_: groot.vla.data.transform.ComposedModalityTransform
|
| 371 |
+
transforms:
|
| 372 |
+
- <<: *totensor_cfg
|
| 373 |
+
apply_to: ${modality_config_libero_sim.video.modality_keys}
|
| 374 |
+
- <<: *crop_cfg
|
| 375 |
+
apply_to: ${modality_config_libero_sim.video.modality_keys}
|
| 376 |
+
- <<: *resize_cfg
|
| 377 |
+
apply_to: ${modality_config_libero_sim.video.modality_keys}
|
| 378 |
+
- <<: *color_jitter_cfg
|
| 379 |
+
apply_to: ${modality_config_libero_sim.video.modality_keys}
|
| 380 |
+
- <<: *to_numpy_cfg
|
| 381 |
+
apply_to: ${modality_config_libero_sim.video.modality_keys}
|
| 382 |
+
- _target_: groot.vla.data.transform.StateActionToTensor
|
| 383 |
+
apply_to: ${modality_config_libero_sim.state.modality_keys}
|
| 384 |
+
- _target_: groot.vla.data.transform.StateActionTransform
|
| 385 |
+
apply_to: ${modality_config_libero_sim.state.modality_keys}
|
| 386 |
+
normalization_modes:
|
| 387 |
+
state.joint_pos: q99
|
| 388 |
+
state.gripper_pos: q99
|
| 389 |
+
- _target_: groot.vla.data.transform.StateActionToTensor
|
| 390 |
+
apply_to: ${modality_config_libero_sim.action.modality_keys}
|
| 391 |
+
- _target_: groot.vla.data.transform.StateActionTransform
|
| 392 |
+
apply_to: ${modality_config_libero_sim.action.modality_keys}
|
| 393 |
+
normalization_modes:
|
| 394 |
+
action.joint_pos: q99
|
| 395 |
+
- _target_: groot.vla.data.transform.ConcatTransform
|
| 396 |
+
video_concat_order: ${modality_config_libero_sim.video.modality_keys}
|
| 397 |
+
state_concat_order: ${modality_config_libero_sim.state.modality_keys}
|
| 398 |
+
action_concat_order: ${modality_config_libero_sim.action.modality_keys}
|
| 399 |
+
- ${model_specific_transform}
|
| 400 |
+
|
| 401 |
+
|
| 402 |
+
|
| 403 |
+
################################################################################
|
| 404 |
+
# Modality Configs
|
| 405 |
+
################################################################################
|
| 406 |
+
|
| 407 |
+
modality_configs:
|
| 408 |
+
oxe_droid: ${modality_config_oxe_droid}
|
| 409 |
+
agibot: ${modality_config_agibot}
|
| 410 |
+
yam: ${modality_config_yam}
|
| 411 |
+
libero_sim: ${modality_config_libero_sim}
|
| 412 |
+
|
| 413 |
+
################################################################################
|
| 414 |
+
# Transforms
|
| 415 |
+
################################################################################
|
| 416 |
+
|
| 417 |
+
transforms:
|
| 418 |
+
oxe_droid: ${transform_oxe_droid}
|
| 419 |
+
agibot: ${transform_agibot}
|
| 420 |
+
yam: ${transform_yam}
|
| 421 |
+
libero_sim: ${transform_libero_sim}
|
| 422 |
+
|
| 423 |
+
################################################################################
|
| 424 |
+
# Metadata Versions
|
| 425 |
+
################################################################################
|
| 426 |
+
|
| 427 |
+
metadata_versions:
|
| 428 |
+
oxe_droid: '0221'
|
| 429 |
+
agibot: '0221'
|
| 430 |
+
yam: '0221'
|
| 431 |
+
libero_sim: 0221
|
| 432 |
+
|
| 433 |
+
################################################################################
|
| 434 |
+
# FPS (per embodiment, null means use dataset default)
|
| 435 |
+
################################################################################
|
| 436 |
+
|
| 437 |
+
fps:
|
| 438 |
+
yam: 30
|
| 439 |
+
libero_sim: 10
|
groot/vla/configs/data/dreamzero/droid_relative.yaml
ADDED
|
@@ -0,0 +1,49 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# @package _global_
|
| 2 |
+
|
| 3 |
+
defaults:
|
| 4 |
+
- dreamzero/base_48_wan_fine_aug_relative
|
| 5 |
+
- _self_ # this file will override the base
|
| 6 |
+
|
| 7 |
+
max_state_dim: 64
|
| 8 |
+
use_global_metadata: false
|
| 9 |
+
relative_action: true
|
| 10 |
+
relative_action_per_horizon: false
|
| 11 |
+
relative_action_keys:
|
| 12 |
+
- joint_position
|
| 13 |
+
max_chunk_size: 5
|
| 14 |
+
# Use 10% of data in shards before moving to next shard
|
| 15 |
+
dataset_shard_sampling_rate: 0.1
|
| 16 |
+
mixture_dataset_cls: groot.vla.data.dataset.lerobot_sharded.ShardedLeRobotMixtureDataset.from_mixture_spec
|
| 17 |
+
single_dataset_cls: groot.vla.data.dataset.lerobot_sharded.ShardedLeRobotSubLangSingleActionChunkDatasetDROID
|
| 18 |
+
|
| 19 |
+
# Set your DROID dataset path here or override via CLI:
|
| 20 |
+
# droid_data_root=/path/to/your/droid_dataset
|
| 21 |
+
droid_data_root: ???
|
| 22 |
+
|
| 23 |
+
train_dataset:
|
| 24 |
+
_target_: ${mixture_dataset_cls}
|
| 25 |
+
_convert_: object
|
| 26 |
+
mixture_spec:
|
| 27 |
+
- dataset_path:
|
| 28 |
+
oxe_droid:
|
| 29 |
+
- ${droid_data_root}
|
| 30 |
+
dataset_weight: 1.0
|
| 31 |
+
distribute_weights: true
|
| 32 |
+
|
| 33 |
+
dataset_class: ${single_dataset_cls}
|
| 34 |
+
all_modality_configs: ${modality_configs}
|
| 35 |
+
all_transforms: ${transforms}
|
| 36 |
+
metadata_versions: ${metadata_versions}
|
| 37 |
+
fps: ${fps}
|
| 38 |
+
dataset_kwargs:
|
| 39 |
+
video_backend: decord
|
| 40 |
+
use_global_metadata: ${use_global_metadata}
|
| 41 |
+
max_chunk_size: ${max_chunk_size}
|
| 42 |
+
relative_action: ${relative_action}
|
| 43 |
+
relative_action_keys: ${relative_action_keys}
|
| 44 |
+
relative_action_per_horizon: ${relative_action_per_horizon}
|
| 45 |
+
mixture_kwargs:
|
| 46 |
+
training: true
|
| 47 |
+
balance_dataset_weights: false
|
| 48 |
+
seed: 42
|
| 49 |
+
shard_sampling_rate: ${dataset_shard_sampling_rate}
|
groot/vla/configs/data/dreamzero/droid_relative_wan22.yaml
ADDED
|
@@ -0,0 +1,51 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# @package _global_
|
| 2 |
+
# DROID data config for Wan 5B (Wan2.2): 320x160 so latent is 20x10 (even H,W) with WanVideoVAE38 (16x), frame_seqlen=50.
|
| 3 |
+
# Extends base and droid settings directly to avoid Hydra nesting (dreamzero/droid_relative would double-resolve defaults).
|
| 4 |
+
|
| 5 |
+
defaults:
|
| 6 |
+
- dreamzero/base_48_wan_fine_aug_relative
|
| 7 |
+
- _self_
|
| 8 |
+
|
| 9 |
+
# Wan 5B: 160x320 (HxW) -> latent 10x20, (10//2)*(20//2)=50. Use H,W divisible by 32 so latent is even (no crop in loss).
|
| 10 |
+
image_resolution_width: 320
|
| 11 |
+
image_resolution_height: 160
|
| 12 |
+
|
| 13 |
+
max_state_dim: 64
|
| 14 |
+
use_global_metadata: false
|
| 15 |
+
relative_action: true
|
| 16 |
+
relative_action_per_horizon: false
|
| 17 |
+
relative_action_keys:
|
| 18 |
+
- joint_position
|
| 19 |
+
max_chunk_size: 5
|
| 20 |
+
dataset_shard_sampling_rate: 0.1
|
| 21 |
+
mixture_dataset_cls: groot.vla.data.dataset.lerobot_sharded.ShardedLeRobotMixtureDataset.from_mixture_spec
|
| 22 |
+
single_dataset_cls: groot.vla.data.dataset.lerobot_sharded.ShardedLeRobotSubLangSingleActionChunkDatasetDROID
|
| 23 |
+
droid_data_root: ???
|
| 24 |
+
|
| 25 |
+
train_dataset:
|
| 26 |
+
_target_: ${mixture_dataset_cls}
|
| 27 |
+
_convert_: object
|
| 28 |
+
mixture_spec:
|
| 29 |
+
- dataset_path:
|
| 30 |
+
oxe_droid:
|
| 31 |
+
- ${droid_data_root}
|
| 32 |
+
dataset_weight: 1.0
|
| 33 |
+
distribute_weights: true
|
| 34 |
+
|
| 35 |
+
dataset_class: ${single_dataset_cls}
|
| 36 |
+
all_modality_configs: ${modality_configs}
|
| 37 |
+
all_transforms: ${transforms}
|
| 38 |
+
metadata_versions: ${metadata_versions}
|
| 39 |
+
fps: ${fps}
|
| 40 |
+
dataset_kwargs:
|
| 41 |
+
video_backend: decord
|
| 42 |
+
use_global_metadata: ${use_global_metadata}
|
| 43 |
+
max_chunk_size: ${max_chunk_size}
|
| 44 |
+
relative_action: ${relative_action}
|
| 45 |
+
relative_action_keys: ${relative_action_keys}
|
| 46 |
+
relative_action_per_horizon: ${relative_action_per_horizon}
|
| 47 |
+
mixture_kwargs:
|
| 48 |
+
training: true
|
| 49 |
+
balance_dataset_weights: false
|
| 50 |
+
seed: 42
|
| 51 |
+
shard_sampling_rate: ${dataset_shard_sampling_rate}
|
groot/vla/configs/data/dreamzero/libero.yaml
ADDED
|
@@ -0,0 +1,85 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# @package _global_
|
| 2 |
+
# LIBERO data config for DreamZero SFT (Wan2.2)
|
| 3 |
+
# Override base config with: model/dreamzero/action_head=wan_flow_matching_action_tf_wan22
|
| 4 |
+
|
| 5 |
+
defaults:
|
| 6 |
+
- dreamzero/base_48_wan_fine_aug_relative
|
| 7 |
+
- _self_
|
| 8 |
+
|
| 9 |
+
libero_dataset_dir: /root/autodl-tmp/data/libero
|
| 10 |
+
|
| 11 |
+
# Wan2.2 resolution
|
| 12 |
+
image_resolution_width: 320
|
| 13 |
+
image_resolution_height: 160
|
| 14 |
+
|
| 15 |
+
# VAE 4x temporal: 8 frames
|
| 16 |
+
num_frames: 12
|
| 17 |
+
action_horizon: 12
|
| 18 |
+
state_horizon: 1
|
| 19 |
+
num_views: 1
|
| 20 |
+
|
| 21 |
+
# Action head alignment
|
| 22 |
+
num_frame_per_block: 2
|
| 23 |
+
num_action_per_block: 12
|
| 24 |
+
num_state_per_block: 1
|
| 25 |
+
max_state_dim: 44
|
| 26 |
+
max_action_dim: 32
|
| 27 |
+
max_chunk_size: 4
|
| 28 |
+
frame_seqlen: 50
|
| 29 |
+
|
| 30 |
+
# Wan2.2 pretrained paths
|
| 31 |
+
dit_version: /root/autodl-tmp/checkpoints/Wan2.2-TI2V-5B
|
| 32 |
+
vae_pretrained_path: /root/autodl-tmp/checkpoints/Wan2.2-TI2V-5B/Wan2.2_VAE.pth
|
| 33 |
+
text_encoder_pretrained_path: /root/autodl-tmp/checkpoints/Wan2.2-TI2V-5B/models_t5_umt5-xxl-enc-bf16.pth
|
| 34 |
+
image_encoder_pretrained_path: /root/autodl-tmp/checkpoints/DreamZero-DROID/models_clip_open-clip-xlm-roberta-large-vit-huge-14.pth
|
| 35 |
+
|
| 36 |
+
train_dataset:
|
| 37 |
+
_target_: groot.vla.data.dataset.libero.LiberoDataset
|
| 38 |
+
_convert_: object
|
| 39 |
+
dataset_dir: ${libero_dataset_dir}
|
| 40 |
+
num_frames: ${num_frames}
|
| 41 |
+
action_horizon: ${action_horizon}
|
| 42 |
+
state_horizon: ${state_horizon}
|
| 43 |
+
num_views: ${num_views}
|
| 44 |
+
video_height: ${image_resolution_height}
|
| 45 |
+
video_width: ${image_resolution_width}
|
| 46 |
+
max_episodes: null
|
| 47 |
+
|
| 48 |
+
transform_libero:
|
| 49 |
+
_target_: groot.vla.data.transform.ComposedModalityTransform
|
| 50 |
+
transforms:
|
| 51 |
+
- ${model_specific_transform}
|
| 52 |
+
|
| 53 |
+
transforms:
|
| 54 |
+
oxe_droid: ${transform_libero}
|
| 55 |
+
|
| 56 |
+
modality_config_libero:
|
| 57 |
+
video:
|
| 58 |
+
_target_: groot.vla.data.dataset.ModalityConfig
|
| 59 |
+
delta_indices: [0, 1, 2, 3, 4, 5, 6, 7]
|
| 60 |
+
modality_keys:
|
| 61 |
+
- video
|
| 62 |
+
state:
|
| 63 |
+
_target_: groot.vla.data.dataset.ModalityConfig
|
| 64 |
+
delta_indices: [0]
|
| 65 |
+
modality_keys:
|
| 66 |
+
- state
|
| 67 |
+
action:
|
| 68 |
+
_target_: groot.vla.data.dataset.ModalityConfig
|
| 69 |
+
delta_indices: [0, 1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11]
|
| 70 |
+
modality_keys:
|
| 71 |
+
- action
|
| 72 |
+
language:
|
| 73 |
+
_target_: groot.vla.data.dataset.ModalityConfig
|
| 74 |
+
delta_indices: [0]
|
| 75 |
+
modality_keys:
|
| 76 |
+
- language
|
| 77 |
+
|
| 78 |
+
modality_configs:
|
| 79 |
+
oxe_droid: ${modality_config_libero}
|
| 80 |
+
|
| 81 |
+
metadata_versions:
|
| 82 |
+
oxe_droid: '0221'
|
| 83 |
+
|
| 84 |
+
fps:
|
| 85 |
+
oxe_droid: 10
|
groot/vla/configs/data/dreamzero/libero_sim_relative.yaml
ADDED
|
@@ -0,0 +1,58 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# @package _global_
|
| 2 |
+
# LIBERO data config for DreamZero GEAR pipeline (Wan2.2, 1 camera)
|
| 3 |
+
# Follows the official DATASET_TO_GEAR_AND_TRAIN.md pattern.
|
| 4 |
+
|
| 5 |
+
defaults:
|
| 6 |
+
- dreamzero/base_48_wan_fine_aug_relative
|
| 7 |
+
- _self_
|
| 8 |
+
|
| 9 |
+
# Wan2.2 resolution: 160x320 -> latent 10x20
|
| 10 |
+
image_resolution_width: 320
|
| 11 |
+
image_resolution_height: 160
|
| 12 |
+
|
| 13 |
+
# LIBERO: 12 frames, 12-step action chunk, 1 camera
|
| 14 |
+
use_global_metadata: false
|
| 15 |
+
relative_action: true
|
| 16 |
+
relative_action_per_horizon: false
|
| 17 |
+
relative_action_keys:
|
| 18 |
+
- joint_position
|
| 19 |
+
max_chunk_size: 5
|
| 20 |
+
dataset_shard_sampling_rate: 0.1
|
| 21 |
+
mixture_dataset_cls: groot.vla.data.dataset.lerobot_sharded.ShardedLeRobotMixtureDataset.from_mixture_spec
|
| 22 |
+
single_dataset_cls: groot.vla.data.dataset.lerobot_sharded.ShardedLeRobotSubLangSingleActionChunkDatasetDROID
|
| 23 |
+
|
| 24 |
+
# Set dataset path via CLI or env var:
|
| 25 |
+
# libero_sim_data_root=/root/autodl-tmp/data/libero_gear
|
| 26 |
+
libero_gear_root: /root/autodl-tmp/data/libero_gear
|
| 27 |
+
|
| 28 |
+
train_dataset:
|
| 29 |
+
_target_: ${mixture_dataset_cls}
|
| 30 |
+
_convert_: object
|
| 31 |
+
mixture_spec:
|
| 32 |
+
- dataset_path:
|
| 33 |
+
libero_sim:
|
| 34 |
+
- ${libero_gear_root}/libero_spatial
|
| 35 |
+
- ${libero_gear_root}/libero_object
|
| 36 |
+
- ${libero_gear_root}/libero_goal
|
| 37 |
+
- ${libero_gear_root}/libero_10
|
| 38 |
+
- ${libero_gear_root}/libero_90
|
| 39 |
+
dataset_weight: 1.0
|
| 40 |
+
distribute_weights: true
|
| 41 |
+
|
| 42 |
+
dataset_class: ${single_dataset_cls}
|
| 43 |
+
all_modality_configs: ${modality_configs}
|
| 44 |
+
all_transforms: ${transforms}
|
| 45 |
+
metadata_versions: ${metadata_versions}
|
| 46 |
+
fps: ${fps}
|
| 47 |
+
dataset_kwargs:
|
| 48 |
+
video_backend: decord
|
| 49 |
+
use_global_metadata: ${use_global_metadata}
|
| 50 |
+
max_chunk_size: ${max_chunk_size}
|
| 51 |
+
relative_action: ${relative_action}
|
| 52 |
+
relative_action_keys: ${relative_action_keys}
|
| 53 |
+
relative_action_per_horizon: ${relative_action_per_horizon}
|
| 54 |
+
mixture_kwargs:
|
| 55 |
+
training: true
|
| 56 |
+
balance_dataset_weights: false
|
| 57 |
+
seed: 42
|
| 58 |
+
shard_sampling_rate: ${dataset_shard_sampling_rate}
|
groot/vla/configs/data/dreamzero/libero_sim_relative.yaml.bak
ADDED
|
@@ -0,0 +1,54 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# @package _global_
|
| 2 |
+
# LIBERO data config for DreamZero GEAR pipeline (Wan2.2, 1 camera)
|
| 3 |
+
# Follows the official DATASET_TO_GEAR_AND_TRAIN.md pattern.
|
| 4 |
+
|
| 5 |
+
defaults:
|
| 6 |
+
- dreamzero/base_48_wan_fine_aug_relative
|
| 7 |
+
- _self_
|
| 8 |
+
|
| 9 |
+
# Wan2.2 resolution: 160x320 -> latent 10x20
|
| 10 |
+
image_resolution_width: 320
|
| 11 |
+
image_resolution_height: 160
|
| 12 |
+
|
| 13 |
+
# LIBERO: 12 frames, 12-step action chunk, 1 camera
|
| 14 |
+
use_global_metadata: false
|
| 15 |
+
relative_action: true
|
| 16 |
+
relative_action_per_horizon: false
|
| 17 |
+
relative_action_keys:
|
| 18 |
+
- joint_pos
|
| 19 |
+
max_chunk_size: 5
|
| 20 |
+
dataset_shard_sampling_rate: 0.1
|
| 21 |
+
mixture_dataset_cls: groot.vla.data.dataset.lerobot_sharded.ShardedLeRobotMixtureDataset.from_mixture_spec
|
| 22 |
+
single_dataset_cls: groot.vla.data.dataset.lerobot_sharded.ShardedLeRobotSubLangSingleActionChunkDatasetDROID
|
| 23 |
+
|
| 24 |
+
# Set dataset path via CLI or env var:
|
| 25 |
+
# libero_sim_data_root=/root/autodl-tmp/data/libero_gear
|
| 26 |
+
libero_sim_data_root: ???
|
| 27 |
+
|
| 28 |
+
train_dataset:
|
| 29 |
+
_target_: ${mixture_dataset_cls}
|
| 30 |
+
_convert_: object
|
| 31 |
+
mixture_spec:
|
| 32 |
+
- dataset_path:
|
| 33 |
+
libero_sim:
|
| 34 |
+
- ${libero_sim_data_root}
|
| 35 |
+
dataset_weight: 1.0
|
| 36 |
+
distribute_weights: true
|
| 37 |
+
|
| 38 |
+
dataset_class: ${single_dataset_cls}
|
| 39 |
+
all_modality_configs: ${modality_configs}
|
| 40 |
+
all_transforms: ${transforms}
|
| 41 |
+
metadata_versions: ${metadata_versions}
|
| 42 |
+
fps: ${fps}
|
| 43 |
+
dataset_kwargs:
|
| 44 |
+
video_backend: decord
|
| 45 |
+
use_global_metadata: ${use_global_metadata}
|
| 46 |
+
max_chunk_size: ${max_chunk_size}
|
| 47 |
+
relative_action: ${relative_action}
|
| 48 |
+
relative_action_keys: ${relative_action_keys}
|
| 49 |
+
relative_action_per_horizon: ${relative_action_per_horizon}
|
| 50 |
+
mixture_kwargs:
|
| 51 |
+
training: true
|
| 52 |
+
balance_dataset_weights: false
|
| 53 |
+
seed: 42
|
| 54 |
+
shard_sampling_rate: ${dataset_shard_sampling_rate}
|
groot/vla/configs/data/dreamzero/manifeel.yaml
ADDED
|
@@ -0,0 +1,86 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# @package _global_
|
| 2 |
+
# ManiFeel data config for DreamZero SFT (Wan2.2)
|
| 3 |
+
# Override base config with: model/dreamzero/action_head=wan_flow_matching_action_tf_wan22
|
| 4 |
+
|
| 5 |
+
defaults:
|
| 6 |
+
- dreamzero/base_48_wan_fine_aug_relative
|
| 7 |
+
- _self_
|
| 8 |
+
|
| 9 |
+
manifeel_dataset_dir: /root/autodl-tmp/tmp/manifeel_extracted
|
| 10 |
+
|
| 11 |
+
# Wan2.2 resolution: 160x320 -> latent 10x20 -> frame_seqlen=50
|
| 12 |
+
image_resolution_width: 320
|
| 13 |
+
image_resolution_height: 160
|
| 14 |
+
|
| 15 |
+
# VAE 4x temporal: 12 frames -> 12//4=3 latent frames
|
| 16 |
+
# With num_frame_per_block=2, state_horizon=1 -> assertion: (3-1)/1 = 2 = 2/1
|
| 17 |
+
num_frames: 12
|
| 18 |
+
action_horizon: 12
|
| 19 |
+
state_horizon: 1
|
| 20 |
+
num_views: 3
|
| 21 |
+
|
| 22 |
+
# Action head alignment (Wan2.2)
|
| 23 |
+
num_frame_per_block: 2
|
| 24 |
+
num_action_per_block: 12
|
| 25 |
+
num_state_per_block: 1
|
| 26 |
+
max_state_dim: 44
|
| 27 |
+
max_action_dim: 32
|
| 28 |
+
max_chunk_size: 4
|
| 29 |
+
frame_seqlen: 50
|
| 30 |
+
|
| 31 |
+
# Wan2.2 pretrained paths
|
| 32 |
+
dit_version: /root/autodl-tmp/checkpoints/Wan2.2-TI2V-5B
|
| 33 |
+
vae_pretrained_path: /root/autodl-tmp/checkpoints/Wan2.2-TI2V-5B/Wan2.2_VAE.pth
|
| 34 |
+
text_encoder_pretrained_path: /root/autodl-tmp/checkpoints/Wan2.2-TI2V-5B/models_t5_umt5-xxl-enc-bf16.pth
|
| 35 |
+
image_encoder_pretrained_path: /root/autodl-tmp/checkpoints/DreamZero-DROID/models_clip_open-clip-xlm-roberta-large-vit-huge-14.pth
|
| 36 |
+
|
| 37 |
+
train_dataset:
|
| 38 |
+
_target_: groot.vla.data.dataset.manifeel.ManiFeelDataset
|
| 39 |
+
_convert_: object
|
| 40 |
+
dataset_dir: ${manifeel_dataset_dir}
|
| 41 |
+
num_frames: ${num_frames}
|
| 42 |
+
action_horizon: ${action_horizon}
|
| 43 |
+
state_horizon: ${state_horizon}
|
| 44 |
+
num_views: ${num_views}
|
| 45 |
+
video_height: ${image_resolution_height}
|
| 46 |
+
video_width: ${image_resolution_width}
|
| 47 |
+
max_episodes: null
|
| 48 |
+
|
| 49 |
+
transform_manifeel:
|
| 50 |
+
_target_: groot.vla.data.transform.ComposedModalityTransform
|
| 51 |
+
transforms:
|
| 52 |
+
- ${model_specific_transform}
|
| 53 |
+
|
| 54 |
+
transforms:
|
| 55 |
+
oxe_droid: ${transform_manifeel}
|
| 56 |
+
|
| 57 |
+
modality_config_manifeel:
|
| 58 |
+
video:
|
| 59 |
+
_target_: groot.vla.data.dataset.ModalityConfig
|
| 60 |
+
delta_indices: [0, 1, 2, 3, 4, 5, 6, 7]
|
| 61 |
+
modality_keys:
|
| 62 |
+
- video
|
| 63 |
+
state:
|
| 64 |
+
_target_: groot.vla.data.dataset.ModalityConfig
|
| 65 |
+
delta_indices: [0]
|
| 66 |
+
modality_keys:
|
| 67 |
+
- state
|
| 68 |
+
action:
|
| 69 |
+
_target_: groot.vla.data.dataset.ModalityConfig
|
| 70 |
+
delta_indices: [0, 1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11]
|
| 71 |
+
modality_keys:
|
| 72 |
+
- action
|
| 73 |
+
language:
|
| 74 |
+
_target_: groot.vla.data.dataset.ModalityConfig
|
| 75 |
+
delta_indices: [0]
|
| 76 |
+
modality_keys:
|
| 77 |
+
- language
|
| 78 |
+
|
| 79 |
+
modality_configs:
|
| 80 |
+
oxe_droid: ${modality_config_manifeel}
|
| 81 |
+
|
| 82 |
+
metadata_versions:
|
| 83 |
+
oxe_droid: '0221'
|
| 84 |
+
|
| 85 |
+
fps:
|
| 86 |
+
oxe_droid: 10
|
groot/vla/configs/data/dreamzero/robotwin.yaml
ADDED
|
@@ -0,0 +1,35 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# @package _global_
|
| 2 |
+
defaults:
|
| 3 |
+
- dreamzero/base_48_wan_fine_aug_relative
|
| 4 |
+
- _self_
|
| 5 |
+
|
| 6 |
+
robotwin_dataset_dir: /root/autodl-tmp/data/robotwin_gear
|
| 7 |
+
image_resolution_width: 320
|
| 8 |
+
image_resolution_height: 160
|
| 9 |
+
num_frames: 12
|
| 10 |
+
action_horizon: 12
|
| 11 |
+
state_horizon: 1
|
| 12 |
+
num_views: 1
|
| 13 |
+
num_frame_per_block: 2
|
| 14 |
+
num_action_per_block: 12
|
| 15 |
+
num_state_per_block: 1
|
| 16 |
+
max_state_dim: 44
|
| 17 |
+
max_action_dim: 32
|
| 18 |
+
max_chunk_size: 4
|
| 19 |
+
frame_seqlen: 50
|
| 20 |
+
dit_version: /root/autodl-tmp/checkpoints/Wan2.2-TI2V-5B
|
| 21 |
+
vae_pretrained_path: /root/autodl-tmp/checkpoints/Wan2.2-TI2V-5B/Wan2.2_VAE.pth
|
| 22 |
+
text_encoder_pretrained_path: /root/autodl-tmp/checkpoints/Wan2.2-TI2V-5B/models_t5_umt5-xxl-enc-bf16.pth
|
| 23 |
+
image_encoder_pretrained_path: /root/autodl-tmp/checkpoints/DreamZero-DROID/models_clip_open-clip-xlm-roberta-large-vit-huge-14.pth
|
| 24 |
+
|
| 25 |
+
train_dataset:
|
| 26 |
+
_target_: groot.vla.data.dataset.robotwin.RobotWinDataset
|
| 27 |
+
_convert_: object
|
| 28 |
+
dataset_dir: ${robotwin_dataset_dir}
|
| 29 |
+
num_frames: ${num_frames}
|
| 30 |
+
action_horizon: ${action_horizon}
|
| 31 |
+
state_horizon: ${state_horizon}
|
| 32 |
+
num_views: ${num_views}
|
| 33 |
+
video_height: ${image_resolution_height}
|
| 34 |
+
video_width: ${image_resolution_width}
|
| 35 |
+
max_episodes: null
|
groot/vla/configs/data/dreamzero/yam_relative.yaml
ADDED
|
@@ -0,0 +1,52 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# @package _global_
|
| 2 |
+
|
| 3 |
+
defaults:
|
| 4 |
+
- dreamzero/base_48_wan_fine_aug_relative
|
| 5 |
+
- _self_ # this file will override the base
|
| 6 |
+
|
| 7 |
+
max_state_dim: 64
|
| 8 |
+
use_global_metadata: false
|
| 9 |
+
relative_action: true
|
| 10 |
+
relative_action_per_horizon: false
|
| 11 |
+
relative_action_keys:
|
| 12 |
+
- left_joint_pos
|
| 13 |
+
- left_gripper_pos
|
| 14 |
+
- right_joint_pos
|
| 15 |
+
- right_gripper_pos
|
| 16 |
+
max_chunk_size: 5
|
| 17 |
+
# Use 10% of data in shards before moving to next shard
|
| 18 |
+
dataset_shard_sampling_rate: 0.1
|
| 19 |
+
mixture_dataset_cls: groot.vla.data.dataset.lerobot_sharded.ShardedLeRobotMixtureDataset.from_mixture_spec
|
| 20 |
+
single_dataset_cls: groot.vla.data.dataset.lerobot_sharded.ShardedLeRobotSubLangSingleActionChunkDatasetDROID
|
| 21 |
+
|
| 22 |
+
# Set your YAM dataset path here or override via CLI:
|
| 23 |
+
# yam_data_root=/path/to/your/yam_dataset
|
| 24 |
+
yam_data_root: ???
|
| 25 |
+
|
| 26 |
+
train_dataset:
|
| 27 |
+
_target_: ${mixture_dataset_cls}
|
| 28 |
+
_convert_: object
|
| 29 |
+
mixture_spec:
|
| 30 |
+
- dataset_path:
|
| 31 |
+
yam:
|
| 32 |
+
- ${yam_data_root}
|
| 33 |
+
dataset_weight: 1.0
|
| 34 |
+
distribute_weights: true
|
| 35 |
+
|
| 36 |
+
dataset_class: ${single_dataset_cls}
|
| 37 |
+
all_modality_configs: ${modality_configs}
|
| 38 |
+
all_transforms: ${transforms}
|
| 39 |
+
metadata_versions: ${metadata_versions}
|
| 40 |
+
fps: ${fps}
|
| 41 |
+
dataset_kwargs:
|
| 42 |
+
video_backend: decord
|
| 43 |
+
use_global_metadata: ${use_global_metadata}
|
| 44 |
+
max_chunk_size: ${max_chunk_size}
|
| 45 |
+
relative_action: ${relative_action}
|
| 46 |
+
relative_action_keys: ${relative_action_keys}
|
| 47 |
+
relative_action_per_horizon: ${relative_action_per_horizon}
|
| 48 |
+
mixture_kwargs:
|
| 49 |
+
training: true
|
| 50 |
+
balance_dataset_weights: false
|
| 51 |
+
seed: 42
|
| 52 |
+
shard_sampling_rate: ${dataset_shard_sampling_rate}
|
groot/vla/configs/deepspeed/zero2.json
ADDED
|
@@ -0,0 +1,28 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
{
|
| 2 |
+
"checkpoint": {
|
| 3 |
+
"load_universal": false
|
| 4 |
+
},
|
| 5 |
+
"train_batch_size": "auto",
|
| 6 |
+
"train_micro_batch_size_per_gpu": "auto",
|
| 7 |
+
"gradient_accumulation_steps": "auto",
|
| 8 |
+
"gradient_clipping": "auto",
|
| 9 |
+
"zero_allow_untested_optimizer": true,
|
| 10 |
+
"fp16": {
|
| 11 |
+
"enabled": "auto",
|
| 12 |
+
"loss_scale": 0,
|
| 13 |
+
"loss_scale_window": 1000,
|
| 14 |
+
"initial_scale_power": 16,
|
| 15 |
+
"hysteresis": 2,
|
| 16 |
+
"min_loss_scale": 1
|
| 17 |
+
},
|
| 18 |
+
"bf16": {
|
| 19 |
+
"enabled": "auto"
|
| 20 |
+
},
|
| 21 |
+
"zero_optimization": {
|
| 22 |
+
"stage": 2,
|
| 23 |
+
"overlap_comm": false,
|
| 24 |
+
"contiguous_gradients": true,
|
| 25 |
+
"sub_group_size": 1e9,
|
| 26 |
+
"reduce_bucket_size": 1e8
|
| 27 |
+
}
|
| 28 |
+
}
|
groot/vla/configs/deepspeed/zero2_offload.json
ADDED
|
@@ -0,0 +1,32 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
{
|
| 2 |
+
"checkpoint": {
|
| 3 |
+
"load_universal": false
|
| 4 |
+
},
|
| 5 |
+
"train_batch_size": "auto",
|
| 6 |
+
"train_micro_batch_size_per_gpu": "auto",
|
| 7 |
+
"gradient_accumulation_steps": "auto",
|
| 8 |
+
"gradient_clipping": "auto",
|
| 9 |
+
"zero_allow_untested_optimizer": true,
|
| 10 |
+
"fp16": {
|
| 11 |
+
"enabled": "auto",
|
| 12 |
+
"loss_scale": 0,
|
| 13 |
+
"loss_scale_window": 1000,
|
| 14 |
+
"initial_scale_power": 16,
|
| 15 |
+
"hysteresis": 2,
|
| 16 |
+
"min_loss_scale": 1
|
| 17 |
+
},
|
| 18 |
+
"bf16": {
|
| 19 |
+
"enabled": "auto"
|
| 20 |
+
},
|
| 21 |
+
"zero_optimization": {
|
| 22 |
+
"stage": 2,
|
| 23 |
+
"offload_optimizer": {
|
| 24 |
+
"device": "cpu",
|
| 25 |
+
"pin_memory": true
|
| 26 |
+
},
|
| 27 |
+
"overlap_comm": false,
|
| 28 |
+
"contiguous_gradients": true,
|
| 29 |
+
"sub_group_size": 1e9,
|
| 30 |
+
"reduce_bucket_size": 1e8
|
| 31 |
+
}
|
| 32 |
+
}
|