ocbench 1.0.0__py3-none-any.whl

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
Files changed (152) hide show
  1. ocbench/__init__.py +440 -0
  2. ocbench/controllers/__init__.py +3 -0
  3. ocbench/controllers/diff_ik.py +115 -0
  4. ocbench/dataset_utils.py +240 -0
  5. ocbench/descriptions/bowling/assets/bowling_pin.blend +0 -0
  6. ocbench/descriptions/bowling/assets/bowling_pin.obj +5465 -0
  7. ocbench/descriptions/bowling/assets/bowling_pin_body.obj +5384 -0
  8. ocbench/descriptions/bowling/assets/bowling_pin_stripe.obj +4614 -0
  9. ocbench/descriptions/bowling_outer.xml +28 -0
  10. ocbench/descriptions/button_inner.xml +26 -0
  11. ocbench/descriptions/button_outer.xml +39 -0
  12. ocbench/descriptions/buttons.xml +84 -0
  13. ocbench/descriptions/cube_inner.xml +12 -0
  14. ocbench/descriptions/cube_outer.xml +12 -0
  15. ocbench/descriptions/drawer.xml +61 -0
  16. ocbench/descriptions/floor.xml +21 -0
  17. ocbench/descriptions/hanoi_disk.xml +12 -0
  18. ocbench/descriptions/hanoi_outer.xml +20 -0
  19. ocbench/descriptions/hanoi_peg.xml +8 -0
  20. ocbench/descriptions/metaworld/button/metal1.png +0 -0
  21. ocbench/descriptions/metaworld/button/stopbot.stl +0 -0
  22. ocbench/descriptions/metaworld/button/stopbutton.stl +0 -0
  23. ocbench/descriptions/metaworld/button/stopbuttonrim.stl +0 -0
  24. ocbench/descriptions/metaworld/button/stopbuttonrod.stl +0 -0
  25. ocbench/descriptions/metaworld/button/stoptop.stl +0 -0
  26. ocbench/descriptions/metaworld/drawer/drawer.stl +0 -0
  27. ocbench/descriptions/metaworld/drawer/drawercase.stl +0 -0
  28. ocbench/descriptions/metaworld/drawer/drawerhandle.stl +0 -0
  29. ocbench/descriptions/metaworld/window/window_base.stl +0 -0
  30. ocbench/descriptions/metaworld/window/window_frame.stl +0 -0
  31. ocbench/descriptions/metaworld/window/window_h_base.stl +0 -0
  32. ocbench/descriptions/metaworld/window/window_h_frame.stl +0 -0
  33. ocbench/descriptions/metaworld/window/windowa_frame.stl +0 -0
  34. ocbench/descriptions/metaworld/window/windowa_glass.stl +0 -0
  35. ocbench/descriptions/metaworld/window/windowa_h_frame.stl +0 -0
  36. ocbench/descriptions/metaworld/window/windowa_h_glass.stl +0 -0
  37. ocbench/descriptions/metaworld/window/windowb_frame.stl +0 -0
  38. ocbench/descriptions/metaworld/window/windowb_glass.stl +0 -0
  39. ocbench/descriptions/metaworld/window/windowb_h_frame.stl +0 -0
  40. ocbench/descriptions/metaworld/window/windowb_h_glass.stl +0 -0
  41. ocbench/descriptions/robotiq_2f85/2f85.png +0 -0
  42. ocbench/descriptions/robotiq_2f85/2f85.xml +192 -0
  43. ocbench/descriptions/robotiq_2f85/LICENSE +23 -0
  44. ocbench/descriptions/robotiq_2f85/README.md +33 -0
  45. ocbench/descriptions/robotiq_2f85/assets/base.stl +0 -0
  46. ocbench/descriptions/robotiq_2f85/assets/base_mount.stl +0 -0
  47. ocbench/descriptions/robotiq_2f85/assets/coupler.stl +0 -0
  48. ocbench/descriptions/robotiq_2f85/assets/driver.stl +0 -0
  49. ocbench/descriptions/robotiq_2f85/assets/follower.stl +0 -0
  50. ocbench/descriptions/robotiq_2f85/assets/pad.stl +0 -0
  51. ocbench/descriptions/robotiq_2f85/assets/silicone_pad.stl +0 -0
  52. ocbench/descriptions/robotiq_2f85/assets/spring_link.stl +0 -0
  53. ocbench/descriptions/robotiq_2f85/scene.xml +40 -0
  54. ocbench/descriptions/universal_robots_ur5e/LICENSE +26 -0
  55. ocbench/descriptions/universal_robots_ur5e/README.md +36 -0
  56. ocbench/descriptions/universal_robots_ur5e/assets/base_0.obj +15306 -0
  57. ocbench/descriptions/universal_robots_ur5e/assets/base_1.obj +17646 -0
  58. ocbench/descriptions/universal_robots_ur5e/assets/forearm_0.obj +45520 -0
  59. ocbench/descriptions/universal_robots_ur5e/assets/forearm_1.obj +2159 -0
  60. ocbench/descriptions/universal_robots_ur5e/assets/forearm_2.obj +25602 -0
  61. ocbench/descriptions/universal_robots_ur5e/assets/forearm_3.obj +29961 -0
  62. ocbench/descriptions/universal_robots_ur5e/assets/shoulder_0.obj +79709 -0
  63. ocbench/descriptions/universal_robots_ur5e/assets/shoulder_1.obj +15676 -0
  64. ocbench/descriptions/universal_robots_ur5e/assets/shoulder_2.obj +65873 -0
  65. ocbench/descriptions/universal_robots_ur5e/assets/upperarm_0.obj +4342 -0
  66. ocbench/descriptions/universal_robots_ur5e/assets/upperarm_1.obj +29711 -0
  67. ocbench/descriptions/universal_robots_ur5e/assets/upperarm_2.obj +99001 -0
  68. ocbench/descriptions/universal_robots_ur5e/assets/upperarm_3.obj +144582 -0
  69. ocbench/descriptions/universal_robots_ur5e/assets/wrist1_0.obj +7761 -0
  70. ocbench/descriptions/universal_robots_ur5e/assets/wrist1_1.obj +66213 -0
  71. ocbench/descriptions/universal_robots_ur5e/assets/wrist1_2.obj +44960 -0
  72. ocbench/descriptions/universal_robots_ur5e/assets/wrist2_0.obj +23375 -0
  73. ocbench/descriptions/universal_robots_ur5e/assets/wrist2_1.obj +57886 -0
  74. ocbench/descriptions/universal_robots_ur5e/assets/wrist2_2.obj +57883 -0
  75. ocbench/descriptions/universal_robots_ur5e/assets/wrist3.obj +5893 -0
  76. ocbench/descriptions/universal_robots_ur5e/scene.xml +23 -0
  77. ocbench/descriptions/universal_robots_ur5e/ur5e.png +0 -0
  78. ocbench/descriptions/universal_robots_ur5e/ur5e.xml +136 -0
  79. ocbench/descriptions/window.xml +92 -0
  80. ocbench/envs/__init__.py +13 -0
  81. ocbench/envs/block_env.py +336 -0
  82. ocbench/envs/bowling_env.py +283 -0
  83. ocbench/envs/chamber_env.py +754 -0
  84. ocbench/envs/env.py +415 -0
  85. ocbench/envs/hanoi_env.py +400 -0
  86. ocbench/envs/manipulation_env.py +423 -0
  87. ocbench/envs/switch_env.py +298 -0
  88. ocbench/lie/__init__.py +12 -0
  89. ocbench/lie/se3.py +152 -0
  90. ocbench/lie/so3.py +190 -0
  91. ocbench/lie/utils.py +36 -0
  92. ocbench/mjcf_utils.py +127 -0
  93. ocbench/mjwarp/__init__.py +53 -0
  94. ocbench/mjwarp/controllers/__init__.py +1 -0
  95. ocbench/mjwarp/controllers/block.py +254 -0
  96. ocbench/mjwarp/controllers/block_kernels.py +626 -0
  97. ocbench/mjwarp/controllers/bowling.py +169 -0
  98. ocbench/mjwarp/controllers/bowling_kernels.py +278 -0
  99. ocbench/mjwarp/controllers/chamber.py +402 -0
  100. ocbench/mjwarp/controllers/chamber_kernels.py +1322 -0
  101. ocbench/mjwarp/controllers/controller_kernels.py +45 -0
  102. ocbench/mjwarp/controllers/hanoi.py +249 -0
  103. ocbench/mjwarp/controllers/hanoi_kernels.py +893 -0
  104. ocbench/mjwarp/controllers/switch.py +223 -0
  105. ocbench/mjwarp/controllers/switch_kernels.py +720 -0
  106. ocbench/mjwarp/envs/__init__.py +1 -0
  107. ocbench/mjwarp/envs/block.py +304 -0
  108. ocbench/mjwarp/envs/block_kernels.py +396 -0
  109. ocbench/mjwarp/envs/bowling.py +341 -0
  110. ocbench/mjwarp/envs/bowling_kernels.py +385 -0
  111. ocbench/mjwarp/envs/chamber.py +596 -0
  112. ocbench/mjwarp/envs/chamber_kernels.py +856 -0
  113. ocbench/mjwarp/envs/hanoi.py +401 -0
  114. ocbench/mjwarp/envs/hanoi_kernels.py +476 -0
  115. ocbench/mjwarp/envs/manipulation.py +693 -0
  116. ocbench/mjwarp/envs/manipulation_kernels.py +307 -0
  117. ocbench/mjwarp/envs/switch.py +342 -0
  118. ocbench/mjwarp/envs/switch_kernels.py +405 -0
  119. ocbench/mjwarp/metadata.py +238 -0
  120. ocbench/mjwarp/primitives/__init__.py +1 -0
  121. ocbench/mjwarp/primitives/bowling.py +54 -0
  122. ocbench/mjwarp/primitives/button.py +18 -0
  123. ocbench/mjwarp/primitives/button_kernels.py +193 -0
  124. ocbench/mjwarp/primitives/cube.py +50 -0
  125. ocbench/mjwarp/primitives/cube_kernels.py +499 -0
  126. ocbench/mjwarp/primitives/drawer.py +18 -0
  127. ocbench/mjwarp/primitives/drawer_kernels.py +196 -0
  128. ocbench/mjwarp/primitives/hanoi.py +44 -0
  129. ocbench/mjwarp/primitives/primitive.py +95 -0
  130. ocbench/mjwarp/primitives/primitive_kernels.py +409 -0
  131. ocbench/mjwarp/primitives/window.py +15 -0
  132. ocbench/mjwarp/primitives/window_kernels.py +199 -0
  133. ocbench/oracles/__init__.py +29 -0
  134. ocbench/oracles/controllers/__init__.py +15 -0
  135. ocbench/oracles/controllers/block.py +259 -0
  136. ocbench/oracles/controllers/bowling.py +68 -0
  137. ocbench/oracles/controllers/chamber.py +394 -0
  138. ocbench/oracles/controllers/controller.py +84 -0
  139. ocbench/oracles/controllers/hanoi.py +194 -0
  140. ocbench/oracles/controllers/switch.py +244 -0
  141. ocbench/oracles/primitives/__init__.py +17 -0
  142. ocbench/oracles/primitives/bowling.py +178 -0
  143. ocbench/oracles/primitives/button.py +133 -0
  144. ocbench/oracles/primitives/cube.py +310 -0
  145. ocbench/oracles/primitives/drawer.py +160 -0
  146. ocbench/oracles/primitives/hanoi.py +261 -0
  147. ocbench/oracles/primitives/primitive.py +297 -0
  148. ocbench/oracles/primitives/window.py +165 -0
  149. ocbench-1.0.0.dist-info/METADATA +28 -0
  150. ocbench-1.0.0.dist-info/RECORD +152 -0
  151. ocbench-1.0.0.dist-info/WHEEL +4 -0
  152. ocbench-1.0.0.dist-info/licenses/LICENSE +21 -0
ocbench/__init__.py ADDED
@@ -0,0 +1,440 @@
1
+ """OCBench manipulation environments."""
2
+
3
+ import gymnasium
4
+ from gymnasium.envs.registration import register, registry
5
+ from ocbench.dataset_utils import download_datasets, get_episode_offsets, load_dataset
6
+
7
+ __all__ = (
8
+ 'download_datasets',
9
+ 'get_episode_offsets',
10
+ 'load_dataset',
11
+ 'make',
12
+ 'make_env_and_datasets',
13
+ 'parse_env_spec',
14
+ )
15
+
16
+ _LITE_KWARGS = dict(
17
+ action_delta_scale=5.0,
18
+ lite=True,
19
+ control_timestep=0.1,
20
+ )
21
+
22
+
23
+ def parse_env_spec(env_name):
24
+ """Return the family, backend, constructor kwargs, and episode limit."""
25
+ visual = env_name.startswith('visual-')
26
+ family, task = env_name.removeprefix('visual-').split('-', 1)
27
+ backend = 'cpu' if task.startswith('cpu-') else 'mjwarp'
28
+ cpu_name = env_name if backend == 'cpu' else f'{family}-cpu-{task}'
29
+ if visual and backend == 'mjwarp':
30
+ cpu_name = f'visual-{cpu_name}'
31
+ spec = gymnasium.spec(cpu_name)
32
+ return family, backend, dict(spec.kwargs), int(spec.max_episode_steps)
33
+
34
+
35
+ def make(env_name, **kwargs):
36
+ """Make a batched MJWarp environment, or a single CPU environment with a `-cpu-` ID.
37
+
38
+ MJWarp uses 1024 worlds by default; pass `nworld` to choose the batch size.
39
+ For MJWarp environments, stop each episode after the step limit returned by `parse_env_spec`.
40
+ CPU environments truncate episodes automatically at their time limit.
41
+ """
42
+ family, backend, env_kwargs, _ = parse_env_spec(env_name)
43
+ if backend == 'cpu':
44
+ return gymnasium.make(env_name, **kwargs)
45
+ from ocbench.mjwarp import make_env
46
+
47
+ nworld = kwargs.pop('nworld', 1024)
48
+ env_kwargs.update(kwargs)
49
+ return make_env(family, env_kwargs, nworld)
50
+
51
+
52
+ def make_env_and_datasets(env_name, dataset_root=None, num_shards=None, **env_kwargs):
53
+ """Make an OCBench environment and load its released datasets.
54
+
55
+ Args:
56
+ env_name: Environment name (e.g., 'block-lite-single-task1-v0').
57
+ dataset_root: Root directory for downloaded datasets. Defaults to '$XDG_CACHE_HOME/ocbench/datasets' if
58
+ XDG_CACHE_HOME is set, or '~/.cache/ocbench/datasets' otherwise.
59
+ num_shards: Number of training shards to load, in filename order. If None, load all shards.
60
+ **env_kwargs: Keyword arguments to pass to make (e.g., nworld=1024 for MJWarp).
61
+
62
+ Returns:
63
+ A tuple of the environment, training dataset, and validation dataset. Each dataset is a dictionary of
64
+ NumPy arrays; see load_dataset for the dataset format.
65
+ """
66
+ train_paths, val_paths = download_datasets(env_name, dataset_root=dataset_root, num_shards=num_shards)
67
+ train_dataset = load_dataset(train_paths)
68
+ val_dataset = load_dataset(val_paths)
69
+ env = make(env_name, **env_kwargs)
70
+ return env, train_dataset, val_dataset
71
+
72
+
73
+ # Block environments.
74
+ register(
75
+ id='block-cpu-single-task1-v0',
76
+ entry_point='ocbench.envs.block_env:BlockEnv',
77
+ max_episode_steps=1250,
78
+ kwargs=dict(env_type='single', task_id=1),
79
+ )
80
+ register(
81
+ id='block-cpu-double-task1-v0',
82
+ entry_point='ocbench.envs.block_env:BlockEnv',
83
+ max_episode_steps=2500,
84
+ kwargs=dict(env_type='double', task_id=1),
85
+ )
86
+ register(
87
+ id='block-cpu-double-task2-v0',
88
+ entry_point='ocbench.envs.block_env:BlockEnv',
89
+ max_episode_steps=2500,
90
+ kwargs=dict(env_type='double', task_id=2),
91
+ )
92
+ register(
93
+ id='block-cpu-triple-task1-v0',
94
+ entry_point='ocbench.envs.block_env:BlockEnv',
95
+ max_episode_steps=3750,
96
+ kwargs=dict(env_type='triple', task_id=1),
97
+ )
98
+ register(
99
+ id='block-cpu-triple-task2-v0',
100
+ entry_point='ocbench.envs.block_env:BlockEnv',
101
+ max_episode_steps=3750,
102
+ kwargs=dict(env_type='triple', task_id=2),
103
+ )
104
+ register(
105
+ id='block-cpu-quadruple-task1-v0',
106
+ entry_point='ocbench.envs.block_env:BlockEnv',
107
+ max_episode_steps=5000,
108
+ kwargs=dict(env_type='quadruple', task_id=1),
109
+ )
110
+ register(
111
+ id='block-cpu-quadruple-task2-v0',
112
+ entry_point='ocbench.envs.block_env:BlockEnv',
113
+ max_episode_steps=5000,
114
+ kwargs=dict(env_type='quadruple', task_id=2),
115
+ )
116
+ register(
117
+ id='block-cpu-quadruple-task3-v0',
118
+ entry_point='ocbench.envs.block_env:BlockEnv',
119
+ max_episode_steps=5000,
120
+ kwargs=dict(env_type='quadruple', task_id=3),
121
+ )
122
+
123
+ # Chamber environments.
124
+ register(
125
+ id='chamber-cpu-easy-task1-v0',
126
+ entry_point='ocbench.envs.chamber_env:ChamberEnv',
127
+ max_episode_steps=4000,
128
+ kwargs=dict(env_type='easy', task_id=1),
129
+ )
130
+ register(
131
+ id='chamber-cpu-easy-task2-v0',
132
+ entry_point='ocbench.envs.chamber_env:ChamberEnv',
133
+ max_episode_steps=4000,
134
+ kwargs=dict(env_type='easy', task_id=2),
135
+ )
136
+ register(
137
+ id='chamber-cpu-easy-task3-v0',
138
+ entry_point='ocbench.envs.chamber_env:ChamberEnv',
139
+ max_episode_steps=4000,
140
+ kwargs=dict(env_type='easy', task_id=3),
141
+ )
142
+ register(
143
+ id='chamber-cpu-medium-task1-v0',
144
+ entry_point='ocbench.envs.chamber_env:ChamberEnv',
145
+ max_episode_steps=5000,
146
+ kwargs=dict(env_type='medium', task_id=1),
147
+ )
148
+ register(
149
+ id='chamber-cpu-medium-task2-v0',
150
+ entry_point='ocbench.envs.chamber_env:ChamberEnv',
151
+ max_episode_steps=5000,
152
+ kwargs=dict(env_type='medium', task_id=2),
153
+ )
154
+ register(
155
+ id='chamber-cpu-hard-task1-v0',
156
+ entry_point='ocbench.envs.chamber_env:ChamberEnv',
157
+ max_episode_steps=5500,
158
+ kwargs=dict(env_type='hard', task_id=1),
159
+ )
160
+ register(
161
+ id='chamber-cpu-hard-task2-v0',
162
+ entry_point='ocbench.envs.chamber_env:ChamberEnv',
163
+ max_episode_steps=5500,
164
+ kwargs=dict(env_type='hard', task_id=2),
165
+ )
166
+
167
+ # Switch environments.
168
+ register(
169
+ id='switch-cpu-3x3-task1-v0',
170
+ entry_point='ocbench.envs.switch_env:SwitchEnv',
171
+ max_episode_steps=1500,
172
+ kwargs=dict(env_type='3x3', task_id=1),
173
+ )
174
+ register(
175
+ id='switch-cpu-3x3-task2-v0',
176
+ entry_point='ocbench.envs.switch_env:SwitchEnv',
177
+ max_episode_steps=1500,
178
+ kwargs=dict(env_type='3x3', task_id=2),
179
+ )
180
+ register(
181
+ id='switch-cpu-4x4-task1-v0',
182
+ entry_point='ocbench.envs.switch_env:SwitchEnv',
183
+ max_episode_steps=1500,
184
+ kwargs=dict(env_type='4x4', task_id=1),
185
+ )
186
+ register(
187
+ id='switch-cpu-4x4-task2-v0',
188
+ entry_point='ocbench.envs.switch_env:SwitchEnv',
189
+ max_episode_steps=1500,
190
+ kwargs=dict(env_type='4x4', task_id=2),
191
+ )
192
+ register(
193
+ id='switch-cpu-5x5-task1-v0',
194
+ entry_point='ocbench.envs.switch_env:SwitchEnv',
195
+ max_episode_steps=4000,
196
+ kwargs=dict(env_type='5x5', task_id=1),
197
+ )
198
+ register(
199
+ id='switch-cpu-5x5-task2-v0',
200
+ entry_point='ocbench.envs.switch_env:SwitchEnv',
201
+ max_episode_steps=4000,
202
+ kwargs=dict(env_type='5x5', task_id=2),
203
+ )
204
+
205
+ # Hanoi environments.
206
+ register(
207
+ id='hanoi-cpu-single-task1-v0',
208
+ entry_point='ocbench.envs.hanoi_env:HanoiEnv',
209
+ max_episode_steps=1000,
210
+ kwargs=dict(env_type='single', task_id=1),
211
+ )
212
+ register(
213
+ id='hanoi-cpu-single-task2-v0',
214
+ entry_point='ocbench.envs.hanoi_env:HanoiEnv',
215
+ max_episode_steps=1000,
216
+ kwargs=dict(env_type='single', task_id=2),
217
+ )
218
+ register(
219
+ id='hanoi-cpu-double-task1-v0',
220
+ entry_point='ocbench.envs.hanoi_env:HanoiEnv',
221
+ max_episode_steps=3000,
222
+ kwargs=dict(env_type='double', task_id=1),
223
+ )
224
+ register(
225
+ id='hanoi-cpu-double-task2-v0',
226
+ entry_point='ocbench.envs.hanoi_env:HanoiEnv',
227
+ max_episode_steps=3000,
228
+ kwargs=dict(env_type='double', task_id=2),
229
+ )
230
+ register(
231
+ id='hanoi-cpu-triple-task1-v0',
232
+ entry_point='ocbench.envs.hanoi_env:HanoiEnv',
233
+ max_episode_steps=7000,
234
+ kwargs=dict(env_type='triple', task_id=1),
235
+ )
236
+ register(
237
+ id='hanoi-cpu-triple-task2-v0',
238
+ entry_point='ocbench.envs.hanoi_env:HanoiEnv',
239
+ max_episode_steps=7000,
240
+ kwargs=dict(env_type='triple', task_id=2),
241
+ )
242
+
243
+ # Bowling environment.
244
+ register(
245
+ id='bowling-cpu-task1-v0',
246
+ entry_point='ocbench.envs.bowling_env:BowlingEnv',
247
+ max_episode_steps=500,
248
+ kwargs=dict(task_id=1),
249
+ )
250
+
251
+ # Lite block environments.
252
+ register(
253
+ id='block-cpu-lite-single-task1-v0',
254
+ entry_point='ocbench.envs.block_env:BlockEnv',
255
+ max_episode_steps=200,
256
+ kwargs=dict(env_type='single', task_id=1, **_LITE_KWARGS),
257
+ )
258
+ register(
259
+ id='block-cpu-lite-double-task1-v0',
260
+ entry_point='ocbench.envs.block_env:BlockEnv',
261
+ max_episode_steps=400,
262
+ kwargs=dict(env_type='double', task_id=1, **_LITE_KWARGS),
263
+ )
264
+ register(
265
+ id='block-cpu-lite-double-task2-v0',
266
+ entry_point='ocbench.envs.block_env:BlockEnv',
267
+ max_episode_steps=400,
268
+ kwargs=dict(env_type='double', task_id=2, **_LITE_KWARGS),
269
+ )
270
+ register(
271
+ id='block-cpu-lite-triple-task1-v0',
272
+ entry_point='ocbench.envs.block_env:BlockEnv',
273
+ max_episode_steps=600,
274
+ kwargs=dict(env_type='triple', task_id=1, **_LITE_KWARGS),
275
+ )
276
+ register(
277
+ id='block-cpu-lite-triple-task2-v0',
278
+ entry_point='ocbench.envs.block_env:BlockEnv',
279
+ max_episode_steps=600,
280
+ kwargs=dict(env_type='triple', task_id=2, **_LITE_KWARGS),
281
+ )
282
+ register(
283
+ id='block-cpu-lite-quadruple-task1-v0',
284
+ entry_point='ocbench.envs.block_env:BlockEnv',
285
+ max_episode_steps=800,
286
+ kwargs=dict(env_type='quadruple', task_id=1, **_LITE_KWARGS),
287
+ )
288
+ register(
289
+ id='block-cpu-lite-quadruple-task2-v0',
290
+ entry_point='ocbench.envs.block_env:BlockEnv',
291
+ max_episode_steps=800,
292
+ kwargs=dict(env_type='quadruple', task_id=2, **_LITE_KWARGS),
293
+ )
294
+ register(
295
+ id='block-cpu-lite-quadruple-task3-v0',
296
+ entry_point='ocbench.envs.block_env:BlockEnv',
297
+ max_episode_steps=800,
298
+ kwargs=dict(env_type='quadruple', task_id=3, **_LITE_KWARGS),
299
+ )
300
+
301
+ # Lite chamber environments.
302
+ register(
303
+ id='chamber-cpu-lite-easy-task1-v0',
304
+ entry_point='ocbench.envs.chamber_env:ChamberEnv',
305
+ max_episode_steps=500,
306
+ kwargs=dict(env_type='easy', task_id=1, **_LITE_KWARGS),
307
+ )
308
+ register(
309
+ id='chamber-cpu-lite-easy-task2-v0',
310
+ entry_point='ocbench.envs.chamber_env:ChamberEnv',
311
+ max_episode_steps=500,
312
+ kwargs=dict(env_type='easy', task_id=2, **_LITE_KWARGS),
313
+ )
314
+ register(
315
+ id='chamber-cpu-lite-easy-task3-v0',
316
+ entry_point='ocbench.envs.chamber_env:ChamberEnv',
317
+ max_episode_steps=500,
318
+ kwargs=dict(env_type='easy', task_id=3, **_LITE_KWARGS),
319
+ )
320
+ register(
321
+ id='chamber-cpu-lite-medium-task1-v0',
322
+ entry_point='ocbench.envs.chamber_env:ChamberEnv',
323
+ max_episode_steps=600,
324
+ kwargs=dict(env_type='medium', task_id=1, **_LITE_KWARGS),
325
+ )
326
+ register(
327
+ id='chamber-cpu-lite-medium-task2-v0',
328
+ entry_point='ocbench.envs.chamber_env:ChamberEnv',
329
+ max_episode_steps=600,
330
+ kwargs=dict(env_type='medium', task_id=2, **_LITE_KWARGS),
331
+ )
332
+ register(
333
+ id='chamber-cpu-lite-hard-task1-v0',
334
+ entry_point='ocbench.envs.chamber_env:ChamberEnv',
335
+ max_episode_steps=700,
336
+ kwargs=dict(env_type='hard', task_id=1, **_LITE_KWARGS),
337
+ )
338
+ register(
339
+ id='chamber-cpu-lite-hard-task2-v0',
340
+ entry_point='ocbench.envs.chamber_env:ChamberEnv',
341
+ max_episode_steps=700,
342
+ kwargs=dict(env_type='hard', task_id=2, **_LITE_KWARGS),
343
+ )
344
+
345
+ # Lite switch environments.
346
+ register(
347
+ id='switch-cpu-lite-3x3-task1-v0',
348
+ entry_point='ocbench.envs.switch_env:SwitchEnv',
349
+ max_episode_steps=300,
350
+ kwargs=dict(env_type='3x3', task_id=1, **_LITE_KWARGS),
351
+ )
352
+ register(
353
+ id='switch-cpu-lite-3x3-task2-v0',
354
+ entry_point='ocbench.envs.switch_env:SwitchEnv',
355
+ max_episode_steps=300,
356
+ kwargs=dict(env_type='3x3', task_id=2, **_LITE_KWARGS),
357
+ )
358
+ register(
359
+ id='switch-cpu-lite-4x4-task1-v0',
360
+ entry_point='ocbench.envs.switch_env:SwitchEnv',
361
+ max_episode_steps=300,
362
+ kwargs=dict(env_type='4x4', task_id=1, **_LITE_KWARGS),
363
+ )
364
+ register(
365
+ id='switch-cpu-lite-4x4-task2-v0',
366
+ entry_point='ocbench.envs.switch_env:SwitchEnv',
367
+ max_episode_steps=300,
368
+ kwargs=dict(env_type='4x4', task_id=2, **_LITE_KWARGS),
369
+ )
370
+ register(
371
+ id='switch-cpu-lite-5x5-task1-v0',
372
+ entry_point='ocbench.envs.switch_env:SwitchEnv',
373
+ max_episode_steps=800,
374
+ kwargs=dict(env_type='5x5', task_id=1, **_LITE_KWARGS),
375
+ )
376
+ register(
377
+ id='switch-cpu-lite-5x5-task2-v0',
378
+ entry_point='ocbench.envs.switch_env:SwitchEnv',
379
+ max_episode_steps=800,
380
+ kwargs=dict(env_type='5x5', task_id=2, **_LITE_KWARGS),
381
+ )
382
+
383
+ # Lite Hanoi environments.
384
+ register(
385
+ id='hanoi-cpu-lite-single-task1-v0',
386
+ entry_point='ocbench.envs.hanoi_env:HanoiEnv',
387
+ max_episode_steps=300,
388
+ kwargs=dict(env_type='single', task_id=1, **_LITE_KWARGS),
389
+ )
390
+ register(
391
+ id='hanoi-cpu-lite-single-task2-v0',
392
+ entry_point='ocbench.envs.hanoi_env:HanoiEnv',
393
+ max_episode_steps=300,
394
+ kwargs=dict(env_type='single', task_id=2, **_LITE_KWARGS),
395
+ )
396
+ register(
397
+ id='hanoi-cpu-lite-double-task1-v0',
398
+ entry_point='ocbench.envs.hanoi_env:HanoiEnv',
399
+ max_episode_steps=900,
400
+ kwargs=dict(env_type='double', task_id=1, **_LITE_KWARGS),
401
+ )
402
+ register(
403
+ id='hanoi-cpu-lite-double-task2-v0',
404
+ entry_point='ocbench.envs.hanoi_env:HanoiEnv',
405
+ max_episode_steps=900,
406
+ kwargs=dict(env_type='double', task_id=2, **_LITE_KWARGS),
407
+ )
408
+ register(
409
+ id='hanoi-cpu-lite-triple-task1-v0',
410
+ entry_point='ocbench.envs.hanoi_env:HanoiEnv',
411
+ max_episode_steps=2100,
412
+ kwargs=dict(env_type='triple', task_id=1, **_LITE_KWARGS),
413
+ )
414
+ register(
415
+ id='hanoi-cpu-lite-triple-task2-v0',
416
+ entry_point='ocbench.envs.hanoi_env:HanoiEnv',
417
+ max_episode_steps=2100,
418
+ kwargs=dict(env_type='triple', task_id=2, **_LITE_KWARGS),
419
+ )
420
+
421
+ # Visual environments.
422
+ for _spec in list(registry.values()):
423
+ if (
424
+ isinstance(_spec.entry_point, str)
425
+ and _spec.entry_point.startswith('ocbench.envs.')
426
+ and not _spec.id.startswith('visual-')
427
+ ):
428
+ register(
429
+ id=f'visual-{_spec.id}',
430
+ entry_point=_spec.entry_point,
431
+ max_episode_steps=_spec.max_episode_steps,
432
+ kwargs=dict(
433
+ **_spec.kwargs,
434
+ ob_type='pixels',
435
+ width=224,
436
+ height=224,
437
+ visualize_info=False,
438
+ pixel_cameras=('front', 'side', 'ur5e/wrist'),
439
+ ),
440
+ )
@@ -0,0 +1,3 @@
1
+ from ocbench.controllers.diff_ik import DiffIKController
2
+
3
+ __all__ = ('DiffIKController',)
@@ -0,0 +1,115 @@
1
+ import mujoco
2
+ import numpy as np
3
+
4
+ PI = np.pi
5
+ PI_2 = 2 * np.pi
6
+
7
+
8
+ def angle_diff(q1: np.ndarray, q2: np.ndarray) -> np.ndarray:
9
+ return np.mod(q1 - q2 + PI, PI_2) - PI
10
+
11
+
12
+ class DiffIKController:
13
+ """Differential inverse kinematics controller."""
14
+
15
+ def __init__(
16
+ self,
17
+ model: mujoco.MjModel,
18
+ sites: list,
19
+ qpos0: np.ndarray = None,
20
+ damping_coeff: float = 1e-12,
21
+ max_angle_change: float = np.radians(45),
22
+ ):
23
+ self._model = model
24
+ self._data = mujoco.MjData(self._model)
25
+ self._qp0 = qpos0
26
+ self._max_angle_change = max_angle_change
27
+
28
+ # Cache references.
29
+ self._ns = len(sites) # Number of sites.
30
+ self._site_ids = np.asarray([self._model.site(s).id for s in sites])
31
+
32
+ # Preallocate arrays.
33
+ self._err = np.empty((self._ns, 6))
34
+ self._site_quat = np.empty((self._ns, 4))
35
+ self._site_quat_inv = np.empty((self._ns, 4))
36
+ self._err_quat = np.empty((self._ns, 4))
37
+ self._jac = np.empty((6 * self._ns, self._model.nv))
38
+ self._damping = damping_coeff * np.eye(6 * self._ns)
39
+ self._eye = np.eye(self._model.nv)
40
+
41
+ def _forward_kinematics(self) -> None:
42
+ """Minimal computation required for forward kinematics."""
43
+ mujoco.mj_kinematics(self._model, self._data)
44
+ mujoco.mj_comPos(self._model, self._data) # Required for mj_jacSite.
45
+
46
+ def _integrate(self, update: np.ndarray) -> None:
47
+ """Integrate the joint velocities in-place."""
48
+ mujoco.mj_integratePos(self._model, self._data.qpos, update, 1.0)
49
+
50
+ def _compute_translational_error(self, pos: np.ndarray) -> None:
51
+ """Compute the error between the desired and current site positions."""
52
+ self._err[:, :3] = pos - self._data.site_xpos[self._site_ids]
53
+
54
+ def _compute_rotational_error(self, quat: np.ndarray) -> None:
55
+ """Compute the error between the desired and current site orientations."""
56
+ for i, site_id in enumerate(self._site_ids):
57
+ mujoco.mju_mat2Quat(self._site_quat[i], self._data.site_xmat[site_id])
58
+ mujoco.mju_negQuat(self._site_quat_inv[i], self._site_quat[i])
59
+ mujoco.mju_mulQuat(self._err_quat[i], quat[i], self._site_quat_inv[i])
60
+ mujoco.mju_quat2Vel(self._err[i, 3:], self._err_quat[i], 1.0)
61
+
62
+ def _compute_jacobian(self) -> None:
63
+ """Update site end-effector Jacobians."""
64
+ for i, site_id in enumerate(self._site_ids):
65
+ jacp = self._jac[6 * i : 6 * i + 3]
66
+ jacr = self._jac[6 * i + 3 : 6 * i + 6]
67
+ mujoco.mj_jacSite(self._model, self._data, jacp, jacr, site_id)
68
+
69
+ def _error_threshold_reached(self, pos_thresh: float, ori_thresh: float) -> bool:
70
+ """Return True if position and rotation errors are below the thresholds."""
71
+ pos_achieved = np.linalg.norm(self._err[:, :3]) <= pos_thresh
72
+ ori_achieved = np.linalg.norm(self._err[:, 3:]) <= ori_thresh
73
+ return pos_achieved and ori_achieved
74
+
75
+ def _solve(self) -> np.ndarray:
76
+ """Solve for joint velocities using damped least squares."""
77
+ H = self._jac @ self._jac.T + self._damping
78
+ x = self._jac.T @ np.linalg.solve(H, self._err.ravel())
79
+ if self._qp0 is not None:
80
+ jac_pinv = np.linalg.pinv(H)
81
+ q_err = angle_diff(self._qp0, self._data.qpos)
82
+ x += (self._eye - (self._jac.T @ jac_pinv) @ self._jac) @ q_err
83
+ return x
84
+
85
+ def _scale_update(self, update: np.ndarray) -> np.ndarray:
86
+ """Scale down update so that the max allowable angle change is not exceeded."""
87
+ update_max = np.max(np.abs(update))
88
+ if update_max > self._max_angle_change:
89
+ update *= self._max_angle_change / update_max
90
+ return update
91
+
92
+ def solve(
93
+ self,
94
+ pos: np.ndarray,
95
+ quat: np.ndarray,
96
+ curr_qpos: np.ndarray,
97
+ max_iters: int = 20,
98
+ pos_thresh: float = 1e-4,
99
+ ori_thresh: float = 1e-4,
100
+ ) -> np.ndarray:
101
+ self._data.qpos = curr_qpos
102
+
103
+ for _ in range(max_iters):
104
+ self._forward_kinematics()
105
+
106
+ self._compute_translational_error(np.atleast_2d(pos))
107
+ self._compute_rotational_error(np.atleast_2d(quat))
108
+ if self._error_threshold_reached(pos_thresh, ori_thresh):
109
+ break
110
+
111
+ self._compute_jacobian()
112
+ update = self._scale_update(self._solve())
113
+ self._integrate(update)
114
+
115
+ return self._data.qpos.copy()