ryanhoangt commited on
Commit
c456c14
·
verified ·
1 Parent(s): 5956b77

Upload folder using huggingface_hub

Browse files
This view is limited to 50 files because it contains too many changes.   See raw diff
Files changed (50) hide show
  1. .gitattributes +16 -0
  2. .gitignore +169 -0
  3. .gradio/certificate.pem +31 -0
  4. LICENSE +21 -0
  5. README.md +158 -7
  6. app.py +34 -0
  7. datasets/.gitignore +5 -0
  8. datasets/data_gen.py +317 -0
  9. datasets/data_gen.yaml +29 -0
  10. requirements.txt +20 -0
  11. scripts/__init__.py +0 -0
  12. scripts/auto_format.sh +15 -0
  13. scripts/benchmark_decomp.py +63 -0
  14. scripts/benchmark_inference.py +163 -0
  15. scripts/examples/microwave-bottom_burner-light_switch-slide_cabinet.mp4 +0 -0
  16. scripts/tsne_visualization.py +110 -0
  17. setup.py +49 -0
  18. uvd/__init__.py +30 -0
  19. uvd/data/__init__.py +3 -0
  20. uvd/data/dataset_aug.py +310 -0
  21. uvd/data/dataset_base.py +19 -0
  22. uvd/data/franka_kitchen_datasets.py +707 -0
  23. uvd/decomp/__init__.py +1 -0
  24. uvd/decomp/decomp.py +636 -0
  25. uvd/decomp/kernel_reg.py +91 -0
  26. uvd/envs/__init__.py +0 -0
  27. uvd/envs/evaluator/__init__.py +3 -0
  28. uvd/envs/evaluator/evaluator.py +571 -0
  29. uvd/envs/evaluator/inference_wrapper.py +435 -0
  30. uvd/envs/evaluator/vec_envs/__init__.py +0 -0
  31. uvd/envs/evaluator/vec_envs/vec_env.py +398 -0
  32. uvd/envs/evaluator/vec_envs/workers.py +409 -0
  33. uvd/envs/evaluator/visualize_wrapper.py +271 -0
  34. uvd/envs/franka_kitchen/__init__.py +20 -0
  35. uvd/envs/franka_kitchen/franka_kitchen_base.py +446 -0
  36. uvd/envs/franka_kitchen/franka_kitchen_constants.py +62 -0
  37. uvd/envs/franka_kitchen/relay-policy-learning/adept_envs/.pylintrc +433 -0
  38. uvd/envs/franka_kitchen/relay-policy-learning/adept_envs/.style.yapf +323 -0
  39. uvd/envs/franka_kitchen/relay-policy-learning/adept_envs/adept_envs/__init__.py +15 -0
  40. uvd/envs/franka_kitchen/relay-policy-learning/adept_envs/adept_envs/base_robot.py +153 -0
  41. uvd/envs/franka_kitchen/relay-policy-learning/adept_envs/adept_envs/franka/__init__.py +24 -0
  42. uvd/envs/franka_kitchen/relay-policy-learning/adept_envs/adept_envs/franka/assets/franka_kitchen_jntpos_act_ab.xml +94 -0
  43. uvd/envs/franka_kitchen/relay-policy-learning/adept_envs/adept_envs/franka/kitchen_multitask_v0.py +234 -0
  44. uvd/envs/franka_kitchen/relay-policy-learning/adept_envs/adept_envs/franka/robot/franka_config.xml +59 -0
  45. uvd/envs/franka_kitchen/relay-policy-learning/adept_envs/adept_envs/franka/robot/franka_robot.py +342 -0
  46. uvd/envs/franka_kitchen/relay-policy-learning/adept_envs/adept_envs/mujoco_env.py +222 -0
  47. uvd/envs/franka_kitchen/relay-policy-learning/adept_envs/adept_envs/robot_env.py +178 -0
  48. uvd/envs/franka_kitchen/relay-policy-learning/adept_envs/adept_envs/simulation/module.py +135 -0
  49. uvd/envs/franka_kitchen/relay-policy-learning/adept_envs/adept_envs/simulation/renderer.py +304 -0
  50. 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
- pinned: false
 
10
  ---
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
11
 
12
- Check out the configuration reference at https://huggingface.co/docs/hub/spaces-config-reference
 
 
 
 
 
 
 
 
 
 
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