Spaces:
Sleeping
Sleeping
Upload folder using huggingface_hub
Browse filesThis view is limited to 50 files because it contains too many changes. See raw diff
- .gitattributes +16 -0
- .gitignore +169 -0
- .gradio/certificate.pem +31 -0
- LICENSE +21 -0
- README.md +158 -7
- app.py +34 -0
- datasets/.gitignore +5 -0
- datasets/data_gen.py +317 -0
- datasets/data_gen.yaml +29 -0
- requirements.txt +20 -0
- scripts/__init__.py +0 -0
- scripts/auto_format.sh +15 -0
- scripts/benchmark_decomp.py +63 -0
- scripts/benchmark_inference.py +163 -0
- scripts/examples/microwave-bottom_burner-light_switch-slide_cabinet.mp4 +0 -0
- scripts/tsne_visualization.py +110 -0
- setup.py +49 -0
- uvd/__init__.py +30 -0
- uvd/data/__init__.py +3 -0
- uvd/data/dataset_aug.py +310 -0
- uvd/data/dataset_base.py +19 -0
- uvd/data/franka_kitchen_datasets.py +707 -0
- uvd/decomp/__init__.py +1 -0
- uvd/decomp/decomp.py +636 -0
- uvd/decomp/kernel_reg.py +91 -0
- uvd/envs/__init__.py +0 -0
- uvd/envs/evaluator/__init__.py +3 -0
- uvd/envs/evaluator/evaluator.py +571 -0
- uvd/envs/evaluator/inference_wrapper.py +435 -0
- uvd/envs/evaluator/vec_envs/__init__.py +0 -0
- uvd/envs/evaluator/vec_envs/vec_env.py +398 -0
- uvd/envs/evaluator/vec_envs/workers.py +409 -0
- uvd/envs/evaluator/visualize_wrapper.py +271 -0
- uvd/envs/franka_kitchen/__init__.py +20 -0
- uvd/envs/franka_kitchen/franka_kitchen_base.py +446 -0
- uvd/envs/franka_kitchen/franka_kitchen_constants.py +62 -0
- uvd/envs/franka_kitchen/relay-policy-learning/adept_envs/.pylintrc +433 -0
- uvd/envs/franka_kitchen/relay-policy-learning/adept_envs/.style.yapf +323 -0
- uvd/envs/franka_kitchen/relay-policy-learning/adept_envs/adept_envs/__init__.py +15 -0
- uvd/envs/franka_kitchen/relay-policy-learning/adept_envs/adept_envs/base_robot.py +153 -0
- uvd/envs/franka_kitchen/relay-policy-learning/adept_envs/adept_envs/franka/__init__.py +24 -0
- uvd/envs/franka_kitchen/relay-policy-learning/adept_envs/adept_envs/franka/assets/franka_kitchen_jntpos_act_ab.xml +94 -0
- uvd/envs/franka_kitchen/relay-policy-learning/adept_envs/adept_envs/franka/kitchen_multitask_v0.py +234 -0
- uvd/envs/franka_kitchen/relay-policy-learning/adept_envs/adept_envs/franka/robot/franka_config.xml +59 -0
- uvd/envs/franka_kitchen/relay-policy-learning/adept_envs/adept_envs/franka/robot/franka_robot.py +342 -0
- uvd/envs/franka_kitchen/relay-policy-learning/adept_envs/adept_envs/mujoco_env.py +222 -0
- uvd/envs/franka_kitchen/relay-policy-learning/adept_envs/adept_envs/robot_env.py +178 -0
- uvd/envs/franka_kitchen/relay-policy-learning/adept_envs/adept_envs/simulation/module.py +135 -0
- uvd/envs/franka_kitchen/relay-policy-learning/adept_envs/adept_envs/simulation/renderer.py +304 -0
- uvd/envs/franka_kitchen/relay-policy-learning/adept_envs/adept_envs/simulation/sim_robot.py +133 -0
.gitattributes
CHANGED
|
@@ -33,3 +33,19 @@ saved_model/**/* filter=lfs diff=lfs merge=lfs -text
|
|
| 33 |
*.zip filter=lfs diff=lfs merge=lfs -text
|
| 34 |
*.zst filter=lfs diff=lfs merge=lfs -text
|
| 35 |
*tfevents* filter=lfs diff=lfs merge=lfs -text
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 33 |
*.zip filter=lfs diff=lfs merge=lfs -text
|
| 34 |
*.zst filter=lfs diff=lfs merge=lfs -text
|
| 35 |
*tfevents* filter=lfs diff=lfs merge=lfs -text
|
| 36 |
+
uvd/envs/franka_kitchen/relay-policy-learning/adept_models/kitchen/textures/marble1.png filter=lfs diff=lfs merge=lfs -text
|
| 37 |
+
uvd/envs/franka_kitchen/relay-policy-learning/adept_models/kitchen/textures/metal1.png filter=lfs diff=lfs merge=lfs -text
|
| 38 |
+
uvd/envs/franka_kitchen/relay-policy-learning/adept_models/kitchen/textures/tile1.png filter=lfs diff=lfs merge=lfs -text
|
| 39 |
+
uvd/envs/franka_kitchen/relay-policy-learning/adept_models/kitchen/textures/wood1.png filter=lfs diff=lfs merge=lfs -text
|
| 40 |
+
uvd/envs/franka_kitchen/relay-policy-learning/adept_models/scenes/textures/white_marble_tile.png filter=lfs diff=lfs merge=lfs -text
|
| 41 |
+
uvd/envs/franka_kitchen/relay-policy-learning/adept_models/scenes/textures/white_marble_tile2.png filter=lfs diff=lfs merge=lfs -text
|
| 42 |
+
uvd/envs/franka_kitchen/relay-policy-learning/third_party/franka/franka_panda.png filter=lfs diff=lfs merge=lfs -text
|
| 43 |
+
uvd/envs/franka_kitchen/relay-policy-learning/third_party/franka/meshes/visual/hand.stl filter=lfs diff=lfs merge=lfs -text
|
| 44 |
+
uvd/envs/franka_kitchen/relay-policy-learning/third_party/franka/meshes/visual/link0.stl filter=lfs diff=lfs merge=lfs -text
|
| 45 |
+
uvd/envs/franka_kitchen/relay-policy-learning/third_party/franka/meshes/visual/link1.stl filter=lfs diff=lfs merge=lfs -text
|
| 46 |
+
uvd/envs/franka_kitchen/relay-policy-learning/third_party/franka/meshes/visual/link2.stl filter=lfs diff=lfs merge=lfs -text
|
| 47 |
+
uvd/envs/franka_kitchen/relay-policy-learning/third_party/franka/meshes/visual/link3.stl filter=lfs diff=lfs merge=lfs -text
|
| 48 |
+
uvd/envs/franka_kitchen/relay-policy-learning/third_party/franka/meshes/visual/link4.stl filter=lfs diff=lfs merge=lfs -text
|
| 49 |
+
uvd/envs/franka_kitchen/relay-policy-learning/third_party/franka/meshes/visual/link5.stl filter=lfs diff=lfs merge=lfs -text
|
| 50 |
+
uvd/envs/franka_kitchen/relay-policy-learning/third_party/franka/meshes/visual/link6.stl filter=lfs diff=lfs merge=lfs -text
|
| 51 |
+
uvd/envs/franka_kitchen/relay-policy-learning/third_party/franka/meshes/visual/link7.stl filter=lfs diff=lfs merge=lfs -text
|
.gitignore
ADDED
|
@@ -0,0 +1,169 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# Byte-compiled / optimized / DLL files
|
| 2 |
+
__pycache__/
|
| 3 |
+
*.py[cod]
|
| 4 |
+
*$py.class
|
| 5 |
+
|
| 6 |
+
# C extensions
|
| 7 |
+
*.so
|
| 8 |
+
|
| 9 |
+
# Distribution / packaging
|
| 10 |
+
.Python
|
| 11 |
+
build/
|
| 12 |
+
develop-eggs/
|
| 13 |
+
dist/
|
| 14 |
+
downloads/
|
| 15 |
+
eggs/
|
| 16 |
+
.eggs/
|
| 17 |
+
lib/
|
| 18 |
+
lib64/
|
| 19 |
+
parts/
|
| 20 |
+
sdist/
|
| 21 |
+
var/
|
| 22 |
+
wheels/
|
| 23 |
+
share/python-wheels/
|
| 24 |
+
*.egg-info/
|
| 25 |
+
.installed.cfg
|
| 26 |
+
*.egg
|
| 27 |
+
MANIFEST
|
| 28 |
+
|
| 29 |
+
# PyInstaller
|
| 30 |
+
# Usually these files are written by a python script from a template
|
| 31 |
+
# before PyInstaller builds the exe, so as to inject date/other infos into it.
|
| 32 |
+
*.manifest
|
| 33 |
+
*.spec
|
| 34 |
+
|
| 35 |
+
# Installer logs
|
| 36 |
+
pip-log.txt
|
| 37 |
+
pip-delete-this-directory.txt
|
| 38 |
+
|
| 39 |
+
# Unit test / coverage reports
|
| 40 |
+
htmlcov/
|
| 41 |
+
.tox/
|
| 42 |
+
.nox/
|
| 43 |
+
.coverage
|
| 44 |
+
.coverage.*
|
| 45 |
+
.cache
|
| 46 |
+
nosetests.xml
|
| 47 |
+
coverage.xml
|
| 48 |
+
*.cover
|
| 49 |
+
*.py,cover
|
| 50 |
+
.hypothesis/
|
| 51 |
+
.pytest_cache/
|
| 52 |
+
cover/
|
| 53 |
+
|
| 54 |
+
# Translations
|
| 55 |
+
*.mo
|
| 56 |
+
*.pot
|
| 57 |
+
|
| 58 |
+
# Django stuff:
|
| 59 |
+
*.log
|
| 60 |
+
local_settings.py
|
| 61 |
+
db.sqlite3
|
| 62 |
+
db.sqlite3-journal
|
| 63 |
+
|
| 64 |
+
# Flask stuff:
|
| 65 |
+
instance/
|
| 66 |
+
.webassets-cache
|
| 67 |
+
|
| 68 |
+
# Scrapy stuff:
|
| 69 |
+
.scrapy
|
| 70 |
+
|
| 71 |
+
# Sphinx documentation
|
| 72 |
+
docs/_build/
|
| 73 |
+
|
| 74 |
+
# PyBuilder
|
| 75 |
+
.pybuilder/
|
| 76 |
+
target/
|
| 77 |
+
|
| 78 |
+
# Jupyter Notebook
|
| 79 |
+
.ipynb_checkpoints
|
| 80 |
+
|
| 81 |
+
# IPython
|
| 82 |
+
profile_default/
|
| 83 |
+
ipython_config.py
|
| 84 |
+
|
| 85 |
+
# pyenv
|
| 86 |
+
# For a library or package, you might want to ignore these files since the code is
|
| 87 |
+
# intended to run in multiple environments; otherwise, check them in:
|
| 88 |
+
# .python-version
|
| 89 |
+
|
| 90 |
+
# pipenv
|
| 91 |
+
# According to pypa/pipenv#598, it is recommended to include Pipfile.lock in version control.
|
| 92 |
+
# However, in case of collaboration, if having platform-specific dependencies or dependencies
|
| 93 |
+
# having no cross-platform support, pipenv may install dependencies that don't work, or not
|
| 94 |
+
# install all needed dependencies.
|
| 95 |
+
#Pipfile.lock
|
| 96 |
+
|
| 97 |
+
# poetry
|
| 98 |
+
# Similar to Pipfile.lock, it is generally recommended to include poetry.lock in version control.
|
| 99 |
+
# This is especially recommended for binary packages to ensure reproducibility, and is more
|
| 100 |
+
# commonly ignored for libraries.
|
| 101 |
+
# https://python-poetry.org/docs/basic-usage/#commit-your-poetrylock-file-to-version-control
|
| 102 |
+
#poetry.lock
|
| 103 |
+
|
| 104 |
+
# pdm
|
| 105 |
+
# Similar to Pipfile.lock, it is generally recommended to include pdm.lock in version control.
|
| 106 |
+
#pdm.lock
|
| 107 |
+
# pdm stores project-wide configurations in .pdm.toml, but it is recommended to not include it
|
| 108 |
+
# in version control.
|
| 109 |
+
# https://pdm.fming.dev/#use-with-ide
|
| 110 |
+
.pdm.toml
|
| 111 |
+
|
| 112 |
+
# PEP 582; used by e.g. github.com/David-OConnor/pyflow and github.com/pdm-project/pdm
|
| 113 |
+
__pypackages__/
|
| 114 |
+
|
| 115 |
+
# Celery stuff
|
| 116 |
+
celerybeat-schedule
|
| 117 |
+
celerybeat.pid
|
| 118 |
+
|
| 119 |
+
# SageMath parsed files
|
| 120 |
+
*.sage.py
|
| 121 |
+
|
| 122 |
+
# Environments
|
| 123 |
+
.env
|
| 124 |
+
.venv
|
| 125 |
+
env/
|
| 126 |
+
venv/
|
| 127 |
+
ENV/
|
| 128 |
+
env.bak/
|
| 129 |
+
venv.bak/
|
| 130 |
+
|
| 131 |
+
# Spyder project settings
|
| 132 |
+
.spyderproject
|
| 133 |
+
.spyproject
|
| 134 |
+
|
| 135 |
+
# Rope project settings
|
| 136 |
+
.ropeproject
|
| 137 |
+
|
| 138 |
+
# mkdocs documentation
|
| 139 |
+
/site
|
| 140 |
+
|
| 141 |
+
# mypy
|
| 142 |
+
.mypy_cache/
|
| 143 |
+
.dmypy.json
|
| 144 |
+
dmypy.json
|
| 145 |
+
|
| 146 |
+
# Pyre type checker
|
| 147 |
+
.pyre/
|
| 148 |
+
|
| 149 |
+
# pytype static type analyzer
|
| 150 |
+
.pytype/
|
| 151 |
+
|
| 152 |
+
# Cython debug symbols
|
| 153 |
+
cython_debug/
|
| 154 |
+
|
| 155 |
+
# PyCharm
|
| 156 |
+
# JetBrains specific template is maintained in a separate JetBrains.gitignore that can
|
| 157 |
+
# be found at https://github.com/github/gitignore/blob/main/Global/JetBrains.gitignore
|
| 158 |
+
# and can be added to the global gitignore or merged into this file. For a more nuclear
|
| 159 |
+
# option (not recommended) you can uncomment the following to ignore the entire idea folder.
|
| 160 |
+
.idea/
|
| 161 |
+
|
| 162 |
+
*.pt
|
| 163 |
+
*.pth
|
| 164 |
+
*pl
|
| 165 |
+
*.patch
|
| 166 |
+
*used_configs
|
| 167 |
+
.allenact_last_start_time_string
|
| 168 |
+
*.lock
|
| 169 |
+
*wandb
|
.gradio/certificate.pem
ADDED
|
@@ -0,0 +1,31 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
-----BEGIN CERTIFICATE-----
|
| 2 |
+
MIIFazCCA1OgAwIBAgIRAIIQz7DSQONZRGPgu2OCiwAwDQYJKoZIhvcNAQELBQAw
|
| 3 |
+
TzELMAkGA1UEBhMCVVMxKTAnBgNVBAoTIEludGVybmV0IFNlY3VyaXR5IFJlc2Vh
|
| 4 |
+
cmNoIEdyb3VwMRUwEwYDVQQDEwxJU1JHIFJvb3QgWDEwHhcNMTUwNjA0MTEwNDM4
|
| 5 |
+
WhcNMzUwNjA0MTEwNDM4WjBPMQswCQYDVQQGEwJVUzEpMCcGA1UEChMgSW50ZXJu
|
| 6 |
+
ZXQgU2VjdXJpdHkgUmVzZWFyY2ggR3JvdXAxFTATBgNVBAMTDElTUkcgUm9vdCBY
|
| 7 |
+
MTCCAiIwDQYJKoZIhvcNAQEBBQADggIPADCCAgoCggIBAK3oJHP0FDfzm54rVygc
|
| 8 |
+
h77ct984kIxuPOZXoHj3dcKi/vVqbvYATyjb3miGbESTtrFj/RQSa78f0uoxmyF+
|
| 9 |
+
0TM8ukj13Xnfs7j/EvEhmkvBioZxaUpmZmyPfjxwv60pIgbz5MDmgK7iS4+3mX6U
|
| 10 |
+
A5/TR5d8mUgjU+g4rk8Kb4Mu0UlXjIB0ttov0DiNewNwIRt18jA8+o+u3dpjq+sW
|
| 11 |
+
T8KOEUt+zwvo/7V3LvSye0rgTBIlDHCNAymg4VMk7BPZ7hm/ELNKjD+Jo2FR3qyH
|
| 12 |
+
B5T0Y3HsLuJvW5iB4YlcNHlsdu87kGJ55tukmi8mxdAQ4Q7e2RCOFvu396j3x+UC
|
| 13 |
+
B5iPNgiV5+I3lg02dZ77DnKxHZu8A/lJBdiB3QW0KtZB6awBdpUKD9jf1b0SHzUv
|
| 14 |
+
KBds0pjBqAlkd25HN7rOrFleaJ1/ctaJxQZBKT5ZPt0m9STJEadao0xAH0ahmbWn
|
| 15 |
+
OlFuhjuefXKnEgV4We0+UXgVCwOPjdAvBbI+e0ocS3MFEvzG6uBQE3xDk3SzynTn
|
| 16 |
+
jh8BCNAw1FtxNrQHusEwMFxIt4I7mKZ9YIqioymCzLq9gwQbooMDQaHWBfEbwrbw
|
| 17 |
+
qHyGO0aoSCqI3Haadr8faqU9GY/rOPNk3sgrDQoo//fb4hVC1CLQJ13hef4Y53CI
|
| 18 |
+
rU7m2Ys6xt0nUW7/vGT1M0NPAgMBAAGjQjBAMA4GA1UdDwEB/wQEAwIBBjAPBgNV
|
| 19 |
+
HRMBAf8EBTADAQH/MB0GA1UdDgQWBBR5tFnme7bl5AFzgAiIyBpY9umbbjANBgkq
|
| 20 |
+
hkiG9w0BAQsFAAOCAgEAVR9YqbyyqFDQDLHYGmkgJykIrGF1XIpu+ILlaS/V9lZL
|
| 21 |
+
ubhzEFnTIZd+50xx+7LSYK05qAvqFyFWhfFQDlnrzuBZ6brJFe+GnY+EgPbk6ZGQ
|
| 22 |
+
3BebYhtF8GaV0nxvwuo77x/Py9auJ/GpsMiu/X1+mvoiBOv/2X/qkSsisRcOj/KK
|
| 23 |
+
NFtY2PwByVS5uCbMiogziUwthDyC3+6WVwW6LLv3xLfHTjuCvjHIInNzktHCgKQ5
|
| 24 |
+
ORAzI4JMPJ+GslWYHb4phowim57iaztXOoJwTdwJx4nLCgdNbOhdjsnvzqvHu7Ur
|
| 25 |
+
TkXWStAmzOVyyghqpZXjFaH3pO3JLF+l+/+sKAIuvtd7u+Nxe5AW0wdeRlN8NwdC
|
| 26 |
+
jNPElpzVmbUq4JUagEiuTDkHzsxHpFKVK7q4+63SM1N95R1NbdWhscdCb+ZAJzVc
|
| 27 |
+
oyi3B43njTOQ5yOf+1CceWxG1bQVs5ZufpsMljq4Ui0/1lvh+wjChP4kqKOJ2qxq
|
| 28 |
+
4RgqsahDYVvTH9w7jXbyLeiNdd8XM2w9U/t7y0Ff/9yi0GE44Za4rF2LN9d11TPA
|
| 29 |
+
mRGunUHBcnWEvgJBQl9nJEiU0Zsnvgc/ubhPgXRR4Xq37Z0j4r7g1SgEEzwxA57d
|
| 30 |
+
emyPxgcYxn/eR44/KJ4EBs+lVDR3veyJm+kXQ99b21/+jh5Xos1AnX5iItreGCc=
|
| 31 |
+
-----END CERTIFICATE-----
|
LICENSE
ADDED
|
@@ -0,0 +1,21 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
MIT License
|
| 2 |
+
|
| 3 |
+
Copyright (c) 2023 UVD
|
| 4 |
+
|
| 5 |
+
Permission is hereby granted, free of charge, to any person obtaining a copy
|
| 6 |
+
of this software and associated documentation files (the "Software"), to deal
|
| 7 |
+
in the Software without restriction, including without limitation the rights
|
| 8 |
+
to use, copy, modify, merge, publish, distribute, sublicense, and/or sell
|
| 9 |
+
copies of the Software, and to permit persons to whom the Software is
|
| 10 |
+
furnished to do so, subject to the following conditions:
|
| 11 |
+
|
| 12 |
+
The above copyright notice and this permission notice shall be included in all
|
| 13 |
+
copies or substantial portions of the Software.
|
| 14 |
+
|
| 15 |
+
THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
|
| 16 |
+
IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY,
|
| 17 |
+
FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE
|
| 18 |
+
AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER
|
| 19 |
+
LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM,
|
| 20 |
+
OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE
|
| 21 |
+
SOFTWARE.
|
README.md
CHANGED
|
@@ -1,12 +1,163 @@
|
|
| 1 |
---
|
| 2 |
title: UVD
|
| 3 |
-
emoji: 🏢
|
| 4 |
-
colorFrom: blue
|
| 5 |
-
colorTo: green
|
| 6 |
-
sdk: gradio
|
| 7 |
-
sdk_version: 6.12.0
|
| 8 |
app_file: app.py
|
| 9 |
-
|
|
|
|
| 10 |
---
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 11 |
|
| 12 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
---
|
| 2 |
title: UVD
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 3 |
app_file: app.py
|
| 4 |
+
sdk: gradio
|
| 5 |
+
sdk_version: 6.11.0
|
| 6 |
---
|
| 7 |
+
# Universal Visual Decomposer: <br>Long-Horizon Manipulation Made Easy
|
| 8 |
+
|
| 9 |
+
<div align="center">
|
| 10 |
+
|
| 11 |
+
[[Website]](https://zcczhang.github.io/UVD/)
|
| 12 |
+
[[arXiv]](https://arxiv.org/abs/2310.08581)
|
| 13 |
+
[[PDF]](https://zcczhang.github.io/UVD/assets/pdf/full_paper.pdf)
|
| 14 |
+
[[Installation]](#Installation)
|
| 15 |
+
[[Usage]](#Usage)
|
| 16 |
+
[[BibTex]](#Citation)
|
| 17 |
+
______________________________________________________________________
|
| 18 |
+
|
| 19 |
+
|
| 20 |
+
|
| 21 |
+
https://github.com/zcczhang/UVD/assets/52727818/5555b99a-76eb-4d76-966f-787af763573a
|
| 22 |
+
|
| 23 |
+
|
| 24 |
+
|
| 25 |
+
|
| 26 |
+
</div>
|
| 27 |
+
|
| 28 |
+
# Installation
|
| 29 |
+
|
| 30 |
+
- Follow the [instruction](https://github.com/openai/mujoco-py#install-mujoco) for installing `mujuco-py` and install the following apt packages if using Ubuntu:
|
| 31 |
+
```commandline
|
| 32 |
+
sudo apt install -y libosmesa6-dev libglfw3 patchelf libgl1 libglx-mesa0
|
| 33 |
+
```
|
| 34 |
+
- create conda env with Python==3.9
|
| 35 |
+
```commandline
|
| 36 |
+
conda create -n uvd python==3.9 -y && conda activate uvd
|
| 37 |
+
```
|
| 38 |
+
- Install any/all standalone visual foundation models from their repos separately *before* setup UVD, in case dependency conflicts, e.g.:
|
| 39 |
+
<details><summary>
|
| 40 |
+
<a href="https://github.com/facebookresearch/vip">VIP</a>
|
| 41 |
+
</summary>
|
| 42 |
+
<p>
|
| 43 |
+
|
| 44 |
+
```commandline
|
| 45 |
+
git clone https://github.com/facebookresearch/vip.git
|
| 46 |
+
cd vip && pip install -e .
|
| 47 |
+
python -c "from vip import load_vip; vip = load_vip()"
|
| 48 |
+
```
|
| 49 |
+
|
| 50 |
+
</p>
|
| 51 |
+
</details>
|
| 52 |
+
|
| 53 |
+
<details><summary>
|
| 54 |
+
<a href="https://github.com/facebookresearch/r3m">R3M</a>
|
| 55 |
+
</summary>
|
| 56 |
+
<p>
|
| 57 |
+
|
| 58 |
+
```commandline
|
| 59 |
+
git clone https://github.com/facebookresearch/r3m.git
|
| 60 |
+
cd r3m && pip install -e .
|
| 61 |
+
python -c "from r3m import load_r3m; r3m = load_r3m('resnet50')"
|
| 62 |
+
```
|
| 63 |
+
|
| 64 |
+
</p>
|
| 65 |
+
</details>
|
| 66 |
+
|
| 67 |
+
<details><summary>
|
| 68 |
+
<a href="https://github.com/penn-pal-lab/LIV">LIV (& CLIP)</a>
|
| 69 |
+
</summary>
|
| 70 |
+
<p>
|
| 71 |
+
|
| 72 |
+
```commandline
|
| 73 |
+
git clone https://github.com/penn-pal-lab/LIV.git
|
| 74 |
+
cd LIV && pip install -e . && cd liv/models/clip && pip install -e .
|
| 75 |
+
python -c "from liv import load_liv; liv = load_liv()"
|
| 76 |
+
```
|
| 77 |
+
|
| 78 |
+
</p>
|
| 79 |
+
</details>
|
| 80 |
+
|
| 81 |
+
|
| 82 |
+
<details><summary>
|
| 83 |
+
<a href="https://github.com/facebookresearch/eai-vc">VC1</a>
|
| 84 |
+
</summary>
|
| 85 |
+
<p>
|
| 86 |
+
|
| 87 |
+
```commandline
|
| 88 |
+
git clone https://github.com/facebookresearch/eai-vc.git
|
| 89 |
+
cd eai-vc && pip install -e vc_models
|
| 90 |
+
```
|
| 91 |
+
|
| 92 |
+
</p>
|
| 93 |
+
</details>
|
| 94 |
+
|
| 95 |
+
<details><summary>
|
| 96 |
+
<a href="https://github.com/facebookresearch/dinov2">DINOv2</a> and <a href="https://pytorch.org/vision/main/models/generated/torchvision.models.resnet50.html">ResNet</a> pretrained with ImageNet-1k are directly loaded via <a href="https://pytorch.org/hub/">torch hub</a> and <a href="https://pytorch.org/vision/main/models/generated/torchvision.models.resnet50.html">torchvision</a>.
|
| 97 |
+
</summary></details>
|
| 98 |
+
|
| 99 |
+
- Under *this* UVD repo directory, install other dependencies
|
| 100 |
+
```commandline
|
| 101 |
+
pip install -e .
|
| 102 |
+
```
|
| 103 |
+
|
| 104 |
+
# Usage
|
| 105 |
+
|
| 106 |
+
We provide a simple API for decompose RGB videos:
|
| 107 |
+
|
| 108 |
+
```python
|
| 109 |
+
import torch
|
| 110 |
+
import uvd
|
| 111 |
+
|
| 112 |
+
# (N sub-goals, *video frame shape)
|
| 113 |
+
subgoals = uvd.get_uvd_subgoals(
|
| 114 |
+
"/PATH/TO/VIDEO.*", # video filename or (L, *video frame shape) video numpy array
|
| 115 |
+
preprocessor_name="vip", # Literal["vip", "r3m", "liv", "clip", "vc1", "dinov2"]
|
| 116 |
+
device="cuda" if torch.cuda.is_available() else "cpu", # device for loading frozen preprocessor
|
| 117 |
+
return_indices=False, # True if only want the list of subgoal timesteps
|
| 118 |
+
)
|
| 119 |
+
```
|
| 120 |
+
|
| 121 |
+
or run
|
| 122 |
+
```commandline
|
| 123 |
+
python demo.py
|
| 124 |
+
```
|
| 125 |
+
to host a Gradio demo locally with different choices of visual representations.
|
| 126 |
+
|
| 127 |
+
## Simulation Data
|
| 128 |
+
|
| 129 |
+
We post-processed the data released from original [Relay-Policy-Learning](https://github.com/google-research/relay-policy-learning/tree/master) that keeps the successful trajectories only and adapt the control and observations used in our paper by:
|
| 130 |
+
```commandline
|
| 131 |
+
python datasets/data_gen.py raw_data_path=/PATH/TO/RAW_DATA
|
| 132 |
+
```
|
| 133 |
+
|
| 134 |
+
Also consider to force set `Builder = LinuxCPUExtensionBuilder` to `Builder = LinuxGPUExtensionBuilder` in `PATH/TO/CONDA/envs/uvd/lib/python3.9/site-packages/mujoco_py/builder.py` to enable (multi-)GPU acceleration.
|
| 135 |
+
|
| 136 |
+
|
| 137 |
+
## Runtime Benchmark
|
| 138 |
+
|
| 139 |
+
Since UVD's goal is to be an off-the-shelf method applying to *any* existing policy learning frameworks and models, across BC and RL, we provide minimal scripts for benchmarking the runtime showing negligible runtime under `./scripts` directory:
|
| 140 |
+
```commandline
|
| 141 |
+
python scripts/benchmark_decomp.py /PATH/TO/VIDEO
|
| 142 |
+
```
|
| 143 |
+
and passing `--preprocessor_name` with other preprocessors (default `vip`) and `--n` for the number of repeated iterations (default `100`).
|
| 144 |
+
|
| 145 |
+
For inference or rollouts, we benchmark the runtime by
|
| 146 |
+
```commandline
|
| 147 |
+
python scripts/benchmark_inference.py
|
| 148 |
+
```
|
| 149 |
+
and passing `--policy` for using MLP or causal GPT policy; `--preprocessor_name` with other preprocessors (default `vip`); `--use_uvd` as boolean arg for whether using UVD or no decomposition (i.e. final goal conditioned); and `--n` for the number of repeated iterations (default `100`). The default episode horizon is set to 300. We found that running in the terminal would be almost 2s slower every episode than directly running with python IDE (e.g. PyCharm, under the script directory and run as script instead of module), but the general trend that including UVD introduces negligible extra runtime still holds true.
|
| 150 |
+
|
| 151 |
+
# Citation
|
| 152 |
+
If you find this project useful in your research, please consider citing:
|
| 153 |
|
| 154 |
+
```bibtex
|
| 155 |
+
@inproceedings{zhang2024universal,
|
| 156 |
+
title={Universal visual decomposer: Long-horizon manipulation made easy},
|
| 157 |
+
author={Zhang, Zichen and Li, Yunshuang and Bastani, Osbert and Gupta, Abhishek and Jayaraman, Dinesh and Ma, Yecheng Jason and Weihs, Luca},
|
| 158 |
+
booktitle={2024 IEEE International Conference on Robotics and Automation (ICRA)},
|
| 159 |
+
pages={6973--6980},
|
| 160 |
+
year={2024},
|
| 161 |
+
organization={IEEE}
|
| 162 |
+
}
|
| 163 |
+
```
|
app.py
ADDED
|
@@ -0,0 +1,34 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
import gradio as gr
|
| 2 |
+
import torch
|
| 3 |
+
import uvd
|
| 4 |
+
import decord
|
| 5 |
+
|
| 6 |
+
|
| 7 |
+
def proc_video(video, preprocessor_name):
|
| 8 |
+
frame_original_res = decord.VideoReader(video)[:].asnumpy()
|
| 9 |
+
indices = uvd.get_uvd_subgoals(
|
| 10 |
+
video, preprocessor_name.lower().replace("-", ""),
|
| 11 |
+
device="cuda" if torch.cuda.is_available() else "cpu",
|
| 12 |
+
return_indices=True,
|
| 13 |
+
)
|
| 14 |
+
subgoals = frame_original_res[indices]
|
| 15 |
+
return [(img, f"No. {i+1} subgoal") for i, img in enumerate(subgoals)]
|
| 16 |
+
|
| 17 |
+
|
| 18 |
+
with gr.Blocks() as demo:
|
| 19 |
+
with gr.Row():
|
| 20 |
+
input_video = gr.Video(height=224, scale=3)
|
| 21 |
+
preprocessor_name = gr.Dropdown(
|
| 22 |
+
["VIP", "R3M", "LIV", "CLIP", "DINO-v2", "VC-1", "ResNet"],
|
| 23 |
+
label="Preprocessor",
|
| 24 |
+
value="VIP",
|
| 25 |
+
scale=1,
|
| 26 |
+
)
|
| 27 |
+
output = gr.Gallery(label="UVD SubGoals", height=224, preview=True, scale=4)
|
| 28 |
+
with gr.Row():
|
| 29 |
+
submit = gr.Button("Submit")
|
| 30 |
+
clr = gr.ClearButton(components=[input_video, output])
|
| 31 |
+
submit.click(proc_video, inputs=[input_video, preprocessor_name], outputs=[output])
|
| 32 |
+
|
| 33 |
+
|
| 34 |
+
demo.queue().launch(share=True, show_error=True)
|
datasets/.gitignore
ADDED
|
@@ -0,0 +1,5 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
*
|
| 2 |
+
!.gitignore
|
| 3 |
+
!data_gen.py
|
| 4 |
+
!generate_in_domain_vip_ft_data.py
|
| 5 |
+
!data_gen.yaml
|
datasets/data_gen.py
ADDED
|
@@ -0,0 +1,317 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
from __future__ import annotations
|
| 2 |
+
|
| 3 |
+
import copy
|
| 4 |
+
import glob
|
| 5 |
+
import multiprocessing as mp
|
| 6 |
+
import shutil
|
| 7 |
+
from multiprocessing import get_logger
|
| 8 |
+
from typing import Any
|
| 9 |
+
|
| 10 |
+
import hydra
|
| 11 |
+
import imageio
|
| 12 |
+
import numpy as np
|
| 13 |
+
import tqdm
|
| 14 |
+
|
| 15 |
+
import uvd.utils as U
|
| 16 |
+
from uvd.envs.franka_kitchen import KitchenBase
|
| 17 |
+
|
| 18 |
+
|
| 19 |
+
def parse_task_name(task_name: str):
|
| 20 |
+
sub_tasks = task_name.split("-")[1].split(",")
|
| 21 |
+
parsed_tasks = []
|
| 22 |
+
for sub_task in sub_tasks:
|
| 23 |
+
if sub_task == "bottomknob":
|
| 24 |
+
parsed_tasks.append("bottom burner")
|
| 25 |
+
elif sub_task == "topknob":
|
| 26 |
+
parsed_tasks.append("top burner")
|
| 27 |
+
elif sub_task == "switch":
|
| 28 |
+
parsed_tasks.append("light switch")
|
| 29 |
+
elif sub_task == "slide":
|
| 30 |
+
parsed_tasks.append("slide cabinet")
|
| 31 |
+
elif sub_task == "hinge":
|
| 32 |
+
parsed_tasks.append("hinge cabinet")
|
| 33 |
+
else:
|
| 34 |
+
parsed_tasks.append(sub_task)
|
| 35 |
+
return "-".join(parsed_tasks)
|
| 36 |
+
|
| 37 |
+
|
| 38 |
+
def _generate_one_task_demo(
|
| 39 |
+
*,
|
| 40 |
+
raw_data_path: str,
|
| 41 |
+
raw_task_dir: str,
|
| 42 |
+
# same seeds for diff tasks for now
|
| 43 |
+
seed_start: int,
|
| 44 |
+
output_path: str,
|
| 45 |
+
max_demos: int | None = None,
|
| 46 |
+
frame_height: int = 256,
|
| 47 |
+
frame_width: int = 256,
|
| 48 |
+
terminal_if_done: bool = True,
|
| 49 |
+
include_no_robot: bool = False,
|
| 50 |
+
max_reject_sampling: int = 10,
|
| 51 |
+
render: bool = False,
|
| 52 |
+
save_mp4: bool = False,
|
| 53 |
+
mp4_fps: int | None = None,
|
| 54 |
+
copy_raw_data: bool = False,
|
| 55 |
+
env: KitchenBase | None = None,
|
| 56 |
+
gpu_id,
|
| 57 |
+
):
|
| 58 |
+
task_path = U.f_join(raw_data_path, raw_task_dir)
|
| 59 |
+
parsed_task_name = parse_task_name(raw_task_dir)
|
| 60 |
+
trajectories = glob.glob(U.f_join(task_path, "*.pkl"))
|
| 61 |
+
assert len(trajectories) > 0, task_path
|
| 62 |
+
# include space, connect goals with "-"
|
| 63 |
+
task_elements = parsed_task_name.split("-")
|
| 64 |
+
n_demos = min(max_demos or len(trajectories), len(trajectories))
|
| 65 |
+
saved_episode_dir = f"{parsed_task_name.replace(' ', '_')}"
|
| 66 |
+
episode_path = U.f_mkdir(output_path, saved_episode_dir)
|
| 67 |
+
|
| 68 |
+
close_env = False
|
| 69 |
+
if env is None:
|
| 70 |
+
env = KitchenBase(
|
| 71 |
+
task_elements=task_elements,
|
| 72 |
+
frame_height=frame_height,
|
| 73 |
+
frame_width=frame_width,
|
| 74 |
+
obs_keys=("rgb", "proprio"),
|
| 75 |
+
gpu_id=gpu_id,
|
| 76 |
+
)
|
| 77 |
+
close_env = True
|
| 78 |
+
|
| 79 |
+
num_demos_gen = 0
|
| 80 |
+
seed = seed_start - 1
|
| 81 |
+
traj_length_map = {}
|
| 82 |
+
unsuccessful_traj = []
|
| 83 |
+
for i, traj in tqdm.tqdm(
|
| 84 |
+
enumerate(trajectories[:n_demos]), total=n_demos, desc=parsed_task_name
|
| 85 |
+
):
|
| 86 |
+
data = U.load_pickle(traj)
|
| 87 |
+
path = data["path"]
|
| 88 |
+
actions = path["actions"]
|
| 89 |
+
ctrls = data["ctrl"]
|
| 90 |
+
raw_data_length, action_dim = actions.shape
|
| 91 |
+
assert action_dim == env.action_space.shape[0]
|
| 92 |
+
|
| 93 |
+
data_dict, info = {}, {} # placeholders
|
| 94 |
+
for _ in range(max_reject_sampling):
|
| 95 |
+
seed += 1
|
| 96 |
+
env.reset(init_qpos=data["qpos"][0], init_qvel=data["qvel"][0])
|
| 97 |
+
obs = env.reset(task_elements=task_elements, seed=seed)
|
| 98 |
+
obs_dict = env.obs_dict
|
| 99 |
+
# slightly difference than directly using
|
| 100 |
+
init_qpos = env.init_qpos.copy()
|
| 101 |
+
init_qvel = env.init_qvel.copy()
|
| 102 |
+
if render:
|
| 103 |
+
env.render()
|
| 104 |
+
data_dict: dict[str, Any] = {
|
| 105 |
+
"reset_kwargs": dict(
|
| 106 |
+
task_elements=task_elements,
|
| 107 |
+
init_qpos=init_qpos,
|
| 108 |
+
init_qvel=init_qvel,
|
| 109 |
+
),
|
| 110 |
+
"obs_full": [obs], # {str: L+1, h, w, 3}
|
| 111 |
+
"actions": [], # L, A
|
| 112 |
+
"rewards": [], # L,
|
| 113 |
+
"dones": [], # L,
|
| 114 |
+
"completed_tasks": [], # L, x
|
| 115 |
+
"scores": [], # L,
|
| 116 |
+
"seed": seed,
|
| 117 |
+
"oracle_milestones": [], # n h, w, 3
|
| 118 |
+
}
|
| 119 |
+
if include_no_robot:
|
| 120 |
+
data_dict["no_robot_rgb"] = [
|
| 121 |
+
env.render(mode="rgb_array", set_robot_alpha=0.0)
|
| 122 |
+
]
|
| 123 |
+
for s in tqdm.trange(raw_data_length):
|
| 124 |
+
# Construct the action
|
| 125 |
+
ctrl = (ctrls[s] - obs_dict["qp"]) / (
|
| 126 |
+
env.frame_skip * env.model.opt.timestep
|
| 127 |
+
)
|
| 128 |
+
act = (ctrl - env.act_mid) / env.act_amp
|
| 129 |
+
act = np.clip(act, -1.0, 1.0)
|
| 130 |
+
obs, reward, done, info = env.step(act)
|
| 131 |
+
if render:
|
| 132 |
+
env.render()
|
| 133 |
+
obs_dict = info["obs_dict"]
|
| 134 |
+
# 1 if complete a new goal this step
|
| 135 |
+
score = info["score"]
|
| 136 |
+
# num goals achieved so far
|
| 137 |
+
completed_tasks = info["completed_tasks"]
|
| 138 |
+
if score == 1:
|
| 139 |
+
data_dict["oracle_milestones"].append(obs["rgb"])
|
| 140 |
+
|
| 141 |
+
data_dict["obs_full"].append(obs)
|
| 142 |
+
data_dict["actions"].append(np.array(act))
|
| 143 |
+
data_dict["rewards"].append(float(reward))
|
| 144 |
+
data_dict["dones"].append(int(done))
|
| 145 |
+
data_dict["scores"].append(float(score))
|
| 146 |
+
data_dict["completed_tasks"].append(np.array(completed_tasks))
|
| 147 |
+
|
| 148 |
+
if include_no_robot:
|
| 149 |
+
data_dict["no_robot_rgb"].append(
|
| 150 |
+
env.render(mode="rgb_array", set_robot_alpha=0.0)
|
| 151 |
+
)
|
| 152 |
+
|
| 153 |
+
if terminal_if_done and done:
|
| 154 |
+
break
|
| 155 |
+
if len(env.tasks_to_complete) != 0:
|
| 156 |
+
get_logger().warning(
|
| 157 |
+
f"reject sampling, with task left: {env.tasks_to_complete}, "
|
| 158 |
+
f"distance left: {info['rewards']['distances_left']}"
|
| 159 |
+
)
|
| 160 |
+
continue
|
| 161 |
+
else:
|
| 162 |
+
break
|
| 163 |
+
if len(data_dict["dones"]) == 0 or data_dict["dones"][-1] != 1:
|
| 164 |
+
# raise ValueError(traj)
|
| 165 |
+
get_logger().warning(
|
| 166 |
+
f"unsuccessful data {traj}, distance left: {info['rewards']['distances_left']}"
|
| 167 |
+
)
|
| 168 |
+
debug_dir = U.f_mkdir(U.f_join(episode_path, "reject_sampling"))
|
| 169 |
+
# noinspection PyUnboundLocalVariable
|
| 170 |
+
imageio.imsave(U.f_join(debug_dir, f"{task_elements}_{i}.png"), obs["rgb"])
|
| 171 |
+
unsuccessful_traj.append(traj)
|
| 172 |
+
continue
|
| 173 |
+
|
| 174 |
+
traj_length = len(data_dict["dones"])
|
| 175 |
+
traj_length_map[i] = traj_length
|
| 176 |
+
# save accurate data
|
| 177 |
+
data_dict["obs_full"] = U.batch_observations(
|
| 178 |
+
data_dict["obs_full"], to_tensor=False
|
| 179 |
+
)
|
| 180 |
+
data_dict["actions"] = np.stack(data_dict["actions"])
|
| 181 |
+
data_dict["rewards"] = np.array(data_dict["rewards"])
|
| 182 |
+
data_dict["dones"] = np.array(data_dict["dones"])
|
| 183 |
+
data_dict["scores"] = np.array(data_dict["scores"])
|
| 184 |
+
data_dict["completed_tasks"] = np.array(data_dict["completed_tasks"])
|
| 185 |
+
data_dict["ctrls"] = ctrls[:traj_length]
|
| 186 |
+
data_dict["last_distances_to_goal"] = info["rewards"]["distances_left"]
|
| 187 |
+
data_dict["oracle_milestones"] = np.stack(data_dict["oracle_milestones"])
|
| 188 |
+
if include_no_robot:
|
| 189 |
+
data_dict["no_robot_rgb"] = np.stack(data_dict["no_robot_rgb"])
|
| 190 |
+
assert data_dict["obs_full"]["rgb"].shape == data_dict["no_robot_rgb"].shape
|
| 191 |
+
|
| 192 |
+
assert (
|
| 193 |
+
len(data_dict["obs_full"]["rgb"]) - 1
|
| 194 |
+
== len(data_dict["obs_full"]["proprio"]) - 1
|
| 195 |
+
== len(data_dict["actions"])
|
| 196 |
+
== len(data_dict["rewards"])
|
| 197 |
+
== len(data_dict["scores"])
|
| 198 |
+
== len(data_dict["completed_tasks"])
|
| 199 |
+
== len(data_dict["ctrls"])
|
| 200 |
+
== traj_length
|
| 201 |
+
)
|
| 202 |
+
assert len(data_dict["oracle_milestones"]) == len(task_elements)
|
| 203 |
+
|
| 204 |
+
num_demos_gen += 1
|
| 205 |
+
U.save_pickle(data_dict, U.f_join(episode_path, f"episode_{i}.pkl"))
|
| 206 |
+
if copy_raw_data:
|
| 207 |
+
out_raw_path = U.f_mkdir(output_path, "raw_data", saved_episode_dir)
|
| 208 |
+
shutil.copy(traj, U.f_join(out_raw_path, traj.split("/")[-1]))
|
| 209 |
+
if save_mp4:
|
| 210 |
+
video_dir = U.f_mkdir(U.f_join(episode_path, "videos"))
|
| 211 |
+
U.save_video(
|
| 212 |
+
video=data_dict["obs_full"]["rgb"],
|
| 213 |
+
fname=U.f_join(video_dir, f"episode_{i}.mp4"),
|
| 214 |
+
fps=mp4_fps,
|
| 215 |
+
compress=True,
|
| 216 |
+
)
|
| 217 |
+
if include_no_robot:
|
| 218 |
+
U.save_video(
|
| 219 |
+
video=data_dict["no_robot_rgb"],
|
| 220 |
+
fname=U.f_join(video_dir, f"episode_{i}_no_robot.mp4"),
|
| 221 |
+
fps=mp4_fps,
|
| 222 |
+
compress=True,
|
| 223 |
+
)
|
| 224 |
+
U.dump_json(
|
| 225 |
+
dict(
|
| 226 |
+
task_elements=task_elements,
|
| 227 |
+
seed_start=seed_start,
|
| 228 |
+
action_dim=env.action_space.shape[0],
|
| 229 |
+
frame_height=frame_height,
|
| 230 |
+
frame_width=frame_width,
|
| 231 |
+
num_demos=num_demos_gen,
|
| 232 |
+
trajectory_lengths=traj_length_map,
|
| 233 |
+
unsuccessful_traj=unsuccessful_traj,
|
| 234 |
+
),
|
| 235 |
+
episode_path,
|
| 236 |
+
f"metadata.json",
|
| 237 |
+
)
|
| 238 |
+
if close_env:
|
| 239 |
+
env.close()
|
| 240 |
+
|
| 241 |
+
|
| 242 |
+
def generate_one_task_demo(kwargs):
|
| 243 |
+
_generate_one_task_demo(**kwargs)
|
| 244 |
+
|
| 245 |
+
|
| 246 |
+
@hydra.main(config_path=".", config_name="data_gen", version_base="1.1")
|
| 247 |
+
def main(cfg):
|
| 248 |
+
if cfg.debug:
|
| 249 |
+
cfg = copy.deepcopy(cfg)
|
| 250 |
+
cfg.num_processes = 1
|
| 251 |
+
cfg.output_path += "_debug"
|
| 252 |
+
# cfg.render = True
|
| 253 |
+
cfg.max_demos = 3
|
| 254 |
+
U.ask_if_overwrite(cfg.output_path)
|
| 255 |
+
if cfg.dm_backend or (cfg.render and (not cfg.mp or cfg.num_processes == 1)):
|
| 256 |
+
import adept_envs.mujoco_env
|
| 257 |
+
|
| 258 |
+
adept_envs.mujoco_env.USE_DM_CONTROL = True
|
| 259 |
+
|
| 260 |
+
raw_data_path = cfg.raw_data_path
|
| 261 |
+
raw_task_dirs = U.f_listdir(raw_data_path)
|
| 262 |
+
assert raw_task_dirs
|
| 263 |
+
if cfg.mp:
|
| 264 |
+
num_processes = min(mp.cpu_count() - 2, cfg.num_processes or len(raw_task_dirs))
|
| 265 |
+
_, num_gpus = U.parse_gpu_devices(cfg.gpus)
|
| 266 |
+
with mp.Pool(num_processes) as pool:
|
| 267 |
+
pool.map(
|
| 268 |
+
generate_one_task_demo,
|
| 269 |
+
[
|
| 270 |
+
dict(
|
| 271 |
+
raw_data_path=cfg.raw_data_path,
|
| 272 |
+
raw_task_dir=raw_task_dirs[i],
|
| 273 |
+
seed_start=cfg.seed_start,
|
| 274 |
+
output_path=cfg.output_path,
|
| 275 |
+
max_demos=cfg.max_demos,
|
| 276 |
+
frame_height=cfg.frame_height,
|
| 277 |
+
frame_width=cfg.frame_width,
|
| 278 |
+
terminal_if_done=cfg.terminal_if_done,
|
| 279 |
+
max_reject_sampling=cfg.max_reject_sampling,
|
| 280 |
+
include_no_robot=cfg.include_no_robot,
|
| 281 |
+
render=cfg.render,
|
| 282 |
+
save_mp4=cfg.save_mp4,
|
| 283 |
+
mp4_fps=cfg.mp4_fps,
|
| 284 |
+
copy_raw_data=cfg.copy_raw_data,
|
| 285 |
+
env=None,
|
| 286 |
+
gpu_id=i % num_gpus,
|
| 287 |
+
)
|
| 288 |
+
for i in range(len(raw_task_dirs))
|
| 289 |
+
],
|
| 290 |
+
)
|
| 291 |
+
else:
|
| 292 |
+
env = KitchenBase(frame_height=cfg.frame_height, frame_width=cfg.frame_width)
|
| 293 |
+
for raw_task_dir in tqdm.tqdm(
|
| 294 |
+
raw_task_dirs, desc="Generate Franka Kitchen Demos"
|
| 295 |
+
):
|
| 296 |
+
_generate_one_task_demo(
|
| 297 |
+
raw_data_path=cfg.raw_data_path,
|
| 298 |
+
raw_task_dir=raw_task_dir,
|
| 299 |
+
seed_start=cfg.seed_start,
|
| 300 |
+
output_path=cfg.output_path,
|
| 301 |
+
max_demos=cfg.max_demos,
|
| 302 |
+
frame_height=cfg.frame_height,
|
| 303 |
+
frame_width=cfg.frame_width,
|
| 304 |
+
terminal_if_done=cfg.terminal_if_done,
|
| 305 |
+
max_reject_sampling=cfg.max_reject_sampling,
|
| 306 |
+
include_no_robot=cfg.include_no_robot,
|
| 307 |
+
render=cfg.render,
|
| 308 |
+
save_mp4=cfg.save_mp4,
|
| 309 |
+
mp4_fps=cfg.mp4_fps,
|
| 310 |
+
copy_raw_data=cfg.copy_raw_data,
|
| 311 |
+
env=env,
|
| 312 |
+
gpu_id=-1,
|
| 313 |
+
)
|
| 314 |
+
|
| 315 |
+
|
| 316 |
+
if __name__ == "__main__":
|
| 317 |
+
main()
|
datasets/data_gen.yaml
ADDED
|
@@ -0,0 +1,29 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
raw_data_path: datasets/franka_kitchen_raw_dataset/full
|
| 2 |
+
|
| 3 |
+
output_path: datasets/franka_kitchen_demos
|
| 4 |
+
max_demos: null # null for all
|
| 5 |
+
frame_height: 224
|
| 6 |
+
frame_width: 224
|
| 7 |
+
terminal_if_done: false
|
| 8 |
+
max_reject_sampling: 1
|
| 9 |
+
include_no_robot: false
|
| 10 |
+
render: false
|
| 11 |
+
save_mp4: true
|
| 12 |
+
mp4_fps: null
|
| 13 |
+
copy_raw_data: false
|
| 14 |
+
|
| 15 |
+
debug: false
|
| 16 |
+
|
| 17 |
+
seed_start: 42
|
| 18 |
+
dm_backend: false
|
| 19 |
+
mp: true
|
| 20 |
+
num_processes: null
|
| 21 |
+
gpus: null
|
| 22 |
+
|
| 23 |
+
|
| 24 |
+
hydra:
|
| 25 |
+
job:
|
| 26 |
+
chdir: true
|
| 27 |
+
run:
|
| 28 |
+
dir: .
|
| 29 |
+
output_subdir: null
|
requirements.txt
ADDED
|
@@ -0,0 +1,20 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
wheel==0.38.4
|
| 2 |
+
setuptools==65.5.0
|
| 3 |
+
torch==2.2.2
|
| 4 |
+
torchvision==0.17.2
|
| 5 |
+
cython<3
|
| 6 |
+
dm_control==1.0.11
|
| 7 |
+
einops==0.6.0
|
| 8 |
+
gdown==4.7.1
|
| 9 |
+
gym==0.25.2
|
| 10 |
+
hydra-core==1.3.1
|
| 11 |
+
mujoco-py<2.2,>=2.1
|
| 12 |
+
numpy==1.23.5
|
| 13 |
+
pytorch_lightning==2.0.0
|
| 14 |
+
scikit-learn
|
| 15 |
+
shimmy==0.2.0
|
| 16 |
+
wandb==0.14.0
|
| 17 |
+
termcolor
|
| 18 |
+
decord==0.6.0
|
| 19 |
+
gradio>=4.0.0
|
| 20 |
+
allenact@git+https://github.com/allenai/allenact.git@usd#egg=allenact&subdirectory=allenact
|
scripts/__init__.py
ADDED
|
File without changes
|
scripts/auto_format.sh
ADDED
|
@@ -0,0 +1,15 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
#!/bin/bash
|
| 2 |
+
|
| 3 |
+
# Move to the directory containing the directory that this file is in
|
| 4 |
+
cd "$( cd "$( dirname "${BASH_SOURCE[0]}/.." )" >/dev/null 2>&1 && pwd )" || exit
|
| 5 |
+
|
| 6 |
+
echo RUNNING BLACK
|
| 7 |
+
black . --exclude src --exclude external_projects
|
| 8 |
+
echo BLACK DONE
|
| 9 |
+
echo ""
|
| 10 |
+
|
| 11 |
+
echo RUNNING DOCFORMATTER
|
| 12 |
+
find . -name "*.py" | grep -v ^./src | grep -v ^./external_projects | grep -v used_configs | xargs docformatter --in-place -r
|
| 13 |
+
echo DOCFORMATTER DONE
|
| 14 |
+
|
| 15 |
+
echo ALL DONE
|
scripts/benchmark_decomp.py
ADDED
|
@@ -0,0 +1,63 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
import argparse
|
| 2 |
+
import warnings
|
| 3 |
+
|
| 4 |
+
warnings.filterwarnings("ignore", category=UserWarning)
|
| 5 |
+
|
| 6 |
+
import torch
|
| 7 |
+
import tqdm
|
| 8 |
+
|
| 9 |
+
import numpy as np
|
| 10 |
+
import time
|
| 11 |
+
|
| 12 |
+
from uvd.decomp.decomp import embedding_decomp, DEFAULT_DECOMP_KWARGS
|
| 13 |
+
from uvd.models import Preprocessor, get_preprocessor
|
| 14 |
+
import uvd.utils as U
|
| 15 |
+
|
| 16 |
+
from decord import VideoReader
|
| 17 |
+
|
| 18 |
+
|
| 19 |
+
def one_run(video_file: str, preprocessor: Preprocessor):
|
| 20 |
+
t = time.time()
|
| 21 |
+
vr = VideoReader(U.f_expand(video_file), height=224, width=224)
|
| 22 |
+
videos = vr[:].asnumpy()
|
| 23 |
+
load_v_t = time.time() - t
|
| 24 |
+
|
| 25 |
+
t = time.time()
|
| 26 |
+
embeddings = preprocessor.process(videos, return_numpy=True)
|
| 27 |
+
preprocess_t = time.time() - t
|
| 28 |
+
|
| 29 |
+
t = time.time()
|
| 30 |
+
_, decomp_meta = embedding_decomp(
|
| 31 |
+
embeddings=embeddings,
|
| 32 |
+
fill_embeddings=False,
|
| 33 |
+
return_intermediate_curves=False,
|
| 34 |
+
window_length=100,
|
| 35 |
+
**DEFAULT_DECOMP_KWARGS["embed"],
|
| 36 |
+
)
|
| 37 |
+
decomp_t = time.time() - t
|
| 38 |
+
return load_v_t, preprocess_t, decomp_t
|
| 39 |
+
|
| 40 |
+
|
| 41 |
+
if __name__ == "__main__":
|
| 42 |
+
parser = argparse.ArgumentParser()
|
| 43 |
+
parser.add_argument("video_file")
|
| 44 |
+
parser.add_argument("--preprocessor_name", default="vip")
|
| 45 |
+
parser.add_argument("--n", type=int, default=100)
|
| 46 |
+
args = parser.parse_args()
|
| 47 |
+
|
| 48 |
+
use_gpu = torch.cuda.is_available()
|
| 49 |
+
if not use_gpu:
|
| 50 |
+
print("NO GPU FOUND")
|
| 51 |
+
preprocessor = get_preprocessor(
|
| 52 |
+
args.preprocessor_name, device="cuda" if use_gpu else None
|
| 53 |
+
)
|
| 54 |
+
one_run(args.video_file, preprocessor)
|
| 55 |
+
|
| 56 |
+
benchmark_times = dict(load=[], preprocess=[], decomp=[])
|
| 57 |
+
for _ in tqdm.trange(args.n):
|
| 58 |
+
load_v_t, preprocess_t, decomp_t = one_run(args.video_file, preprocessor)
|
| 59 |
+
benchmark_times["load"].append(load_v_t)
|
| 60 |
+
benchmark_times["preprocess"].append(preprocess_t)
|
| 61 |
+
benchmark_times["decomp"].append(decomp_t)
|
| 62 |
+
|
| 63 |
+
print({k: (np.mean(v), np.std(v)) for k, v in benchmark_times.items()})
|
scripts/benchmark_inference.py
ADDED
|
@@ -0,0 +1,163 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
import argparse
|
| 2 |
+
import copy
|
| 3 |
+
import time
|
| 4 |
+
|
| 5 |
+
import gym
|
| 6 |
+
import numpy as np
|
| 7 |
+
import torch
|
| 8 |
+
import yaml
|
| 9 |
+
from omegaconf import DictConfig
|
| 10 |
+
|
| 11 |
+
import uvd.utils as U
|
| 12 |
+
from uvd.models.preprocessors import get_preprocessor
|
| 13 |
+
from uvd.decomp.decomp import embedding_decomp, DEFAULT_DECOMP_KWARGS
|
| 14 |
+
from uvd.envs.evaluator.inference_wrapper import InferenceWrapper
|
| 15 |
+
from uvd.envs.franka_kitchen.franka_kitchen_base import KitchenBase
|
| 16 |
+
|
| 17 |
+
MLP_CFG = """\
|
| 18 |
+
policy:
|
| 19 |
+
_target_: uvd.models.policy.MLPPolicy
|
| 20 |
+
observation_space: ???
|
| 21 |
+
action_space: ???
|
| 22 |
+
preprocessor: ???
|
| 23 |
+
obs_encoder:
|
| 24 |
+
__target__: uvd.models.nn.MLP
|
| 25 |
+
hidden_dims: [1024, 512, 256]
|
| 26 |
+
activation: ReLU
|
| 27 |
+
normalization: false
|
| 28 |
+
input_normalization: BatchNorm1d
|
| 29 |
+
input_normalization_full_obs: false
|
| 30 |
+
proprio_output_dim: 512
|
| 31 |
+
proprio_add_layernorm: true
|
| 32 |
+
proprio_activation: Tanh
|
| 33 |
+
proprio_add_noise_eval: false
|
| 34 |
+
actor_act: Tanh
|
| 35 |
+
act_head:
|
| 36 |
+
__target__: uvd.models.distributions.DeterministicHead
|
| 37 |
+
"""
|
| 38 |
+
|
| 39 |
+
GPT_CFG = """\
|
| 40 |
+
policy:
|
| 41 |
+
_target_: uvd.models.policy.GPTPolicy
|
| 42 |
+
observation_space: ???
|
| 43 |
+
action_space: ???
|
| 44 |
+
preprocessor: ???
|
| 45 |
+
use_kv_cache: true
|
| 46 |
+
max_seq_length: 10
|
| 47 |
+
obs_add: false
|
| 48 |
+
proprio_hidden_dim: 512
|
| 49 |
+
obs_encoder:
|
| 50 |
+
__target__: uvd.models.nn.GPT
|
| 51 |
+
use_wte: true
|
| 52 |
+
gpt_config:
|
| 53 |
+
block_size: 10
|
| 54 |
+
vocab_size: null
|
| 55 |
+
n_embd: 768
|
| 56 |
+
n_layer: 8
|
| 57 |
+
n_head: 8
|
| 58 |
+
dropout: 0.1
|
| 59 |
+
bias: false
|
| 60 |
+
use_llama_impl: true
|
| 61 |
+
position_embed: rotary
|
| 62 |
+
act_head:
|
| 63 |
+
__target__: uvd.models.distributions.DeterministicHead
|
| 64 |
+
"""
|
| 65 |
+
|
| 66 |
+
if __name__ == "__main__":
|
| 67 |
+
parser = argparse.ArgumentParser()
|
| 68 |
+
parser.add_argument("--policy", default="gpt")
|
| 69 |
+
parser.add_argument("--preprocessor_name", default="vip")
|
| 70 |
+
parser.add_argument("--use_uvd", action="store_true")
|
| 71 |
+
parser.add_argument("--n", type=int, default=100)
|
| 72 |
+
args = parser.parse_args()
|
| 73 |
+
|
| 74 |
+
use_gpu = torch.cuda.is_available()
|
| 75 |
+
if not use_gpu:
|
| 76 |
+
print("NO GPU FOUND")
|
| 77 |
+
preprocessor = get_preprocessor(
|
| 78 |
+
args.preprocessor_name, device="cuda" if use_gpu else None
|
| 79 |
+
)
|
| 80 |
+
policy_name = args.policy.lower()
|
| 81 |
+
assert policy_name in ["mlp", "gpt"]
|
| 82 |
+
is_causal = policy_name == "gpt"
|
| 83 |
+
|
| 84 |
+
env = KitchenBase(frame_height=224, frame_width=224)
|
| 85 |
+
env = InferenceWrapper(env, dummy_rtn=is_causal)
|
| 86 |
+
env.reset()
|
| 87 |
+
|
| 88 |
+
observation_space = gym.spaces.Dict(
|
| 89 |
+
rgb=gym.spaces.Box(-np.inf, np.inf, preprocessor.output_dim, np.float32),
|
| 90 |
+
proprio=gym.spaces.Box(-1, 1, (9,), np.float32),
|
| 91 |
+
milestones=gym.spaces.Box(
|
| 92 |
+
-np.inf, np.inf, (6,) + preprocessor.output_dim, np.float32
|
| 93 |
+
),
|
| 94 |
+
)
|
| 95 |
+
action_space = env.action_space
|
| 96 |
+
|
| 97 |
+
cfg = yaml.safe_load(MLP_CFG if policy_name == "mlp" else GPT_CFG)
|
| 98 |
+
cfg = DictConfig(cfg)
|
| 99 |
+
policy = U.hydra_instantiate(
|
| 100 |
+
cfg.policy,
|
| 101 |
+
observation_space=observation_space,
|
| 102 |
+
action_space=action_space,
|
| 103 |
+
preprocessor=preprocessor,
|
| 104 |
+
)
|
| 105 |
+
policy = policy.to(preprocessor.device).eval()
|
| 106 |
+
U.debug_model_info(policy)
|
| 107 |
+
if is_causal:
|
| 108 |
+
assert policy.causal and policy.use_kv_cache
|
| 109 |
+
|
| 110 |
+
preprocessor = policy.preprocessor
|
| 111 |
+
# Or load FrankaKitchen dummy datas
|
| 112 |
+
dummy_data = np.random.random((300, 224, 224, 3)).astype(np.float32)
|
| 113 |
+
emb = preprocessor.process(dummy_data, return_numpy=True)
|
| 114 |
+
if args.use_uvd:
|
| 115 |
+
_, decomp_meta = embedding_decomp(
|
| 116 |
+
embeddings=emb,
|
| 117 |
+
fill_embeddings=False,
|
| 118 |
+
return_intermediate_curves=False,
|
| 119 |
+
**DEFAULT_DECOMP_KWARGS["embed"],
|
| 120 |
+
)
|
| 121 |
+
milestones = emb[decomp_meta.milestone_indices] # nhw3
|
| 122 |
+
else:
|
| 123 |
+
milestones = emb[-1][None, ...]
|
| 124 |
+
env.milestones = milestones
|
| 125 |
+
|
| 126 |
+
MAX_HORIZON = 300
|
| 127 |
+
totals = []
|
| 128 |
+
for _ in range(args.n):
|
| 129 |
+
obs = env.reset()
|
| 130 |
+
if is_causal:
|
| 131 |
+
policy.reset_cache()
|
| 132 |
+
|
| 133 |
+
times = []
|
| 134 |
+
for st in range(MAX_HORIZON):
|
| 135 |
+
t = time.time()
|
| 136 |
+
obs = copy.deepcopy(obs)
|
| 137 |
+
batchify_obs = U.batch_observations([obs], device=policy.device)
|
| 138 |
+
if is_causal:
|
| 139 |
+
# B, T, ...
|
| 140 |
+
cur_milestone = env.current_milestone[None, None, ...]
|
| 141 |
+
for k in batchify_obs:
|
| 142 |
+
batchify_obs[k] = batchify_obs[k][:, None, ...]
|
| 143 |
+
else:
|
| 144 |
+
# B, ...
|
| 145 |
+
cur_milestone = env.current_milestone[None, ...]
|
| 146 |
+
with torch.no_grad():
|
| 147 |
+
action, obs_embed, goal_embed = policy(
|
| 148 |
+
batchify_obs,
|
| 149 |
+
goal=torch.as_tensor(cur_milestone, device=policy.device),
|
| 150 |
+
deterministic=True,
|
| 151 |
+
return_embeddings=True,
|
| 152 |
+
input_pos=torch.tensor([st], device=policy.device)
|
| 153 |
+
if is_causal
|
| 154 |
+
else None,
|
| 155 |
+
)
|
| 156 |
+
env.current_obs_embedding = obs_embed[0].cpu().numpy()
|
| 157 |
+
obs, r, done, info = env.step(action[0].cpu().numpy())
|
| 158 |
+
step_t = time.time() - t
|
| 159 |
+
times.append(step_t)
|
| 160 |
+
times = np.sum(times)
|
| 161 |
+
print(times)
|
| 162 |
+
totals.append(times)
|
| 163 |
+
print(np.mean(totals))
|
scripts/examples/microwave-bottom_burner-light_switch-slide_cabinet.mp4
ADDED
|
Binary file (64.7 kB). View file
|
|
|
scripts/tsne_visualization.py
ADDED
|
@@ -0,0 +1,110 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
import argparse
|
| 2 |
+
import os.path
|
| 3 |
+
|
| 4 |
+
import matplotlib.pyplot as plt
|
| 5 |
+
import numpy as np
|
| 6 |
+
import pandas as pd
|
| 7 |
+
import seaborn as sns
|
| 8 |
+
import torch
|
| 9 |
+
from sklearn.manifold import TSNE
|
| 10 |
+
|
| 11 |
+
from uvd.decomp.decomp import (
|
| 12 |
+
embedding_decomp,
|
| 13 |
+
)
|
| 14 |
+
from uvd.models.preprocessors import *
|
| 15 |
+
import uvd.utils as U
|
| 16 |
+
|
| 17 |
+
from decord import VideoReader
|
| 18 |
+
|
| 19 |
+
|
| 20 |
+
def vis_2d_tsne(embeddings: np.ndarray, labels: list):
|
| 21 |
+
tsne = TSNE(n_components=2)
|
| 22 |
+
tsne_result = tsne.fit_transform(embeddings)
|
| 23 |
+
tsne_result_df = pd.DataFrame(
|
| 24 |
+
{"tsne_1": tsne_result[:, 0], "tsne_2": tsne_result[:, 1], "label": labels}
|
| 25 |
+
)
|
| 26 |
+
fig, ax = plt.subplots(1)
|
| 27 |
+
sns.scatterplot(x="tsne_1", y="tsne_2", hue="label", data=tsne_result_df, ax=ax, s=120)
|
| 28 |
+
lim = (tsne_result.min() - 5, tsne_result.max() + 5)
|
| 29 |
+
ax.set_xlim(lim)
|
| 30 |
+
ax.set_ylim(lim)
|
| 31 |
+
ax.set_aspect("equal")
|
| 32 |
+
ax.set_title(f"{preprocessor.__class__.__name__}")
|
| 33 |
+
plt.show()
|
| 34 |
+
|
| 35 |
+
|
| 36 |
+
def vis_3d_tsne(embeddings: np.ndarray, labels: list):
|
| 37 |
+
tsne = TSNE(n_components=3)
|
| 38 |
+
tsne_result = tsne.fit_transform(embeddings)
|
| 39 |
+
tsne_result_df = pd.DataFrame(
|
| 40 |
+
{
|
| 41 |
+
"tsne_1": tsne_result[:, 0],
|
| 42 |
+
"tsne_2": tsne_result[:, 1],
|
| 43 |
+
"tsne_3": tsne_result[:, 2],
|
| 44 |
+
"label": labels,
|
| 45 |
+
}
|
| 46 |
+
)
|
| 47 |
+
|
| 48 |
+
fig = plt.figure()
|
| 49 |
+
ax = fig.add_subplot(111, projection="3d")
|
| 50 |
+
|
| 51 |
+
palette = sns.color_palette("viridis", as_cmap=True)
|
| 52 |
+
unique_labels = tsne_result_df["label"].unique()
|
| 53 |
+
colors = palette(np.linspace(0, 1, len(unique_labels)))
|
| 54 |
+
color_dict = dict(zip(unique_labels, colors))
|
| 55 |
+
|
| 56 |
+
for label in unique_labels:
|
| 57 |
+
subset = tsne_result_df[tsne_result_df["label"] == label]
|
| 58 |
+
ax.scatter(
|
| 59 |
+
subset["tsne_1"],
|
| 60 |
+
subset["tsne_2"],
|
| 61 |
+
subset["tsne_3"],
|
| 62 |
+
c=[color_dict[label]],
|
| 63 |
+
label=label,
|
| 64 |
+
s=120,
|
| 65 |
+
)
|
| 66 |
+
ax.set_title(f"{preprocessor.__class__.__name__}")
|
| 67 |
+
plt.show()
|
| 68 |
+
|
| 69 |
+
|
| 70 |
+
if __name__ == "__main__":
|
| 71 |
+
parser = argparse.ArgumentParser()
|
| 72 |
+
parser.add_argument(
|
| 73 |
+
"--video_file",
|
| 74 |
+
default=U.f_join(
|
| 75 |
+
os.path.dirname(__file__), "examples/microwave-bottom_burner-light_switch-slide_cabinet.mp4"
|
| 76 |
+
)
|
| 77 |
+
)
|
| 78 |
+
parser.add_argument("--preprocessor_name", default="vip")
|
| 79 |
+
args = parser.parse_args()
|
| 80 |
+
|
| 81 |
+
use_gpu = torch.cuda.is_available()
|
| 82 |
+
if not use_gpu:
|
| 83 |
+
print("NO GPU FOUND")
|
| 84 |
+
|
| 85 |
+
frames = VideoReader(args.video_file, height=224, width=224)[:].asnumpy()
|
| 86 |
+
preprocessor = get_preprocessor(
|
| 87 |
+
args.preprocessor_name, device="cuda" if use_gpu else None
|
| 88 |
+
)
|
| 89 |
+
embeddings = preprocessor.process(frames, return_numpy=True)
|
| 90 |
+
_, decomp_meta = embedding_decomp(
|
| 91 |
+
embeddings=embeddings,
|
| 92 |
+
fill_embeddings=False,
|
| 93 |
+
return_intermediate_curves=False,
|
| 94 |
+
normalize_curve=False,
|
| 95 |
+
min_interval=20,
|
| 96 |
+
smooth_method="kernel",
|
| 97 |
+
gamma=0.1,
|
| 98 |
+
)
|
| 99 |
+
milestone_indices = decomp_meta.milestone_indices
|
| 100 |
+
milestone_rgbs = frames[milestone_indices]
|
| 101 |
+
|
| 102 |
+
labels = [
|
| 103 |
+
i
|
| 104 |
+
for i, count in enumerate(milestone_indices)
|
| 105 |
+
for _ in range(count - milestone_indices[i - 1] if i > 0 else count)
|
| 106 |
+
]
|
| 107 |
+
labels = [labels[0]] + labels
|
| 108 |
+
|
| 109 |
+
vis_2d_tsne(embeddings, labels)
|
| 110 |
+
vis_3d_tsne(embeddings, labels)
|
setup.py
ADDED
|
@@ -0,0 +1,49 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
import pathlib
|
| 2 |
+
|
| 3 |
+
from setuptools import setup, find_packages
|
| 4 |
+
|
| 5 |
+
PKG_NAME = "uvd"
|
| 6 |
+
VERSION = "0.0.1"
|
| 7 |
+
|
| 8 |
+
|
| 9 |
+
def _read_file(fname):
|
| 10 |
+
with pathlib.Path(fname).open() as fp:
|
| 11 |
+
return fp.read()
|
| 12 |
+
|
| 13 |
+
|
| 14 |
+
def _read_install_requires():
|
| 15 |
+
with pathlib.Path("requirements.txt").open() as fp:
|
| 16 |
+
return [
|
| 17 |
+
line.strip()
|
| 18 |
+
for line in fp
|
| 19 |
+
if line.strip() and not line.startswith("#")
|
| 20 |
+
]
|
| 21 |
+
|
| 22 |
+
|
| 23 |
+
setup(
|
| 24 |
+
name=PKG_NAME,
|
| 25 |
+
version=VERSION,
|
| 26 |
+
author=f"{PKG_NAME} Developers",
|
| 27 |
+
# url='http://github.com/',
|
| 28 |
+
description="research project",
|
| 29 |
+
long_description=_read_file("README.md"),
|
| 30 |
+
long_description_content_type="text/markdown",
|
| 31 |
+
keywords=["Deep Learning", "Reinforcement Learning"],
|
| 32 |
+
license="MIT License",
|
| 33 |
+
packages=find_packages(include=f"{PKG_NAME}.*"),
|
| 34 |
+
include_package_data=True,
|
| 35 |
+
zip_safe=False,
|
| 36 |
+
entry_points={
|
| 37 |
+
"console_scripts": [
|
| 38 |
+
# 'cmd_tool=mylib.subpkg.module:main',
|
| 39 |
+
]
|
| 40 |
+
},
|
| 41 |
+
install_requires=_read_install_requires(),
|
| 42 |
+
python_requires=">=3.9",
|
| 43 |
+
classifiers=[
|
| 44 |
+
"Development Status :: 3 - Alpha",
|
| 45 |
+
"Topic :: Scientific/Engineering :: Artificial Intelligence",
|
| 46 |
+
"Environment :: Console",
|
| 47 |
+
"Programming Language :: Python :: 3",
|
| 48 |
+
],
|
| 49 |
+
)
|
uvd/__init__.py
ADDED
|
@@ -0,0 +1,30 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
from __future__ import annotations
|
| 2 |
+
|
| 3 |
+
from typing import Literal
|
| 4 |
+
|
| 5 |
+
import numpy as np
|
| 6 |
+
import torch
|
| 7 |
+
|
| 8 |
+
from .decomp import *
|
| 9 |
+
from .models import *
|
| 10 |
+
|
| 11 |
+
|
| 12 |
+
def get_uvd_subgoals(
|
| 13 |
+
frames: np.ndarray | str,
|
| 14 |
+
preprocessor_name: Literal["vip", "r3m", "liv", "clip", "vc1", "dinov2"] = "vip",
|
| 15 |
+
device: torch.device | str | None = "cuda",
|
| 16 |
+
return_indices: bool = False,
|
| 17 |
+
) -> list | np.ndarray:
|
| 18 |
+
"""Quick API for UVD decomposition."""
|
| 19 |
+
if isinstance(frames, str):
|
| 20 |
+
from decord import VideoReader
|
| 21 |
+
|
| 22 |
+
vr = VideoReader(frames, height=224, width=224)
|
| 23 |
+
frames = vr[:].asnumpy()
|
| 24 |
+
preprocessor = get_preprocessor(preprocessor_name, device=device)
|
| 25 |
+
rep = preprocessor.process(frames, return_numpy=True)
|
| 26 |
+
_, decomp_meta = decomp_trajectories("embed", rep)
|
| 27 |
+
indices = decomp_meta.milestone_indices
|
| 28 |
+
if return_indices:
|
| 29 |
+
return indices
|
| 30 |
+
return frames[indices]
|
uvd/data/__init__.py
ADDED
|
@@ -0,0 +1,3 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
from .dataset_aug import *
|
| 2 |
+
from .dataset_base import *
|
| 3 |
+
from .franka_kitchen_datasets import *
|
uvd/data/dataset_aug.py
ADDED
|
@@ -0,0 +1,310 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
from __future__ import annotations
|
| 2 |
+
|
| 3 |
+
import random
|
| 4 |
+
from typing import Literal
|
| 5 |
+
|
| 6 |
+
import gym
|
| 7 |
+
import numpy as np
|
| 8 |
+
import torch
|
| 9 |
+
|
| 10 |
+
import uvd.utils as U
|
| 11 |
+
from uvd.decomp import decomp_trajectories
|
| 12 |
+
from uvd.envs.franka_kitchen import KitchenBase
|
| 13 |
+
from .franka_kitchen_datasets import (
|
| 14 |
+
FrankaKitchenDataset,
|
| 15 |
+
)
|
| 16 |
+
|
| 17 |
+
__all__ = ["DatasetWithAug", "MilestoneRandomSkipDataset"]
|
| 18 |
+
|
| 19 |
+
|
| 20 |
+
class DatasetWithAug(FrankaKitchenDataset):
|
| 21 |
+
def __init__(
|
| 22 |
+
self,
|
| 23 |
+
dataset_path: str | list,
|
| 24 |
+
*,
|
| 25 |
+
specific_tasks: str | list | None,
|
| 26 |
+
obs_keys: str | tuple = ("rgb", "proprio"),
|
| 27 |
+
num_demos: int | None = None,
|
| 28 |
+
shuffle: bool = False,
|
| 29 |
+
replace_last_goal_as_last_frame: bool = True,
|
| 30 |
+
# if `preprocess`, the observation for training is directly embeddings
|
| 31 |
+
preprocess: bool = True,
|
| 32 |
+
preprocess_name: str | None = None,
|
| 33 |
+
preprocess_kwargs: dict | None = None,
|
| 34 |
+
decomp_method: Literal["oracle", "embed", "random", "equally", "near_future"]
|
| 35 |
+
| None = "embed",
|
| 36 |
+
decomp_kwargs: dict | None = None,
|
| 37 |
+
sequential: bool = True,
|
| 38 |
+
include_no_robot: bool = False,
|
| 39 |
+
del_preprocessor: bool = True,
|
| 40 |
+
device: int | str | torch.device | None = None,
|
| 41 |
+
):
|
| 42 |
+
decomp_kwargs = decomp_kwargs or {}
|
| 43 |
+
decomp_kwargs.update(fill_embeddings=False)
|
| 44 |
+
self.idx_to_all_milestones = {}
|
| 45 |
+
self.max_num_milestones = -1
|
| 46 |
+
super().__init__(**U.prepare_locals_for_super(locals()))
|
| 47 |
+
|
| 48 |
+
def prepare_episode_data(
|
| 49 |
+
self,
|
| 50 |
+
episode_idx: int,
|
| 51 |
+
episode: str,
|
| 52 |
+
use_tensor: bool = True,
|
| 53 |
+
sequential: bool = True,
|
| 54 |
+
):
|
| 55 |
+
self.idx_to_episode_path[episode_idx] = episode
|
| 56 |
+
data = U.load_pickle(episode)
|
| 57 |
+
actions = data["actions"]
|
| 58 |
+
episode_length = actions.shape[0]
|
| 59 |
+
full_obs_data = data["obs_full"]
|
| 60 |
+
rgb_traj = full_obs_data["rgb"][:-1, ...] # rm goal
|
| 61 |
+
raw_rgb_no_robot = data.get("no_robot_rgb", None)
|
| 62 |
+
if self.include_no_robot:
|
| 63 |
+
assert raw_rgb_no_robot is not None
|
| 64 |
+
raw_rgb_no_robot = raw_rgb_no_robot[:-1]
|
| 65 |
+
assert raw_rgb_no_robot.shape == rgb_traj.shape
|
| 66 |
+
else:
|
| 67 |
+
raw_rgb_no_robot = None
|
| 68 |
+
assert rgb_traj.shape[0] == episode_length
|
| 69 |
+
self._idx_to_episode_length[episode_idx] = episode_length
|
| 70 |
+
|
| 71 |
+
prepared_data = dict()
|
| 72 |
+
|
| 73 |
+
if use_tensor:
|
| 74 |
+
actions = U.any_to_torch_tensor(actions, device="cpu", dtype=torch.float32)
|
| 75 |
+
else:
|
| 76 |
+
actions = U.any_to_numpy(actions, dtype="float32")
|
| 77 |
+
prepared_data["actions"] = actions
|
| 78 |
+
|
| 79 |
+
# embeddings are either fully preprocessed or transformed
|
| 80 |
+
if self.preprocess or (
|
| 81 |
+
self.decomp_method is not None and "embed" in self.decomp_method
|
| 82 |
+
):
|
| 83 |
+
# preprocess with frozen visual backbones
|
| 84 |
+
embeddings = self.preprocessor.process(
|
| 85 |
+
raw_img_batch=rgb_traj, return_numpy=not use_tensor
|
| 86 |
+
)
|
| 87 |
+
# L, N
|
| 88 |
+
assert embeddings.shape[0] == episode_length, embeddings.shape
|
| 89 |
+
else:
|
| 90 |
+
# L, 3, H, W
|
| 91 |
+
embeddings = self.preprocessor.transform_images(rgb_traj)
|
| 92 |
+
|
| 93 |
+
obs = {}
|
| 94 |
+
for k, v in full_obs_data.items():
|
| 95 |
+
if k in self.obs_keys and k == "rgb":
|
| 96 |
+
# raw rgb, may use for rollout
|
| 97 |
+
obs[k] = rgb_traj # v[:-1]
|
| 98 |
+
elif k in self.obs_keys:
|
| 99 |
+
obs[k] = (
|
| 100 |
+
U.any_to_torch_tensor(
|
| 101 |
+
v[:-1], dtype=torch.float32, device="cpu" # self.device
|
| 102 |
+
)
|
| 103 |
+
if use_tensor
|
| 104 |
+
else U.any_to_numpy(v[:-1], dtype="float32")
|
| 105 |
+
)
|
| 106 |
+
|
| 107 |
+
# obs = {k: v[:-1] for k, v in full_obs_data.items() if k in self.obs_keys}
|
| 108 |
+
prepared_data["obs"] = obs
|
| 109 |
+
rgb_no_robot = None
|
| 110 |
+
if raw_rgb_no_robot is not None:
|
| 111 |
+
if self.preprocess:
|
| 112 |
+
rgb_no_robot = self.preprocessor.process(
|
| 113 |
+
raw_img_batch=raw_rgb_no_robot,
|
| 114 |
+
return_numpy=not use_tensor,
|
| 115 |
+
reconstruct_linear=True,
|
| 116 |
+
)
|
| 117 |
+
else:
|
| 118 |
+
rgb_no_robot = self.preprocessor.transform_images(raw_rgb_no_robot)
|
| 119 |
+
prepared_data["rgb_no_robot"] = rgb_no_robot
|
| 120 |
+
|
| 121 |
+
if self.decomp_method == "oracle":
|
| 122 |
+
num_goals_achieved = data["completed_tasks"]
|
| 123 |
+
_, decomp_meta = decomp_trajectories(
|
| 124 |
+
method_name=self.decomp_method,
|
| 125 |
+
embeddings=None,
|
| 126 |
+
goal_achieved_mask=num_goals_achieved,
|
| 127 |
+
**self.decomp_kwargs or {},
|
| 128 |
+
)
|
| 129 |
+
else:
|
| 130 |
+
decomp_kwargs = self.decomp_kwargs or {}
|
| 131 |
+
if (
|
| 132 |
+
self.decomp_method is not None
|
| 133 |
+
and "embed_no_robot" in self.decomp_method
|
| 134 |
+
):
|
| 135 |
+
decomp_kwargs["no_robot_embeddings"] = rgb_no_robot
|
| 136 |
+
elif self.decomp_method == "embed2":
|
| 137 |
+
decomp_kwargs["no_robot_embeddings"] = embeddings
|
| 138 |
+
decomp_kwargs["task_name"] = str(data["reset_kwargs"]["task_elements"])
|
| 139 |
+
_, decomp_meta = decomp_trajectories(
|
| 140 |
+
method_name=self.decomp_method,
|
| 141 |
+
embeddings=embeddings,
|
| 142 |
+
**decomp_kwargs,
|
| 143 |
+
)
|
| 144 |
+
|
| 145 |
+
if self.preprocessor.use_language_goal:
|
| 146 |
+
raise NotImplementedError
|
| 147 |
+
|
| 148 |
+
if use_tensor:
|
| 149 |
+
embeddings = U.any_to_torch_tensor(embeddings, device="cpu")
|
| 150 |
+
|
| 151 |
+
milestone_indices = decomp_meta.milestone_indices
|
| 152 |
+
prepared_data["milestone_indices"] = milestone_indices
|
| 153 |
+
self.max_num_milestones = max(len(milestone_indices), self.max_num_milestones)
|
| 154 |
+
|
| 155 |
+
diffs = np.diff(milestone_indices, prepend=0)
|
| 156 |
+
milestone_step_mask = np.repeat(np.arange(len(diffs)), diffs)
|
| 157 |
+
milestone_step_mask = np.concatenate(
|
| 158 |
+
[np.array([0]), milestone_step_mask], dtype=np.int32
|
| 159 |
+
)
|
| 160 |
+
assert len(milestone_step_mask) == episode_length, (
|
| 161 |
+
len(milestone_step_mask),
|
| 162 |
+
episode_length,
|
| 163 |
+
)
|
| 164 |
+
|
| 165 |
+
milestone_embeddings = embeddings[milestone_indices]
|
| 166 |
+
|
| 167 |
+
if sequential:
|
| 168 |
+
if self.split_train_eval:
|
| 169 |
+
if episode.split("/")[-2] not in self.train_tasks:
|
| 170 |
+
assert episode.split("/")[-2] in self.eval_tasks, episode.split(
|
| 171 |
+
"/"
|
| 172 |
+
)[-2]
|
| 173 |
+
return prepared_data, episode_length
|
| 174 |
+
|
| 175 |
+
def maybe_tensor(x, dtype):
|
| 176 |
+
return (
|
| 177 |
+
torch.tensor(x, device="cpu", dtype=dtype)
|
| 178 |
+
if use_tensor
|
| 179 |
+
else np.array(x, dtype=dtype)
|
| 180 |
+
)
|
| 181 |
+
|
| 182 |
+
for st in range(episode_length):
|
| 183 |
+
seq_data = dict(
|
| 184 |
+
actions=actions[st],
|
| 185 |
+
milestones=milestone_embeddings,
|
| 186 |
+
milestone_indices=maybe_tensor(milestone_indices, dtype="int32"),
|
| 187 |
+
cur_milestone_idx=maybe_tensor(
|
| 188 |
+
milestone_step_mask[st], dtype="int32"
|
| 189 |
+
),
|
| 190 |
+
episode_idx=maybe_tensor([episode_idx], dtype="int32"),
|
| 191 |
+
timesteps=maybe_tensor([st], dtype="int32"),
|
| 192 |
+
)
|
| 193 |
+
obs_i = {
|
| 194 |
+
k: v[st] if k != "rgb" else embeddings[st] for k, v in obs.items()
|
| 195 |
+
}
|
| 196 |
+
seq_data["obs"] = obs_i
|
| 197 |
+
self.sequence_data.append(seq_data)
|
| 198 |
+
else:
|
| 199 |
+
raise NotImplementedError
|
| 200 |
+
|
| 201 |
+
return prepared_data, episode_length
|
| 202 |
+
|
| 203 |
+
@property
|
| 204 |
+
def dataset_metadata(self) -> dict:
|
| 205 |
+
assert self.max_num_milestones > 0, self.max_num_milestones
|
| 206 |
+
|
| 207 |
+
dummy_env = KitchenBase(obs_keys=self.obs_keys)
|
| 208 |
+
action_space = dummy_env.action_space
|
| 209 |
+
observation_space = dummy_env.observation_space.spaces
|
| 210 |
+
dummy_env.close()
|
| 211 |
+
del dummy_env
|
| 212 |
+
if "rgb" in observation_space and self.preprocess:
|
| 213 |
+
observation_space["rgb"] = gym.spaces.Box(
|
| 214 |
+
low=-np.inf,
|
| 215 |
+
high=np.inf,
|
| 216 |
+
shape=self.preprocessor_output_dim,
|
| 217 |
+
dtype=np.float32,
|
| 218 |
+
)
|
| 219 |
+
observation_space = gym.spaces.Dict(observation_space)
|
| 220 |
+
metadata = dict(action_space=action_space, observation_space=observation_space)
|
| 221 |
+
metadata["observation_space"]["milestones"] = gym.spaces.Box(
|
| 222 |
+
low=-np.inf,
|
| 223 |
+
high=np.inf,
|
| 224 |
+
shape=(self.max_num_milestones,) + self.preprocessor_output_dim,
|
| 225 |
+
dtype=np.float32,
|
| 226 |
+
)
|
| 227 |
+
return metadata
|
| 228 |
+
|
| 229 |
+
def __getitem__(self, idx: int) -> dict:
|
| 230 |
+
data = self.sequence_data[idx]
|
| 231 |
+
milestones = data["milestones"]
|
| 232 |
+
milestone_indices = data["milestone_indices"]
|
| 233 |
+
num_milestones, d = milestones.shape
|
| 234 |
+
assert num_milestones <= self.max_num_milestones
|
| 235 |
+
pad_length = self.max_num_milestones - num_milestones
|
| 236 |
+
if pad_length != 0:
|
| 237 |
+
milestones = U.any_concat(
|
| 238 |
+
# assume not use tensor
|
| 239 |
+
[milestones, np.zeros((pad_length, d), dtype=milestones.dtype)],
|
| 240 |
+
dim=0,
|
| 241 |
+
)
|
| 242 |
+
milestone_indices = U.any_concat(
|
| 243 |
+
[milestone_indices, milestone_indices[-1][None] * pad_length]
|
| 244 |
+
)
|
| 245 |
+
# N, D
|
| 246 |
+
data["milestones"] = milestones # .reshape((self.max_num_milestones * d,))
|
| 247 |
+
data["milestone_masks"] = np.array(
|
| 248 |
+
[1] * num_milestones + [0] * pad_length, dtype=np.bool_
|
| 249 |
+
)[
|
| 250 |
+
:, None
|
| 251 |
+
] # N, 1
|
| 252 |
+
data["milestone_indices"] = milestone_indices
|
| 253 |
+
return data
|
| 254 |
+
|
| 255 |
+
|
| 256 |
+
class MilestoneRandomSkipDataset(DatasetWithAug):
|
| 257 |
+
def __init__(
|
| 258 |
+
self, *, skip_ratio: 0.1, min_skip_n_milestones: int | None = None, **kwargs
|
| 259 |
+
):
|
| 260 |
+
self.skip_ratio = skip_ratio
|
| 261 |
+
self.min_skip_n_milestones = min_skip_n_milestones or 0
|
| 262 |
+
super().__init__(**kwargs)
|
| 263 |
+
|
| 264 |
+
def __getitem__(self, idx: int) -> dict:
|
| 265 |
+
data = self.sequence_data[idx]
|
| 266 |
+
milestones = data["milestones"]
|
| 267 |
+
milestone_indices = data["milestone_indices"].copy()
|
| 268 |
+
num_milestones = len(milestone_indices)
|
| 269 |
+
assert len(milestones) == num_milestones, (
|
| 270 |
+
len(milestones),
|
| 271 |
+
num_milestones,
|
| 272 |
+
milestone_indices,
|
| 273 |
+
)
|
| 274 |
+
cur_milestone_idx = data["cur_milestone_idx"].copy()
|
| 275 |
+
|
| 276 |
+
do_skip = (
|
| 277 |
+
len(milestone_indices) > self.min_skip_n_milestones
|
| 278 |
+
and random.random() > self.skip_ratio
|
| 279 |
+
)
|
| 280 |
+
if do_skip:
|
| 281 |
+
skip_back = (
|
| 282 |
+
cur_milestone_idx > 0 and random.random() > 0.5
|
| 283 |
+
) or cur_milestone_idx == num_milestones - 1
|
| 284 |
+
if skip_back:
|
| 285 |
+
cur_milestone_idx -= 1
|
| 286 |
+
else:
|
| 287 |
+
cur_milestone_idx += 1
|
| 288 |
+
assert 0 <= cur_milestone_idx <= num_milestones - 1, (
|
| 289 |
+
cur_milestone_idx,
|
| 290 |
+
milestone_indices,
|
| 291 |
+
)
|
| 292 |
+
|
| 293 |
+
pad_length = self.max_num_milestones - num_milestones
|
| 294 |
+
if pad_length != 0:
|
| 295 |
+
milestone_indices = U.any_concat(
|
| 296 |
+
[data["milestone_indices"]]
|
| 297 |
+
+ [data["milestone_indices"][-1][None]] * pad_length
|
| 298 |
+
)
|
| 299 |
+
else:
|
| 300 |
+
milestone_indices = data["milestone_indices"]
|
| 301 |
+
|
| 302 |
+
return dict(
|
| 303 |
+
obs=data["obs"],
|
| 304 |
+
milestones=milestones[cur_milestone_idx],
|
| 305 |
+
milestone_indices=milestone_indices,
|
| 306 |
+
cur_milestone_idx=cur_milestone_idx,
|
| 307 |
+
actions=data["actions"],
|
| 308 |
+
episode_idx=data["episode_idx"],
|
| 309 |
+
timesteps=data["timesteps"],
|
| 310 |
+
)
|
uvd/data/dataset_base.py
ADDED
|
@@ -0,0 +1,19 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
import abc
|
| 2 |
+
|
| 3 |
+
from torch.utils.data import Dataset
|
| 4 |
+
|
| 5 |
+
__all__ = ["DatasetBase"]
|
| 6 |
+
|
| 7 |
+
|
| 8 |
+
class DatasetBase(Dataset):
|
| 9 |
+
@abc.abstractproperty
|
| 10 |
+
def dataset_metadata(self) -> dict:
|
| 11 |
+
raise NotImplementedError
|
| 12 |
+
|
| 13 |
+
@abc.abstractmethod
|
| 14 |
+
def __len__(self):
|
| 15 |
+
raise NotImplementedError
|
| 16 |
+
|
| 17 |
+
@abc.abstractmethod
|
| 18 |
+
def __getitem__(self, item: int):
|
| 19 |
+
raise NotImplementedError
|
uvd/data/franka_kitchen_datasets.py
ADDED
|
@@ -0,0 +1,707 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
from __future__ import annotations
|
| 2 |
+
|
| 3 |
+
import collections
|
| 4 |
+
import copy
|
| 5 |
+
import json
|
| 6 |
+
import random
|
| 7 |
+
from collections import OrderedDict
|
| 8 |
+
from typing import Literal
|
| 9 |
+
|
| 10 |
+
import gym
|
| 11 |
+
import numpy as np
|
| 12 |
+
import torch
|
| 13 |
+
import tqdm
|
| 14 |
+
import wandb
|
| 15 |
+
from torch.distributions.utils import lazy_property
|
| 16 |
+
from torch.utils.data import Dataset
|
| 17 |
+
|
| 18 |
+
import uvd.utils as U
|
| 19 |
+
from uvd.decomp import decomp_trajectories
|
| 20 |
+
from uvd.envs.franka_kitchen import KitchenBase
|
| 21 |
+
from uvd.envs.franka_kitchen.franka_kitchen_constants import ELEMENT_TO_IDX
|
| 22 |
+
from uvd.models.preprocessors import get_preprocessor
|
| 23 |
+
from .dataset_base import DatasetBase
|
| 24 |
+
|
| 25 |
+
__all__ = [
|
| 26 |
+
"FrankaKitchenDataset",
|
| 27 |
+
"FrankaKitchenRolloutDataset",
|
| 28 |
+
"ALL_TASKS",
|
| 29 |
+
"task_elements_to_prompt",
|
| 30 |
+
]
|
| 31 |
+
|
| 32 |
+
|
| 33 |
+
ALL_TASKS = [
|
| 34 |
+
"bottom_burner-top_burner-light_switch-slide_cabinet",
|
| 35 |
+
"bottom_burner-top_burner-slide_cabinet-hinge_cabinet",
|
| 36 |
+
"kettle-bottom_burner-light_switch-hinge_cabinet",
|
| 37 |
+
"kettle-bottom_burner-light_switch-slide_cabinet",
|
| 38 |
+
"kettle-bottom_burner-slide_cabinet-hinge_cabinet",
|
| 39 |
+
"kettle-bottom_burner-top_burner-hinge_cabinet",
|
| 40 |
+
"kettle-bottom_burner-top_burner-light_switch",
|
| 41 |
+
"kettle-bottom_burner-top_burner-slide_cabinet",
|
| 42 |
+
"kettle-light_switch-slide_cabinet-hinge_cabinet",
|
| 43 |
+
"kettle-top_burner-light_switch-slide_cabinet",
|
| 44 |
+
"microwave-bottom_burner-light_switch-slide_cabinet",
|
| 45 |
+
"microwave-bottom_burner-slide_cabinet-hinge_cabinet",
|
| 46 |
+
"microwave-bottom_burner-top_burner-hinge_cabinet",
|
| 47 |
+
"microwave-bottom_burner-top_burner-light_switch",
|
| 48 |
+
"microwave-bottom_burner-top_burner-slide_cabinet",
|
| 49 |
+
"microwave-kettle-bottom_burner-hinge_cabinet",
|
| 50 |
+
"microwave-kettle-bottom_burner-slide_cabinet",
|
| 51 |
+
"microwave-kettle-light_switch-hinge_cabinet",
|
| 52 |
+
"microwave-kettle-light_switch-slide_cabinet",
|
| 53 |
+
"microwave-kettle-slide_cabinet-hinge_cabinet",
|
| 54 |
+
"microwave-kettle-top_burner-hinge_cabinet",
|
| 55 |
+
"microwave-kettle-top_burner-light_switch",
|
| 56 |
+
"microwave-light_switch-slide_cabinet-hinge_cabinet",
|
| 57 |
+
"microwave-top_burner-light_switch-hinge_cabinet",
|
| 58 |
+
]
|
| 59 |
+
|
| 60 |
+
|
| 61 |
+
class FrankaKitchenDataset(DatasetBase):
|
| 62 |
+
def __init__(
|
| 63 |
+
self,
|
| 64 |
+
dataset_path: str | list,
|
| 65 |
+
*,
|
| 66 |
+
specific_tasks: str | list | None,
|
| 67 |
+
obs_keys: str | tuple = ("rgb", "proprio"),
|
| 68 |
+
num_demos: int | None = None,
|
| 69 |
+
shuffle: bool = False,
|
| 70 |
+
replace_last_goal_as_last_frame: bool = True,
|
| 71 |
+
# if `preprocess`, the observation for training is directly embeddings
|
| 72 |
+
preprocess: bool = True,
|
| 73 |
+
preprocess_name: str | None = None,
|
| 74 |
+
preprocess_kwargs: dict | None = None,
|
| 75 |
+
decomp_method: Literal["oracle", "embed", "random", "equally", "near_future"]
|
| 76 |
+
| None = "embed",
|
| 77 |
+
decomp_kwargs: dict | None = None,
|
| 78 |
+
sequential: bool = True,
|
| 79 |
+
max_seq_length: int | None = None,
|
| 80 |
+
include_no_robot: bool = False,
|
| 81 |
+
del_preprocessor: bool = True,
|
| 82 |
+
device: int | str | torch.device | None = None,
|
| 83 |
+
):
|
| 84 |
+
super().__init__()
|
| 85 |
+
|
| 86 |
+
self.split_train_eval = None
|
| 87 |
+
if "random_" in specific_tasks:
|
| 88 |
+
n_train = int(specific_tasks.split("_")[-1])
|
| 89 |
+
specific_tasks = ALL_TASKS
|
| 90 |
+
random.shuffle(specific_tasks)
|
| 91 |
+
self.split_train_eval = n_train
|
| 92 |
+
|
| 93 |
+
if isinstance(dataset_path, str):
|
| 94 |
+
assert U.f_exists(dataset_path), dataset_path
|
| 95 |
+
if specific_tasks is not None:
|
| 96 |
+
if isinstance(specific_tasks, str):
|
| 97 |
+
specific_tasks = (specific_tasks,)
|
| 98 |
+
dataset_path = [U.f_join(dataset_path, t) for t in specific_tasks]
|
| 99 |
+
elif "metadata.json" in U.f_listdir(dataset_path):
|
| 100 |
+
# single task
|
| 101 |
+
dataset_path = (dataset_path,)
|
| 102 |
+
else:
|
| 103 |
+
# use all tasks
|
| 104 |
+
dataset_path = [
|
| 105 |
+
U.f_join(dataset_path, d)
|
| 106 |
+
for d in U.f_listdir(dataset_path)
|
| 107 |
+
if "-" in d and "_" in d
|
| 108 |
+
]
|
| 109 |
+
|
| 110 |
+
if self.split_train_eval is not None:
|
| 111 |
+
assert specific_tasks is not None
|
| 112 |
+
self.train_tasks = specific_tasks[: self.split_train_eval]
|
| 113 |
+
self.eval_tasks = specific_tasks[self.split_train_eval :]
|
| 114 |
+
U.rank_zero_print(
|
| 115 |
+
f"{len(self.train_tasks)} training tasks, {len(self.eval_tasks)} unseen tasks",
|
| 116 |
+
color="green",
|
| 117 |
+
)
|
| 118 |
+
if U.is_rank_zero() and wandb.run is not None:
|
| 119 |
+
table = wandb.Table(
|
| 120 |
+
data=[[",\n".join(self.train_tasks), ",\n".join(self.eval_tasks)]],
|
| 121 |
+
columns=["training tasks", "unseen tasks"],
|
| 122 |
+
)
|
| 123 |
+
wandb.log(dict(partition=table))
|
| 124 |
+
|
| 125 |
+
assert sum((U.f_exists(d) for d in dataset_path)) == len(dataset_path)
|
| 126 |
+
self.include_no_robot = include_no_robot and decomp_method is not None
|
| 127 |
+
if decomp_method == "embed_no_robot":
|
| 128 |
+
assert self.include_no_robot and preprocess
|
| 129 |
+
self.cache_scores = dict()
|
| 130 |
+
|
| 131 |
+
episode_for_task = collections.defaultdict(list)
|
| 132 |
+
for d in dataset_path:
|
| 133 |
+
for f in U.f_listdir(d):
|
| 134 |
+
if num_demos is not None and len(episode_for_task[d]) >= num_demos:
|
| 135 |
+
break
|
| 136 |
+
if f.endswith(".pkl") and f.startswith("episode"):
|
| 137 |
+
episode_for_task[d].append(U.f_join(d, f))
|
| 138 |
+
|
| 139 |
+
episodes = sorted([f for v in episode_for_task.values() for f in v])
|
| 140 |
+
|
| 141 |
+
U.rank_zero_print(
|
| 142 |
+
"# Train Demos:\n"
|
| 143 |
+
+ json.dumps(
|
| 144 |
+
{
|
| 145 |
+
k.split("/")[-1]: len(v)
|
| 146 |
+
for k, v in episode_for_task.items()
|
| 147 |
+
if k.split("/")[-1] in self.train_tasks
|
| 148 |
+
},
|
| 149 |
+
indent=2,
|
| 150 |
+
)
|
| 151 |
+
+ "\n# Eval Demos:\n"
|
| 152 |
+
+ json.dumps(
|
| 153 |
+
{
|
| 154 |
+
k.split("/")[-1]: len(v)
|
| 155 |
+
for k, v in episode_for_task.items()
|
| 156 |
+
if k.split("/")[-1] in self.eval_tasks
|
| 157 |
+
},
|
| 158 |
+
indent=2,
|
| 159 |
+
)
|
| 160 |
+
+ f"\nTotal: {len(episodes)}",
|
| 161 |
+
color="blue",
|
| 162 |
+
)
|
| 163 |
+
|
| 164 |
+
if shuffle:
|
| 165 |
+
random.shuffle(episodes)
|
| 166 |
+
|
| 167 |
+
self.obs_keys = (
|
| 168 |
+
tuple(obs_keys) if not isinstance(obs_keys, str) else (obs_keys,)
|
| 169 |
+
)
|
| 170 |
+
self.replace_last_goal_as_last_frame = replace_last_goal_as_last_frame
|
| 171 |
+
self.num_demos = num_demos
|
| 172 |
+
self.idx_to_episode_path = {}
|
| 173 |
+
self._idx_to_episode_length = {}
|
| 174 |
+
|
| 175 |
+
self.preprocess = preprocess
|
| 176 |
+
self.device = device
|
| 177 |
+
if preprocess or decomp_method == "embed":
|
| 178 |
+
assert preprocess_name is not None
|
| 179 |
+
self.preprocessor = get_preprocessor(
|
| 180 |
+
name=preprocess_name, device=device, **preprocess_kwargs or {}
|
| 181 |
+
)
|
| 182 |
+
self.preprocessor_output_dim = self.preprocessor.output_dim
|
| 183 |
+
else:
|
| 184 |
+
# do transform only, with self.preprocess is False
|
| 185 |
+
self.preprocessor = get_preprocessor(
|
| 186 |
+
name=preprocess_name, device=device, **preprocess_kwargs or {}
|
| 187 |
+
)
|
| 188 |
+
self.decomp_method = decomp_method
|
| 189 |
+
self.decomp_kwargs = decomp_kwargs
|
| 190 |
+
self.use_language_goal = self.preprocessor.use_language_goal
|
| 191 |
+
|
| 192 |
+
self.episode_data = {}
|
| 193 |
+
self.sequence_data = []
|
| 194 |
+
self.sequential = sequential
|
| 195 |
+
for episode_idx, episode in tqdm.tqdm(
|
| 196 |
+
enumerate(episodes),
|
| 197 |
+
total=len(episodes),
|
| 198 |
+
desc=f"load {self.__class__.__name__}",
|
| 199 |
+
):
|
| 200 |
+
prepared_data, episode_length = self.prepare_episode_data(
|
| 201 |
+
episode_idx, episode, sequential=sequential, use_tensor=False # True
|
| 202 |
+
)
|
| 203 |
+
self.episode_data[episode_idx] = prepared_data
|
| 204 |
+
|
| 205 |
+
assert len(self._idx_to_episode_length) == len(episodes)
|
| 206 |
+
self._max_episode_length = max(self._idx_to_episode_length.values())
|
| 207 |
+
self.max_seq_length = None
|
| 208 |
+
self.idx_partition = None
|
| 209 |
+
if max_seq_length is not None:
|
| 210 |
+
assert max_seq_length <= self._max_episode_length
|
| 211 |
+
if max_seq_length < self._max_episode_length:
|
| 212 |
+
self.max_seq_length = max_seq_length
|
| 213 |
+
self.idx_partition = [
|
| 214 |
+
(ep_idx, (max(0, step_end - max_seq_length), step_end))
|
| 215 |
+
for ep_idx in range(
|
| 216 |
+
len(self.sequence_data)
|
| 217 |
+
) # only for training data
|
| 218 |
+
for step_end in range(
|
| 219 |
+
1,
|
| 220 |
+
self._idx_to_episode_length[
|
| 221 |
+
int(self.sequence_data[ep_idx]["episode_idx"])
|
| 222 |
+
]
|
| 223 |
+
+ 1,
|
| 224 |
+
)
|
| 225 |
+
]
|
| 226 |
+
|
| 227 |
+
if del_preprocessor:
|
| 228 |
+
# del after using
|
| 229 |
+
del self.preprocessor
|
| 230 |
+
|
| 231 |
+
U.rank_zero_print(
|
| 232 |
+
f"TOTAL TRAINING DATA: {len(self.sequence_data)}", color="blue"
|
| 233 |
+
)
|
| 234 |
+
|
| 235 |
+
@property
|
| 236 |
+
def max_episode_length(self) -> int:
|
| 237 |
+
return self._max_episode_length
|
| 238 |
+
|
| 239 |
+
@lazy_property
|
| 240 |
+
def dataset_metadata(self) -> dict:
|
| 241 |
+
dummy_env = KitchenBase(obs_keys=self.obs_keys)
|
| 242 |
+
action_space = dummy_env.action_space
|
| 243 |
+
observation_space = dummy_env.observation_space.spaces
|
| 244 |
+
dummy_env.close()
|
| 245 |
+
del dummy_env
|
| 246 |
+
if "rgb" in observation_space and self.preprocess:
|
| 247 |
+
observation_space["rgb"] = gym.spaces.Box(
|
| 248 |
+
low=-np.inf,
|
| 249 |
+
high=np.inf,
|
| 250 |
+
shape=(self.preprocessor_output_dim,)
|
| 251 |
+
if isinstance(self.preprocessor_output_dim, int)
|
| 252 |
+
else self.preprocessor_output_dim,
|
| 253 |
+
dtype=np.float32,
|
| 254 |
+
)
|
| 255 |
+
observation_space = gym.spaces.Dict(observation_space)
|
| 256 |
+
metadata = dict(action_space=action_space, observation_space=observation_space)
|
| 257 |
+
if not self.sequential: # causal
|
| 258 |
+
metadata["max_episode_length"] = self.max_episode_length
|
| 259 |
+
if self.max_seq_length is not None:
|
| 260 |
+
metadata["max_seq_length"] = self.max_seq_length
|
| 261 |
+
return metadata
|
| 262 |
+
|
| 263 |
+
def prepare_episode_data(
|
| 264 |
+
self,
|
| 265 |
+
episode_idx: int,
|
| 266 |
+
episode: str,
|
| 267 |
+
use_tensor: bool = True,
|
| 268 |
+
sequential: bool = True,
|
| 269 |
+
):
|
| 270 |
+
self.idx_to_episode_path[episode_idx] = episode
|
| 271 |
+
data = U.load_pickle(episode)
|
| 272 |
+
actions = data["actions"]
|
| 273 |
+
episode_length = actions.shape[0]
|
| 274 |
+
full_obs_data = data["obs_full"]
|
| 275 |
+
rgb_traj = full_obs_data["rgb"][:-1, ...] # rm goal
|
| 276 |
+
task_elements = data["reset_kwargs"]["task_elements"]
|
| 277 |
+
raw_rgb_no_robot = data.get("no_robot_rgb", None)
|
| 278 |
+
if self.include_no_robot:
|
| 279 |
+
assert raw_rgb_no_robot is not None
|
| 280 |
+
raw_rgb_no_robot = raw_rgb_no_robot[:-1]
|
| 281 |
+
assert raw_rgb_no_robot.shape == rgb_traj.shape
|
| 282 |
+
else:
|
| 283 |
+
raw_rgb_no_robot = None
|
| 284 |
+
assert rgb_traj.shape[0] == episode_length
|
| 285 |
+
self._idx_to_episode_length[episode_idx] = episode_length
|
| 286 |
+
|
| 287 |
+
prepared_data = dict()
|
| 288 |
+
|
| 289 |
+
if use_tensor:
|
| 290 |
+
actions = U.any_to_torch_tensor(actions, device="cpu", dtype=torch.float32)
|
| 291 |
+
else:
|
| 292 |
+
actions = U.any_to_numpy(actions, dtype="float32")
|
| 293 |
+
prepared_data["actions"] = actions
|
| 294 |
+
|
| 295 |
+
# embeddings are either fully preprocessed or transformed
|
| 296 |
+
if self.preprocess or (
|
| 297 |
+
self.decomp_method is not None and "embed" in self.decomp_method
|
| 298 |
+
):
|
| 299 |
+
# preprocess with frozen visual backbones
|
| 300 |
+
embeddings = self.preprocessor.process(
|
| 301 |
+
raw_img_batch=rgb_traj, return_numpy=not use_tensor
|
| 302 |
+
)
|
| 303 |
+
# L, N
|
| 304 |
+
assert embeddings.shape[0] == episode_length, embeddings.shape
|
| 305 |
+
else:
|
| 306 |
+
# L, 3, H, W
|
| 307 |
+
embeddings = self.preprocessor.transform_images(rgb_traj)
|
| 308 |
+
|
| 309 |
+
obs = {}
|
| 310 |
+
for k, v in full_obs_data.items():
|
| 311 |
+
if k in self.obs_keys and k == "rgb":
|
| 312 |
+
# raw rgb, may use for rollout
|
| 313 |
+
obs[k] = rgb_traj # v[:-1]
|
| 314 |
+
elif k in self.obs_keys:
|
| 315 |
+
obs[k] = (
|
| 316 |
+
U.any_to_torch_tensor(
|
| 317 |
+
v[:-1], dtype=torch.float32, device="cpu" # self.device
|
| 318 |
+
)
|
| 319 |
+
if use_tensor
|
| 320 |
+
else U.any_to_numpy(v[:-1], dtype="float32")
|
| 321 |
+
)
|
| 322 |
+
|
| 323 |
+
# obs = {k: v[:-1] for k, v in full_obs_data.items() if k in self.obs_keys}
|
| 324 |
+
prepared_data["obs"] = obs
|
| 325 |
+
rgb_no_robot = None
|
| 326 |
+
if raw_rgb_no_robot is not None:
|
| 327 |
+
if self.preprocess:
|
| 328 |
+
rgb_no_robot = self.preprocessor.process(
|
| 329 |
+
raw_img_batch=raw_rgb_no_robot,
|
| 330 |
+
return_numpy=not use_tensor,
|
| 331 |
+
reconstruct_linear=True,
|
| 332 |
+
)
|
| 333 |
+
else:
|
| 334 |
+
rgb_no_robot = self.preprocessor.transform_images(raw_rgb_no_robot)
|
| 335 |
+
prepared_data["rgb_no_robot"] = rgb_no_robot
|
| 336 |
+
|
| 337 |
+
if self.decomp_method == "oracle":
|
| 338 |
+
oracle_milestones = data["oracle_milestones"]
|
| 339 |
+
if self.replace_last_goal_as_last_frame:
|
| 340 |
+
oracle_milestones[-1, ...] = rgb_traj[-1] # full_obs_data["rgb"][-1]
|
| 341 |
+
if self.preprocess:
|
| 342 |
+
oracle_milestone_embeddings = self.preprocessor.process(
|
| 343 |
+
raw_img_batch=oracle_milestones, return_numpy=not use_tensor
|
| 344 |
+
)
|
| 345 |
+
else:
|
| 346 |
+
oracle_milestone_embeddings = self.preprocessor.transform_images(
|
| 347 |
+
oracle_milestones
|
| 348 |
+
)
|
| 349 |
+
assert len(oracle_milestone_embeddings) == len(oracle_milestones)
|
| 350 |
+
num_goals_achieved = data["completed_tasks"]
|
| 351 |
+
milestone_embeddings, decomp_meta = decomp_trajectories(
|
| 352 |
+
method_name=self.decomp_method,
|
| 353 |
+
embeddings=oracle_milestone_embeddings,
|
| 354 |
+
goal_achieved_mask=num_goals_achieved,
|
| 355 |
+
**self.decomp_kwargs or {},
|
| 356 |
+
)
|
| 357 |
+
else:
|
| 358 |
+
decomp_kwargs = self.decomp_kwargs or {}
|
| 359 |
+
if (
|
| 360 |
+
self.decomp_method is not None
|
| 361 |
+
and "embed_no_robot" in self.decomp_method
|
| 362 |
+
):
|
| 363 |
+
decomp_kwargs["no_robot_embeddings"] = rgb_no_robot
|
| 364 |
+
elif self.decomp_method == "embed2":
|
| 365 |
+
decomp_kwargs["no_robot_embeddings"] = embeddings
|
| 366 |
+
decomp_kwargs["task_name"] = str(data["reset_kwargs"]["task_elements"])
|
| 367 |
+
milestone_embeddings, decomp_meta = decomp_trajectories(
|
| 368 |
+
method_name=self.decomp_method,
|
| 369 |
+
embeddings=embeddings,
|
| 370 |
+
**decomp_kwargs,
|
| 371 |
+
)
|
| 372 |
+
|
| 373 |
+
if self.preprocessor.use_language_goal:
|
| 374 |
+
assert self.preprocess
|
| 375 |
+
prompt = task_elements_to_prompt(task_elements)
|
| 376 |
+
prompt_embed = self.preprocessor.encode_text(prompt)
|
| 377 |
+
milestone_embeddings = prompt_embed.repeat(episode_length, 1)
|
| 378 |
+
if not use_tensor:
|
| 379 |
+
milestone_embeddings = U.any_to_numpy(milestone_embeddings)
|
| 380 |
+
prepared_data["lang_embed"] = milestone_embeddings[-1]
|
| 381 |
+
|
| 382 |
+
if use_tensor:
|
| 383 |
+
embeddings = U.any_to_torch_tensor(embeddings, device="cpu")
|
| 384 |
+
milestone_embeddings = U.any_to_torch_tensor(
|
| 385 |
+
milestone_embeddings, device="cpu"
|
| 386 |
+
)
|
| 387 |
+
assert (
|
| 388 |
+
milestone_embeddings.shape == embeddings.shape
|
| 389 |
+
), f"{milestone_embeddings.shape} != {embeddings.shape}"
|
| 390 |
+
prepared_data["milestone_indices"] = decomp_meta.milestone_indices
|
| 391 |
+
|
| 392 |
+
def maybe_tensor(x, dtype):
|
| 393 |
+
return (
|
| 394 |
+
torch.tensor(x, device="cpu", dtype=dtype)
|
| 395 |
+
if use_tensor
|
| 396 |
+
else np.array(x, dtype=dtype)
|
| 397 |
+
)
|
| 398 |
+
|
| 399 |
+
if sequential:
|
| 400 |
+
if self.split_train_eval:
|
| 401 |
+
if episode.split("/")[-2] not in self.train_tasks:
|
| 402 |
+
assert episode.split("/")[-2] in self.eval_tasks, episode.split(
|
| 403 |
+
"/"
|
| 404 |
+
)[-2]
|
| 405 |
+
return prepared_data, episode_length
|
| 406 |
+
for st in range(episode_length):
|
| 407 |
+
seq_data = dict(
|
| 408 |
+
actions=actions[st],
|
| 409 |
+
milestones=milestone_embeddings[st],
|
| 410 |
+
episode_idx=maybe_tensor([episode_idx], dtype="int32"),
|
| 411 |
+
timesteps=maybe_tensor([st], dtype="int32"),
|
| 412 |
+
)
|
| 413 |
+
obs_i = {
|
| 414 |
+
k: v[st] if k != "rgb" else embeddings[st] for k, v in obs.items()
|
| 415 |
+
}
|
| 416 |
+
seq_data["obs"] = obs_i
|
| 417 |
+
self.sequence_data.append(seq_data)
|
| 418 |
+
else:
|
| 419 |
+
if self.split_train_eval:
|
| 420 |
+
if episode.split("/")[-2] not in self.train_tasks:
|
| 421 |
+
assert episode.split("/")[-2] in self.eval_tasks, episode.split(
|
| 422 |
+
"/"
|
| 423 |
+
)[-2]
|
| 424 |
+
return prepared_data, episode_length
|
| 425 |
+
self.sequence_data.append(
|
| 426 |
+
dict(
|
| 427 |
+
actions=actions,
|
| 428 |
+
obs={k: v if k != "rgb" else embeddings for k, v in obs.items()},
|
| 429 |
+
milestones=milestone_embeddings,
|
| 430 |
+
episode_idx=maybe_tensor([episode_idx], dtype="int32"),
|
| 431 |
+
timesteps=maybe_tensor(list(range(episode_length)), dtype="int32"),
|
| 432 |
+
)
|
| 433 |
+
)
|
| 434 |
+
return prepared_data, episode_length
|
| 435 |
+
|
| 436 |
+
def __len__(self) -> int:
|
| 437 |
+
if self.max_seq_length is None:
|
| 438 |
+
return len(self.sequence_data)
|
| 439 |
+
else:
|
| 440 |
+
return len(self.idx_partition)
|
| 441 |
+
|
| 442 |
+
def __getitem__(self, idx: int) -> dict:
|
| 443 |
+
"""Dict(obs, action, milestone, episode_idx, (maybe embedding))"""
|
| 444 |
+
if self.sequential:
|
| 445 |
+
return self.sequence_data[idx]
|
| 446 |
+
|
| 447 |
+
if self.max_seq_length is not None:
|
| 448 |
+
episode_idx, (start_idx, end_idx) = self.idx_partition[idx]
|
| 449 |
+
data = copy.deepcopy(self.sequence_data[episode_idx])
|
| 450 |
+
pad_len = self.max_seq_length - (end_idx - start_idx)
|
| 451 |
+
target_mask = np.array(
|
| 452 |
+
[1] * (end_idx - 1 - start_idx) + [0] * pad_len, dtype=np.int32
|
| 453 |
+
)
|
| 454 |
+
for k in data.keys():
|
| 455 |
+
val = data[k]
|
| 456 |
+
if k == "obs":
|
| 457 |
+
for obs_k, v in val.items():
|
| 458 |
+
if pad_len > 0:
|
| 459 |
+
data[k][obs_k] = U.any_concat(
|
| 460 |
+
[v[start_idx:end_idx, ...]]
|
| 461 |
+
+ [U.any_zeros_like(v[0])[None]] * pad_len
|
| 462 |
+
)
|
| 463 |
+
else:
|
| 464 |
+
data[k][obs_k] = v[start_idx:end_idx, ...]
|
| 465 |
+
elif k != "episode_idx":
|
| 466 |
+
if pad_len > 0:
|
| 467 |
+
data[k] = U.any_concat(
|
| 468 |
+
[val[start_idx:end_idx, ...]]
|
| 469 |
+
+ [U.any_zeros_like(val[0])[None]] * pad_len
|
| 470 |
+
)
|
| 471 |
+
else:
|
| 472 |
+
data[k] = val[start_idx:end_idx]
|
| 473 |
+
data["target_mask"] = target_mask
|
| 474 |
+
return data
|
| 475 |
+
|
| 476 |
+
data = self.sequence_data[idx] # seems no need to deepcopy st only pad once
|
| 477 |
+
actions = data["actions"]
|
| 478 |
+
ep_len = len(actions)
|
| 479 |
+
pad_len = self.max_episode_length - ep_len
|
| 480 |
+
if torch.is_tensor(actions):
|
| 481 |
+
target_mask = torch.tensor(
|
| 482 |
+
[1] * ep_len + [0] * pad_len, dtype=torch.int32, device=actions.device
|
| 483 |
+
)
|
| 484 |
+
else:
|
| 485 |
+
target_mask = np.array([1] * ep_len + [0] * pad_len, dtype=np.int32)
|
| 486 |
+
if pad_len > 0:
|
| 487 |
+
for k in data.keys():
|
| 488 |
+
val = data[k]
|
| 489 |
+
if k == "obs":
|
| 490 |
+
for obs_k, v in val.items():
|
| 491 |
+
if len(v) == ep_len:
|
| 492 |
+
data[k][obs_k] = U.any_concat(
|
| 493 |
+
[v] + [U.any_zeros_like(v[0])[None]] * pad_len
|
| 494 |
+
)
|
| 495 |
+
elif len(val) == ep_len:
|
| 496 |
+
data[k] = U.any_concat(
|
| 497 |
+
[val] + [U.any_zeros_like(val[0])[None]] * pad_len
|
| 498 |
+
)
|
| 499 |
+
data["target_mask"] = target_mask
|
| 500 |
+
return data
|
| 501 |
+
|
| 502 |
+
|
| 503 |
+
class FrankaKitchenRolloutDataset(Dataset):
|
| 504 |
+
"""Reuse training dataset for in-domain rollout evaluation."""
|
| 505 |
+
|
| 506 |
+
def __init__(
|
| 507 |
+
self,
|
| 508 |
+
train_dataset: FrankaKitchenDataset,
|
| 509 |
+
specific_tasks: str | list | None = None,
|
| 510 |
+
# has to be specific tasks
|
| 511 |
+
num_episodes_each_task: int | None = None,
|
| 512 |
+
num_tasks: int | None = None,
|
| 513 |
+
use_no_robot_rgb_milestones: bool = False,
|
| 514 |
+
**kwargs,
|
| 515 |
+
):
|
| 516 |
+
super().__init__()
|
| 517 |
+
assert "rgb" in train_dataset.obs_keys
|
| 518 |
+
|
| 519 |
+
self.is_hybrid = train_dataset.decomp_method == "embed_no_robot_extended"
|
| 520 |
+
|
| 521 |
+
self.max_num_milestones = -1
|
| 522 |
+
episode_data = copy.deepcopy(train_dataset.episode_data)
|
| 523 |
+
del train_dataset.episode_data
|
| 524 |
+
self.dataset_metadata = train_dataset.dataset_metadata
|
| 525 |
+
|
| 526 |
+
self._idx_to_episode_path = train_dataset.idx_to_episode_path
|
| 527 |
+
if specific_tasks is not None:
|
| 528 |
+
if isinstance(specific_tasks, str):
|
| 529 |
+
if "unseen" in specific_tasks:
|
| 530 |
+
specific_tasks = train_dataset.eval_tasks
|
| 531 |
+
elif "all" in specific_tasks:
|
| 532 |
+
specific_tasks = ALL_TASKS
|
| 533 |
+
else:
|
| 534 |
+
specific_tasks = (specific_tasks,)
|
| 535 |
+
# assume all data in the same parent directory
|
| 536 |
+
ds_base_path = "/".join(self._idx_to_episode_path[0].split("/")[:-2])
|
| 537 |
+
spec_episode_paths = [U.f_join(ds_base_path, t) for t in specific_tasks]
|
| 538 |
+
assert sum([U.f_exists(p) for p in spec_episode_paths]) == len(
|
| 539 |
+
spec_episode_paths
|
| 540 |
+
), spec_episode_paths
|
| 541 |
+
|
| 542 |
+
task_counter = {k: 1 for k in specific_tasks}
|
| 543 |
+
_idx_to_episode_path = {}
|
| 544 |
+
_n_episode_each_task = num_episodes_each_task or np.inf
|
| 545 |
+
for i, p in self._idx_to_episode_path.items():
|
| 546 |
+
task_name = p.split("/")[-2]
|
| 547 |
+
if (
|
| 548 |
+
task_name not in specific_tasks
|
| 549 |
+
or task_counter[task_name] > _n_episode_each_task
|
| 550 |
+
):
|
| 551 |
+
if task_name not in specific_tasks:
|
| 552 |
+
U.rank_zero_print(
|
| 553 |
+
f"WARNING: MISSING {task_name} in {specific_tasks}",
|
| 554 |
+
color="red",
|
| 555 |
+
)
|
| 556 |
+
continue
|
| 557 |
+
task_counter[task_name] += 1
|
| 558 |
+
_idx_to_episode_path[i] = p
|
| 559 |
+
self._idx_to_episode_path = _idx_to_episode_path
|
| 560 |
+
use_no_robot_rgb_milestones = (
|
| 561 |
+
use_no_robot_rgb_milestones and train_dataset.decomp_method is not None
|
| 562 |
+
)
|
| 563 |
+
self.use_no_robot_rgb_milestones = use_no_robot_rgb_milestones
|
| 564 |
+
|
| 565 |
+
self.num_demos = len(self._idx_to_episode_path)
|
| 566 |
+
self.num_tasks = min(self.num_demos, num_tasks or self.num_demos)
|
| 567 |
+
|
| 568 |
+
self.episode_data = {}
|
| 569 |
+
idx = -1
|
| 570 |
+
for ds_idx in tqdm.tqdm(
|
| 571 |
+
episode_data.keys(),
|
| 572 |
+
desc=f"load {self.__class__.__name__}",
|
| 573 |
+
total=len(self._idx_to_episode_path),
|
| 574 |
+
):
|
| 575 |
+
if ds_idx not in self._idx_to_episode_path:
|
| 576 |
+
continue
|
| 577 |
+
# else:
|
| 578 |
+
# print(ds_idx)
|
| 579 |
+
idx += 1
|
| 580 |
+
prepared_data = episode_data[ds_idx]
|
| 581 |
+
milestone_indices = prepared_data["milestone_indices"]
|
| 582 |
+
rgb = prepared_data["obs"]["rgb"]
|
| 583 |
+
rgb = (
|
| 584 |
+
U.any_to_torch_tensor(rgb, device="cpu", copy=True)
|
| 585 |
+
if torch.is_tensor(rgb)
|
| 586 |
+
else rgb.copy()
|
| 587 |
+
)
|
| 588 |
+
|
| 589 |
+
data = U.load_pickle(self._idx_to_episode_path[ds_idx])
|
| 590 |
+
reset_kwargs = {
|
| 591 |
+
k: U.any_to_numpy(v, dtype="float32")[None]
|
| 592 |
+
if k != "task_elements"
|
| 593 |
+
else np.array([ELEMENT_TO_IDX[ele] for ele in v], dtype=np.uint8)[None]
|
| 594 |
+
for k, v in data["reset_kwargs"].items()
|
| 595 |
+
}
|
| 596 |
+
|
| 597 |
+
if train_dataset.use_language_goal:
|
| 598 |
+
milestones = prepared_data["lang_embed"][None] # (1, embed_d)
|
| 599 |
+
else:
|
| 600 |
+
milestones = rgb[milestone_indices] # use raw rgb anyway
|
| 601 |
+
if torch.is_tensor(milestones):
|
| 602 |
+
milestones = U.any_to_torch_tensor(milestones, device="cpu")
|
| 603 |
+
if rgb.ndim != milestones.ndim:
|
| 604 |
+
rgb_milestones = U.any_to_numpy(rgb)[milestone_indices]
|
| 605 |
+
else:
|
| 606 |
+
rgb_milestones = None
|
| 607 |
+
rgb_no_robot = prepared_data.get("rgb_no_robot", None)
|
| 608 |
+
|
| 609 |
+
self.max_num_milestones = max(len(milestones), self.max_num_milestones)
|
| 610 |
+
if use_no_robot_rgb_milestones and rgb_no_robot is None:
|
| 611 |
+
rgb_no_robot = data["no_robot_rgb"][:-1]
|
| 612 |
+
|
| 613 |
+
rgb_no_robot_milestones = (
|
| 614 |
+
rgb_no_robot[milestone_indices] if use_no_robot_rgb_milestones else None
|
| 615 |
+
)
|
| 616 |
+
if torch.is_tensor(rgb_no_robot_milestones):
|
| 617 |
+
rgb_no_robot_milestones = U.any_to_torch_tensor(
|
| 618 |
+
rgb_no_robot_milestones, device="cpu"
|
| 619 |
+
)
|
| 620 |
+
if torch.is_tensor(milestones):
|
| 621 |
+
milestone_indices = torch.tensor(
|
| 622 |
+
milestones, dtype=torch.int32, device=milestones.device
|
| 623 |
+
)
|
| 624 |
+
else:
|
| 625 |
+
milestone_indices = np.array(milestone_indices, dtype=np.int32)
|
| 626 |
+
self.episode_data[idx] = dict(
|
| 627 |
+
milestones=milestones,
|
| 628 |
+
reset_kwargs=reset_kwargs,
|
| 629 |
+
rgb_milestones=rgb_milestones,
|
| 630 |
+
rgb_no_robot_milestones=rgb_no_robot_milestones,
|
| 631 |
+
milestone_indices=milestone_indices,
|
| 632 |
+
)
|
| 633 |
+
|
| 634 |
+
del episode_data
|
| 635 |
+
if self.num_tasks != self.num_demos:
|
| 636 |
+
self.episode_data = [
|
| 637 |
+
self.episode_data[i % len(self.episode_data)]
|
| 638 |
+
for i in range(self.num_tasks)
|
| 639 |
+
]
|
| 640 |
+
|
| 641 |
+
def __len__(self) -> int:
|
| 642 |
+
"""Num tasks evaluated."""
|
| 643 |
+
return self.num_tasks
|
| 644 |
+
|
| 645 |
+
def __getitem__(
|
| 646 |
+
self, idx: int
|
| 647 |
+
) -> OrderedDict[str, np.ndarray | dict[str, np.ndarray]]:
|
| 648 |
+
data = self.episode_data[idx]
|
| 649 |
+
milestones = data["milestones"]
|
| 650 |
+
milestone_indices = data.get("milestone_indices")
|
| 651 |
+
rgb_milestones = data["rgb_milestones"]
|
| 652 |
+
rgb_no_robot_milestones = data["rgb_no_robot_milestones"]
|
| 653 |
+
# assert milestones.shape[0] == rgb_milestones.shape[0]
|
| 654 |
+
num_milestones = len(milestones)
|
| 655 |
+
assert num_milestones <= self.max_num_milestones
|
| 656 |
+
if num_milestones < self.max_num_milestones:
|
| 657 |
+
# pad last
|
| 658 |
+
pad_length = self.max_num_milestones - num_milestones
|
| 659 |
+
milestones = U.any_concat(
|
| 660 |
+
[milestones] + [milestones[-1][None]] * pad_length
|
| 661 |
+
)
|
| 662 |
+
milestone_indices = U.any_concat(
|
| 663 |
+
[milestone_indices] + [milestone_indices[-1][None]] * pad_length
|
| 664 |
+
)
|
| 665 |
+
# if milestones.ndim != rgb_milestones.ndim:
|
| 666 |
+
if rgb_milestones is not None:
|
| 667 |
+
rgb_milestones = U.any_concat(
|
| 668 |
+
[rgb_milestones] + [rgb_milestones[-1][None]] * pad_length
|
| 669 |
+
)
|
| 670 |
+
if rgb_no_robot_milestones is not None:
|
| 671 |
+
rgb_no_robot_milestones = U.any_concat(
|
| 672 |
+
[rgb_no_robot_milestones]
|
| 673 |
+
+ [rgb_no_robot_milestones[-1][None]] * pad_length
|
| 674 |
+
)
|
| 675 |
+
rollout_data = OrderedDict(
|
| 676 |
+
milestones=milestones,
|
| 677 |
+
# rgb_milestones=rgb_milestones,
|
| 678 |
+
reset_kwargs=data["reset_kwargs"],
|
| 679 |
+
# rgb_no_robot_milestones=rgb_no_robot_milestones,
|
| 680 |
+
milestone_indices=milestone_indices,
|
| 681 |
+
)
|
| 682 |
+
if rgb_milestones is not None:
|
| 683 |
+
rollout_data["rgb_milestones"] = rgb_milestones
|
| 684 |
+
if rgb_no_robot_milestones is not None:
|
| 685 |
+
rollout_data["rgb_no_robot_milestones"] = rgb_no_robot_milestones
|
| 686 |
+
return rollout_data
|
| 687 |
+
|
| 688 |
+
|
| 689 |
+
PROMPT_DICT = dict(
|
| 690 |
+
microwave="open the microwave",
|
| 691 |
+
kettle="move the kettle to the top left stove",
|
| 692 |
+
light_switch="turn on the light",
|
| 693 |
+
hinge_cabinet="open the left hinge cabinet",
|
| 694 |
+
slide_cabinet="open the right slide cabinet",
|
| 695 |
+
top_burner="turn on the top left burner",
|
| 696 |
+
bottom_burner="turn on the bottom left burner",
|
| 697 |
+
)
|
| 698 |
+
PROMPT_DICT.update({k.replace("_", " "): v for k, v in PROMPT_DICT.items()})
|
| 699 |
+
|
| 700 |
+
|
| 701 |
+
def task_elements_to_prompt(task_elements: list) -> str | list[str]:
|
| 702 |
+
if isinstance(task_elements[0], str):
|
| 703 |
+
prompt = ", ".join([PROMPT_DICT[_] for _ in task_elements])
|
| 704 |
+
return prompt[0].capitalize() + prompt[1:]
|
| 705 |
+
# batch of task elements
|
| 706 |
+
assert isinstance(task_elements[0][0], str), task_elements
|
| 707 |
+
return [task_elements_to_prompt(ele) for ele in task_elements]
|
uvd/decomp/__init__.py
ADDED
|
@@ -0,0 +1 @@
|
|
|
|
|
|
|
| 1 |
+
from .decomp import decomp_trajectories
|
uvd/decomp/decomp.py
ADDED
|
@@ -0,0 +1,636 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
from __future__ import annotations
|
| 2 |
+
|
| 3 |
+
import random
|
| 4 |
+
from typing import Literal, NamedTuple, Callable
|
| 5 |
+
|
| 6 |
+
import numpy as np
|
| 7 |
+
import torch
|
| 8 |
+
import wandb
|
| 9 |
+
from matplotlib import pyplot as plt
|
| 10 |
+
from scipy.signal import medfilt
|
| 11 |
+
from scipy.signal import savgol_filter, argrelextrema
|
| 12 |
+
|
| 13 |
+
import uvd.utils as U
|
| 14 |
+
from uvd.decomp.kernel_reg import KernelRegression
|
| 15 |
+
|
| 16 |
+
|
| 17 |
+
def linear_random_skip(
|
| 18 |
+
cur_step: int, next_goal_step: int, ratio: float = 0.0, progress_lower: float = 0.3
|
| 19 |
+
) -> bool:
|
| 20 |
+
"""Progress toward next goal exceed `progress_lower`, linearly increasing
|
| 21 |
+
temperature for P(skip) <= ratio."""
|
| 22 |
+
if ratio == 0.0 or progress_lower == 1:
|
| 23 |
+
return False
|
| 24 |
+
assert cur_step <= next_goal_step, f"{cur_step} > {next_goal_step}"
|
| 25 |
+
# linearly increase until ratio
|
| 26 |
+
progress = cur_step / next_goal_step
|
| 27 |
+
temp = (
|
| 28 |
+
1.0
|
| 29 |
+
if progress < progress_lower
|
| 30 |
+
else ((progress - progress_lower) / (1 - progress_lower)) * ratio
|
| 31 |
+
)
|
| 32 |
+
if 0 < temp < random.random():
|
| 33 |
+
return True
|
| 34 |
+
return False
|
| 35 |
+
|
| 36 |
+
|
| 37 |
+
class DecompMeta(NamedTuple):
|
| 38 |
+
milestone_indices: list
|
| 39 |
+
milestone_starts: list | None = None
|
| 40 |
+
iter_curves: list[np.ndarray] | None = None
|
| 41 |
+
|
| 42 |
+
|
| 43 |
+
def _debug_plt(
|
| 44 |
+
xs: np.ndarray,
|
| 45 |
+
embed_distances: np.ndarray,
|
| 46 |
+
starts: np.ndarray | list,
|
| 47 |
+
ends: np.ndarray | list,
|
| 48 |
+
return_numpy: bool = False,
|
| 49 |
+
):
|
| 50 |
+
fig = plt.figure()
|
| 51 |
+
for s in starts:
|
| 52 |
+
plt.axvline(x=s, linestyle="--")
|
| 53 |
+
for e in ends:
|
| 54 |
+
plt.axvline(x=e, linestyle="dotted")
|
| 55 |
+
plt.plot(xs, np.gradient(embed_distances), linewidth=1.5, label="1st derivative")
|
| 56 |
+
plt.plot(
|
| 57 |
+
xs,
|
| 58 |
+
np.gradient(np.gradient(embed_distances)),
|
| 59 |
+
linewidth=1.5,
|
| 60 |
+
label="2nd derivative",
|
| 61 |
+
)
|
| 62 |
+
plt.plot(xs, embed_distances, linewidth=1.5, label="embedding distance")
|
| 63 |
+
plt.legend()
|
| 64 |
+
if return_numpy:
|
| 65 |
+
# fig.canvas.draw()
|
| 66 |
+
return U.plt_to_numpy(fig)
|
| 67 |
+
else:
|
| 68 |
+
plt.show()
|
| 69 |
+
|
| 70 |
+
|
| 71 |
+
def embed_decomp_no_robot(
|
| 72 |
+
embeddings: np.ndarray | torch.Tensor,
|
| 73 |
+
no_robot_embeddings: np.ndarray | torch.Tensor,
|
| 74 |
+
window_length: int = 10,
|
| 75 |
+
derivative_order: int = 1,
|
| 76 |
+
derivative_threshold: float = 1e-3,
|
| 77 |
+
threshold_subgoal_passing: float | None = None,
|
| 78 |
+
force_interleave: bool = False,
|
| 79 |
+
debug_plt: bool = False,
|
| 80 |
+
debug_plt_to_wandb: bool = False,
|
| 81 |
+
task_name: str | None = None,
|
| 82 |
+
fill_embeddings: bool = True,
|
| 83 |
+
):
|
| 84 |
+
debug_plt = debug_plt and U.is_rank_zero()
|
| 85 |
+
debug_plt_to_wandb = debug_plt and debug_plt_to_wandb
|
| 86 |
+
if threshold_subgoal_passing is not None:
|
| 87 |
+
assert 0 < threshold_subgoal_passing <= 1.0, threshold_subgoal_passing
|
| 88 |
+
if isinstance(embeddings, torch.Tensor):
|
| 89 |
+
device = embeddings.device
|
| 90 |
+
else:
|
| 91 |
+
device = None
|
| 92 |
+
# clip based preprocessors would have bf16 by default
|
| 93 |
+
embeddings = U.any_to_numpy(embeddings, dtype="float32")
|
| 94 |
+
no_robot_embeddings = U.any_to_numpy(no_robot_embeddings, dtype="float32")
|
| 95 |
+
# L, N (though can be rgb for debugging as well)
|
| 96 |
+
traj_length = embeddings.shape[0]
|
| 97 |
+
|
| 98 |
+
debug_plt_figs = []
|
| 99 |
+
if fill_embeddings:
|
| 100 |
+
milestone_embeddings = []
|
| 101 |
+
else:
|
| 102 |
+
milestone_embeddings = None
|
| 103 |
+
# milestone indices, i.e. end of obj changing
|
| 104 |
+
subgoal_indices = []
|
| 105 |
+
# indices when obj changing starts, i.e. milestone indices for hand reaching
|
| 106 |
+
subgoal_starts = []
|
| 107 |
+
cur_subgoal_idx = traj_length - 1
|
| 108 |
+
# back to start
|
| 109 |
+
while cur_subgoal_idx > 15:
|
| 110 |
+
unnormalized_embed_distance = np.linalg.norm(
|
| 111 |
+
no_robot_embeddings[: cur_subgoal_idx + 1]
|
| 112 |
+
- no_robot_embeddings[cur_subgoal_idx],
|
| 113 |
+
axis=1,
|
| 114 |
+
)
|
| 115 |
+
cur_embed_distance = unnormalized_embed_distance / np.linalg.norm(
|
| 116 |
+
no_robot_embeddings[0] - no_robot_embeddings[cur_subgoal_idx]
|
| 117 |
+
)
|
| 118 |
+
cur_embed_distance = medfilt(cur_embed_distance, kernel_size=None) # smooth
|
| 119 |
+
|
| 120 |
+
if derivative_order == 1:
|
| 121 |
+
slope = np.gradient(cur_embed_distance)
|
| 122 |
+
elif derivative_order == 2:
|
| 123 |
+
slope = np.gradient(np.gradient(cur_embed_distance))
|
| 124 |
+
else:
|
| 125 |
+
raise NotImplementedError(derivative_order)
|
| 126 |
+
valid_slope_indices = np.where(np.abs(slope) >= derivative_threshold)[0]
|
| 127 |
+
# Find the differences between consecutive valid_slope_indices
|
| 128 |
+
diffs = np.diff(valid_slope_indices)
|
| 129 |
+
# Find indices where the difference is greater than window_length + 1
|
| 130 |
+
break_indices = np.where(diffs > window_length + 1)[0]
|
| 131 |
+
# Extract start and end positions
|
| 132 |
+
start_positions = valid_slope_indices[np.concatenate([[0], break_indices + 1])]
|
| 133 |
+
end_positions = valid_slope_indices[
|
| 134 |
+
np.concatenate((break_indices, [len(valid_slope_indices) - 1]))
|
| 135 |
+
]
|
| 136 |
+
|
| 137 |
+
# small tolerance
|
| 138 |
+
start_positions = [max(0, p - 1) for p in start_positions]
|
| 139 |
+
end_positions = [min(traj_length - 1, p + 1) for p in end_positions]
|
| 140 |
+
|
| 141 |
+
if debug_plt:
|
| 142 |
+
x = np.arange(0, len(cur_embed_distance))
|
| 143 |
+
fig = _debug_plt(
|
| 144 |
+
x,
|
| 145 |
+
cur_embed_distance,
|
| 146 |
+
start_positions,
|
| 147 |
+
end_positions,
|
| 148 |
+
return_numpy=debug_plt_to_wandb,
|
| 149 |
+
)
|
| 150 |
+
if debug_plt_to_wandb:
|
| 151 |
+
debug_plt_figs.append(fig)
|
| 152 |
+
|
| 153 |
+
subgoal_indices.append(cur_subgoal_idx)
|
| 154 |
+
subgoal_starts.append(start_positions[-1])
|
| 155 |
+
if len(end_positions) < 2 or end_positions[-2] < 15:
|
| 156 |
+
if threshold_subgoal_passing is not None:
|
| 157 |
+
cur_milestone_dist = None
|
| 158 |
+
if len(subgoal_indices) > 1:
|
| 159 |
+
cur_milestone_dist = np.linalg.norm(
|
| 160 |
+
no_robot_embeddings[cur_subgoal_idx] # cur
|
| 161 |
+
- no_robot_embeddings[subgoal_indices[-2]] # prev
|
| 162 |
+
) # use the order from start to end
|
| 163 |
+
if fill_embeddings:
|
| 164 |
+
for step in reversed(range(cur_subgoal_idx + 1)):
|
| 165 |
+
if (
|
| 166 |
+
len(subgoal_indices) == 1
|
| 167 |
+
or unnormalized_embed_distance[step] / cur_milestone_dist
|
| 168 |
+
> threshold_subgoal_passing
|
| 169 |
+
):
|
| 170 |
+
milestone_embeddings.append(embeddings[cur_subgoal_idx])
|
| 171 |
+
else:
|
| 172 |
+
milestone_embeddings.append(embeddings[subgoal_indices[-2]])
|
| 173 |
+
break
|
| 174 |
+
if threshold_subgoal_passing is not None:
|
| 175 |
+
cur_milestone_dist = None
|
| 176 |
+
if len(subgoal_indices) > 1:
|
| 177 |
+
cur_milestone_dist = np.linalg.norm(
|
| 178 |
+
no_robot_embeddings[cur_subgoal_idx] # cur
|
| 179 |
+
- no_robot_embeddings[subgoal_indices[-2]] # prev
|
| 180 |
+
)
|
| 181 |
+
if fill_embeddings:
|
| 182 |
+
for step in reversed(range(end_positions[-2] + 1, cur_subgoal_idx + 1)):
|
| 183 |
+
# not pass threshold or 1st iter with only last frame contained
|
| 184 |
+
if (
|
| 185 |
+
len(subgoal_indices) == 1
|
| 186 |
+
or unnormalized_embed_distance[step] / cur_milestone_dist
|
| 187 |
+
> threshold_subgoal_passing
|
| 188 |
+
):
|
| 189 |
+
# use the real subgoal
|
| 190 |
+
milestone_embeddings.append(embeddings[cur_subgoal_idx])
|
| 191 |
+
else:
|
| 192 |
+
# skip to the next subgoal if passing threshold
|
| 193 |
+
milestone_embeddings.append(embeddings[subgoal_indices[-2]])
|
| 194 |
+
cur_subgoal_idx = end_positions[-2]
|
| 195 |
+
|
| 196 |
+
subgoal_starts = list(reversed(subgoal_starts))
|
| 197 |
+
subgoal_indices = list(reversed(subgoal_indices))
|
| 198 |
+
|
| 199 |
+
if len(debug_plt_figs) > 0 and wandb.run is not None:
|
| 200 |
+
debug_plt_figs = np.concatenate(debug_plt_figs, axis=1)
|
| 201 |
+
wandb.log(
|
| 202 |
+
{
|
| 203 |
+
f"decomp_curves/{task_name}": wandb.Image(
|
| 204 |
+
debug_plt_figs,
|
| 205 |
+
caption=f"starts: {subgoal_starts}, ends: {subgoal_indices}",
|
| 206 |
+
)
|
| 207 |
+
}
|
| 208 |
+
)
|
| 209 |
+
|
| 210 |
+
assert len(subgoal_starts) == len(
|
| 211 |
+
subgoal_indices
|
| 212 |
+
), f"{subgoal_starts}, {subgoal_indices}"
|
| 213 |
+
|
| 214 |
+
if force_interleave:
|
| 215 |
+
_starts, _ends = [], []
|
| 216 |
+
for i in range(len(subgoal_starts)):
|
| 217 |
+
if subgoal_starts[i] < subgoal_indices[i]:
|
| 218 |
+
_starts.append(subgoal_starts[i])
|
| 219 |
+
_ends.append(subgoal_indices[i])
|
| 220 |
+
else:
|
| 221 |
+
U.get_logger().warning(
|
| 222 |
+
f"{subgoal_starts} & {subgoal_indices} not interleaved"
|
| 223 |
+
)
|
| 224 |
+
|
| 225 |
+
if fill_embeddings:
|
| 226 |
+
if threshold_subgoal_passing is not None:
|
| 227 |
+
milestone_embeddings = np.stack(list(reversed(milestone_embeddings)))
|
| 228 |
+
else:
|
| 229 |
+
# slightly faster to do once here without threshold checking
|
| 230 |
+
milestone_embeddings = np.concatenate(
|
| 231 |
+
[embeddings[subgoal_indices[0], ...][None]]
|
| 232 |
+
+ [
|
| 233 |
+
np.full((end - start, *embeddings.shape[1:]), embeddings[end, ...])
|
| 234 |
+
for start, end in zip([0] + subgoal_indices[:-1], subgoal_indices)
|
| 235 |
+
],
|
| 236 |
+
)
|
| 237 |
+
|
| 238 |
+
if device is not None:
|
| 239 |
+
milestone_embeddings = U.any_to_torch_tensor(
|
| 240 |
+
milestone_embeddings, device=device
|
| 241 |
+
)
|
| 242 |
+
return milestone_embeddings, DecompMeta(
|
| 243 |
+
milestone_indices=subgoal_indices, milestone_starts=subgoal_starts
|
| 244 |
+
)
|
| 245 |
+
|
| 246 |
+
|
| 247 |
+
def embed_decomp_no_robot_extended(
|
| 248 |
+
embeddings: np.ndarray | torch.Tensor,
|
| 249 |
+
no_robot_embeddings: np.ndarray | torch.Tensor,
|
| 250 |
+
threshold_subgoal_passing: float | None = None,
|
| 251 |
+
**kwargs,
|
| 252 |
+
):
|
| 253 |
+
kwargs["fill_embeddings"] = False
|
| 254 |
+
_, decomp_meta = embed_decomp_no_robot(
|
| 255 |
+
embeddings,
|
| 256 |
+
no_robot_embeddings,
|
| 257 |
+
threshold_subgoal_passing=None,
|
| 258 |
+
**kwargs,
|
| 259 |
+
)
|
| 260 |
+
milestone_indices = decomp_meta.milestone_indices
|
| 261 |
+
milestone_starts = decomp_meta.milestone_starts
|
| 262 |
+
norm = (
|
| 263 |
+
np.linalg.norm if isinstance(embeddings[0], np.ndarray) else torch.linalg.norm
|
| 264 |
+
)
|
| 265 |
+
assert len(milestone_starts) == len(milestone_indices)
|
| 266 |
+
|
| 267 |
+
milestone_embeddings = []
|
| 268 |
+
hybrid_indices = list(sorted(milestone_starts + milestone_indices))
|
| 269 |
+
prev_idx = -1
|
| 270 |
+
s = -1
|
| 271 |
+
init_dist = None
|
| 272 |
+
for i, goal_idx in enumerate(hybrid_indices):
|
| 273 |
+
once_passed = False
|
| 274 |
+
for _ in range(goal_idx - prev_idx):
|
| 275 |
+
s += 1
|
| 276 |
+
if threshold_subgoal_passing is None:
|
| 277 |
+
milestone_embeddings.append(embeddings[goal_idx])
|
| 278 |
+
elif once_passed:
|
| 279 |
+
milestone_embeddings.append(embeddings[hybrid_indices[i + 1]])
|
| 280 |
+
else:
|
| 281 |
+
raw_cur_dist = float(norm(embeddings[s] - embeddings[goal_idx]))
|
| 282 |
+
if init_dist is None:
|
| 283 |
+
assert s == 0, s
|
| 284 |
+
cur_dist = 1.0
|
| 285 |
+
init_dist = max(raw_cur_dist, 1e-7)
|
| 286 |
+
else:
|
| 287 |
+
cur_dist = raw_cur_dist / init_dist
|
| 288 |
+
if (
|
| 289 |
+
cur_dist <= threshold_subgoal_passing
|
| 290 |
+
and i < len(hybrid_indices) - 1
|
| 291 |
+
):
|
| 292 |
+
init_dist = float(
|
| 293 |
+
norm(embeddings[s] - embeddings[hybrid_indices[i + 1]])
|
| 294 |
+
)
|
| 295 |
+
init_dist = max(init_dist, 1e-7)
|
| 296 |
+
once_passed = True
|
| 297 |
+
milestone_embeddings.append(embeddings[hybrid_indices[i + 1]])
|
| 298 |
+
else:
|
| 299 |
+
milestone_embeddings.append(embeddings[goal_idx])
|
| 300 |
+
prev_idx = goal_idx
|
| 301 |
+
|
| 302 |
+
milestone_embeddings = U.any_stack(milestone_embeddings)
|
| 303 |
+
U.assert_(milestone_embeddings.shape, embeddings.shape)
|
| 304 |
+
return milestone_embeddings, DecompMeta(milestone_indices=hybrid_indices)
|
| 305 |
+
|
| 306 |
+
|
| 307 |
+
def get_hybrid_milestones(
|
| 308 |
+
start_embeddings: np.ndarray, # w. robot
|
| 309 |
+
end_embeddings: np.ndarray, # w.o robot
|
| 310 |
+
milestone_starts: list,
|
| 311 |
+
milestone_indices: list,
|
| 312 |
+
) -> np.ndarray:
|
| 313 |
+
assert len(milestone_starts) == len(milestone_indices)
|
| 314 |
+
assert len(start_embeddings) == len(end_embeddings)
|
| 315 |
+
milestone_only = len(milestone_starts) == len(start_embeddings)
|
| 316 |
+
hybrid_milestones = np.empty(
|
| 317 |
+
(start_embeddings.shape[0] * 2, *start_embeddings.shape[1:]),
|
| 318 |
+
dtype=start_embeddings.dtype,
|
| 319 |
+
)
|
| 320 |
+
hybrid_milestones[::2] = (
|
| 321 |
+
start_embeddings if milestone_only else start_embeddings[milestone_starts]
|
| 322 |
+
)
|
| 323 |
+
hybrid_milestones[1::2] = (
|
| 324 |
+
end_embeddings if milestone_only else end_embeddings[milestone_indices]
|
| 325 |
+
)
|
| 326 |
+
return hybrid_milestones
|
| 327 |
+
|
| 328 |
+
|
| 329 |
+
def embedding_decomp(
|
| 330 |
+
embeddings: np.ndarray | torch.Tensor,
|
| 331 |
+
normalize_curve: bool = True,
|
| 332 |
+
min_interval: int = 18,
|
| 333 |
+
window_length: int | None = None,
|
| 334 |
+
smooth_method: Literal["kernel", "savgol"] = "kernel",
|
| 335 |
+
extrema_comparator: Callable = np.greater,
|
| 336 |
+
fill_embeddings: bool = True,
|
| 337 |
+
return_intermediate_curves: bool = False,
|
| 338 |
+
**kwargs,
|
| 339 |
+
) -> tuple[torch.Tensor | np.ndarray, DecompMeta]:
|
| 340 |
+
if torch.is_tensor(embeddings):
|
| 341 |
+
device = embeddings.device
|
| 342 |
+
embeddings = U.any_to_numpy(embeddings)
|
| 343 |
+
else:
|
| 344 |
+
device = None
|
| 345 |
+
# L, N
|
| 346 |
+
assert embeddings.ndim == 2, embeddings.shape
|
| 347 |
+
traj_length = embeddings.shape[0]
|
| 348 |
+
|
| 349 |
+
cur_goal_idx = traj_length - 1
|
| 350 |
+
goal_indices = [cur_goal_idx]
|
| 351 |
+
cur_embeddings = embeddings[
|
| 352 |
+
max(0, cur_goal_idx - (window_length or cur_goal_idx)) : cur_goal_idx + 1
|
| 353 |
+
]
|
| 354 |
+
iterate_num = 0
|
| 355 |
+
iter_curves = [] if return_intermediate_curves else None
|
| 356 |
+
while cur_goal_idx > (window_length or min_interval):
|
| 357 |
+
iterate_num += 1
|
| 358 |
+
# get goal embedding
|
| 359 |
+
goal_embedding = cur_embeddings[-1]
|
| 360 |
+
distances = np.linalg.norm(cur_embeddings - goal_embedding, axis=1)
|
| 361 |
+
if normalize_curve:
|
| 362 |
+
distances = distances / np.linalg.norm(cur_embeddings[0] - goal_embedding)
|
| 363 |
+
|
| 364 |
+
x = np.arange(
|
| 365 |
+
max(0, cur_goal_idx - (window_length or cur_goal_idx)), cur_goal_idx + 1
|
| 366 |
+
)
|
| 367 |
+
|
| 368 |
+
if smooth_method == "kernel":
|
| 369 |
+
smooth_kwargs = dict(kernel="rbf", gamma=0.08)
|
| 370 |
+
smooth_kwargs.update(kwargs or {})
|
| 371 |
+
kr = KernelRegression(**smooth_kwargs)
|
| 372 |
+
kr.fit(x.reshape(-1, 1), distances)
|
| 373 |
+
distance_smoothed = kr.predict(x.reshape(-1, 1))
|
| 374 |
+
elif smooth_method == "savgol":
|
| 375 |
+
smooth_kwargs = dict(window_length=85, polyorder=2, mode="nearest")
|
| 376 |
+
smooth_kwargs.update(kwargs or {})
|
| 377 |
+
distance_smoothed = savgol_filter(distances, **smooth_kwargs)
|
| 378 |
+
elif smooth_method is None:
|
| 379 |
+
distance_smoothed = distances
|
| 380 |
+
else:
|
| 381 |
+
raise NotImplementedError(smooth_method)
|
| 382 |
+
|
| 383 |
+
if iter_curves is not None:
|
| 384 |
+
iter_curves.append(distance_smoothed)
|
| 385 |
+
|
| 386 |
+
extrema_indices = argrelextrema(distance_smoothed, extrema_comparator)[0]
|
| 387 |
+
x_extrema = x[extrema_indices]
|
| 388 |
+
|
| 389 |
+
update_goal = False
|
| 390 |
+
for i in range(len(x_extrema) - 1, -1, -1):
|
| 391 |
+
if cur_goal_idx < min_interval:
|
| 392 |
+
break
|
| 393 |
+
if (
|
| 394 |
+
cur_goal_idx - x_extrema[i] > min_interval
|
| 395 |
+
and x_extrema[i] > min_interval
|
| 396 |
+
):
|
| 397 |
+
cur_goal_idx = x_extrema[i]
|
| 398 |
+
update_goal = True
|
| 399 |
+
goal_indices.append(cur_goal_idx)
|
| 400 |
+
break
|
| 401 |
+
|
| 402 |
+
if not update_goal or cur_goal_idx < min_interval:
|
| 403 |
+
break
|
| 404 |
+
cur_embeddings = embeddings[
|
| 405 |
+
max(0, cur_goal_idx - (window_length or cur_goal_idx)) : cur_goal_idx + 1
|
| 406 |
+
]
|
| 407 |
+
|
| 408 |
+
goal_indices = goal_indices[::-1]
|
| 409 |
+
if fill_embeddings:
|
| 410 |
+
milestone_embeddings = np.concatenate(
|
| 411 |
+
[embeddings[goal_indices[0], ...][None]]
|
| 412 |
+
+ [
|
| 413 |
+
np.full((end - start, *embeddings.shape[1:]), embeddings[end, ...])
|
| 414 |
+
for start, end in zip([0] + goal_indices[:-1], goal_indices)
|
| 415 |
+
],
|
| 416 |
+
)
|
| 417 |
+
if device is not None:
|
| 418 |
+
milestone_embeddings = U.any_to_torch_tensor(
|
| 419 |
+
milestone_embeddings, device=device
|
| 420 |
+
)
|
| 421 |
+
else:
|
| 422 |
+
milestone_embeddings = None
|
| 423 |
+
return milestone_embeddings, DecompMeta(
|
| 424 |
+
milestone_indices=goal_indices, iter_curves=iter_curves
|
| 425 |
+
)
|
| 426 |
+
|
| 427 |
+
|
| 428 |
+
def goal_idx_from_mask(goal_achieved_mask):
|
| 429 |
+
diff = np.diff(goal_achieved_mask)
|
| 430 |
+
goal_indices = np.where(diff != 0)[0] + 1
|
| 431 |
+
traj_length = goal_achieved_mask.shape[0]
|
| 432 |
+
goal_indices[-1] = traj_length - 1 # last
|
| 433 |
+
goal_indices = goal_indices.tolist()
|
| 434 |
+
return goal_indices
|
| 435 |
+
|
| 436 |
+
|
| 437 |
+
def oracle_decomp(
|
| 438 |
+
embeddings: np.ndarray | torch.Tensor | None,
|
| 439 |
+
goal_achieved_mask: np.ndarray,
|
| 440 |
+
random_skip_ratio: float | None = None,
|
| 441 |
+
linearly_random_skip_lower: float | None = None,
|
| 442 |
+
fill_embeddings: bool = True,
|
| 443 |
+
) -> tuple[torch.Tensor | np.ndarray, DecompMeta]:
|
| 444 |
+
"""Note: embeddings here only has the oracle subgoals, not full trajectory"""
|
| 445 |
+
goal_indices = goal_idx_from_mask(goal_achieved_mask)
|
| 446 |
+
if not fill_embeddings:
|
| 447 |
+
return None, DecompMeta(milestone_indices=goal_indices)
|
| 448 |
+
|
| 449 |
+
traj_length = goal_achieved_mask.shape[0]
|
| 450 |
+
assert embeddings.shape[0] < traj_length, embeddings.shape
|
| 451 |
+
milestone_embeddings = (
|
| 452 |
+
torch.empty(
|
| 453 |
+
(goal_achieved_mask.shape[0], *embeddings.shape[1:]),
|
| 454 |
+
dtype=embeddings.dtype,
|
| 455 |
+
device=embeddings.device,
|
| 456 |
+
)
|
| 457 |
+
if not isinstance(embeddings, np.ndarray)
|
| 458 |
+
else np.empty(
|
| 459 |
+
(goal_achieved_mask.shape[0], *embeddings.shape[1:]),
|
| 460 |
+
dtype=embeddings.dtype,
|
| 461 |
+
)
|
| 462 |
+
)
|
| 463 |
+
|
| 464 |
+
assert len(goal_indices) == len(embeddings), goal_indices
|
| 465 |
+
|
| 466 |
+
last_embedding = embeddings[-1]
|
| 467 |
+
for i, idx in enumerate(goal_achieved_mask):
|
| 468 |
+
if idx >= embeddings.shape[0]:
|
| 469 |
+
# If the index in the mask is greater than the highest index in the embedding,
|
| 470 |
+
# just use the last row of the embedding
|
| 471 |
+
milestone_embeddings[i] = last_embedding
|
| 472 |
+
else:
|
| 473 |
+
skip = False
|
| 474 |
+
if linearly_random_skip_lower is not None:
|
| 475 |
+
skip = linear_random_skip(
|
| 476 |
+
cur_step=i,
|
| 477 |
+
next_goal_step=goal_indices[idx],
|
| 478 |
+
ratio=random_skip_ratio,
|
| 479 |
+
progress_lower=linearly_random_skip_lower,
|
| 480 |
+
)
|
| 481 |
+
elif (
|
| 482 |
+
random_skip_ratio is not None
|
| 483 |
+
and 0 < random_skip_ratio < random.random()
|
| 484 |
+
):
|
| 485 |
+
skip = True
|
| 486 |
+
if skip:
|
| 487 |
+
milestone_embeddings[i] = embeddings[min(idx + 1, len(embeddings) - 1)]
|
| 488 |
+
else:
|
| 489 |
+
milestone_embeddings[i] = embeddings[idx]
|
| 490 |
+
return milestone_embeddings, DecompMeta(milestone_indices=goal_indices)
|
| 491 |
+
|
| 492 |
+
|
| 493 |
+
def random_decomp(
|
| 494 |
+
embeddings: np.ndarray | torch.Tensor,
|
| 495 |
+
num_milestones: int | tuple[int, int],
|
| 496 |
+
fill_embeddings: bool = True,
|
| 497 |
+
) -> tuple[torch.Tensor | np.ndarray, DecompMeta]:
|
| 498 |
+
if not isinstance(num_milestones, int):
|
| 499 |
+
assert len(num_milestones) == 2, num_milestones
|
| 500 |
+
# by randomly sample from lower and higher bound
|
| 501 |
+
num_milestones = random.randint(*num_milestones)
|
| 502 |
+
traj_length = embeddings.shape[0]
|
| 503 |
+
goal_indices = random.sample(range(traj_length), k=num_milestones)
|
| 504 |
+
goal_indices = list(sorted(goal_indices))
|
| 505 |
+
if fill_embeddings:
|
| 506 |
+
milestone_embeddings = (
|
| 507 |
+
torch.empty_like(
|
| 508 |
+
embeddings, dtype=embeddings.dtype, device=embeddings.device
|
| 509 |
+
)
|
| 510 |
+
if not isinstance(embeddings, np.ndarray)
|
| 511 |
+
else np.empty_like(embeddings, dtype=embeddings.dtype)
|
| 512 |
+
)
|
| 513 |
+
for i, goal_idx in enumerate(goal_indices):
|
| 514 |
+
milestone_embeddings[
|
| 515 |
+
(goal_indices[i - 1] + 1) if i != 0 else 0 : goal_idx + 1
|
| 516 |
+
] = embeddings[goal_idx]
|
| 517 |
+
else:
|
| 518 |
+
milestone_embeddings = None
|
| 519 |
+
return milestone_embeddings, DecompMeta(milestone_indices=goal_indices)
|
| 520 |
+
|
| 521 |
+
|
| 522 |
+
def equally_decomp(
|
| 523 |
+
embeddings: np.ndarray | torch.Tensor,
|
| 524 |
+
num_milestones: int | tuple[int, int],
|
| 525 |
+
fill_embeddings: bool = True,
|
| 526 |
+
) -> tuple[torch.Tensor | np.ndarray, DecompMeta]:
|
| 527 |
+
if not isinstance(num_milestones, int):
|
| 528 |
+
assert len(num_milestones) == 2, num_milestones
|
| 529 |
+
# by randomly sample from lower and higher bound
|
| 530 |
+
num_milestones = random.randint(*num_milestones)
|
| 531 |
+
traj_length = embeddings.shape[0]
|
| 532 |
+
indices = np.linspace(0, traj_length - 1, num_milestones + 1, dtype=int)
|
| 533 |
+
if fill_embeddings:
|
| 534 |
+
milestone_embeddings = (
|
| 535 |
+
torch.empty_like(
|
| 536 |
+
embeddings, dtype=embeddings.dtype, device=embeddings.device
|
| 537 |
+
)
|
| 538 |
+
if not isinstance(embeddings, np.ndarray)
|
| 539 |
+
else np.empty_like(
|
| 540 |
+
embeddings,
|
| 541 |
+
dtype=embeddings.dtype,
|
| 542 |
+
)
|
| 543 |
+
)
|
| 544 |
+
for i, goal_idx in enumerate(indices[1:], start=1):
|
| 545 |
+
milestone_embeddings[
|
| 546 |
+
(indices[i - 1] + 1) if i != 1 else 0 : goal_idx + 1
|
| 547 |
+
] = embeddings[goal_idx]
|
| 548 |
+
else:
|
| 549 |
+
milestone_embeddings = None
|
| 550 |
+
return milestone_embeddings, DecompMeta(milestone_indices=indices[1:].tolist())
|
| 551 |
+
|
| 552 |
+
|
| 553 |
+
def near_future_decomp(
|
| 554 |
+
embeddings: np.ndarray | torch.Tensor, advance_steps: int, **kwargs
|
| 555 |
+
) -> tuple[torch.Tensor | np.ndarray, DecompMeta]:
|
| 556 |
+
return equally_decomp(
|
| 557 |
+
embeddings, num_milestones=embeddings.shape[0] // advance_steps, **kwargs
|
| 558 |
+
)
|
| 559 |
+
|
| 560 |
+
|
| 561 |
+
def no_decomp(
|
| 562 |
+
embeddings: np.ndarray | torch.Tensor, fill_embeddings: bool = True
|
| 563 |
+
) -> tuple[torch.Tensor | np.ndarray, DecompMeta]:
|
| 564 |
+
"""Only conditioned on final goal."""
|
| 565 |
+
if not fill_embeddings:
|
| 566 |
+
return None, DecompMeta(milestone_indices=[-1])
|
| 567 |
+
return embeddings[-1, ...].expand_as(embeddings).clone() if not isinstance(
|
| 568 |
+
embeddings, np.ndarray
|
| 569 |
+
) else np.full(
|
| 570 |
+
embeddings.shape,
|
| 571 |
+
embeddings[-1, ...],
|
| 572 |
+
dtype=embeddings.dtype,
|
| 573 |
+
), DecompMeta(
|
| 574 |
+
milestone_indices=[-1]
|
| 575 |
+
)
|
| 576 |
+
|
| 577 |
+
|
| 578 |
+
def decomp_trajectories(
|
| 579 |
+
method_name: Literal[
|
| 580 |
+
"embed", "embed_no_robot", "oracle", "random", "equally", "near_future"
|
| 581 |
+
]
|
| 582 |
+
| None,
|
| 583 |
+
embeddings: np.ndarray | torch.Tensor,
|
| 584 |
+
**kwargs,
|
| 585 |
+
) -> tuple[torch.Tensor | np.ndarray, DecompMeta]:
|
| 586 |
+
assert embeddings.ndim == 2 or embeddings.ndim == 4, (
|
| 587 |
+
f"input embedding should be either 2 dimensional, "
|
| 588 |
+
f"with (L, feature_dim), or raw rgb with shape (L, H, W, 3), "
|
| 589 |
+
f"but get {embeddings.shape}"
|
| 590 |
+
)
|
| 591 |
+
if method_name is None:
|
| 592 |
+
return no_decomp(embeddings)
|
| 593 |
+
assert method_name in DEFAULT_DECOMP_KWARGS, method_name
|
| 594 |
+
method_kwargs = DEFAULT_DECOMP_KWARGS[method_name]
|
| 595 |
+
method_kwargs.update(kwargs)
|
| 596 |
+
if method_name == "embed":
|
| 597 |
+
return embedding_decomp(embeddings=embeddings, **method_kwargs)
|
| 598 |
+
elif method_name == "embed_no_robot":
|
| 599 |
+
return embed_decomp_no_robot(embeddings=embeddings, **method_kwargs)
|
| 600 |
+
elif method_name == "embed_no_robot_extended":
|
| 601 |
+
return embed_decomp_no_robot_extended(embeddings=embeddings, **method_kwargs)
|
| 602 |
+
elif method_name == "oracle":
|
| 603 |
+
return oracle_decomp(embeddings, **method_kwargs)
|
| 604 |
+
elif method_name == "random":
|
| 605 |
+
return random_decomp(embeddings, **method_kwargs)
|
| 606 |
+
elif method_name == "equally":
|
| 607 |
+
return equally_decomp(embeddings, **method_kwargs)
|
| 608 |
+
elif method_name == "near_future":
|
| 609 |
+
return near_future_decomp(embeddings, **method_kwargs)
|
| 610 |
+
raise NotImplementedError(method_name)
|
| 611 |
+
|
| 612 |
+
|
| 613 |
+
DEFAULT_DECOMP_KWARGS = dict(
|
| 614 |
+
embed=dict(
|
| 615 |
+
normalize_curve=False,
|
| 616 |
+
min_interval=18,
|
| 617 |
+
smooth_method="kernel",
|
| 618 |
+
gamma=0.08,
|
| 619 |
+
),
|
| 620 |
+
embed_no_robot=dict(
|
| 621 |
+
window_length=8,
|
| 622 |
+
derivative_order=1,
|
| 623 |
+
derivative_threshold=1e-3,
|
| 624 |
+
threshold_subgoal_passing=None,
|
| 625 |
+
),
|
| 626 |
+
embed_no_robot_extended=dict(
|
| 627 |
+
window_length=3,
|
| 628 |
+
derivative_order=1,
|
| 629 |
+
derivative_threshold=1e-3,
|
| 630 |
+
threshold_subgoal_passing=None,
|
| 631 |
+
),
|
| 632 |
+
oracle=dict(),
|
| 633 |
+
random=dict(num_milestones=(3, 6)),
|
| 634 |
+
equally=dict(num_milestones=(3, 6)),
|
| 635 |
+
near_future=dict(advance_steps=5),
|
| 636 |
+
)
|
uvd/decomp/kernel_reg.py
ADDED
|
@@ -0,0 +1,91 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""The :mod:`sklearn.kernel_regressor` module implements the Kernel
|
| 2 |
+
Regressor."""
|
| 3 |
+
# Author: Jan Hendrik Metzen <janmetzen@mailbox.de>
|
| 4 |
+
#
|
| 5 |
+
# License: BSD 3 clause
|
| 6 |
+
|
| 7 |
+
import numpy as np
|
| 8 |
+
from sklearn.base import BaseEstimator, RegressorMixin
|
| 9 |
+
from sklearn.metrics.pairwise import pairwise_kernels
|
| 10 |
+
|
| 11 |
+
|
| 12 |
+
class KernelRegression(BaseEstimator, RegressorMixin):
|
| 13 |
+
"""Nadaraya-Watson kernel regression with automatic bandwidth selection.
|
| 14 |
+
This implements Nadaraya-Watson kernel regression with (optional) automatic
|
| 15 |
+
bandwith selection of the kernel via leave-one-out cross-validation. Kernel
|
| 16 |
+
regression is a simple non-parametric kernelized technique for learning a
|
| 17 |
+
non-linear relationship between input variable(s) and a target variable.
|
| 18 |
+
|
| 19 |
+
Parameters
|
| 20 |
+
----------
|
| 21 |
+
kernel : string or callable, default="rbf"
|
| 22 |
+
Kernel map to be approximated. A callable should accept two arguments
|
| 23 |
+
and the keyword arguments passed to this object as kernel_params, and
|
| 24 |
+
should return a floating point number.
|
| 25 |
+
gamma : float, default=None
|
| 26 |
+
Gamma parameter for the RBF ("bandwidth"), polynomial,
|
| 27 |
+
exponential chi2 and sigmoid kernels. Interpretation of the default
|
| 28 |
+
value is left to the kernel; see the documentation for
|
| 29 |
+
sklearn.metrics.pairwise. Ignored by other kernels. If a sequence of
|
| 30 |
+
values is given, one of these values is selected which minimizes
|
| 31 |
+
the mean-squared-error of leave-one-out cross-validation.
|
| 32 |
+
See also
|
| 33 |
+
--------
|
| 34 |
+
sklearn.metrics.pairwise.kernel_metrics : List of built-in kernels.
|
| 35 |
+
"""
|
| 36 |
+
|
| 37 |
+
def __init__(self, kernel="rbf", gamma=None):
|
| 38 |
+
self.kernel = kernel
|
| 39 |
+
self.gamma = gamma
|
| 40 |
+
|
| 41 |
+
def fit(self, X, y):
|
| 42 |
+
"""Fit the model.
|
| 43 |
+
|
| 44 |
+
Parameters
|
| 45 |
+
----------
|
| 46 |
+
X : array-like of shape = [n_samples, n_features]
|
| 47 |
+
The training input samples.
|
| 48 |
+
y : array-like, shape = [n_samples]
|
| 49 |
+
The target values
|
| 50 |
+
Returns
|
| 51 |
+
-------
|
| 52 |
+
self : object
|
| 53 |
+
Returns self.
|
| 54 |
+
"""
|
| 55 |
+
self.X = X
|
| 56 |
+
self.y = y
|
| 57 |
+
|
| 58 |
+
if hasattr(self.gamma, "__iter__"):
|
| 59 |
+
self.gamma = self._optimize_gamma(self.gamma)
|
| 60 |
+
|
| 61 |
+
return self
|
| 62 |
+
|
| 63 |
+
def predict(self, X):
|
| 64 |
+
"""Predict target values for X.
|
| 65 |
+
|
| 66 |
+
Parameters
|
| 67 |
+
----------
|
| 68 |
+
X : array-like of shape = [n_samples, n_features]
|
| 69 |
+
The input samples.
|
| 70 |
+
Returns
|
| 71 |
+
-------
|
| 72 |
+
y : array of shape = [n_samples]
|
| 73 |
+
The predicted target value.
|
| 74 |
+
"""
|
| 75 |
+
K = pairwise_kernels(self.X, X, metric=self.kernel, gamma=self.gamma)
|
| 76 |
+
return (K * self.y[:, None]).sum(axis=0) / K.sum(axis=0)
|
| 77 |
+
|
| 78 |
+
def _optimize_gamma(self, gamma_values):
|
| 79 |
+
# Select specific value of gamma from the range of given gamma_values
|
| 80 |
+
# by minimizing mean-squared error in leave-one-out cross validation
|
| 81 |
+
mse = np.empty_like(gamma_values, dtype=np.float)
|
| 82 |
+
for i, gamma in enumerate(gamma_values):
|
| 83 |
+
K = pairwise_kernels(self.X, self.X, metric=self.kernel, gamma=gamma)
|
| 84 |
+
np.fill_diagonal(K, 0) # leave-one-out
|
| 85 |
+
Ky = K * self.y[:, np.newaxis]
|
| 86 |
+
y_pred = Ky.sum(axis=0) / K.sum(axis=0)
|
| 87 |
+
mse[i] = ((y_pred - self.y) ** 2).mean()
|
| 88 |
+
try:
|
| 89 |
+
return gamma_values[np.nanargmin(mse)]
|
| 90 |
+
except:
|
| 91 |
+
return 0
|
uvd/envs/__init__.py
ADDED
|
File without changes
|
uvd/envs/evaluator/__init__.py
ADDED
|
@@ -0,0 +1,3 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
from .evaluator import *
|
| 2 |
+
from .inference_wrapper import *
|
| 3 |
+
from .visualize_wrapper import *
|
uvd/envs/evaluator/evaluator.py
ADDED
|
@@ -0,0 +1,571 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
from __future__ import annotations
|
| 2 |
+
|
| 3 |
+
import copy
|
| 4 |
+
|
| 5 |
+
import einops
|
| 6 |
+
import gym
|
| 7 |
+
import numpy as np
|
| 8 |
+
import torch
|
| 9 |
+
import torch.nn.functional as F
|
| 10 |
+
import tqdm
|
| 11 |
+
import uvd.utils as U
|
| 12 |
+
import wandb
|
| 13 |
+
from uvd.envs.franka_kitchen import KitchenBase
|
| 14 |
+
from uvd.models.policy import PolicyBase
|
| 15 |
+
|
| 16 |
+
from .inference_wrapper import InferenceWrapper
|
| 17 |
+
from .vec_envs.vec_env import BaseVectorEnv, ShmemVectorEnv, SubprocVectorEnv
|
| 18 |
+
from .visualize_wrapper import VisualizeWrapper
|
| 19 |
+
from ...models import Preprocessor
|
| 20 |
+
|
| 21 |
+
__all__ = ["VectorEnvEvaluator"]
|
| 22 |
+
|
| 23 |
+
|
| 24 |
+
class VectorEnvEvaluator:
|
| 25 |
+
def __init__(
|
| 26 |
+
self,
|
| 27 |
+
*,
|
| 28 |
+
env_name: str,
|
| 29 |
+
reset_kwargs: dict,
|
| 30 |
+
num_tasks: int,
|
| 31 |
+
num_envs: int | None,
|
| 32 |
+
max_horizon: int = 400,
|
| 33 |
+
no_robot_rgb_milestones: np.ndarray | None = None,
|
| 34 |
+
use_no_robot_milestones: bool = False,
|
| 35 |
+
inference_kwargs: dict | None,
|
| 36 |
+
random_skip_inference: bool = False,
|
| 37 |
+
use_milestone_distances_normalization: bool = False,
|
| 38 |
+
decomp_method: str | None = None,
|
| 39 |
+
milestones: np.ndarray,
|
| 40 |
+
rgb_milestones: np.ndarray | None = None,
|
| 41 |
+
seed: int | None = None,
|
| 42 |
+
seed_list: list[int] | None = None,
|
| 43 |
+
save_video: bool = False,
|
| 44 |
+
save_video_kwargs: dict | None = None,
|
| 45 |
+
device: int | str | torch.device | None = None,
|
| 46 |
+
use_milestone_compressor: bool = False,
|
| 47 |
+
milestone_indices: np.ndarray | None = None,
|
| 48 |
+
):
|
| 49 |
+
self.env_name = env_name
|
| 50 |
+
self.reset_kwargs = reset_kwargs
|
| 51 |
+
num_envs = min(num_envs or num_tasks, num_tasks, milestones.shape[0])
|
| 52 |
+
self.num_envs = num_envs
|
| 53 |
+
self.rollout_list = self.generate_rollout_list(num_envs, num_tasks)
|
| 54 |
+
|
| 55 |
+
inference_kwargs = inference_kwargs or {}
|
| 56 |
+
if not random_skip_inference:
|
| 57 |
+
inference_kwargs.update(random_skip_ratio=0.0)
|
| 58 |
+
# inference_kwargs["hybrid"] = decomp_method == "embed_no_robot_extended"
|
| 59 |
+
inference_kwargs["hybrid"] = False
|
| 60 |
+
if decomp_method == "embed_no_robot_extended":
|
| 61 |
+
use_no_robot_milestones = False
|
| 62 |
+
inference_kwargs["use_milestone_compressor"] = use_milestone_compressor
|
| 63 |
+
self.inference_kwargs = inference_kwargs
|
| 64 |
+
self.decomp_method = decomp_method
|
| 65 |
+
self.max_horizon = max_horizon
|
| 66 |
+
self.milestones = milestones
|
| 67 |
+
self.rgb_milestones = rgb_milestones
|
| 68 |
+
self.no_robot_rgb_milestones = no_robot_rgb_milestones
|
| 69 |
+
if not use_no_robot_milestones:
|
| 70 |
+
self.no_robot_rgb_milestones = None
|
| 71 |
+
elif self.no_robot_rgb_milestones is None:
|
| 72 |
+
U.rank_zero_print(
|
| 73 |
+
f"WARNING: set to use no robot milestones {use_no_robot_milestones} but there is no input",
|
| 74 |
+
color="red",
|
| 75 |
+
)
|
| 76 |
+
milestone_distances = None
|
| 77 |
+
if use_milestone_distances_normalization and milestones.shape[1] > 1:
|
| 78 |
+
if self.no_robot_rgb_milestones is not None:
|
| 79 |
+
# assume only using if with linear embedding
|
| 80 |
+
assert (
|
| 81 |
+
self.no_robot_rgb_milestones.ndim == 3
|
| 82 |
+
), self.no_robot_rgb_milestones.shape
|
| 83 |
+
# N_ENVS, N_GOALS, D
|
| 84 |
+
milestone_distances = np.linalg.norm(
|
| 85 |
+
self.no_robot_rgb_milestones[:, :-1]
|
| 86 |
+
- self.no_robot_rgb_milestones[:, 1:],
|
| 87 |
+
axis=-1,
|
| 88 |
+
)
|
| 89 |
+
else:
|
| 90 |
+
assert self.milestones.ndim == 3, self.milestones.shape
|
| 91 |
+
milestone_distances = np.linalg.norm(
|
| 92 |
+
self.milestones[:, :-1] - self.milestones[:, 1:], axis=-1
|
| 93 |
+
)
|
| 94 |
+
assert milestone_distances.shape == (
|
| 95 |
+
milestones.shape[0],
|
| 96 |
+
milestones.shape[1] - 1,
|
| 97 |
+
), milestone_distances.shape
|
| 98 |
+
self.milestone_distances = milestone_distances
|
| 99 |
+
|
| 100 |
+
if seed_list is None:
|
| 101 |
+
_rng = np.random.default_rng(seed=seed)
|
| 102 |
+
seed_list = [_rng.integers(0, 2**31 - 1) for _ in range(num_tasks)]
|
| 103 |
+
assert len(seed_list) == num_tasks
|
| 104 |
+
self._num_tasks = num_tasks
|
| 105 |
+
self.seed_list = seed_list
|
| 106 |
+
self.save_video = save_video
|
| 107 |
+
if save_video:
|
| 108 |
+
assert save_video_kwargs is not None
|
| 109 |
+
self.video_fps = save_video_kwargs["fps"]
|
| 110 |
+
self.add_debug_text = save_video_kwargs["add_debug_text"]
|
| 111 |
+
self.wandb_logging = save_video_kwargs["wandb_logging"]
|
| 112 |
+
self.save_locally = save_video_kwargs["save_locally"]
|
| 113 |
+
if self.save_locally:
|
| 114 |
+
self.save_path = save_video_kwargs["save_path"]
|
| 115 |
+
self.device = device
|
| 116 |
+
|
| 117 |
+
self.batched_env: BaseVectorEnv | None = None
|
| 118 |
+
self.use_milestone_compressor = use_milestone_compressor
|
| 119 |
+
self.milestone_indices = milestone_indices # B, N
|
| 120 |
+
|
| 121 |
+
@staticmethod
|
| 122 |
+
def confirm_embedding(
|
| 123 |
+
embedding: np.ndarray | torch.Tensor, preprocessor: Preprocessor | None
|
| 124 |
+
) -> np.ndarray:
|
| 125 |
+
"""make sure the embedding is linear for calculating embed distance during inference
|
| 126 |
+
Args:
|
| 127 |
+
embedding:
|
| 128 |
+
if embedding.ndim == 2: B, d already;
|
| 129 |
+
if embedding.ndim == 3: B, L, d for milestone embeddings;
|
| 130 |
+
if embedding.ndim == 4: B, H, W, 3 for raw rgb;
|
| 131 |
+
or without pooling, e.g. B, 2048, 7, 7
|
| 132 |
+
if embedding.ndim == 5: B, L, H, W, 3 for L milestones;
|
| 133 |
+
preprocessor: frozen preprocessor
|
| 134 |
+
Return:
|
| 135 |
+
embedding.ndim = 2 for single step obs, ndim=3 for milestone embeddings
|
| 136 |
+
"""
|
| 137 |
+
if preprocessor is None:
|
| 138 |
+
return U.any_to_numpy(embedding)
|
| 139 |
+
preprocessor_output_dim = preprocessor.output_dim
|
| 140 |
+
if embedding.ndim not in [2, 3]:
|
| 141 |
+
input_dims = embedding.shape
|
| 142 |
+
if (
|
| 143 |
+
embedding.ndim in [4, 5]
|
| 144 |
+
and embedding.shape[-3:] != preprocessor_output_dim
|
| 145 |
+
):
|
| 146 |
+
# raw rgb case (B, (L) H, W, 3)
|
| 147 |
+
if embedding.ndim == 5:
|
| 148 |
+
embedding = einops.rearrange(embedding, "b l h w c -> (b l) h w c")
|
| 149 |
+
with torch.no_grad():
|
| 150 |
+
embedding = preprocessor.process(
|
| 151 |
+
embedding,
|
| 152 |
+
# embedding.reshape((np.prod(input_dims[:-3]), *input_dims[-3:])),
|
| 153 |
+
reconstruct_linear=True,
|
| 154 |
+
)
|
| 155 |
+
# if preprocessor.remove_pool:
|
| 156 |
+
if (
|
| 157 |
+
embedding.ndim in [4, 5]
|
| 158 |
+
and embedding.shape[-3:] == preprocessor_output_dim
|
| 159 |
+
):
|
| 160 |
+
U.rank_zero_print("DEPRECATED!", color="red")
|
| 161 |
+
# no pooling case, e.g. (B, (L), 2048, 7, 7)
|
| 162 |
+
embedding = embedding.reshape(
|
| 163 |
+
(np.prod(input_dims[:-3]), *preprocessor.output_dim)
|
| 164 |
+
)
|
| 165 |
+
with torch.no_grad():
|
| 166 |
+
embedding = F.adaptive_avg_pool2d(embedding, output_size=(1, 1))
|
| 167 |
+
embedding = torch.flatten(embedding, 1)
|
| 168 |
+
if preprocessor.preprocessor_fc is not None:
|
| 169 |
+
embedding = preprocessor.preprocessor_fc(embedding)
|
| 170 |
+
embedding = embedding.reshape((*input_dims[:-3], embedding.shape[-1]))
|
| 171 |
+
if torch.is_tensor(embedding):
|
| 172 |
+
embedding = embedding.cpu().numpy()
|
| 173 |
+
assert embedding.ndim in [2, 3], embedding.shape
|
| 174 |
+
return embedding
|
| 175 |
+
|
| 176 |
+
def rollout(self, policy: PolicyBase, mode: str, epoch: int) -> dict:
|
| 177 |
+
use_kv_cache = hasattr(policy, "use_kv_cache") and policy.use_kv_cache
|
| 178 |
+
is_causal = hasattr(policy, "causal") and policy.causal
|
| 179 |
+
cache_obs = is_causal and not use_kv_cache
|
| 180 |
+
if cache_obs: # means causal and not use kv cache
|
| 181 |
+
U.rank_zero_print(
|
| 182 |
+
"WARNING: `use_kv_cache=False` NOT RECOMMEND SINCE SOOOOOO INEFFICIENT!",
|
| 183 |
+
color="red",
|
| 184 |
+
)
|
| 185 |
+
self.inference_kwargs["cache_history"] = True
|
| 186 |
+
elif is_causal and use_kv_cache:
|
| 187 |
+
# keep the bs the same so not terminate the episode
|
| 188 |
+
self.inference_kwargs["dummy_rtn"] = True
|
| 189 |
+
|
| 190 |
+
if self.batched_env is None:
|
| 191 |
+
if cache_obs: # dynamic leading dim obs for history
|
| 192 |
+
self.batched_env = SubprocVectorEnv(self._create_env_fns())
|
| 193 |
+
else:
|
| 194 |
+
self.batched_env = ShmemVectorEnv(self._create_env_fns())
|
| 195 |
+
self.batched_env.reset()
|
| 196 |
+
|
| 197 |
+
results = {k: None for k in range(self._num_tasks)}
|
| 198 |
+
|
| 199 |
+
seed_list = list(self.seed_list)
|
| 200 |
+
logging_videos = self.save_video
|
| 201 |
+
if logging_videos:
|
| 202 |
+
self.batched_env.set_env_attr("recording", True)
|
| 203 |
+
for i, env_ids in enumerate(self.rollout_list):
|
| 204 |
+
if logging_videos and i != 0 and not self.save_locally:
|
| 205 |
+
# logging wandb only 1st iter
|
| 206 |
+
self.batched_env.set_env_attr("recording", False)
|
| 207 |
+
logging_videos = False
|
| 208 |
+
|
| 209 |
+
if is_causal and use_kv_cache:
|
| 210 |
+
policy.reset_cache()
|
| 211 |
+
|
| 212 |
+
num_env_this_batch = len(env_ids)
|
| 213 |
+
self.batched_env.seed(
|
| 214 |
+
seed=[seed_list.pop() for _ in range(num_env_this_batch)]
|
| 215 |
+
)
|
| 216 |
+
self.prepare_states_before_reset(env_ids, global_idx=i)
|
| 217 |
+
if self.milestone_distances is not None:
|
| 218 |
+
milestone_distances_this_batch = self.milestone_distances[
|
| 219 |
+
self.num_envs * i : self.num_envs * i + num_env_this_batch, ...
|
| 220 |
+
].copy()
|
| 221 |
+
self.batched_env.set_env_attr(
|
| 222 |
+
"milestone_distances",
|
| 223 |
+
milestone_distances_this_batch,
|
| 224 |
+
id=env_ids,
|
| 225 |
+
diff_value=True,
|
| 226 |
+
)
|
| 227 |
+
|
| 228 |
+
if self.no_robot_rgb_milestones is not None:
|
| 229 |
+
no_robot_milestones_this_batch = self.no_robot_rgb_milestones[
|
| 230 |
+
self.num_envs * i : self.num_envs * i + num_env_this_batch, ...
|
| 231 |
+
].copy()
|
| 232 |
+
# num_env, num_milestone, (H, W, 3 or D)
|
| 233 |
+
no_robot_milestones_this_batch = self.confirm_embedding(
|
| 234 |
+
no_robot_milestones_this_batch, preprocessor=policy.preprocessor
|
| 235 |
+
)
|
| 236 |
+
|
| 237 |
+
self.batched_env.set_env_attr(
|
| 238 |
+
"no_robot_milestones",
|
| 239 |
+
no_robot_milestones_this_batch,
|
| 240 |
+
id=env_ids,
|
| 241 |
+
diff_value=True,
|
| 242 |
+
)
|
| 243 |
+
# num_env_this_batch, num_goals, ...
|
| 244 |
+
milestones_this_batch = self.milestones[
|
| 245 |
+
self.num_envs * i : self.num_envs * i + num_env_this_batch, ...
|
| 246 |
+
].copy()
|
| 247 |
+
# set milestones for this episode
|
| 248 |
+
self.batched_env.set_env_attr(
|
| 249 |
+
"milestones", milestones_this_batch, id=env_ids, diff_value=True
|
| 250 |
+
)
|
| 251 |
+
|
| 252 |
+
if self.milestone_indices is not None:
|
| 253 |
+
milestone_indices_this_batch = self.milestone_indices[
|
| 254 |
+
self.num_envs * i : self.num_envs * i + num_env_this_batch, ...
|
| 255 |
+
].copy()
|
| 256 |
+
self.batched_env.set_env_attr(
|
| 257 |
+
"milestone_indices",
|
| 258 |
+
milestone_indices_this_batch,
|
| 259 |
+
id=env_ids,
|
| 260 |
+
diff_value=True,
|
| 261 |
+
)
|
| 262 |
+
|
| 263 |
+
if milestones_this_batch.ndim != 3:
|
| 264 |
+
# num_env_this_batch, num_goals, 3, H, W
|
| 265 |
+
if i == 0:
|
| 266 |
+
U.rank_zero_print(f"{milestones_this_batch.shape=}", color="blue")
|
| 267 |
+
self.batched_env.set_env_attr(
|
| 268 |
+
"milestone_embeddings",
|
| 269 |
+
self.confirm_embedding(
|
| 270 |
+
milestones_this_batch.copy(), preprocessor=policy.preprocessor
|
| 271 |
+
),
|
| 272 |
+
id=env_ids,
|
| 273 |
+
diff_value=True,
|
| 274 |
+
)
|
| 275 |
+
if logging_videos:
|
| 276 |
+
rgb_milestones_this_batch = self.rgb_milestones[
|
| 277 |
+
num_env_this_batch * i : num_env_this_batch * (i + 1), ...
|
| 278 |
+
].copy()
|
| 279 |
+
self.batched_env.set_env_attr(
|
| 280 |
+
"rgb_milestones",
|
| 281 |
+
rgb_milestones_this_batch,
|
| 282 |
+
id=env_ids,
|
| 283 |
+
diff_value=True,
|
| 284 |
+
)
|
| 285 |
+
# reset: num_env_this_batch, h, w, 3 (or dict of obs)
|
| 286 |
+
obs = self.batched_env.reset(id=env_ids)
|
| 287 |
+
assert obs.shape[0] == num_env_this_batch, obs.shape
|
| 288 |
+
|
| 289 |
+
running_env_ids = np.arange(num_env_this_batch)
|
| 290 |
+
|
| 291 |
+
for st in tqdm.tqdm(
|
| 292 |
+
range(self.max_horizon),
|
| 293 |
+
initial=1,
|
| 294 |
+
desc=f"rank {U.get_local_rank()}: Rollout {i * self.num_envs + len(env_ids)}/{self._num_tasks}",
|
| 295 |
+
leave=False,
|
| 296 |
+
):
|
| 297 |
+
self.batched_env.set_env_attr(
|
| 298 |
+
"cur_milestone_idx", value=None, id=running_env_ids
|
| 299 |
+
) # for random skipping, sample maybe next goal idx inside env every step
|
| 300 |
+
|
| 301 |
+
if cache_obs:
|
| 302 |
+
current_milestone = self.batched_env.get_env_attr(
|
| 303 |
+
"cached_prev_milestones", id=running_env_ids
|
| 304 |
+
)
|
| 305 |
+
else:
|
| 306 |
+
current_milestone = self.batched_env.get_env_attr(
|
| 307 |
+
"current_milestone", id=running_env_ids
|
| 308 |
+
)
|
| 309 |
+
current_milestone = U.any_to_numpy(current_milestone)
|
| 310 |
+
U.assert_(len(current_milestone), len(obs))
|
| 311 |
+
|
| 312 |
+
with torch.no_grad():
|
| 313 |
+
# only batchify here, keep obs as list as length of running_env_ids anywhere else
|
| 314 |
+
batchify_obs = U.batch_observations(obs, device=self.device)
|
| 315 |
+
if cache_obs:
|
| 316 |
+
b, t, *_ = current_milestone.shape
|
| 317 |
+
if st == 0 and batchify_obs["rgb"].shape[:2] != (b, t):
|
| 318 |
+
for k in batchify_obs:
|
| 319 |
+
batchify_obs[k] = batchify_obs[k][
|
| 320 |
+
:, None, ...
|
| 321 |
+
] # broadcast T dim
|
| 322 |
+
assert batchify_obs["rgb"].shape[:2] == (b, t), (
|
| 323 |
+
batchify_obs["rgb"].shape,
|
| 324 |
+
b,
|
| 325 |
+
t,
|
| 326 |
+
)
|
| 327 |
+
elif is_causal:
|
| 328 |
+
# broadcast T dim, B H W 3 or B D
|
| 329 |
+
assert current_milestone.ndim in [2, 4], current_milestone.shape
|
| 330 |
+
current_milestone = current_milestone[:, None, ...]
|
| 331 |
+
for k in batchify_obs:
|
| 332 |
+
batchify_obs[k] = batchify_obs[k][:, None, ...]
|
| 333 |
+
# L, n
|
| 334 |
+
action, obs_embed, goal_embed = policy(
|
| 335 |
+
batchify_obs,
|
| 336 |
+
goal=current_milestone,
|
| 337 |
+
deterministic=True,
|
| 338 |
+
return_embeddings=True,
|
| 339 |
+
timesteps=torch.tensor(
|
| 340 |
+
[[st]], dtype=torch.int32, device=self.device
|
| 341 |
+
),
|
| 342 |
+
input_pos=torch.tensor([st], device=self.device)
|
| 343 |
+
if is_causal and use_kv_cache
|
| 344 |
+
else None,
|
| 345 |
+
)
|
| 346 |
+
if action.ndim == 3: # has T dim
|
| 347 |
+
action = action[:, -1, :]
|
| 348 |
+
# switch milestones inside the inference wrapper based on current obs embedding
|
| 349 |
+
if self.no_robot_rgb_milestones is not None:
|
| 350 |
+
no_robot_rgbs = self.batched_env.get_env_attr(
|
| 351 |
+
"current_no_robot_frame", id=running_env_ids
|
| 352 |
+
)
|
| 353 |
+
cur_obs_embed = self.confirm_embedding(
|
| 354 |
+
np.stack(no_robot_rgbs), preprocessor=policy.preprocessor
|
| 355 |
+
)
|
| 356 |
+
else:
|
| 357 |
+
cur_obs_embed = self.confirm_embedding(
|
| 358 |
+
batchify_obs["rgb"] if obs_embed.ndim != 2 else obs_embed,
|
| 359 |
+
preprocessor=policy.preprocessor,
|
| 360 |
+
)
|
| 361 |
+
self.batched_env.set_env_attr(
|
| 362 |
+
"current_obs_embedding",
|
| 363 |
+
cur_obs_embed,
|
| 364 |
+
id=running_env_ids,
|
| 365 |
+
diff_value=True,
|
| 366 |
+
)
|
| 367 |
+
|
| 368 |
+
# next_obs (n, dict/(256, 256, 3)) (n,) (n,),list[dict]
|
| 369 |
+
obs, r, done, info = self.batched_env.step(
|
| 370 |
+
action.cpu().numpy(), id=running_env_ids
|
| 371 |
+
)
|
| 372 |
+
prev_running_env_ids = copy.deepcopy(running_env_ids)
|
| 373 |
+
if np.any(done):
|
| 374 |
+
if is_causal and use_kv_cache:
|
| 375 |
+
assert np.all(done), done # keep bs the same for kv cache
|
| 376 |
+
# update running env id & collect results for envs that done
|
| 377 |
+
terminated_env_local_idx = np.where(done)[0]
|
| 378 |
+
masks = np.ones_like(running_env_ids, dtype=bool)
|
| 379 |
+
masks[terminated_env_local_idx] = False
|
| 380 |
+
running_env_ids = running_env_ids[masks]
|
| 381 |
+
obs = obs[masks]
|
| 382 |
+
terminated_local_ids = list(
|
| 383 |
+
set(prev_running_env_ids) - set(running_env_ids)
|
| 384 |
+
)
|
| 385 |
+
metrics = self.batched_env.get_env_attr(
|
| 386 |
+
"metrics", id=terminated_local_ids
|
| 387 |
+
)
|
| 388 |
+
try:
|
| 389 |
+
task_names = self.batched_env.get_env_attr(
|
| 390 |
+
"task_name", id=terminated_local_ids
|
| 391 |
+
)
|
| 392 |
+
except AttributeError:
|
| 393 |
+
task_names = None
|
| 394 |
+
for met_idx, local_id in enumerate(terminated_local_ids):
|
| 395 |
+
assert (
|
| 396 |
+
local_id not in running_env_ids
|
| 397 |
+
), f"{local_id=} is done but still in {running_env_ids=}"
|
| 398 |
+
assert results[local_id + self.num_envs * i] is None, (
|
| 399 |
+
f"{local_id + self.num_envs * i}-th task should be only rollout once, "
|
| 400 |
+
f"{local_id=}, {running_env_ids=}, {prev_running_env_ids=}, {terminated_local_ids=}"
|
| 401 |
+
)
|
| 402 |
+
results[local_id + self.num_envs * i] = {
|
| 403 |
+
(
|
| 404 |
+
k
|
| 405 |
+
if ("success" not in k and "completed_tasks" not in k)
|
| 406 |
+
or task_names is None
|
| 407 |
+
else f"{k}/{task_names[met_idx]}"
|
| 408 |
+
): v
|
| 409 |
+
for k, v in metrics[met_idx].items()
|
| 410 |
+
}
|
| 411 |
+
|
| 412 |
+
if np.all(done):
|
| 413 |
+
assert results is not None
|
| 414 |
+
break
|
| 415 |
+
|
| 416 |
+
# only save first num vec envs videos
|
| 417 |
+
if logging_videos: # gather videos after all env done
|
| 418 |
+
# each video may have diff length among diff envs
|
| 419 |
+
all_frames = self.batched_env.get_env_attr("frames", id=env_ids)
|
| 420 |
+
try:
|
| 421 |
+
task_names = self.batched_env.get_env_attr("task_name", id=env_ids)
|
| 422 |
+
except AttributeError:
|
| 423 |
+
task_names = None
|
| 424 |
+
try:
|
| 425 |
+
self.logging_videos(
|
| 426 |
+
all_frames,
|
| 427 |
+
task_names=task_names,
|
| 428 |
+
global_task_ids=[
|
| 429 |
+
self.num_envs * i + _ for _ in range(num_env_this_batch)
|
| 430 |
+
],
|
| 431 |
+
mode=mode,
|
| 432 |
+
global_step=epoch,
|
| 433 |
+
)
|
| 434 |
+
except:
|
| 435 |
+
pass
|
| 436 |
+
|
| 437 |
+
return self.process_results(results, mode=mode)
|
| 438 |
+
|
| 439 |
+
def prepare_states_before_reset(self, local_env_ids: list[int], global_idx: int):
|
| 440 |
+
if self.env_name == "franka_kitchen":
|
| 441 |
+
reset_states = [
|
| 442 |
+
dict(
|
| 443 |
+
init_qpos=self.reset_kwargs["init_qpos"][
|
| 444 |
+
env_id + self.num_envs * global_idx
|
| 445 |
+
],
|
| 446 |
+
init_qvel=self.reset_kwargs["init_qvel"][
|
| 447 |
+
env_id + self.num_envs * global_idx
|
| 448 |
+
],
|
| 449 |
+
task_elements=self.reset_kwargs["task_elements"][
|
| 450 |
+
env_id + self.num_envs * global_idx
|
| 451 |
+
],
|
| 452 |
+
)
|
| 453 |
+
for env_id in local_env_ids
|
| 454 |
+
]
|
| 455 |
+
self.batched_env.set_env_attr(
|
| 456 |
+
"reset_states", reset_states, id=local_env_ids, diff_value=True
|
| 457 |
+
)
|
| 458 |
+
else:
|
| 459 |
+
raise NotImplementedError(self.env_name)
|
| 460 |
+
|
| 461 |
+
def process_results(self, results: dict, mode: str) -> dict:
|
| 462 |
+
res = {}
|
| 463 |
+
successes = []
|
| 464 |
+
completions = []
|
| 465 |
+
for result in results.values():
|
| 466 |
+
for k, v in result.items():
|
| 467 |
+
if "success" in k:
|
| 468 |
+
successes.append(v)
|
| 469 |
+
elif "completed_tasks" in k:
|
| 470 |
+
completions.append(v)
|
| 471 |
+
res.setdefault(k, []).append(v)
|
| 472 |
+
assert (
|
| 473 |
+
len(successes) == self._num_tasks
|
| 474 |
+
), f"{len(successes)=} != {self._num_tasks}"
|
| 475 |
+
return {
|
| 476 |
+
f"{mode}/success": np.mean(successes),
|
| 477 |
+
f"{mode}/completed_tasks": np.mean(completions),
|
| 478 |
+
f"{mode}/num_tasks": float(self._num_tasks),
|
| 479 |
+
**{f"{mode}/{k}": np.mean(v) for k, v in res.items()},
|
| 480 |
+
}
|
| 481 |
+
|
| 482 |
+
def logging_videos(
|
| 483 |
+
self,
|
| 484 |
+
all_frames: list,
|
| 485 |
+
*,
|
| 486 |
+
task_names: str | None = None,
|
| 487 |
+
global_task_ids: list,
|
| 488 |
+
mode: str,
|
| 489 |
+
global_step: int,
|
| 490 |
+
):
|
| 491 |
+
cur_rank = U.get_local_rank()
|
| 492 |
+
for i, ids in enumerate(global_task_ids):
|
| 493 |
+
# wandb channel-first logging: episode_length, h, w, 3 -> ..., 3, h, w
|
| 494 |
+
cur_video = U.any_to_numpy(all_frames[i])
|
| 495 |
+
if cur_video.ndim != 4 or cur_video.shape[-1] != 3:
|
| 496 |
+
U.rank_zero_print(
|
| 497 |
+
f"cur_video has unexpected shape {cur_video.shape}", color="red"
|
| 498 |
+
)
|
| 499 |
+
continue
|
| 500 |
+
cur_video = cur_video.transpose([0, 3, 1, 2])
|
| 501 |
+
assert cur_video.ndim == 4 and cur_video.shape[1] == 3, cur_video.shape
|
| 502 |
+
if self.wandb_logging and cur_rank == 0:
|
| 503 |
+
log_name = U.f_join(
|
| 504 |
+
f"{mode}", "rollout-videos", f"{task_names[i]}" or "", f"{ids}"
|
| 505 |
+
)
|
| 506 |
+
wandb.log(
|
| 507 |
+
{
|
| 508 |
+
log_name: wandb.Video(
|
| 509 |
+
cur_video, fps=self.video_fps, format="mp4"
|
| 510 |
+
)
|
| 511 |
+
},
|
| 512 |
+
)
|
| 513 |
+
if self.save_locally:
|
| 514 |
+
save_path = U.f_mkdir(
|
| 515 |
+
self.save_path,
|
| 516 |
+
f"{mode}-rank_{cur_rank}"
|
| 517 |
+
+ f"/{task_names[i] if task_names is not None else ''}",
|
| 518 |
+
)
|
| 519 |
+
U.save_video(
|
| 520 |
+
cur_video,
|
| 521 |
+
U.f_join(save_path, f"env_{ids}_epoch_{global_step}.mp4"),
|
| 522 |
+
fps=self.video_fps,
|
| 523 |
+
)
|
| 524 |
+
|
| 525 |
+
def _create_env_fns(self) -> list:
|
| 526 |
+
def _create_env(env_kwargs: dict) -> gym.Env:
|
| 527 |
+
if self.env_name == "franka_kitchen":
|
| 528 |
+
# don't reset before creating vector env
|
| 529 |
+
env = KitchenBase(frame_height=224, frame_width=224, **env_kwargs)
|
| 530 |
+
if self.decomp_method is not None:
|
| 531 |
+
env.COMPLETE_IN_ANY_ORDER = True
|
| 532 |
+
env.TERMINATE_ON_WRONG_COMPLETE = False
|
| 533 |
+
else:
|
| 534 |
+
env.COMPLETE_IN_ANY_ORDER = False
|
| 535 |
+
env.TERMINATE_ON_WRONG_COMPLETE = True
|
| 536 |
+
else:
|
| 537 |
+
raise NotImplementedError(self.env_name)
|
| 538 |
+
env = InferenceWrapper(env, **self.inference_kwargs)
|
| 539 |
+
if self.save_video:
|
| 540 |
+
env = VisualizeWrapper(
|
| 541 |
+
env, add_goal=True, add_debug_text=self.add_debug_text
|
| 542 |
+
)
|
| 543 |
+
env.recording = True
|
| 544 |
+
return env
|
| 545 |
+
|
| 546 |
+
env_kwargs = dict(
|
| 547 |
+
max_horizon=self.max_horizon,
|
| 548 |
+
gpu_id=int(torch.device(self.device).index),
|
| 549 |
+
)
|
| 550 |
+
|
| 551 |
+
return [lambda: _create_env(env_kwargs) for _ in range(self.num_envs)]
|
| 552 |
+
|
| 553 |
+
def close(self):
|
| 554 |
+
self.batched_env.close()
|
| 555 |
+
|
| 556 |
+
@staticmethod
|
| 557 |
+
def generate_rollout_list(num_envs: int, num_tasks: int):
|
| 558 |
+
assert num_envs <= num_tasks
|
| 559 |
+
tasks = []
|
| 560 |
+
k = 0
|
| 561 |
+
for i in range(num_tasks):
|
| 562 |
+
sub_list = []
|
| 563 |
+
j = 0
|
| 564 |
+
while j < num_envs and k < num_tasks:
|
| 565 |
+
sub_list.append(j)
|
| 566 |
+
j += 1
|
| 567 |
+
k += 1
|
| 568 |
+
tasks.append(sub_list)
|
| 569 |
+
if k >= num_tasks:
|
| 570 |
+
break
|
| 571 |
+
return tasks
|
uvd/envs/evaluator/inference_wrapper.py
ADDED
|
@@ -0,0 +1,435 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
from __future__ import annotations
|
| 2 |
+
|
| 3 |
+
import random
|
| 4 |
+
|
| 5 |
+
import gym
|
| 6 |
+
import numpy as np
|
| 7 |
+
|
| 8 |
+
import uvd.utils as U
|
| 9 |
+
from uvd.envs.franka_kitchen import KitchenBase
|
| 10 |
+
|
| 11 |
+
__all__ = ["InferenceWrapper"]
|
| 12 |
+
|
| 13 |
+
|
| 14 |
+
class InferenceWrapper(gym.Wrapper):
|
| 15 |
+
"""Set all milestones => reset => set current_embedding_distance => step =>
|
| 16 |
+
set current_embedding_distance => ...
|
| 17 |
+
|
| 18 |
+
current subgoal milestone should not be set but automatically be
|
| 19 |
+
changed after setting current embedding distance
|
| 20 |
+
"""
|
| 21 |
+
|
| 22 |
+
def __init__(
|
| 23 |
+
self,
|
| 24 |
+
env: gym.Env,
|
| 25 |
+
thresh: float = 0.2,
|
| 26 |
+
hybrid: bool = False,
|
| 27 |
+
use_milestone_compressor: bool = False,
|
| 28 |
+
random_skip_ratio: float = 0.0,
|
| 29 |
+
delay_steps: int | None = 0,
|
| 30 |
+
cache_history: bool = False,
|
| 31 |
+
dummy_rtn: bool = False,
|
| 32 |
+
):
|
| 33 |
+
super().__init__(env=env)
|
| 34 |
+
self._thresh = thresh if thresh is not None else np.inf
|
| 35 |
+
self._random_skip_ratio = random_skip_ratio
|
| 36 |
+
self._cur_milestone_idx = None
|
| 37 |
+
self._delay_steps = delay_steps or 0
|
| 38 |
+
self._n_steps_after_pass_thresh = -1
|
| 39 |
+
|
| 40 |
+
self.num_milestones = None
|
| 41 |
+
self._milestones = None
|
| 42 |
+
self._current_milestone = None
|
| 43 |
+
self._current_obs_embedding = None
|
| 44 |
+
self._achieved_goals = 0
|
| 45 |
+
self._current_embedding_distance = None
|
| 46 |
+
self._unnormalized_embedding_distance = None
|
| 47 |
+
self.milestone_distances = None
|
| 48 |
+
|
| 49 |
+
self._no_robot_milestones = None
|
| 50 |
+
self._current_no_robot_milestone = None
|
| 51 |
+
|
| 52 |
+
self._milestone_embeddings = None
|
| 53 |
+
self._current_milestone_embedding = None
|
| 54 |
+
|
| 55 |
+
self._subgoal_init_distance = None
|
| 56 |
+
|
| 57 |
+
self.hybrid_phase = 0 if hybrid else None
|
| 58 |
+
self._cur_rgb_frame = None
|
| 59 |
+
# U.rank_zero_print(f"{use_milestone_compressor=}", color="blue")
|
| 60 |
+
self.use_milestone_compressor = use_milestone_compressor
|
| 61 |
+
self._n_steps = 0
|
| 62 |
+
self._n_steps_this_milestone = 0
|
| 63 |
+
self._milestone_indices = None
|
| 64 |
+
self.cache_history = cache_history
|
| 65 |
+
if cache_history:
|
| 66 |
+
self.cached_obs = []
|
| 67 |
+
self.cached_prev_milestones = []
|
| 68 |
+
else:
|
| 69 |
+
self.cached_obs = None
|
| 70 |
+
self.cached_prev_milestones = None
|
| 71 |
+
self._has_dummy_reset = True
|
| 72 |
+
|
| 73 |
+
if dummy_rtn:
|
| 74 |
+
self.dummy_rtn = []
|
| 75 |
+
else:
|
| 76 |
+
self.dummy_rtn = None
|
| 77 |
+
|
| 78 |
+
def dummy_step(self) -> tuple:
|
| 79 |
+
self._n_steps += 1
|
| 80 |
+
return self.dummy_rtn
|
| 81 |
+
|
| 82 |
+
def step(self, action: np.ndarray) -> tuple:
|
| 83 |
+
if (
|
| 84 |
+
self.dummy_rtn is not None
|
| 85 |
+
and len(self.dummy_rtn) > 0
|
| 86 |
+
and self._n_steps + 1 >= self.env.max_horizon
|
| 87 |
+
):
|
| 88 |
+
o, r, _, i = self.dummy_rtn
|
| 89 |
+
self.dummy_rtn = []
|
| 90 |
+
return o, r, True, i
|
| 91 |
+
elif self.dummy_rtn is not None and len(self.dummy_rtn) > 0:
|
| 92 |
+
return self.dummy_step()
|
| 93 |
+
|
| 94 |
+
o, r, d, i = super().step(action)
|
| 95 |
+
self._n_steps += 1
|
| 96 |
+
self._n_steps_this_milestone += 1
|
| 97 |
+
o = self.proc_obs(o)
|
| 98 |
+
self._cur_rgb_frame = o.get("rgb", None)
|
| 99 |
+
# switch delay triggered
|
| 100 |
+
if self._n_steps_after_pass_thresh >= 0:
|
| 101 |
+
self._n_steps_after_pass_thresh += 1
|
| 102 |
+
|
| 103 |
+
if d and self.dummy_rtn is not None:
|
| 104 |
+
self.dummy_rtn = (o, r, False, i)
|
| 105 |
+
d = self._n_steps >= self.env.max_horizon
|
| 106 |
+
|
| 107 |
+
if d:
|
| 108 |
+
# force set every episode
|
| 109 |
+
self.milestones = None
|
| 110 |
+
if self.dummy_rtn is not None:
|
| 111 |
+
self.dummy_rtn = []
|
| 112 |
+
return o, r, d, i
|
| 113 |
+
|
| 114 |
+
def reset(self, **kwargs):
|
| 115 |
+
obs = super().reset(**kwargs)
|
| 116 |
+
if self.cached_obs is not None:
|
| 117 |
+
self.cached_obs = []
|
| 118 |
+
self.cached_prev_milestones = []
|
| 119 |
+
obs = self.proc_obs(obs)
|
| 120 |
+
self._achieved_goals = 0
|
| 121 |
+
self._n_steps_after_pass_thresh = -1
|
| 122 |
+
self._subgoal_init_distance = None
|
| 123 |
+
self.hybrid_phase = 0 if self.hybrid_phase is not None else None
|
| 124 |
+
self._cur_rgb_frame = obs.get("rgb", None)
|
| 125 |
+
self._n_steps = 0
|
| 126 |
+
self._n_steps_this_milestone = 0
|
| 127 |
+
self._has_dummy_reset = False
|
| 128 |
+
return obs
|
| 129 |
+
|
| 130 |
+
def proc_obs(self, obs: dict):
|
| 131 |
+
if (
|
| 132 |
+
self.cached_obs is not None and not self._has_dummy_reset
|
| 133 |
+
): # in case dummy reset call
|
| 134 |
+
# this fn should be call only once for every step and reset
|
| 135 |
+
self.cached_obs.append(obs)
|
| 136 |
+
# steps, ...
|
| 137 |
+
obs = U.batch_observations(self.cached_obs, to_tensor=False)
|
| 138 |
+
self.cached_prev_milestones.append(self.current_milestone)
|
| 139 |
+
return obs
|
| 140 |
+
|
| 141 |
+
@property
|
| 142 |
+
def completed_goals(self):
|
| 143 |
+
return self._achieved_goals
|
| 144 |
+
|
| 145 |
+
@property
|
| 146 |
+
def milestones(self) -> np.ndarray:
|
| 147 |
+
assert (
|
| 148 |
+
self._milestones is not None
|
| 149 |
+
), f"must set before using, step: {self.env.episode_length}"
|
| 150 |
+
return self._milestones
|
| 151 |
+
|
| 152 |
+
@milestones.setter
|
| 153 |
+
def milestones(self, milestones: np.ndarray | None):
|
| 154 |
+
"""Set outside.
|
| 155 |
+
|
| 156 |
+
can be rgb or embedding
|
| 157 |
+
"""
|
| 158 |
+
if milestones is not None:
|
| 159 |
+
self._milestones = milestones
|
| 160 |
+
self.num_milestones = len(milestones)
|
| 161 |
+
# set initial sub goal
|
| 162 |
+
self._current_milestone = self._milestones[0]
|
| 163 |
+
else:
|
| 164 |
+
self._milestones = None
|
| 165 |
+
self._current_milestone = None
|
| 166 |
+
self._subgoal_init_distance = None
|
| 167 |
+
|
| 168 |
+
@property
|
| 169 |
+
def milestone_embeddings(self) -> np.ndarray:
|
| 170 |
+
if self.milestones.ndim == 2:
|
| 171 |
+
# preprocessed embedding already
|
| 172 |
+
return self.milestones
|
| 173 |
+
assert self._milestone_embeddings is not None
|
| 174 |
+
return self._milestone_embeddings
|
| 175 |
+
|
| 176 |
+
@milestone_embeddings.setter
|
| 177 |
+
def milestone_embeddings(self, embeddings: np.ndarray | None):
|
| 178 |
+
if embeddings is not None:
|
| 179 |
+
assert embeddings.ndim == 2, embeddings.shape
|
| 180 |
+
self._milestone_embeddings = embeddings
|
| 181 |
+
self._current_milestone_embedding = self.milestone_embeddings[0]
|
| 182 |
+
else:
|
| 183 |
+
self._milestone_embeddings = self._current_milestone_embedding = None
|
| 184 |
+
|
| 185 |
+
@property
|
| 186 |
+
def no_robot_milestones(self) -> np.ndarray:
|
| 187 |
+
return self._no_robot_milestones
|
| 188 |
+
|
| 189 |
+
@no_robot_milestones.setter
|
| 190 |
+
def no_robot_milestones(self, milestones: np.ndarray | None):
|
| 191 |
+
"""If set (outside), sub-goal switches will depend on this, while the
|
| 192 |
+
rollout inference is still based on the goal embed with arm."""
|
| 193 |
+
if milestones is not None:
|
| 194 |
+
self._no_robot_milestones = milestones.copy()
|
| 195 |
+
self._current_no_robot_milestone = milestones[0]
|
| 196 |
+
else:
|
| 197 |
+
self._no_robot_milestones = None
|
| 198 |
+
self._current_no_robot_milestone = None
|
| 199 |
+
self._subgoal_init_distance = None
|
| 200 |
+
|
| 201 |
+
@property
|
| 202 |
+
def cur_milestone_idx(self) -> int | None:
|
| 203 |
+
return self._cur_milestone_idx
|
| 204 |
+
|
| 205 |
+
@cur_milestone_idx.setter
|
| 206 |
+
def cur_milestone_idx(self, idx: int | None):
|
| 207 |
+
self._cur_milestone_idx = idx
|
| 208 |
+
|
| 209 |
+
def maybe_random_skip(self, cur, all_cur) -> np.ndarray:
|
| 210 |
+
if self._random_skip_ratio > 0:
|
| 211 |
+
if (
|
| 212 |
+
self._cur_milestone_idx is None
|
| 213 |
+
and self._random_skip_ratio > random.random()
|
| 214 |
+
):
|
| 215 |
+
self.cur_milestone_idx = min(
|
| 216 |
+
self._achieved_goals + 1, self.num_milestones - 1
|
| 217 |
+
)
|
| 218 |
+
elif self._cur_milestone_idx is None:
|
| 219 |
+
self.cur_milestone_idx = min(
|
| 220 |
+
self._achieved_goals, self.num_milestones - 1
|
| 221 |
+
)
|
| 222 |
+
assert (
|
| 223 |
+
self.cur_milestone_idx
|
| 224 |
+
in [
|
| 225 |
+
self._achieved_goals,
|
| 226 |
+
self._achieved_goals + 1,
|
| 227 |
+
self.num_milestones - 1,
|
| 228 |
+
]
|
| 229 |
+
and self.cur_milestone_idx < self.num_milestones
|
| 230 |
+
), f"WARN: {self.cur_milestone_idx}, {self._achieved_goals}"
|
| 231 |
+
return all_cur[self._cur_milestone_idx]
|
| 232 |
+
self._cur_milestone_idx = min(self._achieved_goals, self.num_milestones - 1)
|
| 233 |
+
return cur
|
| 234 |
+
|
| 235 |
+
@property
|
| 236 |
+
def current_milestone(self) -> np.ndarray:
|
| 237 |
+
if self.use_milestone_compressor:
|
| 238 |
+
return self.milestone_embeddings
|
| 239 |
+
# current sub-goal, switch automatically when set embedding distance
|
| 240 |
+
assert self._current_milestone is not None
|
| 241 |
+
return self.maybe_random_skip(self._current_milestone, self._milestones)
|
| 242 |
+
|
| 243 |
+
@property
|
| 244 |
+
def current_milestone_embedding(self) -> np.ndarray:
|
| 245 |
+
if self.current_milestone.ndim == 1:
|
| 246 |
+
return self.current_milestone
|
| 247 |
+
assert self._current_milestone_embedding is not None
|
| 248 |
+
return self.maybe_random_skip(
|
| 249 |
+
self._current_milestone_embedding, self._milestone_embeddings
|
| 250 |
+
)
|
| 251 |
+
|
| 252 |
+
def get_current_milestone_embedding_without_skip(self) -> np.ndarray:
|
| 253 |
+
if self.current_milestone.ndim == 1:
|
| 254 |
+
return self._current_milestone
|
| 255 |
+
assert self._current_milestone_embedding is not None
|
| 256 |
+
return self._current_milestone_embedding
|
| 257 |
+
|
| 258 |
+
@property
|
| 259 |
+
def current_no_robot_milestone(self) -> np.ndarray | None:
|
| 260 |
+
if self._current_no_robot_milestone is not None:
|
| 261 |
+
assert (
|
| 262 |
+
self._current_no_robot_milestone.ndim == 1
|
| 263 |
+
), self.current_milestone.shape
|
| 264 |
+
# no need for random skip
|
| 265 |
+
return self._current_no_robot_milestone
|
| 266 |
+
|
| 267 |
+
@property
|
| 268 |
+
def current_embedding_distance(self) -> np.ndarray:
|
| 269 |
+
return self._current_embedding_distance
|
| 270 |
+
|
| 271 |
+
@current_embedding_distance.setter
|
| 272 |
+
def current_embedding_distance(self, dist: float | None):
|
| 273 |
+
if dist is None or self._milestones is None:
|
| 274 |
+
self._current_embedding_distance = dist
|
| 275 |
+
return
|
| 276 |
+
switch = False
|
| 277 |
+
if dist is not None and dist < self._thresh:
|
| 278 |
+
switch = True
|
| 279 |
+
if self._delay_steps > 0 and self._n_steps_after_pass_thresh == -1:
|
| 280 |
+
switch = False
|
| 281 |
+
# trigger the delay
|
| 282 |
+
self._n_steps_after_pass_thresh = 0
|
| 283 |
+
if 0 < self._delay_steps <= self._n_steps_after_pass_thresh:
|
| 284 |
+
assert self._delay_steps == self._n_steps_after_pass_thresh
|
| 285 |
+
switch = True
|
| 286 |
+
self._n_steps_after_pass_thresh = -1
|
| 287 |
+
|
| 288 |
+
current_milestone_budget = self.current_milestone_budget()
|
| 289 |
+
if switch and current_milestone_budget is not None:
|
| 290 |
+
if current_milestone_budget > self._n_steps_this_milestone + 3:
|
| 291 |
+
switch = False # make sure not switch too fast by noise of the embed
|
| 292 |
+
elif not switch and current_milestone_budget is not None:
|
| 293 |
+
if current_milestone_budget < self._n_steps_this_milestone - 3:
|
| 294 |
+
switch = True # in case not switch due to noise of the embed
|
| 295 |
+
|
| 296 |
+
if switch:
|
| 297 |
+
self._n_steps_this_milestone = 0
|
| 298 |
+
if self.hybrid_phase is not None:
|
| 299 |
+
self.hybrid_phase += 1
|
| 300 |
+
# switch milestones
|
| 301 |
+
self._achieved_goals += 1
|
| 302 |
+
self._achieved_goals = min(self._achieved_goals, self.num_milestones)
|
| 303 |
+
# self._achieved_goals = self.env.num_goals_achieved
|
| 304 |
+
self._cur_milestone_idx = min(
|
| 305 |
+
self._achieved_goals, self.num_milestones - 1
|
| 306 |
+
) # next idx
|
| 307 |
+
self._current_milestone = self._milestones[self._cur_milestone_idx]
|
| 308 |
+
if self._current_milestone_embedding is not None:
|
| 309 |
+
# switch no next milestone embedding for conditioning
|
| 310 |
+
self._current_milestone_embedding = self.milestone_embeddings[
|
| 311 |
+
self._cur_milestone_idx
|
| 312 |
+
]
|
| 313 |
+
if self.current_no_robot_milestone is not None:
|
| 314 |
+
# switch to next milestone for calculating embed distance
|
| 315 |
+
assert len(self.milestones) == len(self.no_robot_milestones)
|
| 316 |
+
self._current_no_robot_milestone = self.no_robot_milestones[
|
| 317 |
+
self._cur_milestone_idx
|
| 318 |
+
]
|
| 319 |
+
if self.milestone_distances is not None:
|
| 320 |
+
self._subgoal_init_distance = self.milestone_distances[
|
| 321 |
+
min(self._achieved_goals - 1, self.num_milestones - 2)
|
| 322 |
+
]
|
| 323 |
+
else:
|
| 324 |
+
# use no skipping milestone for distance check
|
| 325 |
+
# self._subgoal_init_distance = np.linalg.norm(
|
| 326 |
+
# self.current_obs_embedding - self._current_no_robot_milestone
|
| 327 |
+
# )
|
| 328 |
+
self._subgoal_init_distance = None # set next step back
|
| 329 |
+
else:
|
| 330 |
+
if self.milestone_distances is not None:
|
| 331 |
+
# self._subgoal_init_distance = self.milestone_distances[
|
| 332 |
+
# min(self._achieved_goals - 1, self.num_milestones - 2)
|
| 333 |
+
# ]
|
| 334 |
+
self._subgoal_init_distance = None
|
| 335 |
+
else:
|
| 336 |
+
# use no skipping milestone for distance check
|
| 337 |
+
self._subgoal_init_distance = np.linalg.norm(
|
| 338 |
+
self.current_obs_embedding - self._current_milestone_embedding
|
| 339 |
+
)
|
| 340 |
+
self._current_embedding_distance = 1.0
|
| 341 |
+
else:
|
| 342 |
+
self._current_embedding_distance = dist
|
| 343 |
+
|
| 344 |
+
@property
|
| 345 |
+
def current_obs_embedding(self) -> np.ndarray:
|
| 346 |
+
assert self._current_obs_embedding is not None
|
| 347 |
+
return self._current_obs_embedding
|
| 348 |
+
|
| 349 |
+
@current_obs_embedding.setter
|
| 350 |
+
def current_obs_embedding(self, embedding: np.ndarray):
|
| 351 |
+
"""Set outside, should be embedding already for calculating embed
|
| 352 |
+
distance."""
|
| 353 |
+
assert embedding.ndim == 1, embedding.shape
|
| 354 |
+
self._current_obs_embedding = embedding
|
| 355 |
+
if self.hybrid_phase is None:
|
| 356 |
+
# calculate embed dist automatically after setting embedding
|
| 357 |
+
# use no skipping milestone for distance check
|
| 358 |
+
cur_milestone = (
|
| 359 |
+
self.current_no_robot_milestone
|
| 360 |
+
if self.current_no_robot_milestone is not None
|
| 361 |
+
else self.get_current_milestone_embedding_without_skip()
|
| 362 |
+
)
|
| 363 |
+
else:
|
| 364 |
+
cur_milestone = (
|
| 365 |
+
self.get_current_milestone_embedding_without_skip()
|
| 366 |
+
if self.hybrid_phase % 2 == 0
|
| 367 |
+
else self.current_no_robot_milestone
|
| 368 |
+
)
|
| 369 |
+
assert (
|
| 370 |
+
cur_milestone.ndim == embedding.ndim == 1
|
| 371 |
+
), f"{cur_milestone.shape} != {embedding.shape} != (n,)"
|
| 372 |
+
embed_distance = np.linalg.norm(embedding - cur_milestone)
|
| 373 |
+
self._unnormalized_embedding_distance = embed_distance
|
| 374 |
+
|
| 375 |
+
if self._subgoal_init_distance is None:
|
| 376 |
+
# first step after reset here
|
| 377 |
+
self._subgoal_init_distance = max(embed_distance, 1e-14)
|
| 378 |
+
# property setter
|
| 379 |
+
self.current_embedding_distance = embed_distance / self._subgoal_init_distance
|
| 380 |
+
|
| 381 |
+
@property
|
| 382 |
+
def reset_states(self) -> dict:
|
| 383 |
+
return self.env.reset_states
|
| 384 |
+
|
| 385 |
+
@reset_states.setter
|
| 386 |
+
def reset_states(self, state: dict):
|
| 387 |
+
self.env.reset_states = state
|
| 388 |
+
|
| 389 |
+
@property
|
| 390 |
+
def task_name(self) -> str:
|
| 391 |
+
return self.env.task_name
|
| 392 |
+
|
| 393 |
+
@property
|
| 394 |
+
def metrics(self) -> dict:
|
| 395 |
+
self.env: KitchenBase
|
| 396 |
+
metrics = self.env.metrics
|
| 397 |
+
achieved = (
|
| 398 |
+
self.completed_goals if not metrics["success"] else self.num_milestones
|
| 399 |
+
)
|
| 400 |
+
achieved_percentage = achieved / self.num_milestones
|
| 401 |
+
metrics["achieved_milestones"] = achieved_percentage
|
| 402 |
+
return metrics
|
| 403 |
+
|
| 404 |
+
@property
|
| 405 |
+
def current_no_robot_frame(self) -> np.ndarray:
|
| 406 |
+
if self.hybrid_phase is not None and self.hybrid_phase % 2 == 0:
|
| 407 |
+
# hack for hybrid return rgb again
|
| 408 |
+
assert self._cur_rgb_frame is not None
|
| 409 |
+
return self._cur_rgb_frame
|
| 410 |
+
else:
|
| 411 |
+
return self.env.current_no_robot_frame
|
| 412 |
+
|
| 413 |
+
@property
|
| 414 |
+
def current_rgb_frame(self) -> np.ndarray:
|
| 415 |
+
return self._cur_rgb_frame
|
| 416 |
+
|
| 417 |
+
@property
|
| 418 |
+
def milestone_indices(self):
|
| 419 |
+
return self._milestone_indices
|
| 420 |
+
|
| 421 |
+
@milestone_indices.setter
|
| 422 |
+
def milestone_indices(self, milestone_indices: np.ndarray | None):
|
| 423 |
+
self._milestone_indices = milestone_indices
|
| 424 |
+
|
| 425 |
+
def current_milestone_budget(self) -> int | None:
|
| 426 |
+
if self.cur_milestone_idx is None:
|
| 427 |
+
self.cur_milestone_idx = 0
|
| 428 |
+
if self.milestone_indices is None:
|
| 429 |
+
return None
|
| 430 |
+
return (
|
| 431 |
+
self.milestone_indices[self.cur_milestone_idx] + 1 # indices start from 0
|
| 432 |
+
if self.cur_milestone_idx == 0
|
| 433 |
+
else self.milestone_indices[self.cur_milestone_idx]
|
| 434 |
+
- self.milestone_indices[self.cur_milestone_idx - 1]
|
| 435 |
+
)
|
uvd/envs/evaluator/vec_envs/__init__.py
ADDED
|
File without changes
|
uvd/envs/evaluator/vec_envs/vec_env.py
ADDED
|
@@ -0,0 +1,398 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""Modified from `tianshou`"""
|
| 2 |
+
from typing import Any, Callable, List, Optional, Tuple, Union
|
| 3 |
+
|
| 4 |
+
import gym
|
| 5 |
+
import numpy as np
|
| 6 |
+
|
| 7 |
+
from .workers import EnvWorker, RayEnvWorker, SubprocEnvWorker
|
| 8 |
+
|
| 9 |
+
GYM_RESERVED_KEYS = [
|
| 10 |
+
"metadata",
|
| 11 |
+
"reward_range",
|
| 12 |
+
"spec",
|
| 13 |
+
"action_space",
|
| 14 |
+
"observation_space",
|
| 15 |
+
]
|
| 16 |
+
|
| 17 |
+
|
| 18 |
+
class BaseVectorEnv(object):
|
| 19 |
+
"""Base class for vectorized environments.
|
| 20 |
+
|
| 21 |
+
Usage:
|
| 22 |
+
::
|
| 23 |
+
|
| 24 |
+
env_num = 8
|
| 25 |
+
envs = DummyVectorEnv([lambda: gym.make(task) for _ in range(env_num)])
|
| 26 |
+
assert len(envs) == env_num
|
| 27 |
+
|
| 28 |
+
It accepts a list of environment generators. In other words, an environment
|
| 29 |
+
generator ``efn`` of a specific task means that ``efn()`` returns the
|
| 30 |
+
environment of the given task, for example, ``gym.make(task)``.
|
| 31 |
+
|
| 32 |
+
All of the VectorEnv must inherit :class:`~tianshou.env.BaseVectorEnv`.
|
| 33 |
+
Here are some other usages:
|
| 34 |
+
::
|
| 35 |
+
|
| 36 |
+
envs.seed(2) # which is equal to the next line
|
| 37 |
+
envs.seed([2, 3, 4, 5, 6, 7, 8, 9]) # set specific seed for each env
|
| 38 |
+
obs = envs.reset() # reset all environments
|
| 39 |
+
obs = envs.reset([0, 5, 7]) # reset 3 specific environments
|
| 40 |
+
obs, rew, done, info = envs.step([1] * 8) # step synchronously
|
| 41 |
+
envs.render() # render all environments
|
| 42 |
+
envs.close() # close all environments
|
| 43 |
+
|
| 44 |
+
.. warning::
|
| 45 |
+
|
| 46 |
+
If you use your own environment, please make sure the ``seed`` method
|
| 47 |
+
is set up properly, e.g.,
|
| 48 |
+
::
|
| 49 |
+
|
| 50 |
+
def seed(self, seed):
|
| 51 |
+
np.random.seed(seed)
|
| 52 |
+
|
| 53 |
+
Otherwise, the outputs of these envs may be the same with each other.
|
| 54 |
+
|
| 55 |
+
:param env_fns: a list of callable envs, ``env_fns[i]()`` generates the i-th env.
|
| 56 |
+
:param worker_fn: a callable worker, ``worker_fn(env_fns[i])`` generates a
|
| 57 |
+
worker which contains the i-th env.
|
| 58 |
+
:param int wait_num: use in asynchronous simulation if the time cost of
|
| 59 |
+
``env.step`` varies with time and synchronously waiting for all
|
| 60 |
+
environments to finish a step is time-wasting. In that case, we can
|
| 61 |
+
return when ``wait_num`` environments finish a step and keep on
|
| 62 |
+
simulation in these environments. If ``None``, asynchronous simulation
|
| 63 |
+
is disabled; else, ``1 <= wait_num <= env_num``.
|
| 64 |
+
:param float timeout: use in asynchronous simulation same as above, in each
|
| 65 |
+
vectorized step it only deal with those environments spending time
|
| 66 |
+
within ``timeout`` seconds.
|
| 67 |
+
"""
|
| 68 |
+
|
| 69 |
+
def __init__(
|
| 70 |
+
self,
|
| 71 |
+
env_fns: List[Callable[[], gym.Env]],
|
| 72 |
+
worker_fn: Callable[[Callable[[], gym.Env]], EnvWorker],
|
| 73 |
+
wait_num: Optional[int] = None,
|
| 74 |
+
timeout: Optional[float] = None,
|
| 75 |
+
) -> None:
|
| 76 |
+
self._env_fns = env_fns
|
| 77 |
+
# A VectorEnv contains a pool of EnvWorkers, which corresponds to
|
| 78 |
+
# interact with the given envs (one worker <-> one env).
|
| 79 |
+
self.workers = [worker_fn(fn) for fn in env_fns]
|
| 80 |
+
self.worker_class = type(self.workers[0])
|
| 81 |
+
assert issubclass(self.worker_class, EnvWorker)
|
| 82 |
+
assert all([isinstance(w, self.worker_class) for w in self.workers])
|
| 83 |
+
|
| 84 |
+
self.env_num = len(env_fns)
|
| 85 |
+
self.wait_num = wait_num or len(env_fns)
|
| 86 |
+
assert (
|
| 87 |
+
1 <= self.wait_num <= len(env_fns)
|
| 88 |
+
), f"wait_num should be in [1, {len(env_fns)}], but got {wait_num}"
|
| 89 |
+
self.timeout = timeout
|
| 90 |
+
assert (
|
| 91 |
+
self.timeout is None or self.timeout > 0
|
| 92 |
+
), f"timeout is {timeout}, it should be positive if provided!"
|
| 93 |
+
self.is_async = self.wait_num != len(env_fns) or timeout is not None
|
| 94 |
+
self.waiting_conn: List[EnvWorker] = []
|
| 95 |
+
# environments in self.ready_id is actually ready
|
| 96 |
+
# but environments in self.waiting_id are just waiting when checked,
|
| 97 |
+
# and they may be ready now, but this is not known until we check it
|
| 98 |
+
# in the step() function
|
| 99 |
+
self.waiting_id: List[int] = []
|
| 100 |
+
# all environments are ready in the beginning
|
| 101 |
+
self.ready_id = list(range(self.env_num))
|
| 102 |
+
self.is_closed = False
|
| 103 |
+
|
| 104 |
+
def _assert_is_not_closed(self) -> None:
|
| 105 |
+
assert (
|
| 106 |
+
not self.is_closed
|
| 107 |
+
), f"Methods of {self.__class__.__name__} cannot be called after close."
|
| 108 |
+
|
| 109 |
+
def __len__(self) -> int:
|
| 110 |
+
"""Return len(self), which is the number of environments."""
|
| 111 |
+
return self.env_num
|
| 112 |
+
|
| 113 |
+
def __getattribute__(self, key: str) -> Any:
|
| 114 |
+
"""Switch the attribute getter depending on the key.
|
| 115 |
+
|
| 116 |
+
Any class who inherits ``gym.Env`` will inherit some attributes,
|
| 117 |
+
like ``action_space``. However, we would like the attribute
|
| 118 |
+
lookup to go straight into the worker (in fact, this vector
|
| 119 |
+
env's action_space is always None).
|
| 120 |
+
"""
|
| 121 |
+
if key in GYM_RESERVED_KEYS: # reserved keys in gym.Env
|
| 122 |
+
return self.get_env_attr(key)
|
| 123 |
+
else:
|
| 124 |
+
return super().__getattribute__(key)
|
| 125 |
+
|
| 126 |
+
def get_env_attr(
|
| 127 |
+
self,
|
| 128 |
+
key: str,
|
| 129 |
+
id: Optional[Union[int, List[int], np.ndarray]] = None,
|
| 130 |
+
) -> List[Any]:
|
| 131 |
+
"""Get an attribute from the underlying environments.
|
| 132 |
+
|
| 133 |
+
If id is an int, retrieve the attribute denoted by key from the
|
| 134 |
+
environment underlying the worker at index id. The result is
|
| 135 |
+
returned as a list with one element. Otherwise, retrieve the
|
| 136 |
+
attribute for all workers at indices id and return a list that
|
| 137 |
+
is ordered correspondingly to id.
|
| 138 |
+
|
| 139 |
+
:param str key: The key of the desired attribute.
|
| 140 |
+
:param id: Indice(s) of the desired worker(s). Default to None
|
| 141 |
+
for all env_id.
|
| 142 |
+
:return list: The list of environment attributes.
|
| 143 |
+
"""
|
| 144 |
+
self._assert_is_not_closed()
|
| 145 |
+
id = self._wrap_id(id)
|
| 146 |
+
if self.is_async:
|
| 147 |
+
self._assert_id(id)
|
| 148 |
+
|
| 149 |
+
return [self.workers[j].get_env_attr(key) for j in id]
|
| 150 |
+
|
| 151 |
+
def set_env_attr(
|
| 152 |
+
self,
|
| 153 |
+
key: str,
|
| 154 |
+
value: Any,
|
| 155 |
+
id: Optional[Union[int, List[int], np.ndarray]] = None,
|
| 156 |
+
diff_value: bool = False,
|
| 157 |
+
) -> None:
|
| 158 |
+
"""Set an attribute in the underlying environments.
|
| 159 |
+
|
| 160 |
+
If id is an int, set the attribute denoted by key from the
|
| 161 |
+
environment underlying the worker at index id to value.
|
| 162 |
+
Otherwise, set the attribute for all workers at indices id.
|
| 163 |
+
|
| 164 |
+
:param str key: The key of the desired attribute.
|
| 165 |
+
:param Any value: The new value of the attribute.
|
| 166 |
+
:param id: Indice(s) of the desired worker(s). Default to None
|
| 167 |
+
for all env_id.
|
| 168 |
+
"""
|
| 169 |
+
self._assert_is_not_closed()
|
| 170 |
+
id = self._wrap_id(id)
|
| 171 |
+
if diff_value:
|
| 172 |
+
assert len(value) == len(id)
|
| 173 |
+
if self.is_async:
|
| 174 |
+
self._assert_id(id)
|
| 175 |
+
for i, j in enumerate(id):
|
| 176 |
+
if diff_value:
|
| 177 |
+
self.workers[j].set_env_attr(key, value[i])
|
| 178 |
+
else:
|
| 179 |
+
self.workers[j].set_env_attr(key, value)
|
| 180 |
+
|
| 181 |
+
def _wrap_id(
|
| 182 |
+
self,
|
| 183 |
+
id: Optional[Union[int, List[int], np.ndarray]] = None,
|
| 184 |
+
) -> Union[List[int], np.ndarray]:
|
| 185 |
+
if id is None:
|
| 186 |
+
return list(range(self.env_num))
|
| 187 |
+
return [id] if np.isscalar(id) else id # type: ignore
|
| 188 |
+
|
| 189 |
+
def _assert_id(self, id: Union[List[int], np.ndarray]) -> None:
|
| 190 |
+
for i in id:
|
| 191 |
+
assert (
|
| 192 |
+
i not in self.waiting_id
|
| 193 |
+
), f"Cannot interact with environment {i} which is stepping now."
|
| 194 |
+
assert (
|
| 195 |
+
i in self.ready_id
|
| 196 |
+
), f"Can only interact with ready environments {self.ready_id}."
|
| 197 |
+
|
| 198 |
+
# TODO: compatible issue with reset -> (obs, info)
|
| 199 |
+
def reset(
|
| 200 |
+
self, id: Optional[Union[int, List[int], np.ndarray]] = None
|
| 201 |
+
) -> np.ndarray:
|
| 202 |
+
"""Reset the state of some envs and return initial observations.
|
| 203 |
+
|
| 204 |
+
If id is None, reset the state of all the environments and
|
| 205 |
+
return initial observations, otherwise reset the specific
|
| 206 |
+
environments with the given id, either an int or a list.
|
| 207 |
+
"""
|
| 208 |
+
self._assert_is_not_closed()
|
| 209 |
+
id = self._wrap_id(id)
|
| 210 |
+
if self.is_async:
|
| 211 |
+
self._assert_id(id)
|
| 212 |
+
# send(None) == reset() in worker
|
| 213 |
+
for i in id:
|
| 214 |
+
self.workers[i].send(None)
|
| 215 |
+
obs_list = [self.workers[i].recv() for i in id]
|
| 216 |
+
try:
|
| 217 |
+
obs = np.stack(obs_list)
|
| 218 |
+
except ValueError: # different len(obs)
|
| 219 |
+
obs = np.array(obs_list, dtype=object)
|
| 220 |
+
return obs
|
| 221 |
+
|
| 222 |
+
def step(
|
| 223 |
+
self,
|
| 224 |
+
action: np.ndarray,
|
| 225 |
+
id: Optional[Union[int, List[int], np.ndarray]] = None,
|
| 226 |
+
) -> Tuple[np.ndarray, np.ndarray, np.ndarray, np.ndarray]:
|
| 227 |
+
"""Run one timestep of some environments' dynamics.
|
| 228 |
+
|
| 229 |
+
If id is None, run one timestep of all the environments’ dynamics;
|
| 230 |
+
otherwise run one timestep for some environments with given id, either
|
| 231 |
+
an int or a list. When the end of episode is reached, you are
|
| 232 |
+
responsible for calling reset(id) to reset this environment’s state.
|
| 233 |
+
|
| 234 |
+
Accept a batch of action and return a tuple (batch_obs, batch_rew,
|
| 235 |
+
batch_done, batch_info) in numpy format.
|
| 236 |
+
|
| 237 |
+
:param numpy.ndarray action: a batch of action provided by the agent.
|
| 238 |
+
|
| 239 |
+
:return: A tuple including four items:
|
| 240 |
+
|
| 241 |
+
* ``obs`` a numpy.ndarray, the agent's observation of current environments
|
| 242 |
+
* ``rew`` a numpy.ndarray, the amount of rewards returned after \
|
| 243 |
+
previous actions
|
| 244 |
+
* ``done`` a numpy.ndarray, whether these episodes have ended, in \
|
| 245 |
+
which case further step() calls will return undefined results
|
| 246 |
+
* ``info`` a numpy.ndarray, contains auxiliary diagnostic \
|
| 247 |
+
information (helpful for debugging, and sometimes learning)
|
| 248 |
+
|
| 249 |
+
For the async simulation:
|
| 250 |
+
|
| 251 |
+
Provide the given action to the environments. The action sequence
|
| 252 |
+
should correspond to the ``id`` argument, and the ``id`` argument
|
| 253 |
+
should be a subset of the ``env_id`` in the last returned ``info``
|
| 254 |
+
(initially they are env_ids of all the environments). If action is
|
| 255 |
+
None, fetch unfinished step() calls instead.
|
| 256 |
+
"""
|
| 257 |
+
self._assert_is_not_closed()
|
| 258 |
+
id = self._wrap_id(id)
|
| 259 |
+
if not self.is_async:
|
| 260 |
+
assert len(action) == len(id)
|
| 261 |
+
for i, j in enumerate(id):
|
| 262 |
+
self.workers[j].send(action[i])
|
| 263 |
+
result = []
|
| 264 |
+
for j in id:
|
| 265 |
+
obs, rew, done, info = self.workers[j].recv()
|
| 266 |
+
info["env_id"] = j
|
| 267 |
+
result.append((obs, rew, done, info))
|
| 268 |
+
else:
|
| 269 |
+
if action is not None:
|
| 270 |
+
self._assert_id(id)
|
| 271 |
+
assert len(action) == len(id)
|
| 272 |
+
for act, env_id in zip(action, id):
|
| 273 |
+
self.workers[env_id].send(act)
|
| 274 |
+
self.waiting_conn.append(self.workers[env_id])
|
| 275 |
+
self.waiting_id.append(env_id)
|
| 276 |
+
self.ready_id = [x for x in self.ready_id if x not in id]
|
| 277 |
+
ready_conns: List[EnvWorker] = []
|
| 278 |
+
while not ready_conns:
|
| 279 |
+
ready_conns = self.worker_class.wait(
|
| 280 |
+
self.waiting_conn, self.wait_num, self.timeout
|
| 281 |
+
)
|
| 282 |
+
result = []
|
| 283 |
+
for conn in ready_conns:
|
| 284 |
+
waiting_index = self.waiting_conn.index(conn)
|
| 285 |
+
self.waiting_conn.pop(waiting_index)
|
| 286 |
+
env_id = self.waiting_id.pop(waiting_index)
|
| 287 |
+
obs, rew, done, info = conn.recv()
|
| 288 |
+
info["env_id"] = env_id
|
| 289 |
+
result.append((obs, rew, done, info))
|
| 290 |
+
self.ready_id.append(env_id)
|
| 291 |
+
obs_list, rew_list, done_list, info_list = zip(*result)
|
| 292 |
+
try:
|
| 293 |
+
obs_stack = np.stack(obs_list)
|
| 294 |
+
except ValueError: # different len(obs)
|
| 295 |
+
obs_stack = np.array(obs_list, dtype=object)
|
| 296 |
+
rew_stack, done_stack, info_stack = map(
|
| 297 |
+
np.stack, [rew_list, done_list, info_list]
|
| 298 |
+
)
|
| 299 |
+
return obs_stack, rew_stack, done_stack, info_stack
|
| 300 |
+
|
| 301 |
+
def seed(
|
| 302 |
+
self,
|
| 303 |
+
seed: Optional[Union[int, List[int]]] = None,
|
| 304 |
+
) -> List[Optional[List[int]]]:
|
| 305 |
+
"""Set the seed for all environments.
|
| 306 |
+
|
| 307 |
+
Accept ``None``, an int (which will extend ``i`` to
|
| 308 |
+
``[i, i + 1, i + 2, ...]``) or a list.
|
| 309 |
+
|
| 310 |
+
:return: The list of seeds used in this env's random number generators.
|
| 311 |
+
The first value in the list should be the "main" seed, or the value
|
| 312 |
+
which a reproducer pass to "seed".
|
| 313 |
+
"""
|
| 314 |
+
self._assert_is_not_closed()
|
| 315 |
+
seed_list: Union[List[None], List[int]]
|
| 316 |
+
if seed is None:
|
| 317 |
+
seed_list = [seed] * self.env_num
|
| 318 |
+
elif isinstance(seed, int):
|
| 319 |
+
seed_list = [seed + i for i in range(self.env_num)]
|
| 320 |
+
else:
|
| 321 |
+
seed_list = seed
|
| 322 |
+
return [w.seed(s) for w, s in zip(self.workers, seed_list)]
|
| 323 |
+
|
| 324 |
+
def render(self, **kwargs: Any) -> List[Any]:
|
| 325 |
+
"""Render all of the environments."""
|
| 326 |
+
self._assert_is_not_closed()
|
| 327 |
+
if self.is_async and len(self.waiting_id) > 0:
|
| 328 |
+
raise RuntimeError(
|
| 329 |
+
f"Environments {self.waiting_id} are still stepping, cannot "
|
| 330 |
+
"render them now."
|
| 331 |
+
)
|
| 332 |
+
return [w.render(**kwargs) for w in self.workers]
|
| 333 |
+
|
| 334 |
+
def close(self) -> None:
|
| 335 |
+
"""Close all of the environments.
|
| 336 |
+
|
| 337 |
+
This function will be called only once (if not, it will be
|
| 338 |
+
called during garbage collected). This way, ``close`` of all
|
| 339 |
+
workers can be assured.
|
| 340 |
+
"""
|
| 341 |
+
self._assert_is_not_closed()
|
| 342 |
+
for w in self.workers:
|
| 343 |
+
w.close()
|
| 344 |
+
self.is_closed = True
|
| 345 |
+
|
| 346 |
+
|
| 347 |
+
class SubprocVectorEnv(BaseVectorEnv):
|
| 348 |
+
"""Vectorized environment wrapper based on subprocess.
|
| 349 |
+
|
| 350 |
+
.. seealso::
|
| 351 |
+
|
| 352 |
+
Please refer to :class:`~tianshou.env.BaseVectorEnv` for other APIs' usage.
|
| 353 |
+
"""
|
| 354 |
+
|
| 355 |
+
def __init__(self, env_fns: List[Callable[[], gym.Env]], **kwargs: Any) -> None:
|
| 356 |
+
def worker_fn(fn: Callable[[], gym.Env]) -> SubprocEnvWorker:
|
| 357 |
+
return SubprocEnvWorker(fn, share_memory=False)
|
| 358 |
+
|
| 359 |
+
super().__init__(env_fns, worker_fn, **kwargs)
|
| 360 |
+
|
| 361 |
+
|
| 362 |
+
class ShmemVectorEnv(BaseVectorEnv):
|
| 363 |
+
"""Optimized SubprocVectorEnv with shared buffers to exchange observations.
|
| 364 |
+
|
| 365 |
+
ShmemVectorEnv has exactly the same API as SubprocVectorEnv.
|
| 366 |
+
|
| 367 |
+
.. seealso::
|
| 368 |
+
|
| 369 |
+
Please refer to :class:`~tianshou.env.BaseVectorEnv` for other APIs' usage.
|
| 370 |
+
"""
|
| 371 |
+
|
| 372 |
+
def __init__(self, env_fns: List[Callable[[], gym.Env]], **kwargs: Any) -> None:
|
| 373 |
+
def worker_fn(fn: Callable[[], gym.Env]) -> SubprocEnvWorker:
|
| 374 |
+
return SubprocEnvWorker(fn, share_memory=True)
|
| 375 |
+
|
| 376 |
+
super().__init__(env_fns, worker_fn, **kwargs)
|
| 377 |
+
|
| 378 |
+
|
| 379 |
+
class RayVectorEnv(BaseVectorEnv):
|
| 380 |
+
"""Vectorized environment wrapper based on ray.
|
| 381 |
+
|
| 382 |
+
This is a choice to run distributed environments in a cluster.
|
| 383 |
+
|
| 384 |
+
.. seealso::
|
| 385 |
+
|
| 386 |
+
Please refer to :class:`~tianshou.env.BaseVectorEnv` for other APIs' usage.
|
| 387 |
+
"""
|
| 388 |
+
|
| 389 |
+
def __init__(self, env_fns: List[Callable[[], gym.Env]], **kwargs: Any) -> None:
|
| 390 |
+
try:
|
| 391 |
+
import ray
|
| 392 |
+
except ImportError as exception:
|
| 393 |
+
raise ImportError(
|
| 394 |
+
"Please install ray to support RayVectorEnv: pip install ray"
|
| 395 |
+
) from exception
|
| 396 |
+
if not ray.is_initialized():
|
| 397 |
+
ray.init()
|
| 398 |
+
super().__init__(env_fns, RayEnvWorker, **kwargs)
|
uvd/envs/evaluator/vec_envs/workers.py
ADDED
|
@@ -0,0 +1,409 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
import ctypes
|
| 2 |
+
import time
|
| 3 |
+
import warnings
|
| 4 |
+
from abc import ABC, abstractmethod
|
| 5 |
+
from collections import OrderedDict
|
| 6 |
+
from multiprocessing import Array, Pipe, connection
|
| 7 |
+
from multiprocessing.context import Process
|
| 8 |
+
from typing import Any, Callable, List, Optional, Tuple, Union
|
| 9 |
+
|
| 10 |
+
import cloudpickle
|
| 11 |
+
import gym
|
| 12 |
+
import numpy as np
|
| 13 |
+
|
| 14 |
+
try:
|
| 15 |
+
import ray
|
| 16 |
+
except ImportError:
|
| 17 |
+
pass
|
| 18 |
+
|
| 19 |
+
|
| 20 |
+
warnings.simplefilter("once", DeprecationWarning)
|
| 21 |
+
|
| 22 |
+
_NP_TO_CT = {
|
| 23 |
+
np.bool_: ctypes.c_bool,
|
| 24 |
+
np.uint8: ctypes.c_uint8,
|
| 25 |
+
np.uint16: ctypes.c_uint16,
|
| 26 |
+
np.uint32: ctypes.c_uint32,
|
| 27 |
+
np.uint64: ctypes.c_uint64,
|
| 28 |
+
np.int8: ctypes.c_int8,
|
| 29 |
+
np.int16: ctypes.c_int16,
|
| 30 |
+
np.int32: ctypes.c_int32,
|
| 31 |
+
np.int64: ctypes.c_int64,
|
| 32 |
+
np.float32: ctypes.c_float,
|
| 33 |
+
np.float64: ctypes.c_double,
|
| 34 |
+
}
|
| 35 |
+
|
| 36 |
+
|
| 37 |
+
def deprecation(msg: str) -> None:
|
| 38 |
+
"""Deprecation warning wrapper."""
|
| 39 |
+
warnings.warn(msg, category=DeprecationWarning, stacklevel=2)
|
| 40 |
+
|
| 41 |
+
|
| 42 |
+
class CloudpickleWrapper(object):
|
| 43 |
+
"""A cloudpickle wrapper used in SubprocVectorEnv."""
|
| 44 |
+
|
| 45 |
+
def __init__(self, data: Any) -> None:
|
| 46 |
+
self.data = data
|
| 47 |
+
|
| 48 |
+
def __getstate__(self) -> str:
|
| 49 |
+
return cloudpickle.dumps(self.data)
|
| 50 |
+
|
| 51 |
+
def __setstate__(self, data: str) -> None:
|
| 52 |
+
self.data = cloudpickle.loads(data)
|
| 53 |
+
|
| 54 |
+
|
| 55 |
+
class EnvWorker(ABC):
|
| 56 |
+
"""An abstract worker for an environment."""
|
| 57 |
+
|
| 58 |
+
def __init__(self, env_fn: Callable[[], gym.Env]) -> None:
|
| 59 |
+
self._env_fn = env_fn
|
| 60 |
+
self.is_closed = False
|
| 61 |
+
self.result: Union[
|
| 62 |
+
Tuple[np.ndarray, np.ndarray, np.ndarray, np.ndarray], np.ndarray
|
| 63 |
+
]
|
| 64 |
+
self.action_space = self.get_env_attr("action_space") # noqa: B009
|
| 65 |
+
self.is_reset = False
|
| 66 |
+
|
| 67 |
+
@abstractmethod
|
| 68 |
+
def get_env_attr(self, key: str) -> Any:
|
| 69 |
+
pass
|
| 70 |
+
|
| 71 |
+
@abstractmethod
|
| 72 |
+
def set_env_attr(self, key: str, value: Any) -> None:
|
| 73 |
+
pass
|
| 74 |
+
|
| 75 |
+
def send(self, action: Optional[np.ndarray]) -> None:
|
| 76 |
+
"""Send action signal to low-level worker.
|
| 77 |
+
|
| 78 |
+
When action is None, it indicates sending "reset" signal;
|
| 79 |
+
otherwise it indicates "step" signal. The paired return value
|
| 80 |
+
from "recv" function is determined by such kind of different
|
| 81 |
+
signal.
|
| 82 |
+
"""
|
| 83 |
+
if hasattr(self, "send_action"):
|
| 84 |
+
deprecation(
|
| 85 |
+
"send_action will soon be deprecated. "
|
| 86 |
+
"Please use send and recv for your own EnvWorker."
|
| 87 |
+
)
|
| 88 |
+
if action is None:
|
| 89 |
+
self.is_reset = True
|
| 90 |
+
self.result = self.reset()
|
| 91 |
+
else:
|
| 92 |
+
self.is_reset = False
|
| 93 |
+
self.send_action(action) # type: ignore
|
| 94 |
+
|
| 95 |
+
def recv(
|
| 96 |
+
self,
|
| 97 |
+
) -> Union[Tuple[np.ndarray, np.ndarray, np.ndarray, np.ndarray], np.ndarray]:
|
| 98 |
+
"""Receive result from low-level worker.
|
| 99 |
+
|
| 100 |
+
If the last "send" function sends a NULL action, it only returns
|
| 101 |
+
a single observation; otherwise it returns a tuple of (obs, rew,
|
| 102 |
+
done, info).
|
| 103 |
+
"""
|
| 104 |
+
if hasattr(self, "get_result"):
|
| 105 |
+
deprecation(
|
| 106 |
+
"get_result will soon be deprecated. "
|
| 107 |
+
"Please use send and recv for your own EnvWorker."
|
| 108 |
+
)
|
| 109 |
+
if not self.is_reset:
|
| 110 |
+
self.result = self.get_result() # type: ignore
|
| 111 |
+
return self.result
|
| 112 |
+
|
| 113 |
+
def reset(self) -> np.ndarray:
|
| 114 |
+
self.send(None)
|
| 115 |
+
return self.recv() # type: ignore
|
| 116 |
+
|
| 117 |
+
def step(
|
| 118 |
+
self, action: np.ndarray
|
| 119 |
+
) -> Tuple[np.ndarray, np.ndarray, np.ndarray, np.ndarray]:
|
| 120 |
+
"""Perform one timestep of the environment's dynamic.
|
| 121 |
+
|
| 122 |
+
"send" and "recv" are coupled in sync simulation, so users only
|
| 123 |
+
call "step" function. But they can be called separately in async
|
| 124 |
+
simulation, i.e. someone calls "send" first, and calls "recv"
|
| 125 |
+
later.
|
| 126 |
+
"""
|
| 127 |
+
self.send(action)
|
| 128 |
+
return self.recv() # type: ignore
|
| 129 |
+
|
| 130 |
+
@staticmethod
|
| 131 |
+
def wait(
|
| 132 |
+
workers: List["EnvWorker"], wait_num: int, timeout: Optional[float] = None
|
| 133 |
+
) -> List["EnvWorker"]:
|
| 134 |
+
"""Given a list of workers, return those ready ones."""
|
| 135 |
+
raise NotImplementedError
|
| 136 |
+
|
| 137 |
+
def seed(self, seed: Optional[int] = None) -> Optional[List[int]]:
|
| 138 |
+
return self.action_space.seed(seed) # issue 299
|
| 139 |
+
|
| 140 |
+
@abstractmethod
|
| 141 |
+
def render(self, **kwargs: Any) -> Any:
|
| 142 |
+
"""Render the environment."""
|
| 143 |
+
pass
|
| 144 |
+
|
| 145 |
+
@abstractmethod
|
| 146 |
+
def close_env(self) -> None:
|
| 147 |
+
pass
|
| 148 |
+
|
| 149 |
+
def close(self) -> None:
|
| 150 |
+
if self.is_closed:
|
| 151 |
+
return None
|
| 152 |
+
self.is_closed = True
|
| 153 |
+
self.close_env()
|
| 154 |
+
|
| 155 |
+
|
| 156 |
+
class ShArray:
|
| 157 |
+
"""Wrapper of multiprocessing Array."""
|
| 158 |
+
|
| 159 |
+
def __init__(self, dtype: np.generic, shape: Tuple[int]) -> None:
|
| 160 |
+
self.arr = Array(_NP_TO_CT[dtype.type], int(np.prod(shape))) # type: ignore
|
| 161 |
+
self.dtype = dtype
|
| 162 |
+
self.shape = shape
|
| 163 |
+
|
| 164 |
+
def save(self, ndarray: np.ndarray) -> None:
|
| 165 |
+
assert isinstance(ndarray, np.ndarray)
|
| 166 |
+
dst = self.arr.get_obj()
|
| 167 |
+
dst_np = np.frombuffer(dst, dtype=self.dtype).reshape(self.shape)
|
| 168 |
+
np.copyto(dst_np, ndarray)
|
| 169 |
+
|
| 170 |
+
def get(self) -> np.ndarray:
|
| 171 |
+
obj = self.arr.get_obj()
|
| 172 |
+
return np.frombuffer(obj, dtype=self.dtype).reshape(self.shape)
|
| 173 |
+
|
| 174 |
+
|
| 175 |
+
def _setup_buf(space: gym.Space) -> Union[dict, tuple, ShArray]:
|
| 176 |
+
if isinstance(space, gym.spaces.Dict):
|
| 177 |
+
assert isinstance(space.spaces, OrderedDict)
|
| 178 |
+
return {k: _setup_buf(v) for k, v in space.spaces.items()}
|
| 179 |
+
elif isinstance(space, gym.spaces.Tuple):
|
| 180 |
+
assert isinstance(space.spaces, tuple)
|
| 181 |
+
return tuple([_setup_buf(t) for t in space.spaces])
|
| 182 |
+
else:
|
| 183 |
+
return ShArray(space.dtype, space.shape) # type: ignore
|
| 184 |
+
|
| 185 |
+
|
| 186 |
+
def _worker(
|
| 187 |
+
parent: connection.Connection,
|
| 188 |
+
p: connection.Connection,
|
| 189 |
+
env_fn_wrapper: CloudpickleWrapper,
|
| 190 |
+
obs_bufs: Optional[Union[dict, tuple, ShArray]] = None,
|
| 191 |
+
) -> None:
|
| 192 |
+
def _encode_obs(
|
| 193 |
+
obs: Union[dict, tuple, np.ndarray], buffer: Union[dict, tuple, ShArray]
|
| 194 |
+
) -> None:
|
| 195 |
+
if isinstance(obs, np.ndarray) and isinstance(buffer, ShArray):
|
| 196 |
+
buffer.save(obs)
|
| 197 |
+
elif isinstance(obs, tuple) and isinstance(buffer, tuple):
|
| 198 |
+
for o, b in zip(obs, buffer):
|
| 199 |
+
_encode_obs(o, b)
|
| 200 |
+
elif isinstance(obs, dict) and isinstance(buffer, dict):
|
| 201 |
+
for k in obs.keys():
|
| 202 |
+
_encode_obs(obs[k], buffer[k])
|
| 203 |
+
return None
|
| 204 |
+
|
| 205 |
+
parent.close()
|
| 206 |
+
env = env_fn_wrapper.data()
|
| 207 |
+
try:
|
| 208 |
+
while True:
|
| 209 |
+
try:
|
| 210 |
+
cmd, data = p.recv()
|
| 211 |
+
except EOFError: # the pipe has been closed
|
| 212 |
+
p.close()
|
| 213 |
+
break
|
| 214 |
+
if cmd == "step":
|
| 215 |
+
if data is None: # reset
|
| 216 |
+
obs = env.reset()
|
| 217 |
+
else:
|
| 218 |
+
obs, reward, done, info = env.step(data)
|
| 219 |
+
if obs_bufs is not None:
|
| 220 |
+
_encode_obs(obs, obs_bufs)
|
| 221 |
+
obs = None
|
| 222 |
+
if data is None:
|
| 223 |
+
p.send(obs)
|
| 224 |
+
else:
|
| 225 |
+
p.send((obs, reward, done, info))
|
| 226 |
+
elif cmd == "close":
|
| 227 |
+
p.send(env.close())
|
| 228 |
+
p.close()
|
| 229 |
+
break
|
| 230 |
+
elif cmd == "render":
|
| 231 |
+
p.send(env.render(**data) if hasattr(env, "render") else None)
|
| 232 |
+
elif cmd == "seed":
|
| 233 |
+
p.send(env.seed(data) if hasattr(env, "seed") else None)
|
| 234 |
+
elif cmd == "getattr":
|
| 235 |
+
p.send(getattr(env, data) if hasattr(env, data) else None)
|
| 236 |
+
elif cmd == "setattr":
|
| 237 |
+
setattr(env, data["key"], data["value"])
|
| 238 |
+
else:
|
| 239 |
+
p.close()
|
| 240 |
+
raise NotImplementedError
|
| 241 |
+
except KeyboardInterrupt:
|
| 242 |
+
p.close()
|
| 243 |
+
|
| 244 |
+
|
| 245 |
+
class SubprocEnvWorker(EnvWorker):
|
| 246 |
+
"""Subprocess worker used in SubprocVectorEnv and ShmemVectorEnv."""
|
| 247 |
+
|
| 248 |
+
def __init__(
|
| 249 |
+
self, env_fn: Callable[[], gym.Env], share_memory: bool = False
|
| 250 |
+
) -> None:
|
| 251 |
+
self.parent_remote, self.child_remote = Pipe()
|
| 252 |
+
self.share_memory = share_memory
|
| 253 |
+
self.buffer: Optional[Union[dict, tuple, ShArray]] = None
|
| 254 |
+
if self.share_memory:
|
| 255 |
+
dummy = env_fn()
|
| 256 |
+
obs_space = dummy.observation_space
|
| 257 |
+
dummy.close()
|
| 258 |
+
del dummy
|
| 259 |
+
self.buffer = _setup_buf(obs_space)
|
| 260 |
+
args = (
|
| 261 |
+
self.parent_remote,
|
| 262 |
+
self.child_remote,
|
| 263 |
+
CloudpickleWrapper(env_fn),
|
| 264 |
+
self.buffer,
|
| 265 |
+
)
|
| 266 |
+
self.process = Process(target=_worker, args=args, daemon=True)
|
| 267 |
+
self.process.start()
|
| 268 |
+
self.child_remote.close()
|
| 269 |
+
self.is_reset = False
|
| 270 |
+
super().__init__(env_fn)
|
| 271 |
+
|
| 272 |
+
def get_env_attr(self, key: str) -> Any:
|
| 273 |
+
self.parent_remote.send(["getattr", key])
|
| 274 |
+
return self.parent_remote.recv()
|
| 275 |
+
|
| 276 |
+
def set_env_attr(self, key: str, value: Any) -> None:
|
| 277 |
+
self.parent_remote.send(["setattr", {"key": key, "value": value}])
|
| 278 |
+
|
| 279 |
+
def _decode_obs(self) -> Union[dict, tuple, np.ndarray]:
|
| 280 |
+
def decode_obs(
|
| 281 |
+
buffer: Optional[Union[dict, tuple, ShArray]]
|
| 282 |
+
) -> Union[dict, tuple, np.ndarray]:
|
| 283 |
+
if isinstance(buffer, ShArray):
|
| 284 |
+
return buffer.get()
|
| 285 |
+
elif isinstance(buffer, tuple):
|
| 286 |
+
return tuple([decode_obs(b) for b in buffer])
|
| 287 |
+
elif isinstance(buffer, dict):
|
| 288 |
+
return {k: decode_obs(v) for k, v in buffer.items()}
|
| 289 |
+
else:
|
| 290 |
+
raise NotImplementedError
|
| 291 |
+
|
| 292 |
+
return decode_obs(self.buffer)
|
| 293 |
+
|
| 294 |
+
@staticmethod
|
| 295 |
+
def wait( # type: ignore
|
| 296 |
+
workers: List["SubprocEnvWorker"],
|
| 297 |
+
wait_num: int,
|
| 298 |
+
timeout: Optional[float] = None,
|
| 299 |
+
) -> List["SubprocEnvWorker"]:
|
| 300 |
+
remain_conns = conns = [x.parent_remote for x in workers]
|
| 301 |
+
ready_conns: List[connection.Connection] = []
|
| 302 |
+
remain_time, t1 = timeout, time.time()
|
| 303 |
+
while len(remain_conns) > 0 and len(ready_conns) < wait_num:
|
| 304 |
+
if timeout:
|
| 305 |
+
remain_time = timeout - (time.time() - t1)
|
| 306 |
+
if remain_time <= 0:
|
| 307 |
+
break
|
| 308 |
+
# connection.wait hangs if the list is empty
|
| 309 |
+
new_ready_conns = connection.wait(remain_conns, timeout=remain_time)
|
| 310 |
+
ready_conns.extend(new_ready_conns) # type: ignore
|
| 311 |
+
remain_conns = [conn for conn in remain_conns if conn not in ready_conns]
|
| 312 |
+
return [workers[conns.index(con)] for con in ready_conns]
|
| 313 |
+
|
| 314 |
+
def send(self, action: Optional[np.ndarray]) -> None:
|
| 315 |
+
self.parent_remote.send(["step", action])
|
| 316 |
+
|
| 317 |
+
def recv(
|
| 318 |
+
self,
|
| 319 |
+
) -> Union[Tuple[np.ndarray, np.ndarray, np.ndarray, np.ndarray], np.ndarray]:
|
| 320 |
+
result = self.parent_remote.recv()
|
| 321 |
+
if isinstance(result, tuple):
|
| 322 |
+
obs, rew, done, info = result
|
| 323 |
+
if self.share_memory:
|
| 324 |
+
obs = self._decode_obs()
|
| 325 |
+
return obs, rew, done, info
|
| 326 |
+
else:
|
| 327 |
+
obs = result
|
| 328 |
+
if self.share_memory:
|
| 329 |
+
obs = self._decode_obs()
|
| 330 |
+
return obs
|
| 331 |
+
|
| 332 |
+
def seed(self, seed: Optional[int] = None) -> Optional[List[int]]:
|
| 333 |
+
super().seed(seed)
|
| 334 |
+
self.parent_remote.send(["seed", seed])
|
| 335 |
+
return self.parent_remote.recv()
|
| 336 |
+
|
| 337 |
+
def render(self, **kwargs: Any) -> Any:
|
| 338 |
+
self.parent_remote.send(["render", kwargs])
|
| 339 |
+
return self.parent_remote.recv()
|
| 340 |
+
|
| 341 |
+
def close_env(self) -> None:
|
| 342 |
+
try:
|
| 343 |
+
self.parent_remote.send(["close", None])
|
| 344 |
+
# mp may be deleted so it may raise AttributeError
|
| 345 |
+
self.parent_remote.recv()
|
| 346 |
+
self.process.join()
|
| 347 |
+
except (BrokenPipeError, EOFError, AttributeError):
|
| 348 |
+
pass
|
| 349 |
+
# ensure the subproc is terminated
|
| 350 |
+
self.process.terminate()
|
| 351 |
+
|
| 352 |
+
|
| 353 |
+
class _SetAttrWrapper(gym.Wrapper):
|
| 354 |
+
def set_env_attr(self, key: str, value: Any) -> None:
|
| 355 |
+
setattr(self.env, key, value)
|
| 356 |
+
|
| 357 |
+
def get_env_attr(self, key: str) -> Any:
|
| 358 |
+
return getattr(self.env, key)
|
| 359 |
+
|
| 360 |
+
|
| 361 |
+
class RayEnvWorker(EnvWorker):
|
| 362 |
+
"""Ray worker used in RayVectorEnv."""
|
| 363 |
+
|
| 364 |
+
def __init__(self, env_fn: Callable[[], gym.Env]) -> None:
|
| 365 |
+
self.env = (
|
| 366 |
+
ray.remote(_SetAttrWrapper)
|
| 367 |
+
.options(num_cpus=0) # type: ignore
|
| 368 |
+
.remote(env_fn())
|
| 369 |
+
)
|
| 370 |
+
super().__init__(env_fn)
|
| 371 |
+
|
| 372 |
+
def get_env_attr(self, key: str) -> Any:
|
| 373 |
+
return ray.get(self.env.get_env_attr.remote(key))
|
| 374 |
+
|
| 375 |
+
def set_env_attr(self, key: str, value: Any) -> None:
|
| 376 |
+
ray.get(self.env.set_env_attr.remote(key, value))
|
| 377 |
+
|
| 378 |
+
def reset(self) -> Any:
|
| 379 |
+
return ray.get(self.env.reset.remote())
|
| 380 |
+
|
| 381 |
+
@staticmethod
|
| 382 |
+
def wait( # type: ignore
|
| 383 |
+
workers: List["RayEnvWorker"], wait_num: int, timeout: Optional[float] = None
|
| 384 |
+
) -> List["RayEnvWorker"]:
|
| 385 |
+
results = [x.result for x in workers]
|
| 386 |
+
ready_results, _ = ray.wait(results, num_returns=wait_num, timeout=timeout)
|
| 387 |
+
return [workers[results.index(result)] for result in ready_results]
|
| 388 |
+
|
| 389 |
+
def send(self, action: Optional[np.ndarray]) -> None:
|
| 390 |
+
# self.action is actually a handle
|
| 391 |
+
if action is None:
|
| 392 |
+
self.result = self.env.reset.remote()
|
| 393 |
+
else:
|
| 394 |
+
self.result = self.env.step.remote(action)
|
| 395 |
+
|
| 396 |
+
def recv(
|
| 397 |
+
self,
|
| 398 |
+
) -> Union[Tuple[np.ndarray, np.ndarray, np.ndarray, np.ndarray], np.ndarray]:
|
| 399 |
+
return ray.get(self.result) # type: ignore
|
| 400 |
+
|
| 401 |
+
def seed(self, seed: Optional[int] = None) -> List[int]:
|
| 402 |
+
super().seed(seed)
|
| 403 |
+
return ray.get(self.env.seed.remote(seed))
|
| 404 |
+
|
| 405 |
+
def render(self, **kwargs: Any) -> Any:
|
| 406 |
+
return ray.get(self.env.render.remote(**kwargs))
|
| 407 |
+
|
| 408 |
+
def close_env(self) -> None:
|
| 409 |
+
ray.get(self.env.close.remote())
|
uvd/envs/evaluator/visualize_wrapper.py
ADDED
|
@@ -0,0 +1,271 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
from __future__ import annotations
|
| 2 |
+
|
| 3 |
+
from collections import OrderedDict
|
| 4 |
+
|
| 5 |
+
import gym
|
| 6 |
+
import numpy as np
|
| 7 |
+
|
| 8 |
+
import uvd.utils as U
|
| 9 |
+
from uvd.envs.evaluator.inference_wrapper import InferenceWrapper
|
| 10 |
+
|
| 11 |
+
__all__ = ["VisualizeWrapper"]
|
| 12 |
+
|
| 13 |
+
|
| 14 |
+
class VisualizeWrapper(gym.Wrapper):
|
| 15 |
+
def __init__(
|
| 16 |
+
self, env: InferenceWrapper, add_goal: bool = True, add_debug_text: bool = True
|
| 17 |
+
):
|
| 18 |
+
super().__init__(env=env)
|
| 19 |
+
self.env: InferenceWrapper = env
|
| 20 |
+
self.add_goal = add_goal
|
| 21 |
+
self.add_debug_text = add_debug_text
|
| 22 |
+
self._frames = []
|
| 23 |
+
self._rgb_milestones = None
|
| 24 |
+
self._current_goal_image = None
|
| 25 |
+
self._stepwise_extra_debug_texts = None
|
| 26 |
+
|
| 27 |
+
self._is_init_reset = True
|
| 28 |
+
self._recording = True
|
| 29 |
+
|
| 30 |
+
@property
|
| 31 |
+
def recording(self) -> bool:
|
| 32 |
+
return self._recording
|
| 33 |
+
|
| 34 |
+
@recording.setter
|
| 35 |
+
def recording(self, val: bool):
|
| 36 |
+
self._recording = val
|
| 37 |
+
|
| 38 |
+
@property
|
| 39 |
+
def frames(self):
|
| 40 |
+
return self._frames
|
| 41 |
+
|
| 42 |
+
def clear(self):
|
| 43 |
+
self._frames = []
|
| 44 |
+
|
| 45 |
+
@property
|
| 46 |
+
def video_length(self):
|
| 47 |
+
return len(self._frames)
|
| 48 |
+
|
| 49 |
+
def reset(self, **kwargs) -> np.ndarray:
|
| 50 |
+
obs = super().reset(**kwargs)
|
| 51 |
+
self.clear()
|
| 52 |
+
if self._is_init_reset:
|
| 53 |
+
self._is_init_reset = False
|
| 54 |
+
return obs
|
| 55 |
+
self.add_frame(obs, info=None)
|
| 56 |
+
return obs
|
| 57 |
+
|
| 58 |
+
def step(self, action: np.ndarray) -> tuple[np.ndarray, float, bool, dict]:
|
| 59 |
+
o, r, d, i = super().step(action)
|
| 60 |
+
self.add_frame(o, i)
|
| 61 |
+
self.current_embedding_distance = None
|
| 62 |
+
self.stepwise_extra_debug_texts = None
|
| 63 |
+
return o, r, d, i
|
| 64 |
+
|
| 65 |
+
def add_frame(
|
| 66 |
+
self, obs: np.ndarray | OrderedDict[str, np.ndarray], info: dict | None
|
| 67 |
+
):
|
| 68 |
+
if not self.recording:
|
| 69 |
+
return
|
| 70 |
+
if not isinstance(obs, np.ndarray):
|
| 71 |
+
assert isinstance(obs, (dict, OrderedDict)) and "rgb" in obs, obs
|
| 72 |
+
if self.env.cached_obs is not None:
|
| 73 |
+
rgb = obs["rgb"][-1, ...].copy()
|
| 74 |
+
else:
|
| 75 |
+
rgb = obs["rgb"].copy()
|
| 76 |
+
else:
|
| 77 |
+
rgb = obs.copy()
|
| 78 |
+
if self.add_goal:
|
| 79 |
+
rgb = np.concatenate([rgb, self.current_goal_image], axis=-2)
|
| 80 |
+
if self.add_debug_text:
|
| 81 |
+
if info is not None:
|
| 82 |
+
assert (
|
| 83 |
+
self.current_embedding_distance is not None
|
| 84 |
+
), f"should set before each step"
|
| 85 |
+
txts = dict(
|
| 86 |
+
complete_milestones=self.completed_goals,
|
| 87 |
+
num_milstones=self.env.num_milestones,
|
| 88 |
+
current_embedding_distance=float(self.current_embedding_distance),
|
| 89 |
+
)
|
| 90 |
+
try:
|
| 91 |
+
txts.update(
|
| 92 |
+
current_sub_task=info["current_sub_task"],
|
| 93 |
+
step=info["episode_length"],
|
| 94 |
+
completed_tasks=info["completed_tasks"],
|
| 95 |
+
# success=int(info["success"]),
|
| 96 |
+
)
|
| 97 |
+
if "distances_left" in info["rewards"]:
|
| 98 |
+
distances_left = info["rewards"]["distances_left"]
|
| 99 |
+
if isinstance(distances_left, dict):
|
| 100 |
+
for ele, dist in distances_left.items():
|
| 101 |
+
txts[f"low_dim_distances_left/{ele}"] = float(
|
| 102 |
+
np.sum(dist)
|
| 103 |
+
)
|
| 104 |
+
else:
|
| 105 |
+
txts["low_dim_distances_left"] = float(
|
| 106 |
+
np.sum(distances_left)
|
| 107 |
+
)
|
| 108 |
+
except KeyError:
|
| 109 |
+
pass
|
| 110 |
+
else:
|
| 111 |
+
try:
|
| 112 |
+
current_sub_task = getattr(self.env, "current_subtask")
|
| 113 |
+
txts = dict(current_sub_task=current_sub_task, step=0)
|
| 114 |
+
except AttributeError:
|
| 115 |
+
txts = dict(step=0)
|
| 116 |
+
if self.stepwise_extra_debug_texts is not None:
|
| 117 |
+
txts.update(self.stepwise_extra_debug_texts)
|
| 118 |
+
rgb = U.debug_texts_to_frame(frame=rgb, debug_text=txts)
|
| 119 |
+
self._frames.append(rgb)
|
| 120 |
+
|
| 121 |
+
@property
|
| 122 |
+
def current_goal_image(self):
|
| 123 |
+
assert self._current_goal_image is not None
|
| 124 |
+
if self.cur_milestone_idx is not None:
|
| 125 |
+
return self.rgb_milestones[self.cur_milestone_idx]
|
| 126 |
+
return self._current_goal_image
|
| 127 |
+
|
| 128 |
+
@current_goal_image.setter
|
| 129 |
+
def current_goal_image(self, image: np.ndarray | None):
|
| 130 |
+
if image is not None:
|
| 131 |
+
assert image.ndim == 3, image.shape
|
| 132 |
+
self._current_goal_image = image
|
| 133 |
+
|
| 134 |
+
@property
|
| 135 |
+
def rgb_milestones(self) -> np.ndarray:
|
| 136 |
+
assert self._rgb_milestones is not None, f"must set before using"
|
| 137 |
+
return self._rgb_milestones
|
| 138 |
+
|
| 139 |
+
@rgb_milestones.setter
|
| 140 |
+
def rgb_milestones(self, rgbs: np.ndarray | None):
|
| 141 |
+
"""Set outside."""
|
| 142 |
+
if rgbs is not None:
|
| 143 |
+
# num_goal, h, w, 3
|
| 144 |
+
assert (
|
| 145 |
+
rgbs.ndim == 4 and rgbs.shape[0] == self.milestones.shape[0]
|
| 146 |
+
), rgbs.shape
|
| 147 |
+
self._rgb_milestones = rgbs.copy()
|
| 148 |
+
self._current_goal_image = self._rgb_milestones[0]
|
| 149 |
+
else:
|
| 150 |
+
self._rgb_milestones = None
|
| 151 |
+
self._current_goal_image = None
|
| 152 |
+
|
| 153 |
+
@property
|
| 154 |
+
def stepwise_extra_debug_texts(self) -> dict | None:
|
| 155 |
+
return self._stepwise_extra_debug_texts
|
| 156 |
+
|
| 157 |
+
@stepwise_extra_debug_texts.setter
|
| 158 |
+
def stepwise_extra_debug_texts(self, texts: dict | None):
|
| 159 |
+
self._stepwise_extra_debug_texts = texts
|
| 160 |
+
|
| 161 |
+
@property
|
| 162 |
+
def milestones(self) -> np.ndarray:
|
| 163 |
+
return self.env.milestones
|
| 164 |
+
|
| 165 |
+
@milestones.setter
|
| 166 |
+
def milestones(self, milestones: np.ndarray | None):
|
| 167 |
+
self.env.milestones = milestones
|
| 168 |
+
|
| 169 |
+
@property
|
| 170 |
+
def milestone_embeddings(self) -> np.ndarray:
|
| 171 |
+
return self.env.milestone_embeddings
|
| 172 |
+
|
| 173 |
+
@milestone_embeddings.setter
|
| 174 |
+
def milestone_embeddings(self, embeddings: np.ndarray | None):
|
| 175 |
+
self.env.milestone_embeddings = embeddings
|
| 176 |
+
|
| 177 |
+
@property
|
| 178 |
+
def no_robot_milestones(self) -> np.ndarray:
|
| 179 |
+
return self.env.no_robot_milestones
|
| 180 |
+
|
| 181 |
+
@no_robot_milestones.setter
|
| 182 |
+
def no_robot_milestones(self, milestones: np.ndarray | None):
|
| 183 |
+
self.env.no_robot_milestones = milestones
|
| 184 |
+
|
| 185 |
+
@property
|
| 186 |
+
def current_milestone(self) -> np.ndarray:
|
| 187 |
+
return self.env.current_milestone
|
| 188 |
+
|
| 189 |
+
@property
|
| 190 |
+
def current_no_robot_milestone(self) -> np.ndarray | None:
|
| 191 |
+
return self.env.no_robot_milestones
|
| 192 |
+
|
| 193 |
+
@property
|
| 194 |
+
def current_embedding_distance(self) -> np.ndarray:
|
| 195 |
+
return self.env.current_embedding_distance
|
| 196 |
+
|
| 197 |
+
@current_embedding_distance.setter
|
| 198 |
+
def current_embedding_distance(self, dist: float | None):
|
| 199 |
+
self.env.current_embedding_distance = dist
|
| 200 |
+
|
| 201 |
+
@property
|
| 202 |
+
def current_obs_embedding(self) -> np.ndarray:
|
| 203 |
+
return self.env.current_obs_embedding
|
| 204 |
+
|
| 205 |
+
@current_obs_embedding.setter
|
| 206 |
+
def current_obs_embedding(self, embedding: np.ndarray):
|
| 207 |
+
before_set_achieved = self.completed_goals
|
| 208 |
+
self.env.current_obs_embedding = embedding
|
| 209 |
+
after_set_achieved = self.completed_goals
|
| 210 |
+
# if after_set_achieved > before_set_achieved and self.recording:
|
| 211 |
+
if self.recording:
|
| 212 |
+
# switch rgb current goal image
|
| 213 |
+
self.current_goal_image = self.rgb_milestones[
|
| 214 |
+
min(after_set_achieved, len(self.rgb_milestones) - 1)
|
| 215 |
+
]
|
| 216 |
+
|
| 217 |
+
@property
|
| 218 |
+
def cur_milestone_idx(self) -> int | None:
|
| 219 |
+
return self.env.cur_milestone_idx
|
| 220 |
+
|
| 221 |
+
@cur_milestone_idx.setter
|
| 222 |
+
def cur_milestone_idx(self, idx: int | None):
|
| 223 |
+
self.env.cur_milestone_idx = idx
|
| 224 |
+
|
| 225 |
+
@property
|
| 226 |
+
def completed_goals(self) -> int:
|
| 227 |
+
return self.env.completed_goals
|
| 228 |
+
|
| 229 |
+
@property
|
| 230 |
+
def reset_states(self) -> dict:
|
| 231 |
+
return self.env.reset_states
|
| 232 |
+
|
| 233 |
+
@reset_states.setter
|
| 234 |
+
def reset_states(self, state: dict):
|
| 235 |
+
self.env.reset_states = state
|
| 236 |
+
|
| 237 |
+
@property
|
| 238 |
+
def task_name(self) -> str:
|
| 239 |
+
return self.env.task_name
|
| 240 |
+
|
| 241 |
+
@property
|
| 242 |
+
def metrics(self) -> dict:
|
| 243 |
+
return self.env.metrics
|
| 244 |
+
|
| 245 |
+
@property
|
| 246 |
+
def milestone_distances(self) -> np.ndarray | None:
|
| 247 |
+
return self.env.milestone_distances
|
| 248 |
+
|
| 249 |
+
@milestone_distances.setter
|
| 250 |
+
def milestone_distances(self, milestone_distances: np.ndarray | None):
|
| 251 |
+
self.env.milestone_distances = milestone_distances
|
| 252 |
+
|
| 253 |
+
@property
|
| 254 |
+
def current_no_robot_frame(self) -> np.ndarray:
|
| 255 |
+
return self.env.current_no_robot_frame
|
| 256 |
+
|
| 257 |
+
@property
|
| 258 |
+
def current_rgb_frame(self) -> np.ndarray:
|
| 259 |
+
return self.env.current_rgb_frame
|
| 260 |
+
|
| 261 |
+
@property
|
| 262 |
+
def num_milestones(self):
|
| 263 |
+
return self.env.num_milestones
|
| 264 |
+
|
| 265 |
+
@property
|
| 266 |
+
def milestone_indices(self):
|
| 267 |
+
return self.env.milestone_indices
|
| 268 |
+
|
| 269 |
+
@milestone_indices.setter
|
| 270 |
+
def milestone_indices(self, milestone_indices: np.ndarray | None):
|
| 271 |
+
self.env.milestone_indices = milestone_indices
|
uvd/envs/franka_kitchen/__init__.py
ADDED
|
@@ -0,0 +1,20 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
import os
|
| 2 |
+
import sys
|
| 3 |
+
|
| 4 |
+
os.environ[
|
| 5 |
+
"LD_LIBRARY_PATH"
|
| 6 |
+
] = f":{os.environ['HOME']}/.mujoco/mujoco210/bin:/usr/lib/nvidia"
|
| 7 |
+
|
| 8 |
+
# workaround to import adept envs
|
| 9 |
+
ADEPT_DIR = os.path.join(
|
| 10 |
+
os.path.dirname(__file__), "relay-policy-learning", "adept_envs"
|
| 11 |
+
)
|
| 12 |
+
assert os.path.exists(ADEPT_DIR), ADEPT_DIR
|
| 13 |
+
sys.path.append(ADEPT_DIR)
|
| 14 |
+
|
| 15 |
+
from .franka_kitchen_base import *
|
| 16 |
+
from .franka_kitchen_constants import *
|
| 17 |
+
|
| 18 |
+
import adept_envs.mujoco_env
|
| 19 |
+
|
| 20 |
+
adept_envs.mujoco_env.USE_DM_CONTROL = USE_DM_CONTROL
|
uvd/envs/franka_kitchen/franka_kitchen_base.py
ADDED
|
@@ -0,0 +1,446 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""Environments using kitchen and Franka robot."""
|
| 2 |
+
from __future__ import annotations
|
| 3 |
+
|
| 4 |
+
import copy
|
| 5 |
+
from collections import OrderedDict
|
| 6 |
+
|
| 7 |
+
import gym
|
| 8 |
+
import numpy as np
|
| 9 |
+
|
| 10 |
+
|
| 11 |
+
from uvd.envs.franka_kitchen.franka_kitchen_constants import *
|
| 12 |
+
from mujoco_py import MjRenderContextOffscreen
|
| 13 |
+
|
| 14 |
+
from adept_envs.franka.kitchen_multitask_v0 import KitchenV0
|
| 15 |
+
from adept_envs.simulation.renderer import RenderMode, DMRenderer
|
| 16 |
+
|
| 17 |
+
__all__ = ["KitchenBase"]
|
| 18 |
+
|
| 19 |
+
|
| 20 |
+
class KitchenBase(KitchenV0):
|
| 21 |
+
ALL_TASKS = FRANKA_KITCHEN_ALL_TASKS
|
| 22 |
+
REMOVE_TASKS_WHEN_COMPLETE = True
|
| 23 |
+
TERMINATE_ON_TASK_COMPLETE = True
|
| 24 |
+
TERMINATE_ON_WRONG_COMPLETE = False
|
| 25 |
+
COMPLETE_IN_ANY_ORDER = False
|
| 26 |
+
|
| 27 |
+
frame_height = 224
|
| 28 |
+
frame_width = 224
|
| 29 |
+
|
| 30 |
+
def __init__(
|
| 31 |
+
self,
|
| 32 |
+
*,
|
| 33 |
+
task_elements: list | None = None,
|
| 34 |
+
reward_config: dict | None = FRANKA_KITCHEN_REWARD_CONFIG,
|
| 35 |
+
goal_masking: bool = True,
|
| 36 |
+
obs_keys: str | tuple[str, ...] = ("rgb", "proprio"),
|
| 37 |
+
max_horizon: int = 1000,
|
| 38 |
+
frame_height: int = 224,
|
| 39 |
+
frame_width: int = 224,
|
| 40 |
+
robot_params: dict | None = None,
|
| 41 |
+
frame_skip: int = 40,
|
| 42 |
+
gpu_id: int = -1,
|
| 43 |
+
):
|
| 44 |
+
task_elements = (
|
| 45 |
+
list(task_elements) if task_elements is not None else list(self.ALL_TASKS)
|
| 46 |
+
)
|
| 47 |
+
self._task_elements = [
|
| 48 |
+
e if isinstance(e, str) else IDX_TO_ELEMENT[e] for e in task_elements
|
| 49 |
+
]
|
| 50 |
+
self.goal_masking = goal_masking
|
| 51 |
+
# workaround for vector env preventing render
|
| 52 |
+
self._init_step = True
|
| 53 |
+
super(KitchenBase, self).__init__(
|
| 54 |
+
robot_params=robot_params or {}, frame_skip=frame_skip
|
| 55 |
+
)
|
| 56 |
+
self._num_goals_achieved = 0
|
| 57 |
+
self.tasks_to_complete = copy.deepcopy(self._task_elements)
|
| 58 |
+
self.reward_config = reward_config
|
| 59 |
+
self.obs_keys = (
|
| 60 |
+
tuple(obs_keys) if not isinstance(obs_keys, str) else (obs_keys,)
|
| 61 |
+
)
|
| 62 |
+
dict_obs_space = OrderedDict()
|
| 63 |
+
for key in obs_keys:
|
| 64 |
+
if key == "rgb":
|
| 65 |
+
dict_obs_space[key] = gym.spaces.Box(
|
| 66 |
+
low=0,
|
| 67 |
+
high=255,
|
| 68 |
+
shape=(frame_height, frame_width, 3),
|
| 69 |
+
dtype=np.uint8,
|
| 70 |
+
)
|
| 71 |
+
elif key == "proprio":
|
| 72 |
+
dict_obs_space[key] = gym.spaces.Box(
|
| 73 |
+
low=-np.inf,
|
| 74 |
+
high=np.inf,
|
| 75 |
+
shape=(self.obs_dict["qp"].shape[0],),
|
| 76 |
+
dtype=np.float32,
|
| 77 |
+
)
|
| 78 |
+
else:
|
| 79 |
+
raise NotImplementedError(key)
|
| 80 |
+
self.observation_space = gym.spaces.Dict(dict_obs_space)
|
| 81 |
+
self.max_horizon = max_horizon
|
| 82 |
+
self.frame_height = frame_height
|
| 83 |
+
self.frame_width = frame_width
|
| 84 |
+
self.gpu_id = gpu_id
|
| 85 |
+
self._renderer_setup = False
|
| 86 |
+
self.env_info = None
|
| 87 |
+
self.episode_length = 0
|
| 88 |
+
|
| 89 |
+
@property
|
| 90 |
+
def task_elements(self) -> list[str]:
|
| 91 |
+
return self._task_elements
|
| 92 |
+
|
| 93 |
+
@property
|
| 94 |
+
def task_name(self) -> str:
|
| 95 |
+
return "-".join(self.task_elements).replace(" ", "_")
|
| 96 |
+
|
| 97 |
+
@property
|
| 98 |
+
def current_subtask(self) -> str:
|
| 99 |
+
if self.num_goals_achieved == len(self.task_elements):
|
| 100 |
+
return "null" # success
|
| 101 |
+
return self.task_elements[self.num_goals_achieved]
|
| 102 |
+
|
| 103 |
+
@property
|
| 104 |
+
def metrics(self):
|
| 105 |
+
"""Accessible metrics for vector env."""
|
| 106 |
+
assert self.env_info is not None, f"try to call metrics without stepping"
|
| 107 |
+
metrics = dict(
|
| 108 |
+
episode_length=self.episode_length,
|
| 109 |
+
success=int(self.env_info["success"]),
|
| 110 |
+
completed_tasks=self.env_info["completed_tasks"],
|
| 111 |
+
)
|
| 112 |
+
return metrics
|
| 113 |
+
|
| 114 |
+
def _get_task_goal(
|
| 115 |
+
self, task: list[str] | None = None, actually_return_goal: bool = True
|
| 116 |
+
) -> np.ndarray:
|
| 117 |
+
if task is None:
|
| 118 |
+
task = self.task_elements
|
| 119 |
+
new_goal = np.zeros_like(self.goal)
|
| 120 |
+
if self.goal_masking and not actually_return_goal:
|
| 121 |
+
return new_goal
|
| 122 |
+
for element in task:
|
| 123 |
+
element_idx = OBS_ELEMENT_INDICES[element]
|
| 124 |
+
element_goal = OBS_ELEMENT_GOALS[element]
|
| 125 |
+
new_goal[element_idx] = element_goal
|
| 126 |
+
return new_goal
|
| 127 |
+
|
| 128 |
+
def reset(
|
| 129 |
+
self,
|
| 130 |
+
init_qpos: np.ndarray | None = None,
|
| 131 |
+
init_qvel: np.ndarray | None = None,
|
| 132 |
+
task_elements: list[str] | None = None,
|
| 133 |
+
seed: int | None = None,
|
| 134 |
+
**kwargs,
|
| 135 |
+
) -> OrderedDict[str, np.ndarray]:
|
| 136 |
+
"""Return visual obs."""
|
| 137 |
+
if seed is not None:
|
| 138 |
+
self.seed(seed)
|
| 139 |
+
# set initial states if not None
|
| 140 |
+
self.reset_states = dict(
|
| 141 |
+
init_qpos=init_qpos, init_qvel=init_qvel, task_elements=task_elements
|
| 142 |
+
)
|
| 143 |
+
assert len(self.task_elements) > 0
|
| 144 |
+
self.tasks_to_complete = copy.deepcopy(self.task_elements)
|
| 145 |
+
self._num_goals_achieved = 0
|
| 146 |
+
|
| 147 |
+
self.env_info = None
|
| 148 |
+
self.episode_length = 0
|
| 149 |
+
|
| 150 |
+
super().reset()
|
| 151 |
+
return self.get_dict_obs(**kwargs)
|
| 152 |
+
|
| 153 |
+
def get_dict_obs(self, **kwargs):
|
| 154 |
+
"""Observation sensor for dict observation space."""
|
| 155 |
+
obs = OrderedDict()
|
| 156 |
+
for k in self.obs_keys:
|
| 157 |
+
if k == "rgb":
|
| 158 |
+
if "frame_height" in kwargs:
|
| 159 |
+
self.frame_height = kwargs.pop("frame_height")
|
| 160 |
+
if "frame_width" in kwargs:
|
| 161 |
+
self.frame_width = kwargs.pop("frame_width")
|
| 162 |
+
obs[k] = self.render(
|
| 163 |
+
mode="rgb_array", height=self.frame_height, width=self.frame_width
|
| 164 |
+
)
|
| 165 |
+
elif k == "proprio":
|
| 166 |
+
obs[k] = self.obs_dict["qp"]
|
| 167 |
+
else:
|
| 168 |
+
raise NotImplementedError(k)
|
| 169 |
+
return obs
|
| 170 |
+
|
| 171 |
+
@property
|
| 172 |
+
def reset_states(self) -> dict:
|
| 173 |
+
return dict(
|
| 174 |
+
init_qpos=self.init_qpos,
|
| 175 |
+
init_qvel=self.init_qvel,
|
| 176 |
+
task_elements=self.task_elements,
|
| 177 |
+
)
|
| 178 |
+
|
| 179 |
+
@reset_states.setter
|
| 180 |
+
def reset_states(self, state: dict | None):
|
| 181 |
+
"""Set state in vector env for convenience."""
|
| 182 |
+
if state.get("init_qpos", None) is not None:
|
| 183 |
+
assert (
|
| 184 |
+
state.get("init_qvel", None) is not None
|
| 185 |
+
), "qpos and qvel set together"
|
| 186 |
+
qpos = self._squeeze_state_vec(state["init_qpos"].copy())
|
| 187 |
+
qvel = self._squeeze_state_vec(state["init_qvel"].copy())
|
| 188 |
+
self.set_state(qpos=qpos, qvel=qvel)
|
| 189 |
+
self.init_qpos = self.data.qpos.ravel().copy()
|
| 190 |
+
self.init_qvel = self.data.qvel.ravel().copy()
|
| 191 |
+
if state.get("task_elements", None) is not None:
|
| 192 |
+
self._task_elements = [
|
| 193 |
+
e if isinstance(e, str) else IDX_TO_ELEMENT[e]
|
| 194 |
+
for e in self._squeeze_state_vec(
|
| 195 |
+
state["task_elements"].copy(), to_list=True
|
| 196 |
+
)
|
| 197 |
+
]
|
| 198 |
+
|
| 199 |
+
def set_state(self, qpos: np.ndarray, qvel: np.ndarray):
|
| 200 |
+
if isinstance(self.sim_robot.renderer, DMRenderer):
|
| 201 |
+
assert qpos.shape == (self.model.nq,) and qvel.shape == (self.model.nv,)
|
| 202 |
+
state = np.concatenate([qpos, qvel])
|
| 203 |
+
self.sim.set_state(state)
|
| 204 |
+
self.sim.forward()
|
| 205 |
+
else:
|
| 206 |
+
super().set_state(qpos=qpos, qvel=qvel)
|
| 207 |
+
|
| 208 |
+
def state_vector(self) -> np.ndarray:
|
| 209 |
+
if isinstance(self.sim_robot.renderer, DMRenderer):
|
| 210 |
+
return self.sim.get_state()
|
| 211 |
+
else:
|
| 212 |
+
return super().state_vector()
|
| 213 |
+
|
| 214 |
+
@staticmethod
|
| 215 |
+
def _squeeze_state_vec(vec: np.ndarray, to_list: bool = False) -> np.ndarray | list:
|
| 216 |
+
if not isinstance(vec, np.ndarray):
|
| 217 |
+
return vec
|
| 218 |
+
if vec.ndim != 1:
|
| 219 |
+
vec = vec.reshape((vec.shape[-1],))
|
| 220 |
+
if to_list:
|
| 221 |
+
return vec.tolist()
|
| 222 |
+
return vec
|
| 223 |
+
|
| 224 |
+
def _get_reward_n_score(self, obs_dict: dict) -> tuple[dict, float]:
|
| 225 |
+
"""Score here means whether complete a new goal this step."""
|
| 226 |
+
reward_dict = {"true_reward": 0.0}
|
| 227 |
+
|
| 228 |
+
next_q_obs = obs_dict["qp"]
|
| 229 |
+
next_obj_obs = obs_dict["obj_qp"]
|
| 230 |
+
next_goal = self._get_task_goal(
|
| 231 |
+
task=self.task_elements, actually_return_goal=True
|
| 232 |
+
) # obs_dict['goal']
|
| 233 |
+
idx_offset = len(next_q_obs)
|
| 234 |
+
completions = []
|
| 235 |
+
all_completed_so_far = True
|
| 236 |
+
distances_dict = {}
|
| 237 |
+
# # for element in self.tasks_to_complete:
|
| 238 |
+
# make metrics consistent for sync in DDP rollout
|
| 239 |
+
for i, element in enumerate(self.task_elements):
|
| 240 |
+
element_idx = OBS_ELEMENT_INDICES[element]
|
| 241 |
+
# distance = np.linalg.norm(
|
| 242 |
+
# next_obj_obs[..., element_idx - idx_offset] - next_goal[element_idx]
|
| 243 |
+
# )
|
| 244 |
+
distance = abs(
|
| 245 |
+
next_obj_obs[..., element_idx - idx_offset] - next_goal[element_idx]
|
| 246 |
+
) # keep state dims
|
| 247 |
+
distances_dict[element] = distance
|
| 248 |
+
# complete = distance < BONUS_THRESH
|
| 249 |
+
complete = distance < OBS_ELEMENT_THRESH[element] # diff criteria
|
| 250 |
+
if isinstance(complete, np.ndarray):
|
| 251 |
+
complete = np.all(complete)
|
| 252 |
+
|
| 253 |
+
if (
|
| 254 |
+
self.REMOVE_TASKS_WHEN_COMPLETE
|
| 255 |
+
and element not in self.tasks_to_complete
|
| 256 |
+
):
|
| 257 |
+
# already achieved task(s), and
|
| 258 |
+
# edge case for knobs that distances increasing later though rendered states not changed
|
| 259 |
+
condition = True
|
| 260 |
+
else:
|
| 261 |
+
condition = (
|
| 262 |
+
complete and all_completed_so_far
|
| 263 |
+
if not self.COMPLETE_IN_ANY_ORDER
|
| 264 |
+
else complete
|
| 265 |
+
)
|
| 266 |
+
if condition:
|
| 267 |
+
completions.append(element)
|
| 268 |
+
all_completed_so_far = all_completed_so_far and condition
|
| 269 |
+
prev_tasks_to_complete = list(self.tasks_to_complete)
|
| 270 |
+
if self.REMOVE_TASKS_WHEN_COMPLETE:
|
| 271 |
+
[
|
| 272 |
+
self.tasks_to_complete.remove(element)
|
| 273 |
+
for element in completions
|
| 274 |
+
if element in self.tasks_to_complete
|
| 275 |
+
]
|
| 276 |
+
num_goals_achieved = len(self.task_elements) - len(self.tasks_to_complete)
|
| 277 |
+
else:
|
| 278 |
+
num_goals_achieved = len(completions)
|
| 279 |
+
|
| 280 |
+
complete_new_goal_this_step = False
|
| 281 |
+
if num_goals_achieved > self._num_goals_achieved:
|
| 282 |
+
self._num_goals_achieved += 1
|
| 283 |
+
# edge case for knobs
|
| 284 |
+
if num_goals_achieved != self._num_goals_achieved:
|
| 285 |
+
self._num_goals_achieved = len(self.task_elements) - len(
|
| 286 |
+
self.tasks_to_complete
|
| 287 |
+
)
|
| 288 |
+
|
| 289 |
+
reward_dict["true_reward"] += self.reward_config["progress"]
|
| 290 |
+
complete_new_goal_this_step = True
|
| 291 |
+
|
| 292 |
+
if self.num_goals_achieved == len(self.task_elements):
|
| 293 |
+
reward_dict["true_reward"] += self.reward_config["terminal"]
|
| 294 |
+
|
| 295 |
+
reward_dict["distances_left"] = distances_dict
|
| 296 |
+
score = int(complete_new_goal_this_step)
|
| 297 |
+
return reward_dict, score
|
| 298 |
+
|
| 299 |
+
@property
|
| 300 |
+
def num_goals_achieved(self):
|
| 301 |
+
if self.REMOVE_TASKS_WHEN_COMPLETE:
|
| 302 |
+
assert self._num_goals_achieved == len(self.task_elements) - len(
|
| 303 |
+
self.tasks_to_complete
|
| 304 |
+
), f"{self._num_goals_achieved}, {self.task_elements}, {self.tasks_to_complete}"
|
| 305 |
+
return self._num_goals_achieved
|
| 306 |
+
|
| 307 |
+
def step(
|
| 308 |
+
self, action: np.ndarray, **kwargs
|
| 309 |
+
) -> tuple[OrderedDict[str, np.ndarray], float, bool, dict]:
|
| 310 |
+
action = np.clip(action, -1.0, 1.0)
|
| 311 |
+
if not self.initializing:
|
| 312 |
+
action = self.act_mid + action * self.act_amp # mean center and scale
|
| 313 |
+
else:
|
| 314 |
+
self.goal = self._get_task_goal() # update goal if init
|
| 315 |
+
self.robot.step(self, action, step_duration=self.skip * self.model.opt.timestep)
|
| 316 |
+
|
| 317 |
+
low_dim_obs = self._get_obs()
|
| 318 |
+
# workaround for vector env
|
| 319 |
+
if self._init_step:
|
| 320 |
+
self._init_step = False
|
| 321 |
+
return low_dim_obs, 0.0, False, {}
|
| 322 |
+
|
| 323 |
+
self.episode_length += 1
|
| 324 |
+
|
| 325 |
+
reward_dict, score = self._get_reward_n_score(self.obs_dict)
|
| 326 |
+
success = (
|
| 327 |
+
not self.tasks_to_complete
|
| 328 |
+
if self.REMOVE_TASKS_WHEN_COMPLETE
|
| 329 |
+
else self.num_goals_achieved == len(self.tasks_to_complete)
|
| 330 |
+
)
|
| 331 |
+
|
| 332 |
+
env_info = {
|
| 333 |
+
"time": self.obs_dict["t"],
|
| 334 |
+
"obs_dict": self.obs_dict,
|
| 335 |
+
"rewards": reward_dict,
|
| 336 |
+
"score": score,
|
| 337 |
+
"success": success,
|
| 338 |
+
# self.task_elements[:self.num_goals_achieved - 1]
|
| 339 |
+
"completed_tasks": self.num_goals_achieved,
|
| 340 |
+
"episode_length": self.episode_length,
|
| 341 |
+
"current_sub_task": self.current_subtask,
|
| 342 |
+
}
|
| 343 |
+
done = self.episode_length >= self.max_horizon
|
| 344 |
+
if self.TERMINATE_ON_TASK_COMPLETE and success:
|
| 345 |
+
done = True
|
| 346 |
+
if self.TERMINATE_ON_WRONG_COMPLETE:
|
| 347 |
+
all_goal = self._get_task_goal(task=self.ALL_TASKS)
|
| 348 |
+
for wrong_task in list(set(self.ALL_TASKS) - set(self.task_elements)):
|
| 349 |
+
element_idx = OBS_ELEMENT_INDICES[wrong_task]
|
| 350 |
+
distance = np.linalg.norm(
|
| 351 |
+
low_dim_obs[..., element_idx] - all_goal[element_idx]
|
| 352 |
+
)
|
| 353 |
+
complete = distance < BONUS_THRESH
|
| 354 |
+
if complete:
|
| 355 |
+
done = True
|
| 356 |
+
break
|
| 357 |
+
|
| 358 |
+
obs = self.get_dict_obs(**kwargs)
|
| 359 |
+
env_info["done"] = done
|
| 360 |
+
self.env_info = env_info
|
| 361 |
+
return obs, reward_dict["true_reward"], done, env_info
|
| 362 |
+
|
| 363 |
+
def close(self):
|
| 364 |
+
super().close()
|
| 365 |
+
self.close_env()
|
| 366 |
+
|
| 367 |
+
def render(
|
| 368 |
+
self,
|
| 369 |
+
mode: str = "human",
|
| 370 |
+
height: int | None = None,
|
| 371 |
+
width: int | None = None,
|
| 372 |
+
camera_id: int = -1,
|
| 373 |
+
set_robot_alpha: float | None = None,
|
| 374 |
+
) -> np.ndarray | None:
|
| 375 |
+
height = height or self.frame_height
|
| 376 |
+
width = width or self.frame_width
|
| 377 |
+
if not self._renderer_setup:
|
| 378 |
+
self.set_gpu_id(gpu_id=self.gpu_id)
|
| 379 |
+
camera_settings = dict(
|
| 380 |
+
distance=2.2, lookat=[-0.2, 0.5, 2.0], azimuth=70, elevation=-35
|
| 381 |
+
)
|
| 382 |
+
self.sim_robot.renderer._camera_settings = camera_settings
|
| 383 |
+
alpha_map = None
|
| 384 |
+
if set_robot_alpha is not None:
|
| 385 |
+
alpha_map = self.set_robot_alpha(alpha=set_robot_alpha)
|
| 386 |
+
|
| 387 |
+
height = height or self.frame_height
|
| 388 |
+
width = width or self.frame_width
|
| 389 |
+
if mode in ["rgb_array", "rgb"]:
|
| 390 |
+
frame = self.sim_robot.renderer.render_offscreen(
|
| 391 |
+
width=width, height=height, mode=RenderMode.RGB, camera_id=camera_id
|
| 392 |
+
)
|
| 393 |
+
elif mode == "human":
|
| 394 |
+
frame = self.sim_robot.renderer.render_to_window()
|
| 395 |
+
else:
|
| 396 |
+
frame = super(KitchenV0, self).render(
|
| 397 |
+
mode=mode, height=height, width=width, camera_id=camera_id
|
| 398 |
+
)
|
| 399 |
+
if set_robot_alpha is not None:
|
| 400 |
+
self.unset_robot_alpha(alpha_map)
|
| 401 |
+
return frame
|
| 402 |
+
|
| 403 |
+
def set_robot_alpha(self, alpha: float) -> dict[int, float]:
|
| 404 |
+
alpha_map = dict()
|
| 405 |
+
if not USE_DM_CONTROL:
|
| 406 |
+
for name in self.sim.model.site_names:
|
| 407 |
+
if "end_effector" in name:
|
| 408 |
+
orig_val = self.sim.model.site_rgba[
|
| 409 |
+
self.sim.model.site_name2id(name)
|
| 410 |
+
][-1]
|
| 411 |
+
alpha_map[name] = max(orig_val, alpha_map.get(name, -1))
|
| 412 |
+
self.sim.model.site_rgba[self.sim.model.site_name2id(name)][
|
| 413 |
+
-1
|
| 414 |
+
] = alpha
|
| 415 |
+
for dof_id in self.sim.model.dof_jntid:
|
| 416 |
+
alpha_map[dof_id] = max(
|
| 417 |
+
self.sim.model.geom_rgba[dof_id][-1], alpha_map.get(dof_id, -1)
|
| 418 |
+
)
|
| 419 |
+
self.sim.model.geom_rgba[dof_id][-1] = alpha
|
| 420 |
+
for dof_id in self.sim.model.dof_bodyid:
|
| 421 |
+
alpha_map[dof_id] = max(
|
| 422 |
+
self.sim.model.geom_rgba[dof_id][-1], alpha_map.get(dof_id, -1)
|
| 423 |
+
)
|
| 424 |
+
self.sim.model.geom_rgba[dof_id][-1] = alpha
|
| 425 |
+
return alpha_map
|
| 426 |
+
|
| 427 |
+
def unset_robot_alpha(self, alpha_map: dict[int, float]):
|
| 428 |
+
for dof_id, alpha in alpha_map.items():
|
| 429 |
+
if isinstance(dof_id, str) and "end_effector" in dof_id:
|
| 430 |
+
self.sim.model.site_rgba[self.sim.model.site_name2id(dof_id)][
|
| 431 |
+
-1
|
| 432 |
+
] = alpha
|
| 433 |
+
else:
|
| 434 |
+
self.sim.model.geom_rgba[dof_id][-1] = alpha
|
| 435 |
+
|
| 436 |
+
@property
|
| 437 |
+
def current_no_robot_frame(self) -> np.ndarray:
|
| 438 |
+
return self.render(mode="rgb_array", set_robot_alpha=0.0)
|
| 439 |
+
|
| 440 |
+
def set_gpu_id(self, gpu_id: int):
|
| 441 |
+
self._renderer_setup = True
|
| 442 |
+
# if not USE_DM_CONTROL:
|
| 443 |
+
if not isinstance(self.sim_robot.renderer, DMRenderer):
|
| 444 |
+
self.sim_robot.renderer._offscreen_renderer = MjRenderContextOffscreen(
|
| 445 |
+
self.sim_robot.renderer._sim, device_id=gpu_id
|
| 446 |
+
)
|
uvd/envs/franka_kitchen/franka_kitchen_constants.py
ADDED
|
@@ -0,0 +1,62 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
import numpy as np
|
| 2 |
+
|
| 3 |
+
USE_DM_CONTROL = False
|
| 4 |
+
BONUS_THRESH = 0.3
|
| 5 |
+
|
| 6 |
+
FRANKA_KITCHEN_ALL_TASKS = [
|
| 7 |
+
"bottom burner",
|
| 8 |
+
"top burner",
|
| 9 |
+
"light switch",
|
| 10 |
+
"slide cabinet",
|
| 11 |
+
"hinge cabinet",
|
| 12 |
+
"microwave",
|
| 13 |
+
"kettle",
|
| 14 |
+
]
|
| 15 |
+
|
| 16 |
+
OBS_ELEMENT_INDICES = {
|
| 17 |
+
# rotation & opening for bottom left burner
|
| 18 |
+
"bottom burner": np.array([11, 12]),
|
| 19 |
+
# rotation & opening for top left burner
|
| 20 |
+
"top burner": np.array([15, 16]),
|
| 21 |
+
# joint angle & opening
|
| 22 |
+
"light switch": np.array([17, 18]),
|
| 23 |
+
# Translation of the slide cabinet joint
|
| 24 |
+
"slide cabinet": np.array([19]),
|
| 25 |
+
# Rotation of the joint in the (left, right) hinge cabinet
|
| 26 |
+
"hinge cabinet": np.array([20, 21]),
|
| 27 |
+
# Rotation of the joint in the microwave door
|
| 28 |
+
"microwave": np.array([22]),
|
| 29 |
+
# x, y, z, qx, qy, qz, qw
|
| 30 |
+
"kettle": np.array([23, 24, 25, 26, 27, 28, 29]),
|
| 31 |
+
}
|
| 32 |
+
|
| 33 |
+
OBS_ELEMENT_GOALS = {
|
| 34 |
+
"bottom burner": np.array([-0.88, -0.01]),
|
| 35 |
+
"top burner": np.array([-0.92, -0.01]),
|
| 36 |
+
"light switch": np.array([-0.69, -0.05]),
|
| 37 |
+
"slide cabinet": np.array([0.37]),
|
| 38 |
+
# right hinge, left should be (-1.45, 0)
|
| 39 |
+
"hinge cabinet": np.array([0.0, 1.45]),
|
| 40 |
+
"microwave": np.array([-0.75]),
|
| 41 |
+
"kettle": np.array([-0.23, 0.75, 1.62, 0.99, 0.0, 0.0, -0.06]),
|
| 42 |
+
}
|
| 43 |
+
|
| 44 |
+
OBS_ELEMENT_THRESH = {
|
| 45 |
+
"bottom burner": 0.31,
|
| 46 |
+
"top burner": 0.31,
|
| 47 |
+
"light switch": 0.3, # 0.1
|
| 48 |
+
"slide cabinet": 0.2,
|
| 49 |
+
"hinge cabinet": 0.2,
|
| 50 |
+
"microwave": 0.2,
|
| 51 |
+
# "kettle": np.array([0.1, 0.1, 0.1, 0.2, 0.2, 0.2, 0.2]),
|
| 52 |
+
"kettle": np.array([0.1, 0.1, 0.1, 0.2, 0.2, 0.2, 0.3]),
|
| 53 |
+
}
|
| 54 |
+
|
| 55 |
+
ELEMENT_TO_IDX = {k: i for i, k in enumerate(OBS_ELEMENT_GOALS.keys())}
|
| 56 |
+
IDX_TO_ELEMENT = {i: k for k, i in ELEMENT_TO_IDX.items()}
|
| 57 |
+
|
| 58 |
+
FRANKA_KITCHEN_REWARD_CONFIG = {
|
| 59 |
+
"progress": 0.0, # complete sub-goal
|
| 60 |
+
"terminal": 10.0, # complete all goals
|
| 61 |
+
"intrinsic_weight": 1.0, # intrinsic reward if using embedding-dist-diff reward
|
| 62 |
+
}
|
uvd/envs/franka_kitchen/relay-policy-learning/adept_envs/.pylintrc
ADDED
|
@@ -0,0 +1,433 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
[MASTER]
|
| 2 |
+
|
| 3 |
+
# A comma-separated list of package or module names from where C extensions may
|
| 4 |
+
# be loaded. Extensions are loading into the active Python interpreter and may
|
| 5 |
+
# run arbitrary code.
|
| 6 |
+
extension-pkg-whitelist=
|
| 7 |
+
|
| 8 |
+
# Add files or directories to the blacklist. They should be base names, not
|
| 9 |
+
# paths.
|
| 10 |
+
ignore=CVS
|
| 11 |
+
|
| 12 |
+
# Add files or directories matching the regex patterns to the blacklist. The
|
| 13 |
+
# regex matches against base names, not paths.
|
| 14 |
+
ignore-patterns=
|
| 15 |
+
|
| 16 |
+
# Python code to execute, usually for sys.path manipulation such as
|
| 17 |
+
# pygtk.require().
|
| 18 |
+
#init-hook=
|
| 19 |
+
|
| 20 |
+
# Use multiple processes to speed up Pylint. Specifying 0 will auto-detect the
|
| 21 |
+
# number of processors available to use.
|
| 22 |
+
jobs=1
|
| 23 |
+
|
| 24 |
+
# Control the amount of potential inferred values when inferring a single
|
| 25 |
+
# object. This can help the performance when dealing with large functions or
|
| 26 |
+
# complex, nested conditions.
|
| 27 |
+
limit-inference-results=100
|
| 28 |
+
|
| 29 |
+
# List of plugins (as comma separated values of python modules names) to load,
|
| 30 |
+
# usually to register additional checkers.
|
| 31 |
+
load-plugins=
|
| 32 |
+
|
| 33 |
+
# Pickle collected data for later comparisons.
|
| 34 |
+
persistent=yes
|
| 35 |
+
|
| 36 |
+
# Specify a configuration file.
|
| 37 |
+
#rcfile=
|
| 38 |
+
|
| 39 |
+
# When enabled, pylint would attempt to guess common misconfiguration and emit
|
| 40 |
+
# user-friendly hints instead of false-positive error messages.
|
| 41 |
+
suggestion-mode=yes
|
| 42 |
+
|
| 43 |
+
# Allow loading of arbitrary C extensions. Extensions are imported into the
|
| 44 |
+
# active Python interpreter and may run arbitrary code.
|
| 45 |
+
unsafe-load-any-extension=no
|
| 46 |
+
|
| 47 |
+
|
| 48 |
+
[MESSAGES CONTROL]
|
| 49 |
+
|
| 50 |
+
# Only show warnings with the listed confidence levels. Leave empty to show
|
| 51 |
+
# all. Valid levels: HIGH, INFERENCE, INFERENCE_FAILURE, UNDEFINED.
|
| 52 |
+
confidence=
|
| 53 |
+
|
| 54 |
+
# Disable the message, report, category or checker with the given id(s). You
|
| 55 |
+
# can either give multiple identifiers separated by comma (,) or put this
|
| 56 |
+
# option multiple times (only on the command line, not in the configuration
|
| 57 |
+
# file where it should appear only once). You can also use "--disable=all" to
|
| 58 |
+
# disable everything first and then reenable specific checks. For example, if
|
| 59 |
+
# you want to run only the similarities checker, you can use "--disable=all
|
| 60 |
+
# --enable=similarities". If you want to run only the classes checker, but have
|
| 61 |
+
# no Warning level messages displayed, use "--disable=all --enable=classes
|
| 62 |
+
# --disable=W".
|
| 63 |
+
disable=relative-beyond-top-level
|
| 64 |
+
|
| 65 |
+
|
| 66 |
+
[REPORTS]
|
| 67 |
+
|
| 68 |
+
# Python expression which should return a note less than 10 (10 is the highest
|
| 69 |
+
# note). You have access to the variables errors warning, statement which
|
| 70 |
+
# respectively contain the number of errors / warnings messages and the total
|
| 71 |
+
# number of statements analyzed. This is used by the global evaluation report
|
| 72 |
+
# (RP0004).
|
| 73 |
+
evaluation=10.0 - ((float(5 * error + warning + refactor + convention) / statement) * 10)
|
| 74 |
+
|
| 75 |
+
# Template used to display messages. This is a python new-style format string
|
| 76 |
+
# used to format the message information. See doc for all details.
|
| 77 |
+
#msg-template=
|
| 78 |
+
|
| 79 |
+
# Set the output format. Available formats are text, parseable, colorized, json
|
| 80 |
+
# and msvs (visual studio). You can also give a reporter class, e.g.
|
| 81 |
+
# mypackage.mymodule.MyReporterClass.
|
| 82 |
+
output-format=text
|
| 83 |
+
|
| 84 |
+
# Tells whether to display a full report or only the messages.
|
| 85 |
+
reports=no
|
| 86 |
+
|
| 87 |
+
# Activate the evaluation score.
|
| 88 |
+
score=yes
|
| 89 |
+
|
| 90 |
+
|
| 91 |
+
[REFACTORING]
|
| 92 |
+
|
| 93 |
+
# Maximum number of nested blocks for function / method body
|
| 94 |
+
max-nested-blocks=5
|
| 95 |
+
|
| 96 |
+
# Complete name of functions that never returns. When checking for
|
| 97 |
+
# inconsistent-return-statements if a never returning function is called then
|
| 98 |
+
# it will be considered as an explicit return statement and no message will be
|
| 99 |
+
# printed.
|
| 100 |
+
never-returning-functions=sys.exit
|
| 101 |
+
|
| 102 |
+
|
| 103 |
+
[LOGGING]
|
| 104 |
+
|
| 105 |
+
# Format style used to check logging format string. `old` means using %
|
| 106 |
+
# formatting, while `new` is for `{}` formatting.
|
| 107 |
+
logging-format-style=old
|
| 108 |
+
|
| 109 |
+
# Logging modules to check that the string format arguments are in logging
|
| 110 |
+
# function parameter format.
|
| 111 |
+
logging-modules=logging
|
| 112 |
+
|
| 113 |
+
|
| 114 |
+
[VARIABLES]
|
| 115 |
+
|
| 116 |
+
# List of additional names supposed to be defined in builtins. Remember that
|
| 117 |
+
# you should avoid defining new builtins when possible.
|
| 118 |
+
additional-builtins=
|
| 119 |
+
|
| 120 |
+
# Tells whether unused global variables should be treated as a violation.
|
| 121 |
+
allow-global-unused-variables=yes
|
| 122 |
+
|
| 123 |
+
# List of strings which can identify a callback function by name. A callback
|
| 124 |
+
# name must start or end with one of those strings.
|
| 125 |
+
callbacks=cb_,
|
| 126 |
+
_cb
|
| 127 |
+
|
| 128 |
+
# A regular expression matching the name of dummy variables (i.e. expected to
|
| 129 |
+
# not be used).
|
| 130 |
+
dummy-variables-rgx=_+$|(_[a-zA-Z0-9_]*[a-zA-Z0-9]+?$)|dummy|^ignored_|^unused_
|
| 131 |
+
|
| 132 |
+
# Argument names that match this expression will be ignored. Default to name
|
| 133 |
+
# with leading underscore.
|
| 134 |
+
ignored-argument-names=_.*|^ignored_|^unused_
|
| 135 |
+
|
| 136 |
+
# Tells whether we should check for unused import in __init__ files.
|
| 137 |
+
init-import=no
|
| 138 |
+
|
| 139 |
+
# List of qualified module names which can have objects that can redefine
|
| 140 |
+
# builtins.
|
| 141 |
+
redefining-builtins-modules=six.moves,past.builtins,future.builtins,builtins,io
|
| 142 |
+
|
| 143 |
+
|
| 144 |
+
[FORMAT]
|
| 145 |
+
|
| 146 |
+
# Expected format of line ending, e.g. empty (any line ending), LF or CRLF.
|
| 147 |
+
expected-line-ending-format=
|
| 148 |
+
|
| 149 |
+
# Regexp for a line that is allowed to be longer than the limit.
|
| 150 |
+
ignore-long-lines=^\s*(# )?<?https?://\S+>?$
|
| 151 |
+
|
| 152 |
+
# Number of spaces of indent required inside a hanging or continued line.
|
| 153 |
+
indent-after-paren=4
|
| 154 |
+
|
| 155 |
+
# String used as indentation unit. This is usually " " (4 spaces) or "\t" (1
|
| 156 |
+
# tab).
|
| 157 |
+
indent-string=' '
|
| 158 |
+
|
| 159 |
+
# Maximum number of characters on a single line.
|
| 160 |
+
max-line-length=80
|
| 161 |
+
|
| 162 |
+
# Maximum number of lines in a module
|
| 163 |
+
max-module-lines=99999
|
| 164 |
+
|
| 165 |
+
# List of optional constructs for which whitespace checking is disabled. `dict-
|
| 166 |
+
# separator` is used to allow tabulation in dicts, etc.: {1 : 1,\n222: 2}.
|
| 167 |
+
# `trailing-comma` allows a space between comma and closing bracket: (a, ).
|
| 168 |
+
# `empty-line` allows space-only lines.
|
| 169 |
+
no-space-check=trailing-comma,
|
| 170 |
+
dict-separator
|
| 171 |
+
|
| 172 |
+
# Allow the body of a class to be on the same line as the declaration if body
|
| 173 |
+
# contains single statement.
|
| 174 |
+
single-line-class-stmt=no
|
| 175 |
+
|
| 176 |
+
# Allow the body of an if to be on the same line as the test if there is no
|
| 177 |
+
# else.
|
| 178 |
+
single-line-if-stmt=no
|
| 179 |
+
|
| 180 |
+
|
| 181 |
+
[TYPECHECK]
|
| 182 |
+
|
| 183 |
+
# List of decorators that produce context managers, such as
|
| 184 |
+
# contextlib.contextmanager. Add to this list to register other decorators that
|
| 185 |
+
# produce valid context managers.
|
| 186 |
+
contextmanager-decorators=contextlib.contextmanager
|
| 187 |
+
|
| 188 |
+
# List of members which are set dynamically and missed by pylint inference
|
| 189 |
+
# system, and so shouldn't trigger E1101 when accessed. Python regular
|
| 190 |
+
# expressions are accepted.
|
| 191 |
+
generated-members=
|
| 192 |
+
|
| 193 |
+
# Tells whether missing members accessed in mixin class should be ignored. A
|
| 194 |
+
# mixin class is detected if its name ends with "mixin" (case insensitive).
|
| 195 |
+
ignore-mixin-members=yes
|
| 196 |
+
|
| 197 |
+
# Tells whether to warn about missing members when the owner of the attribute
|
| 198 |
+
# is inferred to be None.
|
| 199 |
+
ignore-none=yes
|
| 200 |
+
|
| 201 |
+
# This flag controls whether pylint should warn about no-member and similar
|
| 202 |
+
# checks whenever an opaque object is returned when inferring. The inference
|
| 203 |
+
# can return multiple potential results while evaluating a Python object, but
|
| 204 |
+
# some branches might not be evaluated, which results in partial inference. In
|
| 205 |
+
# that case, it might be useful to still emit no-member and other checks for
|
| 206 |
+
# the rest of the inferred objects.
|
| 207 |
+
ignore-on-opaque-inference=yes
|
| 208 |
+
|
| 209 |
+
# List of class names for which member attributes should not be checked (useful
|
| 210 |
+
# for classes with dynamically set attributes). This supports the use of
|
| 211 |
+
# qualified names.
|
| 212 |
+
ignored-classes=optparse.Values,thread._local,_thread._local
|
| 213 |
+
|
| 214 |
+
# List of module names for which member attributes should not be checked
|
| 215 |
+
# (useful for modules/projects where namespaces are manipulated during runtime
|
| 216 |
+
# and thus existing member attributes cannot be deduced by static analysis. It
|
| 217 |
+
# supports qualified module names, as well as Unix pattern matching.
|
| 218 |
+
ignored-modules=
|
| 219 |
+
|
| 220 |
+
# Show a hint with possible names when a member name was not found. The aspect
|
| 221 |
+
# of finding the hint is based on edit distance.
|
| 222 |
+
missing-member-hint=yes
|
| 223 |
+
|
| 224 |
+
# The minimum edit distance a name should have in order to be considered a
|
| 225 |
+
# similar match for a missing member name.
|
| 226 |
+
missing-member-hint-distance=1
|
| 227 |
+
|
| 228 |
+
# The total number of similar names that should be taken in consideration when
|
| 229 |
+
# showing a hint for a missing member.
|
| 230 |
+
missing-member-max-choices=1
|
| 231 |
+
|
| 232 |
+
|
| 233 |
+
[SIMILARITIES]
|
| 234 |
+
|
| 235 |
+
# Ignore comments when computing similarities.
|
| 236 |
+
ignore-comments=yes
|
| 237 |
+
|
| 238 |
+
# Ignore docstrings when computing similarities.
|
| 239 |
+
ignore-docstrings=yes
|
| 240 |
+
|
| 241 |
+
# Ignore imports when computing similarities.
|
| 242 |
+
ignore-imports=no
|
| 243 |
+
|
| 244 |
+
# Minimum lines number of a similarity.
|
| 245 |
+
min-similarity-lines=4
|
| 246 |
+
|
| 247 |
+
|
| 248 |
+
[BASIC]
|
| 249 |
+
|
| 250 |
+
# Naming style matching correct argument names
|
| 251 |
+
argument-naming-style=snake_case
|
| 252 |
+
|
| 253 |
+
# Regular expression matching correct argument names. Overrides argument-
|
| 254 |
+
# naming-style
|
| 255 |
+
argument-rgx=^[a-z][a-z0-9_]*$
|
| 256 |
+
|
| 257 |
+
# Naming style matching correct attribute names
|
| 258 |
+
attr-naming-style=snake_case
|
| 259 |
+
|
| 260 |
+
# Regular expression matching correct attribute names. Overrides attr-naming-
|
| 261 |
+
# style
|
| 262 |
+
attr-rgx=^_{0,2}[a-z][a-z0-9_]*$
|
| 263 |
+
|
| 264 |
+
# Bad variable names which should always be refused, separated by a comma
|
| 265 |
+
bad-names=
|
| 266 |
+
|
| 267 |
+
# Naming style matching correct class attribute names
|
| 268 |
+
class-attribute-naming-style=any
|
| 269 |
+
|
| 270 |
+
# Regular expression matching correct class attribute names. Overrides class-
|
| 271 |
+
# attribute-naming-style
|
| 272 |
+
class-attribute-rgx=^(_?[A-Z][A-Z0-9_]*|__[a-z0-9_]+__|_?[a-z][a-z0-9_]*)$
|
| 273 |
+
|
| 274 |
+
# Naming style matching correct class names
|
| 275 |
+
class-naming-style=PascalCase
|
| 276 |
+
|
| 277 |
+
# Regular expression matching correct class names. Overrides class-naming-style
|
| 278 |
+
class-rgx=^_?[A-Z][a-zA-Z0-9]*$
|
| 279 |
+
|
| 280 |
+
# Naming style matching correct constant names
|
| 281 |
+
const-naming-style=UPPER_CASE
|
| 282 |
+
|
| 283 |
+
# Regular expression matching correct constant names. Overrides const-naming-
|
| 284 |
+
# style
|
| 285 |
+
const-rgx=^(_?[A-Z][A-Z0-9_]*|__[a-z0-9_]+__|_?[a-z][a-z0-9_]*)$
|
| 286 |
+
|
| 287 |
+
# Minimum line length for functions/classes that require docstrings, shorter
|
| 288 |
+
# ones are exempt.
|
| 289 |
+
docstring-min-length=10
|
| 290 |
+
|
| 291 |
+
# Naming style matching correct function names
|
| 292 |
+
function-naming-style=snake_case
|
| 293 |
+
|
| 294 |
+
# Regular expression matching correct function names. Overrides function-
|
| 295 |
+
# naming-style
|
| 296 |
+
function-rgx=^(?:(?P<exempt>setUp|tearDown|setUpModule|tearDownModule)|(?P<camel_case>_?[A-Z][a-zA-Z0-9]*)|(?P<snake_case>_?[a-z][a-z0-9_]*))$
|
| 297 |
+
|
| 298 |
+
# Good variable names which should always be accepted, separated by a comma
|
| 299 |
+
good-names=main,
|
| 300 |
+
_
|
| 301 |
+
|
| 302 |
+
# Include a hint for the correct naming format with invalid-name
|
| 303 |
+
include-naming-hint=no
|
| 304 |
+
|
| 305 |
+
# Naming style matching correct inline iteration names
|
| 306 |
+
inlinevar-naming-style=any
|
| 307 |
+
|
| 308 |
+
# Regular expression matching correct inline iteration names. Overrides
|
| 309 |
+
# inlinevar-naming-style
|
| 310 |
+
inlinevar-rgx=^[a-z][a-z0-9_]*$
|
| 311 |
+
|
| 312 |
+
# Naming style matching correct method names
|
| 313 |
+
method-naming-style=snake_case
|
| 314 |
+
|
| 315 |
+
# Regular expression matching correct method names. Overrides method-naming-
|
| 316 |
+
# style
|
| 317 |
+
method-rgx=(?x)^(?:(?P<exempt>_[a-z0-9_]+__|runTest|setUp|tearDown|setUpTestCase|tearDownTestCase|setupSelf|tearDownClass|setUpClass|(test|assert)_*[A-Z0-9][a-zA-Z0-9_]*|next)|(?P<camel_case>_{0,2}[A-Z][a-zA-Z0-9_]*)|(?P<snake_case>_{0,2}[a-z][a-z0-9_]*))$
|
| 318 |
+
|
| 319 |
+
# Naming style matching correct module names
|
| 320 |
+
module-naming-style=snake_case
|
| 321 |
+
|
| 322 |
+
# Regular expression matching correct module names. Overrides module-naming-
|
| 323 |
+
# style
|
| 324 |
+
module-rgx=^(_?[a-z][a-z0-9_]*)|__init__|PRESUBMIT|PRESUBMIT_unittest$
|
| 325 |
+
|
| 326 |
+
# Colon-delimited sets of names that determine each other's naming style when
|
| 327 |
+
# the name regexes allow several styles.
|
| 328 |
+
name-group=function:method
|
| 329 |
+
|
| 330 |
+
# Regular expression which should only match function or class names that do
|
| 331 |
+
# not require a docstring.
|
| 332 |
+
no-docstring-rgx=(__.*__|main)
|
| 333 |
+
|
| 334 |
+
# List of decorators that produce properties, such as abc.abstractproperty. Add
|
| 335 |
+
# to this list to register other decorators that produce valid properties.
|
| 336 |
+
property-classes=abc.abstractproperty,google3.pyglib.function_utils.cached.property
|
| 337 |
+
|
| 338 |
+
# Naming style matching correct variable names
|
| 339 |
+
variable-naming-style=snake_case
|
| 340 |
+
|
| 341 |
+
# Regular expression matching correct variable names. Overrides variable-
|
| 342 |
+
# naming-style
|
| 343 |
+
variable-rgx=^[a-z][a-z0-9_]*$
|
| 344 |
+
|
| 345 |
+
|
| 346 |
+
[SPELLING]
|
| 347 |
+
|
| 348 |
+
# Limits count of emitted suggestions for spelling mistakes.
|
| 349 |
+
max-spelling-suggestions=4
|
| 350 |
+
|
| 351 |
+
# Spelling dictionary name. Available dictionaries: none. To make it working
|
| 352 |
+
# install python-enchant package..
|
| 353 |
+
spelling-dict=
|
| 354 |
+
|
| 355 |
+
# List of comma separated words that should not be checked.
|
| 356 |
+
spelling-ignore-words=
|
| 357 |
+
|
| 358 |
+
# A path to a file that contains private dictionary; one word per line.
|
| 359 |
+
spelling-private-dict-file=
|
| 360 |
+
|
| 361 |
+
# Tells whether to store unknown words to indicated private dictionary in
|
| 362 |
+
# --spelling-private-dict-file option instead of raising a message.
|
| 363 |
+
spelling-store-unknown-words=no
|
| 364 |
+
|
| 365 |
+
|
| 366 |
+
[MISCELLANEOUS]
|
| 367 |
+
|
| 368 |
+
# List of note tags to take in consideration, separated by a comma.
|
| 369 |
+
notes=FIXME,
|
| 370 |
+
XXX,
|
| 371 |
+
TODO
|
| 372 |
+
|
| 373 |
+
|
| 374 |
+
[IMPORTS]
|
| 375 |
+
|
| 376 |
+
# Allow wildcard imports from modules that define __all__.
|
| 377 |
+
allow-wildcard-with-all=no
|
| 378 |
+
|
| 379 |
+
# Analyse import fallback blocks. This can be used to support both Python 2 and
|
| 380 |
+
# 3 compatible code, which means that the block might have code that exists
|
| 381 |
+
# only in one or another interpreter, leading to false positives when analysed.
|
| 382 |
+
analyse-fallback-blocks=no
|
| 383 |
+
|
| 384 |
+
# Deprecated modules which should not be used, separated by a comma.
|
| 385 |
+
deprecated-modules=optparse,tkinter.tix
|
| 386 |
+
|
| 387 |
+
# Create a graph of external dependencies in the given file (report RP0402 must
|
| 388 |
+
# not be disabled).
|
| 389 |
+
ext-import-graph=
|
| 390 |
+
|
| 391 |
+
# Create a graph of every (i.e. internal and external) dependencies in the
|
| 392 |
+
# given file (report RP0402 must not be disabled).
|
| 393 |
+
import-graph=
|
| 394 |
+
|
| 395 |
+
# Create a graph of internal dependencies in the given file (report RP0402 must
|
| 396 |
+
# not be disabled).
|
| 397 |
+
int-import-graph=
|
| 398 |
+
|
| 399 |
+
# Force import order to recognize a module as part of the standard
|
| 400 |
+
# compatibility libraries.
|
| 401 |
+
known-standard-library=
|
| 402 |
+
|
| 403 |
+
# Force import order to recognize a module as part of a third party library.
|
| 404 |
+
known-third-party=enchant
|
| 405 |
+
|
| 406 |
+
|
| 407 |
+
[CLASSES]
|
| 408 |
+
|
| 409 |
+
# List of method names used to declare (i.e. assign) instance attributes.
|
| 410 |
+
defining-attr-methods=__init__,
|
| 411 |
+
__new__,
|
| 412 |
+
setUp
|
| 413 |
+
|
| 414 |
+
# List of member names, which should be excluded from the protected access
|
| 415 |
+
# warning.
|
| 416 |
+
exclude-protected=_asdict,
|
| 417 |
+
_fields,
|
| 418 |
+
_replace,
|
| 419 |
+
_source,
|
| 420 |
+
_make
|
| 421 |
+
|
| 422 |
+
# List of valid names for the first argument in a class method.
|
| 423 |
+
valid-classmethod-first-arg=cls
|
| 424 |
+
|
| 425 |
+
# List of valid names for the first argument in a metaclass class method.
|
| 426 |
+
valid-metaclass-classmethod-first-arg=cls
|
| 427 |
+
|
| 428 |
+
|
| 429 |
+
[EXCEPTIONS]
|
| 430 |
+
|
| 431 |
+
# Exceptions that will emit a warning when being caught. Defaults to
|
| 432 |
+
# "Exception".
|
| 433 |
+
overgeneral-exceptions=Exception
|
uvd/envs/franka_kitchen/relay-policy-learning/adept_envs/.style.yapf
ADDED
|
@@ -0,0 +1,323 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
[style]
|
| 2 |
+
# Align closing bracket with visual indentation.
|
| 3 |
+
align_closing_bracket_with_visual_indent=False
|
| 4 |
+
|
| 5 |
+
# Allow dictionary keys to exist on multiple lines. For example:
|
| 6 |
+
#
|
| 7 |
+
# x = {
|
| 8 |
+
# ('this is the first element of a tuple',
|
| 9 |
+
# 'this is the second element of a tuple'):
|
| 10 |
+
# value,
|
| 11 |
+
# }
|
| 12 |
+
allow_multiline_dictionary_keys=False
|
| 13 |
+
|
| 14 |
+
# Allow lambdas to be formatted on more than one line.
|
| 15 |
+
allow_multiline_lambdas=False
|
| 16 |
+
|
| 17 |
+
# Allow splitting before a default / named assignment in an argument list.
|
| 18 |
+
allow_split_before_default_or_named_assigns=True
|
| 19 |
+
|
| 20 |
+
# Allow splits before the dictionary value.
|
| 21 |
+
allow_split_before_dict_value=True
|
| 22 |
+
|
| 23 |
+
# Let spacing indicate operator precedence. For example:
|
| 24 |
+
#
|
| 25 |
+
# a = 1 * 2 + 3 / 4
|
| 26 |
+
# b = 1 / 2 - 3 * 4
|
| 27 |
+
# c = (1 + 2) * (3 - 4)
|
| 28 |
+
# d = (1 - 2) / (3 + 4)
|
| 29 |
+
# e = 1 * 2 - 3
|
| 30 |
+
# f = 1 + 2 + 3 + 4
|
| 31 |
+
#
|
| 32 |
+
# will be formatted as follows to indicate precedence:
|
| 33 |
+
#
|
| 34 |
+
# a = 1*2 + 3/4
|
| 35 |
+
# b = 1/2 - 3*4
|
| 36 |
+
# c = (1+2) * (3-4)
|
| 37 |
+
# d = (1-2) / (3+4)
|
| 38 |
+
# e = 1*2 - 3
|
| 39 |
+
# f = 1 + 2 + 3 + 4
|
| 40 |
+
#
|
| 41 |
+
arithmetic_precedence_indication=False
|
| 42 |
+
|
| 43 |
+
# Number of blank lines surrounding top-level function and class
|
| 44 |
+
# definitions.
|
| 45 |
+
blank_lines_around_top_level_definition=2
|
| 46 |
+
|
| 47 |
+
# Insert a blank line before a class-level docstring.
|
| 48 |
+
blank_line_before_class_docstring=False
|
| 49 |
+
|
| 50 |
+
# Insert a blank line before a module docstring.
|
| 51 |
+
blank_line_before_module_docstring=False
|
| 52 |
+
|
| 53 |
+
# Insert a blank line before a 'def' or 'class' immediately nested
|
| 54 |
+
# within another 'def' or 'class'. For example:
|
| 55 |
+
#
|
| 56 |
+
# class Foo:
|
| 57 |
+
# # <------ this blank line
|
| 58 |
+
# def method():
|
| 59 |
+
# ...
|
| 60 |
+
blank_line_before_nested_class_or_def=True
|
| 61 |
+
|
| 62 |
+
# Do not split consecutive brackets. Only relevant when
|
| 63 |
+
# dedent_closing_brackets is set. For example:
|
| 64 |
+
#
|
| 65 |
+
# call_func_that_takes_a_dict(
|
| 66 |
+
# {
|
| 67 |
+
# 'key1': 'value1',
|
| 68 |
+
# 'key2': 'value2',
|
| 69 |
+
# }
|
| 70 |
+
# )
|
| 71 |
+
#
|
| 72 |
+
# would reformat to:
|
| 73 |
+
#
|
| 74 |
+
# call_func_that_takes_a_dict({
|
| 75 |
+
# 'key1': 'value1',
|
| 76 |
+
# 'key2': 'value2',
|
| 77 |
+
# })
|
| 78 |
+
coalesce_brackets=False
|
| 79 |
+
|
| 80 |
+
# The column limit.
|
| 81 |
+
column_limit=80
|
| 82 |
+
|
| 83 |
+
# The style for continuation alignment. Possible values are:
|
| 84 |
+
#
|
| 85 |
+
# - SPACE: Use spaces for continuation alignment. This is default behavior.
|
| 86 |
+
# - FIXED: Use fixed number (CONTINUATION_INDENT_WIDTH) of columns
|
| 87 |
+
# (ie: CONTINUATION_INDENT_WIDTH/INDENT_WIDTH tabs) for continuation
|
| 88 |
+
# alignment.
|
| 89 |
+
# - LESS: Slightly left if cannot vertically align continuation lines with
|
| 90 |
+
# indent characters.
|
| 91 |
+
# - VALIGN-RIGHT: Vertically align continuation lines with indent
|
| 92 |
+
# characters. Slightly right (one more indent character) if cannot
|
| 93 |
+
# vertically align continuation lines with indent characters.
|
| 94 |
+
#
|
| 95 |
+
# For options FIXED, and VALIGN-RIGHT are only available when USE_TABS is
|
| 96 |
+
# enabled.
|
| 97 |
+
continuation_align_style=SPACE
|
| 98 |
+
|
| 99 |
+
# Indent width used for line continuations.
|
| 100 |
+
continuation_indent_width=4
|
| 101 |
+
|
| 102 |
+
# Put closing brackets on a separate line, dedented, if the bracketed
|
| 103 |
+
# expression can't fit in a single line. Applies to all kinds of brackets,
|
| 104 |
+
# including function definitions and calls. For example:
|
| 105 |
+
#
|
| 106 |
+
# config = {
|
| 107 |
+
# 'key1': 'value1',
|
| 108 |
+
# 'key2': 'value2',
|
| 109 |
+
# } # <--- this bracket is dedented and on a separate line
|
| 110 |
+
#
|
| 111 |
+
# time_series = self.remote_client.query_entity_counters(
|
| 112 |
+
# entity='dev3246.region1',
|
| 113 |
+
# key='dns.query_latency_tcp',
|
| 114 |
+
# transform=Transformation.AVERAGE(window=timedelta(seconds=60)),
|
| 115 |
+
# start_ts=now()-timedelta(days=3),
|
| 116 |
+
# end_ts=now(),
|
| 117 |
+
# ) # <--- this bracket is dedented and on a separate line
|
| 118 |
+
dedent_closing_brackets=False
|
| 119 |
+
|
| 120 |
+
# Disable the heuristic which places each list element on a separate line
|
| 121 |
+
# if the list is comma-terminated.
|
| 122 |
+
disable_ending_comma_heuristic=False
|
| 123 |
+
|
| 124 |
+
# Place each dictionary entry onto its own line.
|
| 125 |
+
each_dict_entry_on_separate_line=True
|
| 126 |
+
|
| 127 |
+
# The regex for an i18n comment. The presence of this comment stops
|
| 128 |
+
# reformatting of that line, because the comments are required to be
|
| 129 |
+
# next to the string they translate.
|
| 130 |
+
i18n_comment=#\..*
|
| 131 |
+
|
| 132 |
+
# The i18n function call names. The presence of this function stops
|
| 133 |
+
# reformattting on that line, because the string it has cannot be moved
|
| 134 |
+
# away from the i18n comment.
|
| 135 |
+
i18n_function_call=N_, _
|
| 136 |
+
|
| 137 |
+
# Indent blank lines.
|
| 138 |
+
indent_blank_lines=False
|
| 139 |
+
|
| 140 |
+
# Indent the dictionary value if it cannot fit on the same line as the
|
| 141 |
+
# dictionary key. For example:
|
| 142 |
+
#
|
| 143 |
+
# config = {
|
| 144 |
+
# 'key1':
|
| 145 |
+
# 'value1',
|
| 146 |
+
# 'key2': value1 +
|
| 147 |
+
# value2,
|
| 148 |
+
# }
|
| 149 |
+
indent_dictionary_value=False
|
| 150 |
+
|
| 151 |
+
# The number of columns to use for indentation.
|
| 152 |
+
indent_width=4
|
| 153 |
+
|
| 154 |
+
# Join short lines into one line. E.g., single line 'if' statements.
|
| 155 |
+
join_multiple_lines=True
|
| 156 |
+
|
| 157 |
+
# Do not include spaces around selected binary operators. For example:
|
| 158 |
+
#
|
| 159 |
+
# 1 + 2 * 3 - 4 / 5
|
| 160 |
+
#
|
| 161 |
+
# will be formatted as follows when configured with "*,/":
|
| 162 |
+
#
|
| 163 |
+
# 1 + 2*3 - 4/5
|
| 164 |
+
#
|
| 165 |
+
no_spaces_around_selected_binary_operators=
|
| 166 |
+
|
| 167 |
+
# Use spaces around default or named assigns.
|
| 168 |
+
spaces_around_default_or_named_assign=False
|
| 169 |
+
|
| 170 |
+
# Use spaces around the power operator.
|
| 171 |
+
spaces_around_power_operator=False
|
| 172 |
+
|
| 173 |
+
# The number of spaces required before a trailing comment.
|
| 174 |
+
# This can be a single value (representing the number of spaces
|
| 175 |
+
# before each trailing comment) or list of values (representing
|
| 176 |
+
# alignment column values; trailing comments within a block will
|
| 177 |
+
# be aligned to the first column value that is greater than the maximum
|
| 178 |
+
# line length within the block). For example:
|
| 179 |
+
#
|
| 180 |
+
# With spaces_before_comment=5:
|
| 181 |
+
#
|
| 182 |
+
# 1 + 1 # Adding values
|
| 183 |
+
#
|
| 184 |
+
# will be formatted as:
|
| 185 |
+
#
|
| 186 |
+
# 1 + 1 # Adding values <-- 5 spaces between the end of the statement and comment
|
| 187 |
+
#
|
| 188 |
+
# With spaces_before_comment=15, 20:
|
| 189 |
+
#
|
| 190 |
+
# 1 + 1 # Adding values
|
| 191 |
+
# two + two # More adding
|
| 192 |
+
#
|
| 193 |
+
# longer_statement # This is a longer statement
|
| 194 |
+
# short # This is a shorter statement
|
| 195 |
+
#
|
| 196 |
+
# a_very_long_statement_that_extends_beyond_the_final_column # Comment
|
| 197 |
+
# short # This is a shorter statement
|
| 198 |
+
#
|
| 199 |
+
# will be formatted as:
|
| 200 |
+
#
|
| 201 |
+
# 1 + 1 # Adding values <-- end of line comments in block aligned to col 15
|
| 202 |
+
# two + two # More adding
|
| 203 |
+
#
|
| 204 |
+
# longer_statement # This is a longer statement <-- end of line comments in block aligned to col 20
|
| 205 |
+
# short # This is a shorter statement
|
| 206 |
+
#
|
| 207 |
+
# a_very_long_statement_that_extends_beyond_the_final_column # Comment <-- the end of line comments are aligned based on the line length
|
| 208 |
+
# short # This is a shorter statement
|
| 209 |
+
#
|
| 210 |
+
spaces_before_comment=2
|
| 211 |
+
|
| 212 |
+
# Insert a space between the ending comma and closing bracket of a list,
|
| 213 |
+
# etc.
|
| 214 |
+
space_between_ending_comma_and_closing_bracket=False
|
| 215 |
+
|
| 216 |
+
# Split before arguments
|
| 217 |
+
split_all_comma_separated_values=False
|
| 218 |
+
|
| 219 |
+
# Split before arguments if the argument list is terminated by a
|
| 220 |
+
# comma.
|
| 221 |
+
split_arguments_when_comma_terminated=False
|
| 222 |
+
|
| 223 |
+
# Set to True to prefer splitting before '&', '|' or '^' rather than
|
| 224 |
+
# after.
|
| 225 |
+
split_before_bitwise_operator=False
|
| 226 |
+
|
| 227 |
+
# Split before the closing bracket if a list or dict literal doesn't fit on
|
| 228 |
+
# a single line.
|
| 229 |
+
split_before_closing_bracket=True
|
| 230 |
+
|
| 231 |
+
# Split before a dictionary or set generator (comp_for). For example, note
|
| 232 |
+
# the split before the 'for':
|
| 233 |
+
#
|
| 234 |
+
# foo = {
|
| 235 |
+
# variable: 'Hello world, have a nice day!'
|
| 236 |
+
# for variable in bar if variable != 42
|
| 237 |
+
# }
|
| 238 |
+
split_before_dict_set_generator=False
|
| 239 |
+
|
| 240 |
+
# Split before the '.' if we need to split a longer expression:
|
| 241 |
+
#
|
| 242 |
+
# foo = ('This is a really long string: {}, {}, {}, {}'.format(a, b, c, d))
|
| 243 |
+
#
|
| 244 |
+
# would reformat to something like:
|
| 245 |
+
#
|
| 246 |
+
# foo = ('This is a really long string: {}, {}, {}, {}'
|
| 247 |
+
# .format(a, b, c, d))
|
| 248 |
+
split_before_dot=False
|
| 249 |
+
|
| 250 |
+
# Split after the opening paren which surrounds an expression if it doesn't
|
| 251 |
+
# fit on a single line.
|
| 252 |
+
split_before_expression_after_opening_paren=False
|
| 253 |
+
|
| 254 |
+
# If an argument / parameter list is going to be split, then split before
|
| 255 |
+
# the first argument.
|
| 256 |
+
split_before_first_argument=False
|
| 257 |
+
|
| 258 |
+
# Set to True to prefer splitting before 'and' or 'or' rather than
|
| 259 |
+
# after.
|
| 260 |
+
split_before_logical_operator=False
|
| 261 |
+
|
| 262 |
+
# Split named assignments onto individual lines.
|
| 263 |
+
split_before_named_assigns=True
|
| 264 |
+
|
| 265 |
+
# Set to True to split list comprehensions and generators that have
|
| 266 |
+
# non-trivial expressions and multiple clauses before each of these
|
| 267 |
+
# clauses. For example:
|
| 268 |
+
#
|
| 269 |
+
# result = [
|
| 270 |
+
# a_long_var + 100 for a_long_var in xrange(1000)
|
| 271 |
+
# if a_long_var % 10]
|
| 272 |
+
#
|
| 273 |
+
# would reformat to something like:
|
| 274 |
+
#
|
| 275 |
+
# result = [
|
| 276 |
+
# a_long_var + 100
|
| 277 |
+
# for a_long_var in xrange(1000)
|
| 278 |
+
# if a_long_var % 10]
|
| 279 |
+
split_complex_comprehension=True
|
| 280 |
+
|
| 281 |
+
# The penalty for splitting right after the opening bracket.
|
| 282 |
+
split_penalty_after_opening_bracket=30
|
| 283 |
+
|
| 284 |
+
# The penalty for splitting the line after a unary operator.
|
| 285 |
+
split_penalty_after_unary_operator=10000
|
| 286 |
+
|
| 287 |
+
# The penalty for splitting right before an if expression.
|
| 288 |
+
split_penalty_before_if_expr=0
|
| 289 |
+
|
| 290 |
+
# The penalty of splitting the line around the '&', '|', and '^'
|
| 291 |
+
# operators.
|
| 292 |
+
split_penalty_bitwise_operator=300
|
| 293 |
+
|
| 294 |
+
# The penalty for splitting a list comprehension or generator
|
| 295 |
+
# expression.
|
| 296 |
+
split_penalty_comprehension=2100
|
| 297 |
+
|
| 298 |
+
# The penalty for characters over the column limit.
|
| 299 |
+
split_penalty_excess_character=7000
|
| 300 |
+
|
| 301 |
+
# The penalty incurred by adding a line split to the unwrapped line. The
|
| 302 |
+
# more line splits added the higher the penalty.
|
| 303 |
+
split_penalty_for_added_line_split=30
|
| 304 |
+
|
| 305 |
+
# The penalty of splitting a list of "import as" names. For example:
|
| 306 |
+
#
|
| 307 |
+
# from a_very_long_or_indented_module_name_yada_yad import (long_argument_1,
|
| 308 |
+
# long_argument_2,
|
| 309 |
+
# long_argument_3)
|
| 310 |
+
#
|
| 311 |
+
# would reformat to something like:
|
| 312 |
+
#
|
| 313 |
+
# from a_very_long_or_indented_module_name_yada_yad import (
|
| 314 |
+
# long_argument_1, long_argument_2, long_argument_3)
|
| 315 |
+
split_penalty_import_names=0
|
| 316 |
+
|
| 317 |
+
# The penalty of splitting the line around the 'and' and 'or'
|
| 318 |
+
# operators.
|
| 319 |
+
split_penalty_logical_operator=300
|
| 320 |
+
|
| 321 |
+
# Use the Tab character for indentation.
|
| 322 |
+
use_tabs=False
|
| 323 |
+
|
uvd/envs/franka_kitchen/relay-policy-learning/adept_envs/adept_envs/__init__.py
ADDED
|
@@ -0,0 +1,15 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
#!/usr/bin/python
|
| 2 |
+
#
|
| 3 |
+
# Copyright 2020 Google LLC
|
| 4 |
+
#
|
| 5 |
+
# Licensed under the Apache License, Version 2.0 (the "License");
|
| 6 |
+
# you may not use this file except in compliance with the License.
|
| 7 |
+
# You may obtain a copy of the License at
|
| 8 |
+
#
|
| 9 |
+
# http://www.apache.org/licenses/LICENSE-2.0
|
| 10 |
+
#
|
| 11 |
+
# Unless required by applicable law or agreed to in writing, software
|
| 12 |
+
# distributed under the License is distributed on an "AS IS" BASIS,
|
| 13 |
+
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
| 14 |
+
# See the License for the specific language governing permissions and
|
| 15 |
+
# limitations under the License.
|
uvd/envs/franka_kitchen/relay-policy-learning/adept_envs/adept_envs/base_robot.py
ADDED
|
@@ -0,0 +1,153 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
#!/usr/bin/python
|
| 2 |
+
#
|
| 3 |
+
# Copyright 2020 Google LLC
|
| 4 |
+
#
|
| 5 |
+
# Licensed under the Apache License, Version 2.0 (the "License");
|
| 6 |
+
# you may not use this file except in compliance with the License.
|
| 7 |
+
# You may obtain a copy of the License at
|
| 8 |
+
#
|
| 9 |
+
# http://www.apache.org/licenses/LICENSE-2.0
|
| 10 |
+
#
|
| 11 |
+
# Unless required by applicable law or agreed to in writing, software
|
| 12 |
+
# distributed under the License is distributed on an "AS IS" BASIS,
|
| 13 |
+
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
| 14 |
+
# See the License for the specific language governing permissions and
|
| 15 |
+
# limitations under the License.
|
| 16 |
+
|
| 17 |
+
from collections import deque
|
| 18 |
+
|
| 19 |
+
import numpy as np
|
| 20 |
+
|
| 21 |
+
|
| 22 |
+
class BaseRobot(object):
|
| 23 |
+
"""Base class for all robot classes."""
|
| 24 |
+
|
| 25 |
+
def __init__(
|
| 26 |
+
self,
|
| 27 |
+
n_jnt,
|
| 28 |
+
n_obj,
|
| 29 |
+
pos_bounds=None,
|
| 30 |
+
vel_bounds=None,
|
| 31 |
+
calibration_path=None,
|
| 32 |
+
is_hardware=False,
|
| 33 |
+
device_name=None,
|
| 34 |
+
overlay=False,
|
| 35 |
+
calibration_mode=False,
|
| 36 |
+
observation_cache_maxsize=5,
|
| 37 |
+
):
|
| 38 |
+
"""Create a new robot.
|
| 39 |
+
|
| 40 |
+
Args:
|
| 41 |
+
n_jnt: The number of dofs in the robot.
|
| 42 |
+
n_obj: The number of dofs in the object.
|
| 43 |
+
pos_bounds: (n_jnt, 2)-shape matrix denoting the min and max joint
|
| 44 |
+
position for each joint.
|
| 45 |
+
vel_bounds: (n_jnt, 2)-shape matrix denoting the min and max joint
|
| 46 |
+
velocity for each joint.
|
| 47 |
+
calibration_path: File path to the calibration configuration file to
|
| 48 |
+
use.
|
| 49 |
+
is_hardware: Whether to run on hardware or not.
|
| 50 |
+
device_name: The device path for the robot hardware. Only required
|
| 51 |
+
in legacy mode.
|
| 52 |
+
overlay: Whether to show a simulation overlay of the hardware.
|
| 53 |
+
calibration_mode: Start with motors disengaged.
|
| 54 |
+
"""
|
| 55 |
+
|
| 56 |
+
assert n_jnt > 0
|
| 57 |
+
assert n_obj >= 0
|
| 58 |
+
|
| 59 |
+
self._n_jnt = n_jnt
|
| 60 |
+
self._n_obj = n_obj
|
| 61 |
+
self._n_dofs = n_jnt + n_obj
|
| 62 |
+
|
| 63 |
+
self._pos_bounds = None
|
| 64 |
+
if pos_bounds is not None:
|
| 65 |
+
pos_bounds = np.array(pos_bounds, dtype=np.float32)
|
| 66 |
+
assert pos_bounds.shape == (self._n_dofs, 2)
|
| 67 |
+
for low, high in pos_bounds:
|
| 68 |
+
assert low < high
|
| 69 |
+
self._pos_bounds = pos_bounds
|
| 70 |
+
self._vel_bounds = None
|
| 71 |
+
if vel_bounds is not None:
|
| 72 |
+
vel_bounds = np.array(vel_bounds, dtype=np.float32)
|
| 73 |
+
assert vel_bounds.shape == (self._n_dofs, 2)
|
| 74 |
+
for low, high in vel_bounds:
|
| 75 |
+
assert low < high
|
| 76 |
+
self._vel_bounds = vel_bounds
|
| 77 |
+
|
| 78 |
+
self._is_hardware = is_hardware
|
| 79 |
+
self._device_name = device_name
|
| 80 |
+
self._calibration_path = calibration_path
|
| 81 |
+
self._overlay = overlay
|
| 82 |
+
self._calibration_mode = calibration_mode
|
| 83 |
+
self._observation_cache_maxsize = observation_cache_maxsize
|
| 84 |
+
|
| 85 |
+
# Gets updated
|
| 86 |
+
self._observation_cache = deque([], maxlen=self._observation_cache_maxsize)
|
| 87 |
+
|
| 88 |
+
@property
|
| 89 |
+
def n_jnt(self):
|
| 90 |
+
return self._n_jnt
|
| 91 |
+
|
| 92 |
+
@property
|
| 93 |
+
def n_obj(self):
|
| 94 |
+
return self._n_obj
|
| 95 |
+
|
| 96 |
+
@property
|
| 97 |
+
def n_dofs(self):
|
| 98 |
+
return self._n_dofs
|
| 99 |
+
|
| 100 |
+
@property
|
| 101 |
+
def pos_bounds(self):
|
| 102 |
+
return self._pos_bounds
|
| 103 |
+
|
| 104 |
+
@property
|
| 105 |
+
def vel_bounds(self):
|
| 106 |
+
return self._vel_bounds
|
| 107 |
+
|
| 108 |
+
@property
|
| 109 |
+
def is_hardware(self):
|
| 110 |
+
return self._is_hardware
|
| 111 |
+
|
| 112 |
+
@property
|
| 113 |
+
def device_name(self):
|
| 114 |
+
return self._device_name
|
| 115 |
+
|
| 116 |
+
@property
|
| 117 |
+
def calibration_path(self):
|
| 118 |
+
return self._calibration_path
|
| 119 |
+
|
| 120 |
+
@property
|
| 121 |
+
def overlay(self):
|
| 122 |
+
return self._overlay
|
| 123 |
+
|
| 124 |
+
@property
|
| 125 |
+
def has_obj(self):
|
| 126 |
+
return self._n_obj > 0
|
| 127 |
+
|
| 128 |
+
@property
|
| 129 |
+
def calibration_mode(self):
|
| 130 |
+
return self._calibration_mode
|
| 131 |
+
|
| 132 |
+
@property
|
| 133 |
+
def observation_cache_maxsize(self):
|
| 134 |
+
return self._observation_cache_maxsize
|
| 135 |
+
|
| 136 |
+
@property
|
| 137 |
+
def observation_cache(self):
|
| 138 |
+
return self._observation_cache
|
| 139 |
+
|
| 140 |
+
def clip_positions(self, positions):
|
| 141 |
+
"""Clips the given joint positions to the position bounds.
|
| 142 |
+
|
| 143 |
+
Args:
|
| 144 |
+
positions: The joint positions.
|
| 145 |
+
|
| 146 |
+
Returns:
|
| 147 |
+
The bounded joint positions.
|
| 148 |
+
"""
|
| 149 |
+
if self.pos_bounds is None:
|
| 150 |
+
return positions
|
| 151 |
+
assert len(positions) == self.n_jnt or len(positions) == self.n_dofs
|
| 152 |
+
pos_bounds = self.pos_bounds[: len(positions)]
|
| 153 |
+
return np.clip(positions, pos_bounds[:, 0], pos_bounds[:, 1])
|
uvd/envs/franka_kitchen/relay-policy-learning/adept_envs/adept_envs/franka/__init__.py
ADDED
|
@@ -0,0 +1,24 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
#!/usr/bin/python
|
| 2 |
+
#
|
| 3 |
+
# Copyright 2020 Google LLC
|
| 4 |
+
#
|
| 5 |
+
# Licensed under the Apache License, Version 2.0 (the "License");
|
| 6 |
+
# you may not use this file except in compliance with the License.
|
| 7 |
+
# You may obtain a copy of the License at
|
| 8 |
+
#
|
| 9 |
+
# http://www.apache.org/licenses/LICENSE-2.0
|
| 10 |
+
#
|
| 11 |
+
# Unless required by applicable law or agreed to in writing, software
|
| 12 |
+
# distributed under the License is distributed on an "AS IS" BASIS,
|
| 13 |
+
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
| 14 |
+
# See the License for the specific language governing permissions and
|
| 15 |
+
# limitations under the License.
|
| 16 |
+
|
| 17 |
+
from gym.envs.registration import register
|
| 18 |
+
|
| 19 |
+
# Relax the robot
|
| 20 |
+
register(
|
| 21 |
+
id="kitchen_relax-v1",
|
| 22 |
+
entry_point="adept_envs.franka.kitchen_multitask_v0:KitchenTaskRelaxV1",
|
| 23 |
+
max_episode_steps=280,
|
| 24 |
+
)
|
uvd/envs/franka_kitchen/relay-policy-learning/adept_envs/adept_envs/franka/assets/franka_kitchen_jntpos_act_ab.xml
ADDED
|
@@ -0,0 +1,94 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
<!--Copyright 2020 Google LLC-->
|
| 2 |
+
|
| 3 |
+
<!--Licensed under the Apache License, Version 2.0 (the "License");-->
|
| 4 |
+
<!--you may not use this file except in compliance with the License.-->
|
| 5 |
+
<!--You may obtain a copy of the License at-->
|
| 6 |
+
|
| 7 |
+
<!--https://www.apache.org/licenses/LICENSE-2.0-->
|
| 8 |
+
|
| 9 |
+
<!--Unless required by applicable law or agreed to in writing, software-->
|
| 10 |
+
<!--distributed under the License is distributed on an "AS IS" BASIS,-->
|
| 11 |
+
<!--WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.-->
|
| 12 |
+
<!--See the License for the specific language governing permissions and-->
|
| 13 |
+
<!--limitations under the License.-->
|
| 14 |
+
|
| 15 |
+
<mujoco model="franka_mocap_studyTable_buttons">
|
| 16 |
+
|
| 17 |
+
<size njmax='1000' nconmax='1000'/>
|
| 18 |
+
|
| 19 |
+
<include file="../../../../adept_models/scenes/basic_scene.xml"/>
|
| 20 |
+
<include file="../../../../third_party/franka/assets/assets.xml"/>
|
| 21 |
+
<include file="../../../../third_party/franka/assets/actuator0.xml"/>
|
| 22 |
+
<include file="../../../../adept_models/kitchen/assets/oven_asset.xml"/>
|
| 23 |
+
<include file="../../../../adept_models/kitchen/assets/counters_asset.xml"/>
|
| 24 |
+
<include file="../../../../adept_models/kitchen/assets/backwall_asset.xml"/>
|
| 25 |
+
<include file="../../../../adept_models/kitchen/assets/slidecabinet_asset.xml"/>
|
| 26 |
+
<include file="../../../../adept_models/kitchen/assets/hingecabinet_asset.xml"/>
|
| 27 |
+
<include file="../../../../adept_models/kitchen/assets/microwave_asset.xml"/>
|
| 28 |
+
<include file="../../../../adept_models/kitchen/assets/kettle_asset.xml"/>
|
| 29 |
+
|
| 30 |
+
<visual>
|
| 31 |
+
<global offwidth="2560" offheight="1920" />
|
| 32 |
+
<quality shadowsize="4096" offsamples="8" />
|
| 33 |
+
<map force="0.1" fogend="5" />
|
| 34 |
+
</visual>
|
| 35 |
+
|
| 36 |
+
<compiler inertiafromgeom='auto' inertiagrouprange='3 5' angle="radian"
|
| 37 |
+
meshdir="../../../../adept_models/kitchen"
|
| 38 |
+
texturedir="../../../../adept_models/kitchen"/>
|
| 39 |
+
|
| 40 |
+
<equality>
|
| 41 |
+
<weld body1="vive_controller" body2="world" solref="0.02 1" solimp=".7 .95 0.050"/>
|
| 42 |
+
</equality>
|
| 43 |
+
|
| 44 |
+
<worldbody>
|
| 45 |
+
|
| 46 |
+
<!-- Mocap -->
|
| 47 |
+
<body name="vive_controller" mocap="true" pos="-0.440 -0.092 2.026" euler="-1.57 0 -.785">
|
| 48 |
+
<geom type="box" group="2" pos='0 0 .142' size="0.02 0.10 0.03" contype="0" conaffinity="0" rgba=".9 .7 .95 0" euler="0 0 -.785"/>
|
| 49 |
+
</body>
|
| 50 |
+
|
| 51 |
+
<site name='target' pos='0 0 0' size='0.1' rgba='0 2 0 .2'/>
|
| 52 |
+
<camera name='left_cap' pos='-1.2 -0.5 1.8' quat='0.78 0.49 -0.22 -0.32' />
|
| 53 |
+
<camera name='right_cap' pos='1.2 -0.5 1.8' quat='0.76 0.5 0.21 0.35'/>
|
| 54 |
+
|
| 55 |
+
<!-- Robot -->
|
| 56 |
+
<body pos='0. 0 1.8' euler='0 0 1.57'>
|
| 57 |
+
<geom type='cylinder' size='.120 .90' pos='-.04 0 -0.90' class='panda_viz'/>
|
| 58 |
+
<include file="../../../../third_party/franka/assets/chain0.xml"/>
|
| 59 |
+
</body>
|
| 60 |
+
|
| 61 |
+
<body name='desk' pos='-0.1 0.75 0'>
|
| 62 |
+
|
| 63 |
+
<body name="counters1" pos="0 0 0" >
|
| 64 |
+
<include file="../../../../adept_models/kitchen/assets/counters_chain.xml"/>
|
| 65 |
+
</body>
|
| 66 |
+
<body name="oven" pos="0 0 0" >
|
| 67 |
+
<include file="../../../../adept_models/kitchen/assets/oven_chain.xml"/>
|
| 68 |
+
</body>
|
| 69 |
+
<body name="backwall" pos="0 0 0" >
|
| 70 |
+
<include file="../../../../adept_models/kitchen/assets/backwall_chain.xml"/>
|
| 71 |
+
</body>
|
| 72 |
+
<body name="slidecabinet" pos="0.4 0.3 2.6" >
|
| 73 |
+
<include file="../../../../adept_models/kitchen/assets/slidecabinet_chain.xml"/>
|
| 74 |
+
</body>
|
| 75 |
+
<body name="hingecabinet" pos="-0.504 0.28 2.6" >
|
| 76 |
+
<include file="../../../../adept_models/kitchen/assets/hingecabinet_chain.xml"/>
|
| 77 |
+
</body>
|
| 78 |
+
<body name="microwave" pos="-0.750 -0.025 1.6" euler="0 0 0.3">
|
| 79 |
+
<include file="../../../../adept_models/kitchen/assets/microwave_chain.xml"/>
|
| 80 |
+
</body>
|
| 81 |
+
</body>
|
| 82 |
+
<body name="kettle" pos="-0.269 0.35 1.626">
|
| 83 |
+
<freejoint/>
|
| 84 |
+
<include file="../../../../adept_models/kitchen/assets/kettle_chain.xml"/>
|
| 85 |
+
</body>
|
| 86 |
+
|
| 87 |
+
</worldbody>
|
| 88 |
+
|
| 89 |
+
|
| 90 |
+
<keyframe>
|
| 91 |
+
<key qpos='0.16 -1.76 1.84 -2.51 0.36 0.79 1.55 0.00 0.0 1.25561e-05 1.57437e-07 1.25561e-05 1.57437e-07 1.25561e-05 1.57437e-07 1.25561e-05 1.57437e-07 8.24417e-05 9.48283e-05 0 0 0 0 -0.269 0.35 1.61523 1 1.34939e-19 -3.51612e-05 -7.50168e-19'/>
|
| 92 |
+
</keyframe>
|
| 93 |
+
|
| 94 |
+
</mujoco>
|
uvd/envs/franka_kitchen/relay-policy-learning/adept_envs/adept_envs/franka/kitchen_multitask_v0.py
ADDED
|
@@ -0,0 +1,234 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""Kitchen environment for long horizon manipulation."""
|
| 2 |
+
#!/usr/bin/python
|
| 3 |
+
#
|
| 4 |
+
# Copyright 2020 Google LLC
|
| 5 |
+
#
|
| 6 |
+
# Licensed under the Apache License, Version 2.0 (the "License");
|
| 7 |
+
# you may not use this file except in compliance with the License.
|
| 8 |
+
# You may obtain a copy of the License at
|
| 9 |
+
#
|
| 10 |
+
# http://www.apache.org/licenses/LICENSE-2.0
|
| 11 |
+
#
|
| 12 |
+
# Unless required by applicable law or agreed to in writing, software
|
| 13 |
+
# distributed under the License is distributed on an "AS IS" BASIS,
|
| 14 |
+
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
| 15 |
+
# See the License for the specific language governing permissions and
|
| 16 |
+
# limitations under the License.
|
| 17 |
+
|
| 18 |
+
import os
|
| 19 |
+
|
| 20 |
+
import numpy as np
|
| 21 |
+
from dm_control.mujoco import engine
|
| 22 |
+
from gym import spaces
|
| 23 |
+
|
| 24 |
+
from adept_envs import robot_env
|
| 25 |
+
from adept_envs.utils.configurable import configurable
|
| 26 |
+
|
| 27 |
+
|
| 28 |
+
@configurable(pickleable=True)
|
| 29 |
+
class KitchenV0(robot_env.RobotEnv):
|
| 30 |
+
CALIBRATION_PATHS = {
|
| 31 |
+
"default": os.path.join(os.path.dirname(__file__), "robot/franka_config.xml")
|
| 32 |
+
}
|
| 33 |
+
# Converted to velocity actuation
|
| 34 |
+
ROBOTS = {"robot": "adept_envs.franka.robot.franka_robot:Robot_VelAct"}
|
| 35 |
+
MODEl = os.path.join(
|
| 36 |
+
os.path.dirname(__file__), "../franka/assets/franka_kitchen_jntpos_act_ab.xml"
|
| 37 |
+
)
|
| 38 |
+
N_DOF_ROBOT = 9
|
| 39 |
+
N_DOF_OBJECT = 21
|
| 40 |
+
|
| 41 |
+
def __init__(self, robot_params={}, frame_skip=40):
|
| 42 |
+
self.goal_concat = True
|
| 43 |
+
self.obs_dict = {}
|
| 44 |
+
self.robot_noise_ratio = 0.1 # 10% as per robot_config specs
|
| 45 |
+
self.goal = np.zeros((30,))
|
| 46 |
+
|
| 47 |
+
super().__init__(
|
| 48 |
+
self.MODEl,
|
| 49 |
+
robot=self.make_robot(
|
| 50 |
+
n_jnt=self.N_DOF_ROBOT, # root+robot_jnts
|
| 51 |
+
n_obj=self.N_DOF_OBJECT,
|
| 52 |
+
**robot_params
|
| 53 |
+
),
|
| 54 |
+
frame_skip=frame_skip,
|
| 55 |
+
camera_settings=dict(
|
| 56 |
+
distance=4.5,
|
| 57 |
+
azimuth=-66,
|
| 58 |
+
elevation=-65,
|
| 59 |
+
),
|
| 60 |
+
)
|
| 61 |
+
self.init_qpos = self.sim.model.key_qpos[0].copy()
|
| 62 |
+
|
| 63 |
+
# For the microwave kettle slide hinge
|
| 64 |
+
self.init_qpos = np.array(
|
| 65 |
+
[
|
| 66 |
+
1.48388023e-01,
|
| 67 |
+
-1.76848573e00,
|
| 68 |
+
1.84390296e00,
|
| 69 |
+
-2.47685760e00,
|
| 70 |
+
2.60252026e-01,
|
| 71 |
+
7.12533105e-01,
|
| 72 |
+
1.59515394e00,
|
| 73 |
+
4.79267505e-02,
|
| 74 |
+
3.71350919e-02,
|
| 75 |
+
-2.66279850e-04,
|
| 76 |
+
-5.18043486e-05,
|
| 77 |
+
3.12877220e-05,
|
| 78 |
+
-4.51199853e-05,
|
| 79 |
+
-3.90842156e-06,
|
| 80 |
+
-4.22629655e-05,
|
| 81 |
+
6.28065475e-05,
|
| 82 |
+
4.04984708e-05,
|
| 83 |
+
4.62730939e-04,
|
| 84 |
+
-2.26906415e-04,
|
| 85 |
+
-4.65501369e-04,
|
| 86 |
+
-6.44129196e-03,
|
| 87 |
+
-1.77048263e-03,
|
| 88 |
+
1.08009684e-03,
|
| 89 |
+
-2.69397440e-01,
|
| 90 |
+
3.50383255e-01,
|
| 91 |
+
1.61944683e00,
|
| 92 |
+
1.00618764e00,
|
| 93 |
+
4.06395120e-03,
|
| 94 |
+
-6.62095997e-03,
|
| 95 |
+
-2.68278933e-04,
|
| 96 |
+
]
|
| 97 |
+
)
|
| 98 |
+
|
| 99 |
+
self.init_qvel = self.sim.model.key_qvel[0].copy()
|
| 100 |
+
|
| 101 |
+
self.act_mid = np.zeros(self.N_DOF_ROBOT)
|
| 102 |
+
self.act_amp = 2.0 * np.ones(self.N_DOF_ROBOT)
|
| 103 |
+
|
| 104 |
+
act_lower = -1 * np.ones((self.N_DOF_ROBOT,))
|
| 105 |
+
act_upper = 1 * np.ones((self.N_DOF_ROBOT,))
|
| 106 |
+
self.action_space = spaces.Box(act_lower, act_upper)
|
| 107 |
+
|
| 108 |
+
obs_upper = 8.0 * np.ones(self.obs_dim)
|
| 109 |
+
obs_lower = -obs_upper
|
| 110 |
+
self.observation_space = spaces.Box(obs_lower, obs_upper)
|
| 111 |
+
|
| 112 |
+
def _get_reward_n_score(self, obs_dict):
|
| 113 |
+
raise NotImplementedError()
|
| 114 |
+
|
| 115 |
+
def step(self, a, b=None):
|
| 116 |
+
a = np.clip(a, -1.0, 1.0)
|
| 117 |
+
|
| 118 |
+
if not self.initializing:
|
| 119 |
+
a = self.act_mid + a * self.act_amp # mean center and scale
|
| 120 |
+
else:
|
| 121 |
+
self.goal = self._get_task_goal() # update goal if init
|
| 122 |
+
|
| 123 |
+
self.robot.step(self, a, step_duration=self.skip * self.model.opt.timestep)
|
| 124 |
+
|
| 125 |
+
# observations
|
| 126 |
+
obs = self._get_obs()
|
| 127 |
+
|
| 128 |
+
# rewards
|
| 129 |
+
reward_dict, score = self._get_reward_n_score(self.obs_dict)
|
| 130 |
+
|
| 131 |
+
# termination
|
| 132 |
+
done = False
|
| 133 |
+
|
| 134 |
+
# finalize step
|
| 135 |
+
env_info = {
|
| 136 |
+
"time": self.obs_dict["t"],
|
| 137 |
+
"obs_dict": self.obs_dict,
|
| 138 |
+
"rewards": reward_dict,
|
| 139 |
+
"score": score,
|
| 140 |
+
"images": np.asarray(self.render(mode="rgb_array")),
|
| 141 |
+
}
|
| 142 |
+
# self.render()
|
| 143 |
+
return obs, reward_dict["r_total"], done, env_info
|
| 144 |
+
|
| 145 |
+
def _get_obs(self):
|
| 146 |
+
t, qp, qv, obj_qp, obj_qv = self.robot.get_obs(
|
| 147 |
+
self, robot_noise_ratio=self.robot_noise_ratio
|
| 148 |
+
)
|
| 149 |
+
|
| 150 |
+
self.obs_dict = {}
|
| 151 |
+
self.obs_dict["t"] = t
|
| 152 |
+
self.obs_dict["qp"] = qp
|
| 153 |
+
self.obs_dict["qv"] = qv
|
| 154 |
+
self.obs_dict["obj_qp"] = obj_qp
|
| 155 |
+
self.obs_dict["obj_qv"] = obj_qv
|
| 156 |
+
self.obs_dict["goal"] = self.goal
|
| 157 |
+
if self.goal_concat:
|
| 158 |
+
return np.concatenate(
|
| 159 |
+
[self.obs_dict["qp"], self.obs_dict["obj_qp"], self.obs_dict["goal"]]
|
| 160 |
+
)
|
| 161 |
+
|
| 162 |
+
def reset_model(self):
|
| 163 |
+
reset_pos = self.init_qpos[:].copy()
|
| 164 |
+
reset_vel = self.init_qvel[:].copy()
|
| 165 |
+
self.robot.reset(self, reset_pos, reset_vel)
|
| 166 |
+
self.sim.forward()
|
| 167 |
+
self.goal = self._get_task_goal() # sample a new goal on reset
|
| 168 |
+
return self._get_obs()
|
| 169 |
+
|
| 170 |
+
def evaluate_success(self, paths):
|
| 171 |
+
# score
|
| 172 |
+
mean_score_per_rollout = np.zeros(shape=len(paths))
|
| 173 |
+
for idx, path in enumerate(paths):
|
| 174 |
+
mean_score_per_rollout[idx] = np.mean(path["env_infos"]["score"])
|
| 175 |
+
mean_score = np.mean(mean_score_per_rollout)
|
| 176 |
+
|
| 177 |
+
# success percentage
|
| 178 |
+
num_success = 0
|
| 179 |
+
num_paths = len(paths)
|
| 180 |
+
for path in paths:
|
| 181 |
+
num_success += bool(path["env_infos"]["rewards"]["bonus"][-1])
|
| 182 |
+
success_percentage = num_success * 100.0 / num_paths
|
| 183 |
+
|
| 184 |
+
# fuse results
|
| 185 |
+
return np.sign(mean_score) * (
|
| 186 |
+
1e6 * round(success_percentage, 2) + abs(mean_score)
|
| 187 |
+
)
|
| 188 |
+
|
| 189 |
+
def close_env(self):
|
| 190 |
+
self.robot.close()
|
| 191 |
+
|
| 192 |
+
def set_goal(self, goal):
|
| 193 |
+
self.goal = goal
|
| 194 |
+
|
| 195 |
+
def _get_task_goal(self):
|
| 196 |
+
return self.goal
|
| 197 |
+
|
| 198 |
+
# Only include goal
|
| 199 |
+
@property
|
| 200 |
+
def goal_space(self):
|
| 201 |
+
len_obs = self.observation_space.low.shape[0]
|
| 202 |
+
env_lim = np.abs(self.observation_space.low[0])
|
| 203 |
+
return spaces.Box(
|
| 204 |
+
low=-env_lim, high=env_lim, shape=(len_obs // 2,), dtype=np.float32
|
| 205 |
+
)
|
| 206 |
+
|
| 207 |
+
def convert_to_active_observation(self, observation):
|
| 208 |
+
return observation
|
| 209 |
+
|
| 210 |
+
|
| 211 |
+
class KitchenTaskRelaxV1(KitchenV0):
|
| 212 |
+
"""Kitchen environment with proper camera and goal setup."""
|
| 213 |
+
|
| 214 |
+
def __init__(self):
|
| 215 |
+
super(KitchenTaskRelaxV1, self).__init__()
|
| 216 |
+
|
| 217 |
+
def _get_reward_n_score(self, obs_dict):
|
| 218 |
+
reward_dict = {}
|
| 219 |
+
reward_dict["true_reward"] = 0.0
|
| 220 |
+
reward_dict["bonus"] = 0.0
|
| 221 |
+
reward_dict["r_total"] = 0.0
|
| 222 |
+
score = 0.0
|
| 223 |
+
return reward_dict, score
|
| 224 |
+
|
| 225 |
+
def render(self, mode="human"):
|
| 226 |
+
if mode == "rgb_array":
|
| 227 |
+
camera = engine.MovableCamera(self.sim, 1920, 2560)
|
| 228 |
+
camera.set_pose(
|
| 229 |
+
distance=2.2, lookat=[-0.2, 0.5, 2.0], azimuth=70, elevation=-35
|
| 230 |
+
)
|
| 231 |
+
img = camera.render()
|
| 232 |
+
return img
|
| 233 |
+
else:
|
| 234 |
+
super(KitchenTaskRelaxV1, self).render()
|
uvd/envs/franka_kitchen/relay-policy-learning/adept_envs/adept_envs/franka/robot/franka_config.xml
ADDED
|
@@ -0,0 +1,59 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
<!--Copyright 2020 Google LLC-->
|
| 2 |
+
|
| 3 |
+
<!--Licensed under the Apache License, Version 2.0 (the "License");-->
|
| 4 |
+
<!--you may not use this file except in compliance with the License.-->
|
| 5 |
+
<!--You may obtain a copy of the License at-->
|
| 6 |
+
|
| 7 |
+
<!--https://www.apache.org/licenses/LICENSE-2.0-->
|
| 8 |
+
|
| 9 |
+
<!--Unless required by applicable law or agreed to in writing, software-->
|
| 10 |
+
<!--distributed under the License is distributed on an "AS IS" BASIS,-->
|
| 11 |
+
<!--WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.-->
|
| 12 |
+
<!--See the License for the specific language governing permissions and-->
|
| 13 |
+
<!--limitations under the License.-->
|
| 14 |
+
<config name='Franka'>
|
| 15 |
+
|
| 16 |
+
<!-- Franka -->
|
| 17 |
+
<qpos0 name='q0' mode='1' mj_dof='0' hardware_dof='0' scale='1' offset='0' pos_bound='-2.9 2.9' vel_bound='-10 10' pos_noise_amp='0.1' vel_noise_amp='0.1' />
|
| 18 |
+
<qpos1 name='q1' mode='1' mj_dof='1' hardware_dof='1' scale='1' offset='0' pos_bound='-1.8 1.8' vel_bound='-10 10' pos_noise_amp='0.1' vel_noise_amp='0.1' />
|
| 19 |
+
<qpos2 name='q2' mode='1' mj_dof='2' hardware_dof='2' scale='1' offset='0' pos_bound='-2.9 2.9' vel_bound='-10 10' pos_noise_amp='0.1' vel_noise_amp='0.1' />
|
| 20 |
+
<qpos3 name='q3' mode='1' mj_dof='3' hardware_dof='3' scale='1' offset='0' pos_bound='-3.1 0.0' vel_bound='-10 10' pos_noise_amp='0.1' vel_noise_amp='0.1' />
|
| 21 |
+
<qpos4 name='q4' mode='1' mj_dof='4' hardware_dof='4' scale='1' offset='0' pos_bound='-2.9 2.9' vel_bound='-10 10' pos_noise_amp='0.1' vel_noise_amp='0.1' />
|
| 22 |
+
<qpos5 name='q5' mode='1' mj_dof='5' hardware_dof='5' scale='1' offset='0' pos_bound='00.0 3.8' vel_bound='-10 10' pos_noise_amp='0.1' vel_noise_amp='0.1' />
|
| 23 |
+
<qpos6 name='q6' mode='1' mj_dof='6' hardware_dof='6' scale='1' offset='0' pos_bound='-2.9 2.9' vel_bound='-10 10' pos_noise_amp='0.1' vel_noise_amp='0.1' />
|
| 24 |
+
<qpos7 name='q7' mode='1' mj_dof='7' hardware_dof='7' scale='1' offset='0' pos_bound='00.0 0.04' vel_bound='-10 10' pos_noise_amp='0.1' vel_noise_amp='0.1' />
|
| 25 |
+
<qpos8 name='q8' mode='1' mj_dof='8' hardware_dof='8' scale='1' offset='0' pos_bound='0.0 0.04' vel_bound='-10 10' pos_noise_amp='0.1' vel_noise_amp='0.1' />
|
| 26 |
+
|
| 27 |
+
<!-- Desk -->
|
| 28 |
+
<qpos9 name='deskSlideB' mode='1' mj_dof='9' hardware_dof='9' scale='1' offset='0' pos_bound='-.5 0.0' vel_bound='-5 5' pos_noise_amp='0.005' vel_noise_amp='0.005' />
|
| 29 |
+
<qpos10 name='deskSlideT' mode='1' mj_dof='10' hardware_dof='10' scale='1' offset='0' pos_bound='-.5 0.0' vel_bound='-5 5' pos_noise_amp='0.005' vel_noise_amp='0.005' />
|
| 30 |
+
|
| 31 |
+
<!-- Buttons -->
|
| 32 |
+
<qpos11 name='rBotton' mode='1' mj_dof='11' hardware_dof='11' scale='1' offset='0' pos_bound='-.005 0.0' vel_bound='-5 5' pos_noise_amp='0.0005' vel_noise_amp='0.005' />
|
| 33 |
+
<qpos12 name='gButton' mode='1' mj_dof='12' hardware_dof='12' scale='1' offset='0' pos_bound='-.005 0.0' vel_bound='-5 5' pos_noise_amp='0.0005' vel_noise_amp='0.005' />
|
| 34 |
+
<qpos13 name='bButton' mode='1' mj_dof='13' hardware_dof='13' scale='1' offset='0' pos_bound='-.005 0.0' vel_bound='-5 5' pos_noise_amp='0.0005' vel_noise_amp='0.005' />
|
| 35 |
+
<qpos14 name='rLight' mode='1' mj_dof='14' hardware_dof='14' scale='1' offset='0' pos_bound='-.005 0.0' vel_bound='-5 5' pos_noise_amp='0.0005' vel_noise_amp='0.005' />
|
| 36 |
+
<qpos15 name='bLight' mode='1' mj_dof='15' hardware_dof='15' scale='1' offset='0' pos_bound='-.005 0.0' vel_bound='-5 5' pos_noise_amp='0.0005' vel_noise_amp='0.005' />
|
| 37 |
+
<qpos16 name='gLight' mode='1' mj_dof='16' hardware_dof='16' scale='1' offset='0' pos_bound='-.005 0.0' vel_bound='-5 5' pos_noise_amp='0.0005' vel_noise_amp='0.005' />
|
| 38 |
+
|
| 39 |
+
<!-- Blocks -->
|
| 40 |
+
<qpos17 name='q17' mode='1' mj_dof='17' hardware_dof='17' scale='1' offset='0' pos_bound='-1.5 1.5' vel_bound='-5 5' pos_noise_amp='0.005' vel_noise_amp='0.005' />
|
| 41 |
+
<qpos18 name='q18' mode='1' mj_dof='18' hardware_dof='18' scale='1' offset='0' pos_bound='-1.5 1.5' vel_bound='-5 5' pos_noise_amp='0.005' vel_noise_amp='0.005' />
|
| 42 |
+
<qpos19 name='q19' mode='1' mj_dof='19' hardware_dof='19' scale='1' offset='0' pos_bound='-1.5 1.5' vel_bound='-5 5' pos_noise_amp='0.005' vel_noise_amp='0.005' />
|
| 43 |
+
<qpos20 name='q20' mode='1' mj_dof='20' hardware_dof='20' scale='1' offset='0' pos_bound='-10.57 10.57' vel_bound='-.5 .5' pos_noise_amp='0.1' vel_noise_amp='0.1' />
|
| 44 |
+
<qpos21 name='q21' mode='1' mj_dof='21' hardware_dof='21' scale='1' offset='0' pos_bound='-10.57 10.57' vel_bound='-.5 .5' pos_noise_amp='0.1' vel_noise_amp='0.1' />
|
| 45 |
+
<qpos22 name='q22' mode='1' mj_dof='22' hardware_dof='22' scale='1' offset='0' pos_bound='-10.57 10.57' vel_bound='-.5 .5' pos_noise_amp='0.1' vel_noise_amp='0.1' />
|
| 46 |
+
<qpos23 name='q23' mode='1' mj_dof='23' hardware_dof='23' scale='1' offset='0' pos_bound='-1.5 1.5' vel_bound='-5 5' pos_noise_amp='0.005' vel_noise_amp='0.005' />
|
| 47 |
+
<qpos24 name='q24' mode='1' mj_dof='24' hardware_dof='24' scale='1' offset='0' pos_bound='-1.5 1.5' vel_bound='-5 5' pos_noise_amp='0.005' vel_noise_amp='0.005' />
|
| 48 |
+
<qpos25 name='q25' mode='1' mj_dof='25' hardware_dof='25' scale='1' offset='0' pos_bound='-1.5 1.5' vel_bound='-5 5' pos_noise_amp='0.005' vel_noise_amp='0.005' />
|
| 49 |
+
<qpos26 name='q26' mode='1' mj_dof='26' hardware_dof='26' scale='1' offset='0' pos_bound='-10.57 10.57' vel_bound='-.5 .5' pos_noise_amp='0.1' vel_noise_amp='0.1' />
|
| 50 |
+
<qpos27 name='q27' mode='1' mj_dof='27' hardware_dof='27' scale='1' offset='0' pos_bound='-10.57 10.57' vel_bound='-.5 .5' pos_noise_amp='0.1' vel_noise_amp='0.1' />
|
| 51 |
+
<qpos28 name='q28' mode='1' mj_dof='28' hardware_dof='28' scale='1' offset='0' pos_bound='-10.57 10.57' vel_bound='-.5 .5' pos_noise_amp='0.1' vel_noise_amp='0.1' />
|
| 52 |
+
<qpos29 name='q29' mode='1' mj_dof='29' hardware_dof='29' scale='1' offset='0' pos_bound='-1.5 1.5' vel_bound='-5 5' pos_noise_amp='0.005' vel_noise_amp='0.005' />
|
| 53 |
+
<qpos30 name='q30' mode='1' mj_dof='30' hardware_dof='30' scale='1' offset='0' pos_bound='-1.5 1.5' vel_bound='-5 5' pos_noise_amp='0.005' vel_noise_amp='0.005' />
|
| 54 |
+
<qpos31 name='q31' mode='1' mj_dof='31' hardware_dof='31' scale='1' offset='0' pos_bound='-1.5 1.5' vel_bound='-5 5' pos_noise_amp='0.005' vel_noise_amp='0.005' />
|
| 55 |
+
<qpos32 name='q32' mode='1' mj_dof='32' hardware_dof='32' scale='1' offset='0' pos_bound='-10.57 10.57' vel_bound='-.5 .5' pos_noise_amp='0.1' vel_noise_amp='0.1' />
|
| 56 |
+
<qpos33 name='q33' mode='1' mj_dof='33' hardware_dof='33' scale='1' offset='0' pos_bound='-10.57 10.57' vel_bound='-.5 .5' pos_noise_amp='0.1' vel_noise_amp='0.1' />
|
| 57 |
+
<qpos34 name='q34' mode='1' mj_dof='34' hardware_dof='34' scale='1' offset='0' pos_bound='-10.57 10.57' vel_bound='-.5 .5' pos_noise_amp='0.1' vel_noise_amp='0.1' />
|
| 58 |
+
|
| 59 |
+
</config>
|
uvd/envs/franka_kitchen/relay-policy-learning/adept_envs/adept_envs/franka/robot/franka_robot.py
ADDED
|
@@ -0,0 +1,342 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
#!/usr/bin/python
|
| 2 |
+
#
|
| 3 |
+
# Copyright 2020 Google LLC
|
| 4 |
+
#
|
| 5 |
+
# Licensed under the Apache License, Version 2.0 (the "License");
|
| 6 |
+
# you may not use this file except in compliance with the License.
|
| 7 |
+
# You may obtain a copy of the License at
|
| 8 |
+
#
|
| 9 |
+
# http://www.apache.org/licenses/LICENSE-2.0
|
| 10 |
+
#
|
| 11 |
+
# Unless required by applicable law or agreed to in writing, software
|
| 12 |
+
# distributed under the License is distributed on an "AS IS" BASIS,
|
| 13 |
+
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
| 14 |
+
# See the License for the specific language governing permissions and
|
| 15 |
+
# limitations under the License.
|
| 16 |
+
|
| 17 |
+
import time
|
| 18 |
+
|
| 19 |
+
# obervations structure
|
| 20 |
+
from collections import namedtuple
|
| 21 |
+
|
| 22 |
+
import numpy as np
|
| 23 |
+
from termcolor import cprint
|
| 24 |
+
|
| 25 |
+
from adept_envs import base_robot
|
| 26 |
+
from adept_envs.utils.config import get_config_root_node, read_config_from_node
|
| 27 |
+
|
| 28 |
+
observation = namedtuple(
|
| 29 |
+
"observation", ["time", "qpos_robot", "qvel_robot", "qpos_object", "qvel_object"]
|
| 30 |
+
)
|
| 31 |
+
|
| 32 |
+
|
| 33 |
+
franka_interface = ""
|
| 34 |
+
|
| 35 |
+
|
| 36 |
+
class Robot(base_robot.BaseRobot):
|
| 37 |
+
"""Abstracts away the differences between the robot_simulation and
|
| 38 |
+
robot_hardware."""
|
| 39 |
+
|
| 40 |
+
def __init__(self, *args, **kwargs):
|
| 41 |
+
super(Robot, self).__init__(*args, **kwargs)
|
| 42 |
+
global franka_interface
|
| 43 |
+
|
| 44 |
+
# Read robot configurations
|
| 45 |
+
self._read_specs_from_config(robot_configs=self.calibration_path)
|
| 46 |
+
|
| 47 |
+
# Robot: Handware
|
| 48 |
+
if self.is_hardware:
|
| 49 |
+
if franka_interface is "":
|
| 50 |
+
raise NotImplementedError()
|
| 51 |
+
from handware.franka import franka
|
| 52 |
+
|
| 53 |
+
# initialize franka
|
| 54 |
+
self.franka_interface = franka()
|
| 55 |
+
franka_interface = self.franka_interface
|
| 56 |
+
cprint(
|
| 57 |
+
"Initializing %s Hardware (Status:%d)"
|
| 58 |
+
% (self.robot_name, self.franka.okay(self.robot_hardware_dof)),
|
| 59 |
+
"white",
|
| 60 |
+
"on_grey",
|
| 61 |
+
)
|
| 62 |
+
else:
|
| 63 |
+
self.franka_interface = franka_interface
|
| 64 |
+
cprint("Reusing previours Franka session", "white", "on_grey")
|
| 65 |
+
|
| 66 |
+
# Robot: Simulation
|
| 67 |
+
else:
|
| 68 |
+
self.robot_name = "Franka"
|
| 69 |
+
cprint("Initializing %s sim" % self.robot_name, "white", "on_grey")
|
| 70 |
+
|
| 71 |
+
# Robot's time
|
| 72 |
+
self.time_start = time.time()
|
| 73 |
+
self.time = time.time() - self.time_start
|
| 74 |
+
self.time_render = -1 # time of rendering
|
| 75 |
+
|
| 76 |
+
# read specs from the calibration file
|
| 77 |
+
def _read_specs_from_config(self, robot_configs):
|
| 78 |
+
root, root_name = get_config_root_node(config_file_name=robot_configs)
|
| 79 |
+
self.robot_name = root_name[0]
|
| 80 |
+
self.robot_mode = np.zeros(self.n_dofs, dtype=int)
|
| 81 |
+
self.robot_mj_dof = np.zeros(self.n_dofs, dtype=int)
|
| 82 |
+
self.robot_hardware_dof = np.zeros(self.n_dofs, dtype=int)
|
| 83 |
+
self.robot_scale = np.zeros(self.n_dofs, dtype=float)
|
| 84 |
+
self.robot_offset = np.zeros(self.n_dofs, dtype=float)
|
| 85 |
+
self.robot_pos_bound = np.zeros([self.n_dofs, 2], dtype=float)
|
| 86 |
+
self.robot_vel_bound = np.zeros([self.n_dofs, 2], dtype=float)
|
| 87 |
+
self.robot_pos_noise_amp = np.zeros(self.n_dofs, dtype=float)
|
| 88 |
+
self.robot_vel_noise_amp = np.zeros(self.n_dofs, dtype=float)
|
| 89 |
+
|
| 90 |
+
print("Reading configurations for %s" % self.robot_name)
|
| 91 |
+
for i in range(self.n_dofs):
|
| 92 |
+
self.robot_mode[i] = read_config_from_node(
|
| 93 |
+
root, "qpos" + str(i), "mode", int
|
| 94 |
+
)
|
| 95 |
+
self.robot_mj_dof[i] = read_config_from_node(
|
| 96 |
+
root, "qpos" + str(i), "mj_dof", int
|
| 97 |
+
)
|
| 98 |
+
self.robot_hardware_dof[i] = read_config_from_node(
|
| 99 |
+
root, "qpos" + str(i), "hardware_dof", int
|
| 100 |
+
)
|
| 101 |
+
self.robot_scale[i] = read_config_from_node(
|
| 102 |
+
root, "qpos" + str(i), "scale", float
|
| 103 |
+
)
|
| 104 |
+
self.robot_offset[i] = read_config_from_node(
|
| 105 |
+
root, "qpos" + str(i), "offset", float
|
| 106 |
+
)
|
| 107 |
+
self.robot_pos_bound[i] = read_config_from_node(
|
| 108 |
+
root, "qpos" + str(i), "pos_bound", float
|
| 109 |
+
)
|
| 110 |
+
self.robot_vel_bound[i] = read_config_from_node(
|
| 111 |
+
root, "qpos" + str(i), "vel_bound", float
|
| 112 |
+
)
|
| 113 |
+
self.robot_pos_noise_amp[i] = read_config_from_node(
|
| 114 |
+
root, "qpos" + str(i), "pos_noise_amp", float
|
| 115 |
+
)
|
| 116 |
+
self.robot_vel_noise_amp[i] = read_config_from_node(
|
| 117 |
+
root, "qpos" + str(i), "vel_noise_amp", float
|
| 118 |
+
)
|
| 119 |
+
|
| 120 |
+
# convert to hardware space
|
| 121 |
+
def _de_calib(self, qp_mj, qv_mj=None):
|
| 122 |
+
qp_ad = (qp_mj - self.robot_offset) / self.robot_scale
|
| 123 |
+
if qv_mj is not None:
|
| 124 |
+
qv_ad = qv_mj / self.robot_scale
|
| 125 |
+
return qp_ad, qv_ad
|
| 126 |
+
else:
|
| 127 |
+
return qp_ad
|
| 128 |
+
|
| 129 |
+
# convert to mujoco space
|
| 130 |
+
def _calib(self, qp_ad, qv_ad):
|
| 131 |
+
qp_mj = qp_ad * self.robot_scale + self.robot_offset
|
| 132 |
+
qv_mj = qv_ad * self.robot_scale
|
| 133 |
+
return qp_mj, qv_mj
|
| 134 |
+
|
| 135 |
+
# refresh the observation cache
|
| 136 |
+
def _observation_cache_refresh(self, env):
|
| 137 |
+
for _ in range(self.observation_cache_maxsize):
|
| 138 |
+
self.get_obs(env, sim_mimic_hardware=False)
|
| 139 |
+
|
| 140 |
+
# get past observation
|
| 141 |
+
def get_obs_from_cache(self, env, index=-1):
|
| 142 |
+
assert (index >= 0 and index < self.observation_cache_maxsize) or (
|
| 143 |
+
index < 0 and index >= -self.observation_cache_maxsize
|
| 144 |
+
), (
|
| 145 |
+
"cache index out of bound. (cache size is %2d)"
|
| 146 |
+
% self.observation_cache_maxsize
|
| 147 |
+
)
|
| 148 |
+
obs = self.observation_cache[index]
|
| 149 |
+
if self.has_obj:
|
| 150 |
+
return (
|
| 151 |
+
obs.time,
|
| 152 |
+
obs.qpos_robot,
|
| 153 |
+
obs.qvel_robot,
|
| 154 |
+
obs.qpos_object,
|
| 155 |
+
obs.qvel_object,
|
| 156 |
+
)
|
| 157 |
+
else:
|
| 158 |
+
return obs.time, obs.qpos_robot, obs.qvel_robot
|
| 159 |
+
|
| 160 |
+
# get observation
|
| 161 |
+
def get_obs(
|
| 162 |
+
self, env, robot_noise_ratio=1, object_noise_ratio=1, sim_mimic_hardware=True
|
| 163 |
+
):
|
| 164 |
+
if self.is_hardware:
|
| 165 |
+
raise NotImplementedError()
|
| 166 |
+
|
| 167 |
+
else:
|
| 168 |
+
# Gather simulated observation
|
| 169 |
+
qp = env.sim.data.qpos[: self.n_jnt].copy()
|
| 170 |
+
qv = env.sim.data.qvel[: self.n_jnt].copy()
|
| 171 |
+
if self.has_obj:
|
| 172 |
+
qp_obj = env.sim.data.qpos[-self.n_obj :].copy()
|
| 173 |
+
qv_obj = env.sim.data.qvel[-self.n_obj :].copy()
|
| 174 |
+
else:
|
| 175 |
+
qp_obj = None
|
| 176 |
+
qv_obj = None
|
| 177 |
+
self.time = env.sim.data.time
|
| 178 |
+
|
| 179 |
+
# Simulate observation noise
|
| 180 |
+
if not env.initializing:
|
| 181 |
+
qp += (
|
| 182 |
+
robot_noise_ratio
|
| 183 |
+
* self.robot_pos_noise_amp[: self.n_jnt]
|
| 184 |
+
* env.np_random.uniform(low=-1.0, high=1.0, size=self.n_jnt)
|
| 185 |
+
)
|
| 186 |
+
qv += (
|
| 187 |
+
robot_noise_ratio
|
| 188 |
+
* self.robot_vel_noise_amp[: self.n_jnt]
|
| 189 |
+
* env.np_random.uniform(low=-1.0, high=1.0, size=self.n_jnt)
|
| 190 |
+
)
|
| 191 |
+
if self.has_obj:
|
| 192 |
+
qp_obj += (
|
| 193 |
+
robot_noise_ratio
|
| 194 |
+
* self.robot_pos_noise_amp[-self.n_obj :]
|
| 195 |
+
* env.np_random.uniform(low=-1.0, high=1.0, size=self.n_obj)
|
| 196 |
+
)
|
| 197 |
+
qv_obj += (
|
| 198 |
+
robot_noise_ratio
|
| 199 |
+
* self.robot_vel_noise_amp[-self.n_obj :]
|
| 200 |
+
* env.np_random.uniform(low=-1.0, high=1.0, size=self.n_obj)
|
| 201 |
+
)
|
| 202 |
+
|
| 203 |
+
# cache observations
|
| 204 |
+
obs = observation(
|
| 205 |
+
time=self.time,
|
| 206 |
+
qpos_robot=qp,
|
| 207 |
+
qvel_robot=qv,
|
| 208 |
+
qpos_object=qp_obj,
|
| 209 |
+
qvel_object=qv_obj,
|
| 210 |
+
)
|
| 211 |
+
self.observation_cache.append(obs)
|
| 212 |
+
|
| 213 |
+
if self.has_obj:
|
| 214 |
+
return (
|
| 215 |
+
obs.time,
|
| 216 |
+
obs.qpos_robot,
|
| 217 |
+
obs.qvel_robot,
|
| 218 |
+
obs.qpos_object,
|
| 219 |
+
obs.qvel_object,
|
| 220 |
+
)
|
| 221 |
+
else:
|
| 222 |
+
return obs.time, obs.qpos_robot, obs.qvel_robot
|
| 223 |
+
|
| 224 |
+
# enforce position specs.
|
| 225 |
+
def ctrl_position_limits(self, ctrl_position):
|
| 226 |
+
ctrl_feasible_position = np.clip(
|
| 227 |
+
ctrl_position,
|
| 228 |
+
self.robot_pos_bound[: self.n_jnt, 0],
|
| 229 |
+
self.robot_pos_bound[: self.n_jnt, 1],
|
| 230 |
+
)
|
| 231 |
+
return ctrl_feasible_position
|
| 232 |
+
|
| 233 |
+
# step the robot env
|
| 234 |
+
def step(self, env, ctrl_desired, step_duration, sim_override=False):
|
| 235 |
+
# Populate observation cache during startup
|
| 236 |
+
if env.initializing:
|
| 237 |
+
self._observation_cache_refresh(env)
|
| 238 |
+
|
| 239 |
+
# enforce velocity limits
|
| 240 |
+
ctrl_feasible = self.ctrl_velocity_limits(ctrl_desired, step_duration)
|
| 241 |
+
|
| 242 |
+
# enforce position limits
|
| 243 |
+
ctrl_feasible = self.ctrl_position_limits(ctrl_feasible)
|
| 244 |
+
|
| 245 |
+
# Send controls to the robot
|
| 246 |
+
if self.is_hardware and (not sim_override):
|
| 247 |
+
raise NotImplementedError()
|
| 248 |
+
else:
|
| 249 |
+
env.do_simulation(
|
| 250 |
+
ctrl_feasible, int(step_duration / env.sim.model.opt.timestep)
|
| 251 |
+
) # render is folded in here
|
| 252 |
+
|
| 253 |
+
# Update current robot state on the overlay
|
| 254 |
+
if self.overlay:
|
| 255 |
+
env.sim.data.qpos[self.n_jnt : 2 * self.n_jnt] = env.desired_pose.copy()
|
| 256 |
+
env.sim.forward()
|
| 257 |
+
|
| 258 |
+
# synchronize time
|
| 259 |
+
if self.is_hardware:
|
| 260 |
+
time_now = time.time() - self.time_start
|
| 261 |
+
time_left_in_step = step_duration - (time_now - self.time)
|
| 262 |
+
if time_left_in_step > 0.0001:
|
| 263 |
+
time.sleep(time_left_in_step)
|
| 264 |
+
return 1
|
| 265 |
+
|
| 266 |
+
def reset(
|
| 267 |
+
self,
|
| 268 |
+
env,
|
| 269 |
+
reset_pose,
|
| 270 |
+
reset_vel,
|
| 271 |
+
overlay_mimic_reset_pose=True,
|
| 272 |
+
sim_override=False,
|
| 273 |
+
):
|
| 274 |
+
reset_pose = self.clip_positions(reset_pose)
|
| 275 |
+
|
| 276 |
+
if self.is_hardware:
|
| 277 |
+
raise NotImplementedError()
|
| 278 |
+
else:
|
| 279 |
+
env.sim.reset()
|
| 280 |
+
env.sim.data.qpos[: self.n_jnt] = reset_pose[: self.n_jnt].copy()
|
| 281 |
+
env.sim.data.qvel[: self.n_jnt] = reset_vel[: self.n_jnt].copy()
|
| 282 |
+
if self.has_obj:
|
| 283 |
+
env.sim.data.qpos[-self.n_obj :] = reset_pose[-self.n_obj :].copy()
|
| 284 |
+
env.sim.data.qvel[-self.n_obj :] = reset_vel[-self.n_obj :].copy()
|
| 285 |
+
env.sim.forward()
|
| 286 |
+
|
| 287 |
+
if self.overlay:
|
| 288 |
+
env.sim.data.qpos[self.n_jnt : 2 * self.n_jnt] = env.desired_pose[
|
| 289 |
+
: self.n_jnt
|
| 290 |
+
].copy()
|
| 291 |
+
env.sim.forward()
|
| 292 |
+
|
| 293 |
+
# refresh observation cache before exit
|
| 294 |
+
self._observation_cache_refresh(env)
|
| 295 |
+
|
| 296 |
+
def close(self):
|
| 297 |
+
if self.is_hardware:
|
| 298 |
+
cprint(
|
| 299 |
+
"Closing Franka hardware... ", "white", "on_grey", end="", flush=True
|
| 300 |
+
)
|
| 301 |
+
status = 0
|
| 302 |
+
raise NotImplementedError()
|
| 303 |
+
cprint("Closed (Status: {})".format(status), "white", "on_grey", flush=True)
|
| 304 |
+
else:
|
| 305 |
+
cprint("Closing Franka sim", "white", "on_grey", flush=True)
|
| 306 |
+
|
| 307 |
+
|
| 308 |
+
class Robot_PosAct(Robot):
|
| 309 |
+
# enforce velocity sepcs.
|
| 310 |
+
# ALERT: This depends on previous observation. This is not ideal as it breaks MDP addumptions. Be careful
|
| 311 |
+
def ctrl_velocity_limits(self, ctrl_position, step_duration):
|
| 312 |
+
last_obs = self.observation_cache[-1]
|
| 313 |
+
ctrl_desired_vel = (
|
| 314 |
+
ctrl_position - last_obs.qpos_robot[: self.n_jnt]
|
| 315 |
+
) / step_duration
|
| 316 |
+
|
| 317 |
+
ctrl_feasible_vel = np.clip(
|
| 318 |
+
ctrl_desired_vel,
|
| 319 |
+
self.robot_vel_bound[: self.n_jnt, 0],
|
| 320 |
+
self.robot_vel_bound[: self.n_jnt, 1],
|
| 321 |
+
)
|
| 322 |
+
ctrl_feasible_position = (
|
| 323 |
+
last_obs.qpos_robot[: self.n_jnt] + ctrl_feasible_vel * step_duration
|
| 324 |
+
)
|
| 325 |
+
return ctrl_feasible_position
|
| 326 |
+
|
| 327 |
+
|
| 328 |
+
class Robot_VelAct(Robot):
|
| 329 |
+
# enforce velocity sepcs.
|
| 330 |
+
# ALERT: This depends on previous observation. This is not ideal as it breaks MDP addumptions. Be careful
|
| 331 |
+
def ctrl_velocity_limits(self, ctrl_velocity, step_duration):
|
| 332 |
+
last_obs = self.observation_cache[-1]
|
| 333 |
+
|
| 334 |
+
ctrl_feasible_vel = np.clip(
|
| 335 |
+
ctrl_velocity,
|
| 336 |
+
self.robot_vel_bound[: self.n_jnt, 0],
|
| 337 |
+
self.robot_vel_bound[: self.n_jnt, 1],
|
| 338 |
+
)
|
| 339 |
+
ctrl_feasible_position = (
|
| 340 |
+
last_obs.qpos_robot[: self.n_jnt] + ctrl_feasible_vel * step_duration
|
| 341 |
+
)
|
| 342 |
+
return ctrl_feasible_position
|
uvd/envs/franka_kitchen/relay-policy-learning/adept_envs/adept_envs/mujoco_env.py
ADDED
|
@@ -0,0 +1,222 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""Base environment for MuJoCo-based environments."""
|
| 2 |
+
|
| 3 |
+
#!/usr/bin/python
|
| 4 |
+
#
|
| 5 |
+
# Copyright 2020 Google LLC
|
| 6 |
+
#
|
| 7 |
+
# Licensed under the Apache License, Version 2.0 (the "License");
|
| 8 |
+
# you may not use this file except in compliance with the License.
|
| 9 |
+
# You may obtain a copy of the License at
|
| 10 |
+
#
|
| 11 |
+
# http://www.apache.org/licenses/LICENSE-2.0
|
| 12 |
+
#
|
| 13 |
+
# Unless required by applicable law or agreed to in writing, software
|
| 14 |
+
# distributed under the License is distributed on an "AS IS" BASIS,
|
| 15 |
+
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
| 16 |
+
# See the License for the specific language governing permissions and
|
| 17 |
+
# limitations under the License.
|
| 18 |
+
|
| 19 |
+
|
| 20 |
+
import collections
|
| 21 |
+
from collections import abc
|
| 22 |
+
import os
|
| 23 |
+
from typing import Dict, Optional
|
| 24 |
+
|
| 25 |
+
import gym
|
| 26 |
+
import numpy as np
|
| 27 |
+
from gym import spaces
|
| 28 |
+
from gym.utils import seeding
|
| 29 |
+
|
| 30 |
+
from adept_envs.simulation.sim_robot import MujocoSimRobot, RenderMode
|
| 31 |
+
|
| 32 |
+
# from simulation.renderer import RenderMode
|
| 33 |
+
|
| 34 |
+
DEFAULT_RENDER_SIZE = 480
|
| 35 |
+
|
| 36 |
+
USE_DM_CONTROL = True
|
| 37 |
+
|
| 38 |
+
|
| 39 |
+
class MujocoEnv(gym.Env):
|
| 40 |
+
"""Superclass for all MuJoCo environments."""
|
| 41 |
+
|
| 42 |
+
def __init__(
|
| 43 |
+
self,
|
| 44 |
+
model_path: str,
|
| 45 |
+
frame_skip: int,
|
| 46 |
+
camera_settings: Optional[Dict] = None,
|
| 47 |
+
use_dm_backend: Optional[bool] = None,
|
| 48 |
+
):
|
| 49 |
+
"""Initializes a new MuJoCo environment.
|
| 50 |
+
|
| 51 |
+
Args:
|
| 52 |
+
model_path: The path to the MuJoCo XML file.
|
| 53 |
+
frame_skip: The number of simulation steps per environment step. On
|
| 54 |
+
hardware this influences the duration of each environment step.
|
| 55 |
+
camera_settings: Settings to initialize the simulation camera. This
|
| 56 |
+
can contain the keys `distance`, `azimuth`, and `elevation`.
|
| 57 |
+
use_dm_backend: A boolean to switch between mujoco-py and dm_control.
|
| 58 |
+
"""
|
| 59 |
+
self._seed()
|
| 60 |
+
if not os.path.isfile(model_path):
|
| 61 |
+
raise IOError(
|
| 62 |
+
"[MujocoEnv]: Model path does not exist: {}".format(model_path)
|
| 63 |
+
)
|
| 64 |
+
self.frame_skip = frame_skip
|
| 65 |
+
|
| 66 |
+
self.sim_robot = MujocoSimRobot(
|
| 67 |
+
model_path,
|
| 68 |
+
use_dm_backend=use_dm_backend or USE_DM_CONTROL,
|
| 69 |
+
camera_settings=camera_settings,
|
| 70 |
+
)
|
| 71 |
+
self.sim = self.sim_robot.sim
|
| 72 |
+
self.model = self.sim_robot.model
|
| 73 |
+
self.data = self.sim_robot.data
|
| 74 |
+
|
| 75 |
+
self.metadata = {
|
| 76 |
+
"render.modes": ["human", "rgb_array", "depth_array"],
|
| 77 |
+
"video.frames_per_second": int(np.round(1.0 / self.dt)),
|
| 78 |
+
}
|
| 79 |
+
self.mujoco_render_frames = False
|
| 80 |
+
|
| 81 |
+
self.init_qpos = self.data.qpos.ravel().copy()
|
| 82 |
+
self.init_qvel = self.data.qvel.ravel().copy()
|
| 83 |
+
observation, _reward, done, _info = self.step(np.zeros(self.model.nu))
|
| 84 |
+
assert not done
|
| 85 |
+
|
| 86 |
+
bounds = self.model.actuator_ctrlrange.copy()
|
| 87 |
+
act_upper = bounds[:, 1]
|
| 88 |
+
act_lower = bounds[:, 0]
|
| 89 |
+
|
| 90 |
+
# Define the action and observation spaces.
|
| 91 |
+
# HACK: MJRL is still using gym 0.9.x so we can't provide a dtype.
|
| 92 |
+
try:
|
| 93 |
+
self.action_space = spaces.Box(act_lower, act_upper, dtype=np.float32)
|
| 94 |
+
if isinstance(observation, abc.Mapping):
|
| 95 |
+
self.observation_space = spaces.Dict(
|
| 96 |
+
{
|
| 97 |
+
k: spaces.Box(-np.inf, np.inf, shape=v.shape, dtype=np.float32)
|
| 98 |
+
for k, v in observation.items()
|
| 99 |
+
}
|
| 100 |
+
)
|
| 101 |
+
else:
|
| 102 |
+
self.obs_dim = (
|
| 103 |
+
np.sum([o.size for o in observation])
|
| 104 |
+
if type(observation) is tuple
|
| 105 |
+
else observation.size
|
| 106 |
+
)
|
| 107 |
+
self.observation_space = spaces.Box(
|
| 108 |
+
-np.inf, np.inf, observation.shape, dtype=np.float32
|
| 109 |
+
)
|
| 110 |
+
|
| 111 |
+
except TypeError:
|
| 112 |
+
# Fallback case for gym 0.9.x
|
| 113 |
+
self.action_space = spaces.Box(act_lower, act_upper)
|
| 114 |
+
assert not isinstance(
|
| 115 |
+
observation, collections.Mapping
|
| 116 |
+
), "gym 0.9.x does not support dictionary observation."
|
| 117 |
+
self.obs_dim = (
|
| 118 |
+
np.sum([o.size for o in observation])
|
| 119 |
+
if type(observation) is tuple
|
| 120 |
+
else observation.size
|
| 121 |
+
)
|
| 122 |
+
self.observation_space = spaces.Box(-np.inf, np.inf, observation.shape)
|
| 123 |
+
|
| 124 |
+
def seed(self, seed=None): # Compatibility with new gym
|
| 125 |
+
return self._seed(seed)
|
| 126 |
+
|
| 127 |
+
def _seed(self, seed=None):
|
| 128 |
+
self.np_random, seed = seeding.np_random(seed)
|
| 129 |
+
return [seed]
|
| 130 |
+
|
| 131 |
+
# methods to override:
|
| 132 |
+
# ----------------------------
|
| 133 |
+
|
| 134 |
+
def reset_model(self):
|
| 135 |
+
"""Reset the robot degrees of freedom (qpos and qvel).
|
| 136 |
+
|
| 137 |
+
Implement this in each subclass.
|
| 138 |
+
"""
|
| 139 |
+
raise NotImplementedError
|
| 140 |
+
|
| 141 |
+
# -----------------------------
|
| 142 |
+
|
| 143 |
+
def reset(self): # compatibility with new gym
|
| 144 |
+
return self._reset()
|
| 145 |
+
|
| 146 |
+
def _reset(self):
|
| 147 |
+
self.sim.reset()
|
| 148 |
+
self.sim.forward()
|
| 149 |
+
ob = self.reset_model()
|
| 150 |
+
return ob
|
| 151 |
+
|
| 152 |
+
def set_state(self, qpos, qvel):
|
| 153 |
+
assert qpos.shape == (self.model.nq,) and qvel.shape == (self.model.nv,)
|
| 154 |
+
state = self.sim.get_state()
|
| 155 |
+
for i in range(self.model.nq):
|
| 156 |
+
state.qpos[i] = qpos[i]
|
| 157 |
+
for i in range(self.model.nv):
|
| 158 |
+
state.qvel[i] = qvel[i]
|
| 159 |
+
self.sim.set_state(state)
|
| 160 |
+
self.sim.forward()
|
| 161 |
+
|
| 162 |
+
@property
|
| 163 |
+
def dt(self):
|
| 164 |
+
return self.model.opt.timestep * self.frame_skip
|
| 165 |
+
|
| 166 |
+
def do_simulation(self, ctrl, n_frames):
|
| 167 |
+
for i in range(self.model.nu):
|
| 168 |
+
self.sim.data.ctrl[i] = ctrl[i]
|
| 169 |
+
|
| 170 |
+
for _ in range(n_frames):
|
| 171 |
+
self.sim.step()
|
| 172 |
+
|
| 173 |
+
# TODO(michaelahn): Remove this; render should be called separately.
|
| 174 |
+
if self.mujoco_render_frames is True:
|
| 175 |
+
self.mj_render()
|
| 176 |
+
|
| 177 |
+
def render(
|
| 178 |
+
self,
|
| 179 |
+
mode="human",
|
| 180 |
+
width=DEFAULT_RENDER_SIZE,
|
| 181 |
+
height=DEFAULT_RENDER_SIZE,
|
| 182 |
+
camera_id=-1,
|
| 183 |
+
):
|
| 184 |
+
"""Renders the environment.
|
| 185 |
+
|
| 186 |
+
Args:
|
| 187 |
+
mode: The type of rendering to use.
|
| 188 |
+
- 'human': Renders to a graphical window.
|
| 189 |
+
- 'rgb_array': Returns the RGB image as an np.ndarray.
|
| 190 |
+
- 'depth_array': Returns the depth image as an np.ndarray.
|
| 191 |
+
width: The width of the rendered image. This only affects offscreen
|
| 192 |
+
rendering.
|
| 193 |
+
height: The height of the rendered image. This only affects
|
| 194 |
+
offscreen rendering.
|
| 195 |
+
camera_id: The ID of the camera to use. By default, this is the free
|
| 196 |
+
camera. If specified, only affects offscreen rendering.
|
| 197 |
+
"""
|
| 198 |
+
if mode == "human":
|
| 199 |
+
self.sim_robot.renderer.render_to_window()
|
| 200 |
+
elif mode == "rgb_array":
|
| 201 |
+
assert width and height
|
| 202 |
+
return self.sim_robot.renderer.render_offscreen(
|
| 203 |
+
width, height, mode=RenderMode.RGB, camera_id=camera_id
|
| 204 |
+
)
|
| 205 |
+
elif mode == "depth_array":
|
| 206 |
+
assert width and height
|
| 207 |
+
return self.sim_robot.renderer.render_offscreen(
|
| 208 |
+
width, height, mode=RenderMode.DEPTH, camera_id=camera_id
|
| 209 |
+
)
|
| 210 |
+
else:
|
| 211 |
+
raise NotImplementedError(mode)
|
| 212 |
+
|
| 213 |
+
def close(self):
|
| 214 |
+
self.sim_robot.close()
|
| 215 |
+
|
| 216 |
+
def mj_render(self):
|
| 217 |
+
"""Backwards compatibility with MJRL."""
|
| 218 |
+
self.render(mode="human")
|
| 219 |
+
|
| 220 |
+
def state_vector(self):
|
| 221 |
+
state = self.sim.get_state()
|
| 222 |
+
return np.concatenate([state.qpos.flat, state.qvel.flat])
|
uvd/envs/franka_kitchen/relay-policy-learning/adept_envs/adept_envs/robot_env.py
ADDED
|
@@ -0,0 +1,178 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""Base class for robotics environments."""
|
| 2 |
+
|
| 3 |
+
#!/usr/bin/python
|
| 4 |
+
#
|
| 5 |
+
# Copyright 2020 Google LLC
|
| 6 |
+
#
|
| 7 |
+
# Licensed under the Apache License, Version 2.0 (the "License");
|
| 8 |
+
# you may not use this file except in compliance with the License.
|
| 9 |
+
# You may obtain a copy of the License at
|
| 10 |
+
#
|
| 11 |
+
# http://www.apache.org/licenses/LICENSE-2.0
|
| 12 |
+
#
|
| 13 |
+
# Unless required by applicable law or agreed to in writing, software
|
| 14 |
+
# distributed under the License is distributed on an "AS IS" BASIS,
|
| 15 |
+
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
| 16 |
+
# See the License for the specific language governing permissions and
|
| 17 |
+
# limitations under the License.
|
| 18 |
+
|
| 19 |
+
import os
|
| 20 |
+
from typing import Dict, Optional
|
| 21 |
+
|
| 22 |
+
import numpy as np
|
| 23 |
+
|
| 24 |
+
from adept_envs import mujoco_env
|
| 25 |
+
from adept_envs.base_robot import BaseRobot
|
| 26 |
+
from adept_envs.utils.configurable import import_class_from_path
|
| 27 |
+
from adept_envs.utils.constants import MODELS_PATH
|
| 28 |
+
|
| 29 |
+
|
| 30 |
+
class RobotEnv(mujoco_env.MujocoEnv):
|
| 31 |
+
"""Base environment for all adept robots."""
|
| 32 |
+
|
| 33 |
+
# Mapping of robot name to fully qualified class path.
|
| 34 |
+
# e.g. 'robot': 'adept_envs.dclaw.robot.Robot'
|
| 35 |
+
# Subclasses should override this to specify the Robot classes they support.
|
| 36 |
+
ROBOTS = {}
|
| 37 |
+
|
| 38 |
+
# Mapping of device path to the calibration file to use. If the device path
|
| 39 |
+
# is not found, the 'default' key is used.
|
| 40 |
+
# This can be overriden by subclasses.
|
| 41 |
+
CALIBRATION_PATHS = {}
|
| 42 |
+
|
| 43 |
+
def __init__(
|
| 44 |
+
self,
|
| 45 |
+
model_path: str,
|
| 46 |
+
robot: BaseRobot,
|
| 47 |
+
frame_skip: int,
|
| 48 |
+
camera_settings: Optional[Dict] = None,
|
| 49 |
+
):
|
| 50 |
+
"""Initializes a robotics environment.
|
| 51 |
+
|
| 52 |
+
Args:
|
| 53 |
+
model_path: The path to the model to run. Relative paths will be
|
| 54 |
+
interpreted as relative to the 'adept_models' folder.
|
| 55 |
+
robot: The Robot object to use.
|
| 56 |
+
frame_skip: The number of simulation steps per environment step. On
|
| 57 |
+
hardware this influences the duration of each environment step.
|
| 58 |
+
camera_settings: Settings to initialize the simulation camera. This
|
| 59 |
+
can contain the keys `distance`, `azimuth`, and `elevation`.
|
| 60 |
+
"""
|
| 61 |
+
self._robot = robot
|
| 62 |
+
|
| 63 |
+
# Initial pose for first step.
|
| 64 |
+
self.desired_pose = np.zeros(self.n_jnt)
|
| 65 |
+
|
| 66 |
+
if not model_path.startswith("/"):
|
| 67 |
+
model_path = os.path.abspath(os.path.join(MODELS_PATH, model_path))
|
| 68 |
+
|
| 69 |
+
self.remote_viz = None
|
| 70 |
+
|
| 71 |
+
try:
|
| 72 |
+
from adept_envs.utils.remote_viz import RemoteViz
|
| 73 |
+
|
| 74 |
+
self.remote_viz = RemoteViz(model_path)
|
| 75 |
+
except ImportError:
|
| 76 |
+
pass
|
| 77 |
+
|
| 78 |
+
self._initializing = True
|
| 79 |
+
super(RobotEnv, self).__init__(
|
| 80 |
+
model_path, frame_skip, camera_settings=camera_settings
|
| 81 |
+
)
|
| 82 |
+
self._initializing = False
|
| 83 |
+
|
| 84 |
+
@property
|
| 85 |
+
def robot(self):
|
| 86 |
+
return self._robot
|
| 87 |
+
|
| 88 |
+
@property
|
| 89 |
+
def n_jnt(self):
|
| 90 |
+
return self._robot.n_jnt
|
| 91 |
+
|
| 92 |
+
@property
|
| 93 |
+
def n_obj(self):
|
| 94 |
+
return self._robot.n_obj
|
| 95 |
+
|
| 96 |
+
@property
|
| 97 |
+
def skip(self):
|
| 98 |
+
"""Alias for frame_skip.
|
| 99 |
+
|
| 100 |
+
Needed for MJRL.
|
| 101 |
+
"""
|
| 102 |
+
return self.frame_skip
|
| 103 |
+
|
| 104 |
+
@property
|
| 105 |
+
def initializing(self):
|
| 106 |
+
return self._initializing
|
| 107 |
+
|
| 108 |
+
def close_env(self):
|
| 109 |
+
if self._robot is not None:
|
| 110 |
+
self._robot.close()
|
| 111 |
+
|
| 112 |
+
def make_robot(
|
| 113 |
+
self,
|
| 114 |
+
n_jnt,
|
| 115 |
+
n_obj=0,
|
| 116 |
+
is_hardware=False,
|
| 117 |
+
device_name=None,
|
| 118 |
+
legacy=False,
|
| 119 |
+
**kwargs
|
| 120 |
+
):
|
| 121 |
+
"""Creates a new robot for the environment.
|
| 122 |
+
|
| 123 |
+
Args:
|
| 124 |
+
n_jnt: The number of joints in the robot.
|
| 125 |
+
n_obj: The number of object joints in the robot environment.
|
| 126 |
+
is_hardware: Whether to run on hardware or not.
|
| 127 |
+
device_name: The device path for the robot hardware.
|
| 128 |
+
legacy: If true, runs using direct dynamixel communication rather
|
| 129 |
+
than DDS.
|
| 130 |
+
kwargs: See BaseRobot for other parameters.
|
| 131 |
+
|
| 132 |
+
Returns:
|
| 133 |
+
A Robot object.
|
| 134 |
+
"""
|
| 135 |
+
if not self.ROBOTS:
|
| 136 |
+
raise NotImplementedError("Subclasses must override ROBOTS.")
|
| 137 |
+
|
| 138 |
+
if is_hardware and not device_name:
|
| 139 |
+
raise ValueError("Must provide device name if running on hardware.")
|
| 140 |
+
|
| 141 |
+
robot_name = "dds_robot" if not legacy and is_hardware else "robot"
|
| 142 |
+
if robot_name not in self.ROBOTS:
|
| 143 |
+
raise KeyError(
|
| 144 |
+
"Unsupported robot '{}', available: {}".format(
|
| 145 |
+
robot_name, list(self.ROBOTS.keys())
|
| 146 |
+
)
|
| 147 |
+
)
|
| 148 |
+
|
| 149 |
+
cls = import_class_from_path(self.ROBOTS[robot_name])
|
| 150 |
+
|
| 151 |
+
calibration_path = None
|
| 152 |
+
if self.CALIBRATION_PATHS:
|
| 153 |
+
if not device_name:
|
| 154 |
+
calibration_name = "default"
|
| 155 |
+
elif device_name not in self.CALIBRATION_PATHS:
|
| 156 |
+
print(
|
| 157 |
+
'Device "{}" not in CALIBRATION_PATHS; using default.'.format(
|
| 158 |
+
device_name
|
| 159 |
+
)
|
| 160 |
+
)
|
| 161 |
+
calibration_name = "default"
|
| 162 |
+
else:
|
| 163 |
+
calibration_name = device_name
|
| 164 |
+
|
| 165 |
+
calibration_path = self.CALIBRATION_PATHS[calibration_name]
|
| 166 |
+
if not os.path.isfile(calibration_path):
|
| 167 |
+
raise OSError(
|
| 168 |
+
"Could not find calibration file at: {}".format(calibration_path)
|
| 169 |
+
)
|
| 170 |
+
|
| 171 |
+
return cls(
|
| 172 |
+
n_jnt,
|
| 173 |
+
n_obj,
|
| 174 |
+
is_hardware=is_hardware,
|
| 175 |
+
device_name=device_name,
|
| 176 |
+
calibration_path=calibration_path,
|
| 177 |
+
**kwargs
|
| 178 |
+
)
|
uvd/envs/franka_kitchen/relay-policy-learning/adept_envs/adept_envs/simulation/module.py
ADDED
|
@@ -0,0 +1,135 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
#!/usr/bin/python
|
| 2 |
+
#
|
| 3 |
+
# Copyright 2020 Google LLC
|
| 4 |
+
#
|
| 5 |
+
# Licensed under the Apache License, Version 2.0 (the "License");
|
| 6 |
+
# you may not use this file except in compliance with the License.
|
| 7 |
+
# You may obtain a copy of the License at
|
| 8 |
+
#
|
| 9 |
+
# http://www.apache.org/licenses/LICENSE-2.0
|
| 10 |
+
#
|
| 11 |
+
# Unless required by applicable law or agreed to in writing, software
|
| 12 |
+
# distributed under the License is distributed on an "AS IS" BASIS,
|
| 13 |
+
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
| 14 |
+
# See the License for the specific language governing permissions and
|
| 15 |
+
# limitations under the License.
|
| 16 |
+
"""Module for caching Python modules related to simulation."""
|
| 17 |
+
|
| 18 |
+
import sys
|
| 19 |
+
|
| 20 |
+
_MUJOCO_PY_MODULE = None
|
| 21 |
+
|
| 22 |
+
_DM_MUJOCO_MODULE = None
|
| 23 |
+
_DM_VIEWER_MODULE = None
|
| 24 |
+
_DM_RENDER_MODULE = None
|
| 25 |
+
|
| 26 |
+
_GLFW_MODULE = None
|
| 27 |
+
|
| 28 |
+
|
| 29 |
+
def get_mujoco_py():
|
| 30 |
+
"""Returns the mujoco_py module."""
|
| 31 |
+
global _MUJOCO_PY_MODULE
|
| 32 |
+
if _MUJOCO_PY_MODULE:
|
| 33 |
+
return _MUJOCO_PY_MODULE
|
| 34 |
+
try:
|
| 35 |
+
import mujoco_py
|
| 36 |
+
|
| 37 |
+
# Override the warning function.
|
| 38 |
+
from mujoco_py.builder import cymj
|
| 39 |
+
|
| 40 |
+
cymj.set_warning_callback(_mj_warning_fn)
|
| 41 |
+
except ImportError:
|
| 42 |
+
print(
|
| 43 |
+
"Failed to import mujoco_py. Ensure that mujoco_py (using MuJoCo "
|
| 44 |
+
"v1.50) is installed.",
|
| 45 |
+
file=sys.stderr,
|
| 46 |
+
)
|
| 47 |
+
sys.exit(1)
|
| 48 |
+
_MUJOCO_PY_MODULE = mujoco_py
|
| 49 |
+
return mujoco_py
|
| 50 |
+
|
| 51 |
+
|
| 52 |
+
def get_mujoco_py_mjlib():
|
| 53 |
+
"""Returns the mujoco_py mjlib module."""
|
| 54 |
+
|
| 55 |
+
class MjlibDelegate:
|
| 56 |
+
"""Wrapper that forwards mjlib calls."""
|
| 57 |
+
|
| 58 |
+
def __init__(self, lib):
|
| 59 |
+
self._lib = lib
|
| 60 |
+
|
| 61 |
+
def __getattr__(self, name: str):
|
| 62 |
+
if name.startswith("mj"):
|
| 63 |
+
return getattr(self._lib, "_" + name)
|
| 64 |
+
raise AttributeError(name)
|
| 65 |
+
|
| 66 |
+
return MjlibDelegate(get_mujoco_py().cymj)
|
| 67 |
+
|
| 68 |
+
|
| 69 |
+
def get_dm_mujoco():
|
| 70 |
+
"""Returns the DM Control mujoco module."""
|
| 71 |
+
global _DM_MUJOCO_MODULE
|
| 72 |
+
if _DM_MUJOCO_MODULE:
|
| 73 |
+
return _DM_MUJOCO_MODULE
|
| 74 |
+
try:
|
| 75 |
+
from dm_control import mujoco
|
| 76 |
+
except ImportError:
|
| 77 |
+
print(
|
| 78 |
+
"Failed to import dm_control.mujoco. Ensure that dm_control (using "
|
| 79 |
+
"MuJoCo v2.00) is installed.",
|
| 80 |
+
file=sys.stderr,
|
| 81 |
+
)
|
| 82 |
+
sys.exit(1)
|
| 83 |
+
_DM_MUJOCO_MODULE = mujoco
|
| 84 |
+
return mujoco
|
| 85 |
+
|
| 86 |
+
|
| 87 |
+
def get_dm_viewer():
|
| 88 |
+
"""Returns the DM Control viewer module."""
|
| 89 |
+
global _DM_VIEWER_MODULE
|
| 90 |
+
if _DM_VIEWER_MODULE:
|
| 91 |
+
return _DM_VIEWER_MODULE
|
| 92 |
+
try:
|
| 93 |
+
from dm_control import viewer
|
| 94 |
+
except ImportError:
|
| 95 |
+
print(
|
| 96 |
+
"Failed to import dm_control.viewer. Ensure that dm_control (using "
|
| 97 |
+
"MuJoCo v2.00) is installed.",
|
| 98 |
+
file=sys.stderr,
|
| 99 |
+
)
|
| 100 |
+
sys.exit(1)
|
| 101 |
+
_DM_VIEWER_MODULE = viewer
|
| 102 |
+
return viewer
|
| 103 |
+
|
| 104 |
+
|
| 105 |
+
def get_dm_render():
|
| 106 |
+
"""Returns the DM Control render module."""
|
| 107 |
+
global _DM_RENDER_MODULE
|
| 108 |
+
if _DM_RENDER_MODULE:
|
| 109 |
+
return _DM_RENDER_MODULE
|
| 110 |
+
try:
|
| 111 |
+
try:
|
| 112 |
+
from dm_control import _render
|
| 113 |
+
|
| 114 |
+
render = _render
|
| 115 |
+
except ImportError:
|
| 116 |
+
print("Warning: DM Control is out of date.")
|
| 117 |
+
from dm_control import render
|
| 118 |
+
except ImportError:
|
| 119 |
+
print(
|
| 120 |
+
"Failed to import dm_control.render. Ensure that dm_control (using "
|
| 121 |
+
"MuJoCo v2.00) is installed.",
|
| 122 |
+
file=sys.stderr,
|
| 123 |
+
)
|
| 124 |
+
sys.exit(1)
|
| 125 |
+
_DM_RENDER_MODULE = render
|
| 126 |
+
return render
|
| 127 |
+
|
| 128 |
+
|
| 129 |
+
def _mj_warning_fn(warn_data: bytes):
|
| 130 |
+
"""Warning function override for mujoco_py."""
|
| 131 |
+
print(
|
| 132 |
+
"WARNING: Mujoco simulation is unstable (has NaNs): {}".format(
|
| 133 |
+
warn_data.decode()
|
| 134 |
+
)
|
| 135 |
+
)
|
uvd/envs/franka_kitchen/relay-policy-learning/adept_envs/adept_envs/simulation/renderer.py
ADDED
|
@@ -0,0 +1,304 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
#!/usr/bin/python
|
| 2 |
+
#
|
| 3 |
+
# Copyright 2020 Google LLC
|
| 4 |
+
#
|
| 5 |
+
# Licensed under the Apache License, Version 2.0 (the "License");
|
| 6 |
+
# you may not use this file except in compliance with the License.
|
| 7 |
+
# You may obtain a copy of the License at
|
| 8 |
+
#
|
| 9 |
+
# http://www.apache.org/licenses/LICENSE-2.0
|
| 10 |
+
#
|
| 11 |
+
# Unless required by applicable law or agreed to in writing, software
|
| 12 |
+
# distributed under the License is distributed on an "AS IS" BASIS,
|
| 13 |
+
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
| 14 |
+
# See the License for the specific language governing permissions and
|
| 15 |
+
# limitations under the License.
|
| 16 |
+
"""Module for viewing Physics objects in the DM Control viewer."""
|
| 17 |
+
|
| 18 |
+
import abc
|
| 19 |
+
import enum
|
| 20 |
+
import sys
|
| 21 |
+
from typing import Dict, Optional
|
| 22 |
+
|
| 23 |
+
import numpy as np
|
| 24 |
+
|
| 25 |
+
from adept_envs.simulation import module
|
| 26 |
+
|
| 27 |
+
# Default window dimensions.
|
| 28 |
+
DEFAULT_WINDOW_WIDTH = 1024
|
| 29 |
+
DEFAULT_WINDOW_HEIGHT = 768
|
| 30 |
+
|
| 31 |
+
DEFAULT_WINDOW_TITLE = "MuJoCo Viewer"
|
| 32 |
+
|
| 33 |
+
_MAX_RENDERBUFFER_SIZE = 2048
|
| 34 |
+
|
| 35 |
+
|
| 36 |
+
class RenderMode(enum.Enum):
|
| 37 |
+
"""Rendering modes for offscreen rendering."""
|
| 38 |
+
|
| 39 |
+
RGB = 0
|
| 40 |
+
DEPTH = 1
|
| 41 |
+
SEGMENTATION = 2
|
| 42 |
+
|
| 43 |
+
|
| 44 |
+
class Renderer(abc.ABC):
|
| 45 |
+
"""Base interface for rendering simulations."""
|
| 46 |
+
|
| 47 |
+
def __init__(self, camera_settings: Optional[Dict] = None):
|
| 48 |
+
self._camera_settings = camera_settings
|
| 49 |
+
|
| 50 |
+
@abc.abstractmethod
|
| 51 |
+
def close(self):
|
| 52 |
+
"""Cleans up any resources being used by the renderer."""
|
| 53 |
+
|
| 54 |
+
@abc.abstractmethod
|
| 55 |
+
def render_to_window(self):
|
| 56 |
+
"""Renders the simulation to a window."""
|
| 57 |
+
|
| 58 |
+
@abc.abstractmethod
|
| 59 |
+
def render_offscreen(
|
| 60 |
+
self,
|
| 61 |
+
width: int,
|
| 62 |
+
height: int,
|
| 63 |
+
mode: RenderMode = RenderMode.RGB,
|
| 64 |
+
camera_id: int = -1,
|
| 65 |
+
) -> np.ndarray:
|
| 66 |
+
"""Renders the camera view as a NumPy array of pixels.
|
| 67 |
+
|
| 68 |
+
Args:
|
| 69 |
+
width: The viewport width (pixels).
|
| 70 |
+
height: The viewport height (pixels).
|
| 71 |
+
mode: The rendering mode.
|
| 72 |
+
camera_id: The ID of the camera to render from. By default, uses
|
| 73 |
+
the free camera.
|
| 74 |
+
|
| 75 |
+
Returns:
|
| 76 |
+
A NumPy array of the pixels.
|
| 77 |
+
"""
|
| 78 |
+
|
| 79 |
+
def _update_camera(self, camera):
|
| 80 |
+
"""Updates the given camera to move to the initial settings."""
|
| 81 |
+
if not self._camera_settings:
|
| 82 |
+
return
|
| 83 |
+
distance = self._camera_settings.get("distance")
|
| 84 |
+
azimuth = self._camera_settings.get("azimuth")
|
| 85 |
+
elevation = self._camera_settings.get("elevation")
|
| 86 |
+
lookat = self._camera_settings.get("lookat")
|
| 87 |
+
|
| 88 |
+
if distance is not None:
|
| 89 |
+
camera.distance = distance
|
| 90 |
+
if azimuth is not None:
|
| 91 |
+
camera.azimuth = azimuth
|
| 92 |
+
if elevation is not None:
|
| 93 |
+
camera.elevation = elevation
|
| 94 |
+
if lookat is not None:
|
| 95 |
+
camera.lookat[:] = lookat
|
| 96 |
+
|
| 97 |
+
|
| 98 |
+
class MjPyRenderer(Renderer):
|
| 99 |
+
"""Class for rendering mujoco_py simulations."""
|
| 100 |
+
|
| 101 |
+
def __init__(self, sim, **kwargs):
|
| 102 |
+
assert isinstance(
|
| 103 |
+
sim, module.get_mujoco_py().MjSim
|
| 104 |
+
), "MjPyRenderer takes a mujoco_py MjSim object."
|
| 105 |
+
super().__init__(**kwargs)
|
| 106 |
+
self._sim = sim
|
| 107 |
+
self._onscreen_renderer = None
|
| 108 |
+
self._offscreen_renderer = None
|
| 109 |
+
|
| 110 |
+
def render_to_window(self):
|
| 111 |
+
"""Renders the simulation to a window."""
|
| 112 |
+
if not self._onscreen_renderer:
|
| 113 |
+
self._onscreen_renderer = module.get_mujoco_py().MjViewer(self._sim)
|
| 114 |
+
self._update_camera(self._onscreen_renderer.cam)
|
| 115 |
+
|
| 116 |
+
self._onscreen_renderer.render()
|
| 117 |
+
|
| 118 |
+
def render_offscreen(
|
| 119 |
+
self,
|
| 120 |
+
width: int,
|
| 121 |
+
height: int,
|
| 122 |
+
mode: RenderMode = RenderMode.RGB,
|
| 123 |
+
camera_id: int = -1,
|
| 124 |
+
) -> np.ndarray:
|
| 125 |
+
"""Renders the camera view as a NumPy array of pixels.
|
| 126 |
+
|
| 127 |
+
Args:
|
| 128 |
+
width: The viewport width (pixels).
|
| 129 |
+
height: The viewport height (pixels).
|
| 130 |
+
mode: The rendering mode.
|
| 131 |
+
camera_id: The ID of the camera to render from. By default, uses
|
| 132 |
+
the free camera.
|
| 133 |
+
|
| 134 |
+
Returns:
|
| 135 |
+
A NumPy array of the pixels.
|
| 136 |
+
"""
|
| 137 |
+
if not self._offscreen_renderer:
|
| 138 |
+
self._offscreen_renderer = module.get_mujoco_py().MjRenderContextOffscreen(
|
| 139 |
+
self._sim
|
| 140 |
+
)
|
| 141 |
+
|
| 142 |
+
# Update the camera configuration for the free-camera.
|
| 143 |
+
if camera_id == -1:
|
| 144 |
+
self._update_camera(self._offscreen_renderer.cam)
|
| 145 |
+
|
| 146 |
+
self._offscreen_renderer.render(width, height, camera_id)
|
| 147 |
+
if mode == RenderMode.RGB:
|
| 148 |
+
data = self._offscreen_renderer.read_pixels(width, height, depth=False)
|
| 149 |
+
# Original image is upside-down, so flip it
|
| 150 |
+
return data[::-1, :, :]
|
| 151 |
+
elif mode == RenderMode.DEPTH:
|
| 152 |
+
data = self._offscreen_renderer.read_pixels(width, height, depth=True)[1]
|
| 153 |
+
# Original image is upside-down, so flip it
|
| 154 |
+
return data[::-1, :]
|
| 155 |
+
else:
|
| 156 |
+
raise NotImplementedError(mode)
|
| 157 |
+
|
| 158 |
+
def close(self):
|
| 159 |
+
"""Cleans up any resources being used by the renderer."""
|
| 160 |
+
|
| 161 |
+
|
| 162 |
+
class DMRenderer(Renderer):
|
| 163 |
+
"""Class for rendering DM Control Physics objects."""
|
| 164 |
+
|
| 165 |
+
def __init__(self, physics, **kwargs):
|
| 166 |
+
assert isinstance(
|
| 167 |
+
physics, module.get_dm_mujoco().Physics
|
| 168 |
+
), "DMRenderer takes a DM Control Physics object."
|
| 169 |
+
super().__init__(**kwargs)
|
| 170 |
+
self._physics = physics
|
| 171 |
+
self._window = None
|
| 172 |
+
|
| 173 |
+
# Set the camera to lookat the center of the geoms. (mujoco_py does
|
| 174 |
+
# this automatically.
|
| 175 |
+
if "lookat" not in self._camera_settings:
|
| 176 |
+
self._camera_settings["lookat"] = [
|
| 177 |
+
np.median(self._physics.data.geom_xpos[:, i]) for i in range(3)
|
| 178 |
+
]
|
| 179 |
+
|
| 180 |
+
def render_to_window(self):
|
| 181 |
+
"""Renders the Physics object to a window.
|
| 182 |
+
|
| 183 |
+
The window continuously renders the Physics in a separate
|
| 184 |
+
thread.
|
| 185 |
+
|
| 186 |
+
This function is a no-op if the window was already created.
|
| 187 |
+
"""
|
| 188 |
+
if not self._window:
|
| 189 |
+
self._window = DMRenderWindow()
|
| 190 |
+
self._window.load_model(self._physics)
|
| 191 |
+
self._update_camera(self._window.camera)
|
| 192 |
+
self._window.run_frame()
|
| 193 |
+
|
| 194 |
+
def render_offscreen(
|
| 195 |
+
self,
|
| 196 |
+
width: int,
|
| 197 |
+
height: int,
|
| 198 |
+
mode: RenderMode = RenderMode.RGB,
|
| 199 |
+
camera_id: int = -1,
|
| 200 |
+
) -> np.ndarray:
|
| 201 |
+
"""Renders the camera view as a NumPy array of pixels.
|
| 202 |
+
|
| 203 |
+
Args:
|
| 204 |
+
width: The viewport width (pixels).
|
| 205 |
+
height: The viewport height (pixels).
|
| 206 |
+
mode: The rendering mode.
|
| 207 |
+
camera_id: The ID of the camera to render from. By default, uses
|
| 208 |
+
the free camera.
|
| 209 |
+
|
| 210 |
+
Returns:
|
| 211 |
+
A NumPy array of the pixels.
|
| 212 |
+
"""
|
| 213 |
+
mujoco = module.get_dm_mujoco()
|
| 214 |
+
# TODO(michaelahn): Consider caching the camera.
|
| 215 |
+
camera = mujoco.Camera(
|
| 216 |
+
physics=self._physics, height=height, width=width, camera_id=camera_id
|
| 217 |
+
)
|
| 218 |
+
|
| 219 |
+
# Update the camera configuration for the free-camera.
|
| 220 |
+
if camera_id == -1:
|
| 221 |
+
self._update_camera(
|
| 222 |
+
camera._render_camera, # pylint: disable=protected-access
|
| 223 |
+
)
|
| 224 |
+
|
| 225 |
+
image = camera.render(
|
| 226 |
+
depth=(mode == RenderMode.DEPTH),
|
| 227 |
+
segmentation=(mode == RenderMode.SEGMENTATION),
|
| 228 |
+
)
|
| 229 |
+
camera._scene.free() # pylint: disable=protected-access
|
| 230 |
+
return image
|
| 231 |
+
|
| 232 |
+
def close(self):
|
| 233 |
+
"""Cleans up any resources being used by the renderer."""
|
| 234 |
+
if self._window:
|
| 235 |
+
self._window.close()
|
| 236 |
+
self._window = None
|
| 237 |
+
|
| 238 |
+
|
| 239 |
+
class DMRenderWindow:
|
| 240 |
+
"""Class that encapsulates a graphical window."""
|
| 241 |
+
|
| 242 |
+
def __init__(
|
| 243 |
+
self,
|
| 244 |
+
width: int = DEFAULT_WINDOW_WIDTH,
|
| 245 |
+
height: int = DEFAULT_WINDOW_HEIGHT,
|
| 246 |
+
title: str = DEFAULT_WINDOW_TITLE,
|
| 247 |
+
):
|
| 248 |
+
"""Creates a graphical render window.
|
| 249 |
+
|
| 250 |
+
Args:
|
| 251 |
+
width: The width of the window.
|
| 252 |
+
height: The height of the window.
|
| 253 |
+
title: The title of the window.
|
| 254 |
+
"""
|
| 255 |
+
dmv = module.get_dm_viewer()
|
| 256 |
+
self._viewport = dmv.renderer.Viewport(width, height)
|
| 257 |
+
self._window = dmv.gui.RenderWindow(width, height, title)
|
| 258 |
+
self._viewer = dmv.viewer.Viewer(
|
| 259 |
+
self._viewport, self._window.mouse, self._window.keyboard
|
| 260 |
+
)
|
| 261 |
+
self._draw_surface = None
|
| 262 |
+
self._renderer = dmv.renderer.NullRenderer()
|
| 263 |
+
|
| 264 |
+
@property
|
| 265 |
+
def camera(self):
|
| 266 |
+
return self._viewer._camera._camera
|
| 267 |
+
|
| 268 |
+
def close(self):
|
| 269 |
+
self._viewer.deinitialize()
|
| 270 |
+
self._renderer.release()
|
| 271 |
+
self._draw_surface.free()
|
| 272 |
+
self._window.close()
|
| 273 |
+
|
| 274 |
+
def load_model(self, physics):
|
| 275 |
+
"""Loads the given Physics object to render."""
|
| 276 |
+
self._viewer.deinitialize()
|
| 277 |
+
|
| 278 |
+
self._draw_surface = module.get_dm_render().Renderer(
|
| 279 |
+
max_width=_MAX_RENDERBUFFER_SIZE, max_height=_MAX_RENDERBUFFER_SIZE
|
| 280 |
+
)
|
| 281 |
+
self._renderer = module.get_dm_viewer().renderer.OffScreenRenderer(
|
| 282 |
+
physics.model, self._draw_surface
|
| 283 |
+
)
|
| 284 |
+
|
| 285 |
+
self._viewer.initialize(physics, self._renderer, touchpad=False)
|
| 286 |
+
|
| 287 |
+
def run_frame(self):
|
| 288 |
+
"""Renders one frame of the simulation.
|
| 289 |
+
|
| 290 |
+
NOTE: This is extremely slow at the moment.
|
| 291 |
+
"""
|
| 292 |
+
glfw = module.get_dm_viewer().gui.glfw_gui.glfw
|
| 293 |
+
glfw_window = self._window._context.window
|
| 294 |
+
if glfw.window_should_close(glfw_window):
|
| 295 |
+
sys.exit(0)
|
| 296 |
+
|
| 297 |
+
self._viewport.set_size(*self._window.shape)
|
| 298 |
+
self._viewer.render()
|
| 299 |
+
pixels = self._renderer.pixels
|
| 300 |
+
|
| 301 |
+
with self._window._context.make_current() as ctx:
|
| 302 |
+
ctx.call(self._window._update_gui_on_render_thread, glfw_window, pixels)
|
| 303 |
+
self._window._mouse.process_events()
|
| 304 |
+
self._window._keyboard.process_events()
|
uvd/envs/franka_kitchen/relay-policy-learning/adept_envs/adept_envs/simulation/sim_robot.py
ADDED
|
@@ -0,0 +1,133 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
#!/usr/bin/python
|
| 2 |
+
#
|
| 3 |
+
# Copyright 2020 Google LLC
|
| 4 |
+
#
|
| 5 |
+
# Licensed under the Apache License, Version 2.0 (the "License");
|
| 6 |
+
# you may not use this file except in compliance with the License.
|
| 7 |
+
# You may obtain a copy of the License at
|
| 8 |
+
#
|
| 9 |
+
# http://www.apache.org/licenses/LICENSE-2.0
|
| 10 |
+
#
|
| 11 |
+
# Unless required by applicable law or agreed to in writing, software
|
| 12 |
+
# distributed under the License is distributed on an "AS IS" BASIS,
|
| 13 |
+
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
| 14 |
+
# See the License for the specific language governing permissions and
|
| 15 |
+
# limitations under the License.
|
| 16 |
+
"""Module for loading MuJoCo models."""
|
| 17 |
+
|
| 18 |
+
import os
|
| 19 |
+
from typing import Dict, Optional
|
| 20 |
+
|
| 21 |
+
from adept_envs.simulation import module
|
| 22 |
+
from adept_envs.simulation.renderer import DMRenderer, MjPyRenderer, RenderMode
|
| 23 |
+
|
| 24 |
+
|
| 25 |
+
class MujocoSimRobot:
|
| 26 |
+
"""Class that encapsulates a MuJoCo simulation.
|
| 27 |
+
|
| 28 |
+
This class exposes methods that are agnostic to the simulation backend.
|
| 29 |
+
Two backends are supported:
|
| 30 |
+
1. mujoco_py - MuJoCo v1.50
|
| 31 |
+
2. dm_control - MuJoCo v2.00
|
| 32 |
+
"""
|
| 33 |
+
|
| 34 |
+
def __init__(
|
| 35 |
+
self,
|
| 36 |
+
model_file: str,
|
| 37 |
+
use_dm_backend: bool = False,
|
| 38 |
+
camera_settings: Optional[Dict] = None,
|
| 39 |
+
):
|
| 40 |
+
"""Initializes a new simulation.
|
| 41 |
+
|
| 42 |
+
Args:
|
| 43 |
+
model_file: The MuJoCo XML model file to load.
|
| 44 |
+
use_dm_backend: If True, uses DM Control's Physics (MuJoCo v2.0) as
|
| 45 |
+
the backend for the simulation. Otherwise, uses mujoco_py (MuJoCo
|
| 46 |
+
v1.5) as the backend.
|
| 47 |
+
camera_settings: Settings to initialize the renderer's camera. This
|
| 48 |
+
can contain the keys `distance`, `azimuth`, and `elevation`.
|
| 49 |
+
"""
|
| 50 |
+
self._use_dm_backend = use_dm_backend
|
| 51 |
+
|
| 52 |
+
if not os.path.isfile(model_file):
|
| 53 |
+
raise ValueError(
|
| 54 |
+
"[MujocoSimRobot] Invalid model file path: {}".format(model_file)
|
| 55 |
+
)
|
| 56 |
+
|
| 57 |
+
if self._use_dm_backend:
|
| 58 |
+
dm_mujoco = module.get_dm_mujoco()
|
| 59 |
+
if model_file.endswith(".mjb"):
|
| 60 |
+
self.sim = dm_mujoco.Physics.from_binary_path(model_file)
|
| 61 |
+
else:
|
| 62 |
+
self.sim = dm_mujoco.Physics.from_xml_path(model_file)
|
| 63 |
+
self.model = self.sim.model
|
| 64 |
+
self._patch_mjlib_accessors(self.model, self.sim.data)
|
| 65 |
+
self.renderer = DMRenderer(self.sim, camera_settings=camera_settings)
|
| 66 |
+
else: # Use mujoco_py
|
| 67 |
+
mujoco_py = module.get_mujoco_py()
|
| 68 |
+
self.model = mujoco_py.load_model_from_path(model_file)
|
| 69 |
+
self.sim = mujoco_py.MjSim(self.model)
|
| 70 |
+
self.renderer = MjPyRenderer(self.sim, camera_settings=camera_settings)
|
| 71 |
+
|
| 72 |
+
self.data = self.sim.data
|
| 73 |
+
|
| 74 |
+
def close(self):
|
| 75 |
+
"""Cleans up any resources being used by the simulation."""
|
| 76 |
+
self.renderer.close()
|
| 77 |
+
|
| 78 |
+
def save_binary(self, path: str):
|
| 79 |
+
"""Saves the loaded model to a binary .mjb file."""
|
| 80 |
+
if os.path.exists(path):
|
| 81 |
+
raise ValueError("[MujocoSimRobot] Path already exists: {}".format(path))
|
| 82 |
+
if not path.endswith(".mjb"):
|
| 83 |
+
path = path + ".mjb"
|
| 84 |
+
if self._use_dm_backend:
|
| 85 |
+
self.model.save_binary(path)
|
| 86 |
+
else:
|
| 87 |
+
with open(path, "wb") as f:
|
| 88 |
+
f.write(self.model.get_mjb())
|
| 89 |
+
|
| 90 |
+
def get_mjlib(self):
|
| 91 |
+
"""Returns an object that exposes the low-level MuJoCo API."""
|
| 92 |
+
if self._use_dm_backend:
|
| 93 |
+
return module.get_dm_mujoco().wrapper.mjbindings.mjlib
|
| 94 |
+
else:
|
| 95 |
+
return module.get_mujoco_py_mjlib()
|
| 96 |
+
|
| 97 |
+
def _patch_mjlib_accessors(self, model, data):
|
| 98 |
+
"""Adds accessors to the DM Control objects to support mujoco_py
|
| 99 |
+
API."""
|
| 100 |
+
assert self._use_dm_backend
|
| 101 |
+
mjlib = self.get_mjlib()
|
| 102 |
+
|
| 103 |
+
def name2id(type_name, name):
|
| 104 |
+
obj_id = mjlib.mj_name2id(
|
| 105 |
+
model.ptr, mjlib.mju_str2Type(type_name.encode()), name.encode()
|
| 106 |
+
)
|
| 107 |
+
if obj_id < 0:
|
| 108 |
+
raise ValueError('No {} with name "{}" exists.'.format(type_name, name))
|
| 109 |
+
return obj_id
|
| 110 |
+
|
| 111 |
+
if not hasattr(model, "body_name2id"):
|
| 112 |
+
model.body_name2id = lambda name: name2id("body", name)
|
| 113 |
+
|
| 114 |
+
if not hasattr(model, "geom_name2id"):
|
| 115 |
+
model.geom_name2id = lambda name: name2id("geom", name)
|
| 116 |
+
|
| 117 |
+
if not hasattr(model, "site_name2id"):
|
| 118 |
+
model.site_name2id = lambda name: name2id("site", name)
|
| 119 |
+
|
| 120 |
+
if not hasattr(model, "joint_name2id"):
|
| 121 |
+
model.joint_name2id = lambda name: name2id("joint", name)
|
| 122 |
+
|
| 123 |
+
if not hasattr(model, "actuator_name2id"):
|
| 124 |
+
model.actuator_name2id = lambda name: name2id("actuator", name)
|
| 125 |
+
|
| 126 |
+
if not hasattr(model, "camera_name2id"):
|
| 127 |
+
model.camera_name2id = lambda name: name2id("camera", name)
|
| 128 |
+
|
| 129 |
+
if not hasattr(data, "body_xpos"):
|
| 130 |
+
data.body_xpos = data.xpos
|
| 131 |
+
|
| 132 |
+
if not hasattr(data, "body_xquat"):
|
| 133 |
+
data.body_xquat = data.xquat
|