1010 - stack_gym_space: Helper function to repeat/stack Gym spaces
1111"""
1212
13+ import time
1314from typing import Any , Dict , Optional , Tuple
1415
1516import gymnasium as gym
@@ -42,7 +43,9 @@ def stack_gym_space(space: gym.Space, repeat: int) -> gym.Space:
4243 dtype = dtype ,
4344 )
4445 elif isinstance (space , gym .spaces .Dict ):
45- return gym .spaces .Dict ({k : stack_gym_space (v , repeat ) for k , v in space .spaces .items ()})
46+ return gym .spaces .Dict (
47+ {k : stack_gym_space (v , repeat ) for k , v in space .spaces .items ()}
48+ )
4649 else :
4750 raise ValueError (f"Space { space } is not supported." )
4851
@@ -71,7 +74,9 @@ def __init__(self, env: ManipulatorBaseEnv, window_size: int) -> None:
7174 super ().__init__ (env )
7275 self .window_size = window_size
7376 self .window = []
74- self .observation_space = stack_gym_space (self .env .observation_space , self .window_size )
77+ self .observation_space = stack_gym_space (
78+ self .env .observation_space , self .window_size
79+ )
7580
7681 def step (
7782 self , action : NDArray [np .float32 ], ** kwargs : Any
@@ -94,7 +99,8 @@ def step(
9499 self .window .append (obs )
95100 self .window = self .window [- self .window_size :]
96101 obs = {
97- key : np .stack ([frame [key ] for frame in self .window ]) for key in self .window [0 ].keys ()
102+ key : np .stack ([frame [key ] for frame in self .window ])
103+ for key in self .window [0 ].keys ()
98104 }
99105 return obs , float (reward ), terminated , truncated , info
100106
@@ -115,7 +121,8 @@ def reset(
115121 obs , info = self .env .reset (seed = seed , options = options )
116122 self .window = [obs ] * self .window_size
117123 obs = {
118- key : np .stack ([frame [key ] for frame in self .window ]) for key in self .window [0 ].keys ()
124+ key : np .stack ([frame [key ] for frame in self .window ])
125+ for key in self .window [0 ].keys ()
119126 }
120127 return obs , info
121128
@@ -186,7 +193,9 @@ def step(
186193 assert action .shape [0 ] >= self .horizon_length
187194
188195 for i in range (self .horizon_length ):
189- obs , reward , terminated , truncated , info = self .env .step (action [i ], ** kwargs )
196+ obs , reward , terminated , truncated , info = self .env .step (
197+ action [i ], ** kwargs
198+ )
190199 rewards .append (reward )
191200 if terminated or truncated :
192201 break
@@ -212,7 +221,7 @@ def reset(
212221
213222 def close (self ) -> None :
214223 """Clean up the environment's resources."""
215- if rclpy .ok ():
224+ if rclpy .ok (): # pyright: ignore[reportPrivateImportUsage]
216225 rclpy .shutdown ()
217226 self .env .close ()
218227
@@ -226,3 +235,264 @@ def __getattr__(self, name: str) -> Any:
226235 Any: The value of the requested attribute.
227236 """
228237 return getattr (self .env , name )
238+
239+
240+ class ActionTimeStampWrapper (gym .Wrapper ):
241+ def __init__ (self , env ):
242+ super ().__init__ (env )
243+
244+ def step (self , action ) -> tuple [Any , float , bool , bool , dict [str , Any ]]:
245+ time_stamp = time .time ()
246+ observation , reward , terminated , truncated , info = self .env .step (action )
247+ info ["action_t" ] = time_stamp
248+ return observation , float (reward ), terminated , truncated , info
249+
250+
251+ class NoRotationNoGripperActionWrapper (gym .ActionWrapper ):
252+ def __init__ (self , env ):
253+ super ().__init__ (env )
254+ self .action_space = gym .spaces .Box (- np .inf , np .inf , (3 ,))
255+
256+ def action (self , action ):
257+ return np .concatenate ((action , np .zeros (4 )))
258+
259+
260+ class LastObservationWrapper (gym .Wrapper ):
261+ def __init__ (self , env ):
262+ super ().__init__ (env )
263+ self .last_cartesian_state = None
264+ self .last_angular_state = None
265+ self .last_gripper_state = None
266+ self .last_cartesian_error = None
267+ self .last_angular_error = None
268+ self .last_gripper_error = None
269+ self .last_t_obs = None
270+
271+ @staticmethod
272+ def _skew (v : NDArray [np .float64 ]) -> NDArray [np .float64 ]:
273+ return np .array (
274+ [[0.0 , - v [2 ], v [1 ]], [v [2 ], 0.0 , - v [0 ]], [- v [1 ], v [0 ], 0.0 ]],
275+ dtype = np .float64 ,
276+ )
277+
278+ @classmethod
279+ def _rotvec_to_matrix (cls , rotvec : NDArray [np .float64 ]) -> NDArray [np .float64 ]:
280+ theta = float (np .linalg .norm (rotvec ))
281+ if theta < 1e-12 :
282+ return np .eye (3 , dtype = np .float64 ) + cls ._skew (rotvec )
283+
284+ axis = rotvec / theta
285+ k = cls ._skew (axis )
286+ return (
287+ np .eye (3 , dtype = np .float64 )
288+ + np .sin (theta ) * k
289+ + (1.0 - np .cos (theta )) * (k @ k )
290+ )
291+
292+ @staticmethod
293+ def _matrix_to_rotvec (rotation : NDArray [np .float64 ]) -> NDArray [np .float64 ]:
294+ cos_theta = float (np .clip ((np .trace (rotation ) - 1.0 ) / 2.0 , - 1.0 , 1.0 ))
295+ theta = float (np .arccos (cos_theta ))
296+
297+ if theta < 1e-7 :
298+ return 0.5 * np .array (
299+ [
300+ rotation [2 , 1 ] - rotation [1 , 2 ],
301+ rotation [0 , 2 ] - rotation [2 , 0 ],
302+ rotation [1 , 0 ] - rotation [0 , 1 ],
303+ ],
304+ dtype = np .float64 ,
305+ )
306+
307+ sin_theta = float (np .sin (theta ))
308+ if abs (sin_theta ) > 1e-7 :
309+ axis = np .array (
310+ [
311+ rotation [2 , 1 ] - rotation [1 , 2 ],
312+ rotation [0 , 2 ] - rotation [2 , 0 ],
313+ rotation [1 , 0 ] - rotation [0 , 1 ],
314+ ],
315+ dtype = np .float64 ,
316+ ) / (2.0 * sin_theta )
317+ return axis * theta
318+
319+ # Near pi, infer axis from diagonal terms for numerical stability.
320+ axis = np .sqrt (np .maximum ((np .diag (rotation ) + 1.0 ) / 2.0 , 0.0 ))
321+ axis [0 ] = np .copysign (axis [0 ], rotation [2 , 1 ] - rotation [1 , 2 ])
322+ axis [1 ] = np .copysign (axis [1 ], rotation [0 , 2 ] - rotation [2 , 0 ])
323+ axis [2 ] = np .copysign (axis [2 ], rotation [1 , 0 ] - rotation [0 , 1 ])
324+ axis_norm = float (np .linalg .norm (axis ))
325+ if axis_norm < 1e-12 :
326+ return np .array ([theta , 0.0 , 0.0 ], dtype = np .float64 )
327+ return (axis / axis_norm ) * theta
328+
329+ @classmethod
330+ def _relative_rotation_error (
331+ cls ,
332+ target_rotvec : NDArray [np .float64 ],
333+ current_rotvec : NDArray [np .float64 ],
334+ ) -> NDArray [np .float64 ]:
335+ rotation_target = cls ._rotvec_to_matrix (target_rotvec )
336+ rotation_current = cls ._rotvec_to_matrix (current_rotvec )
337+ rotation_error = rotation_target @ rotation_current .T
338+ return cls ._matrix_to_rotvec (rotation_error )
339+
340+ @classmethod
341+ def _relative_angular_velocity (
342+ cls ,
343+ current_rotvec : NDArray [np .float64 ],
344+ previous_rotvec : NDArray [np .float64 ],
345+ dt : float ,
346+ ) -> NDArray [np .float64 ]:
347+ if dt <= 0.0 :
348+ return np .zeros_like (current_rotvec )
349+
350+ rotation_current = cls ._rotvec_to_matrix (current_rotvec )
351+ rotation_previous = cls ._rotvec_to_matrix (previous_rotvec )
352+ delta_rotation = rotation_current @ rotation_previous .T
353+ return cls ._matrix_to_rotvec (delta_rotation ) / dt
354+
355+ def reset (self , * , seed = None , options = None ):
356+ observation , info = self .env .reset (seed = seed , options = options )
357+ current_cartesian_state = observation ["observation.state.cartesian" ][:3 ]
358+ current_angular_state = observation ["observation.state.cartesian" ][3 :].astype (
359+ np .float64
360+ )
361+ target_angular_state = observation ["observation.state.target" ][3 :].astype (
362+ np .float64
363+ )
364+ observation ["observation.previous.action" ] = np .zeros (
365+ self .env .action_space .shape or 7
366+ )
367+ observation ["observation.previous.error.cartesian" ] = np .zeros (3 )
368+ observation ["observation.previous.error.angular" ] = np .zeros (3 )
369+
370+ observation ["observation.velocity.cartesian" ] = np .zeros_like (
371+ current_cartesian_state
372+ )
373+ observation ["observation.velocity.angular" ] = np .zeros_like (
374+ current_angular_state
375+ )
376+ observation ["observation.error.cartesian" ] = (
377+ observation ["observation.state.target" ][:3 ] - current_cartesian_state
378+ )
379+ observation ["observation.error.angular" ] = self ._relative_rotation_error (
380+ target_angular_state , current_angular_state
381+ )
382+
383+ self .last_cartesian_state = current_cartesian_state
384+ self .last_angular_state = current_angular_state
385+ self .last_cartesian_error = observation ["observation.error.cartesian" ]
386+ self .last_angular_error = observation ["observation.error.angular" ]
387+ self .last_t_obs = time .perf_counter ()
388+ return observation , info
389+
390+ def step (self , action ):
391+ observation , reward , terminated , truncated , info = self .env .step (action )
392+ t_obs = time .perf_counter ()
393+ current_cartesian_state = observation ["observation.state.cartesian" ][:3 ]
394+ current_angular_state = observation ["observation.state.cartesian" ][3 :].astype (
395+ np .float64
396+ )
397+ target_angular_state = observation ["observation.state.target" ][3 :].astype (
398+ np .float64
399+ )
400+ dt = (
401+ t_obs - self .last_t_obs
402+ if self .last_cartesian_state is not None and self .last_t_obs is not None
403+ else 0.0
404+ )
405+
406+ observation ["observation.previous.action" ] = action
407+ observation ["observation.previous.error.cartesian" ] = self .last_cartesian_error
408+ observation ["observation.previous.error.angular" ] = self .last_angular_error
409+
410+ observation ["observation.velocity.cartesian" ] = (
411+ (current_cartesian_state - self .last_cartesian_state ) / dt
412+ if self .last_cartesian_state is not None and self .last_t_obs is not None
413+ else np .zeros_like (current_cartesian_state )
414+ )
415+ observation ["observation.velocity.angular" ] = (
416+ self ._relative_angular_velocity (
417+ current_angular_state ,
418+ self .last_angular_state ,
419+ dt ,
420+ )
421+ if self .last_angular_state is not None and self .last_t_obs is not None
422+ else np .zeros_like (current_angular_state )
423+ )
424+ observation ["observation.error.cartesian" ] = (
425+ observation ["observation.state.target" ][:3 ] - current_cartesian_state
426+ )
427+ observation ["observation.error.angular" ] = self ._relative_rotation_error (
428+ target_angular_state , current_angular_state
429+ )
430+
431+ self .last_cartesian_state = current_cartesian_state
432+ self .last_angular_state = current_angular_state
433+ self .last_cartesian_error = observation ["observation.error.cartesian" ]
434+ self .last_angular_error = observation ["observation.error.angular" ]
435+ self .last_t_obs = t_obs
436+ return observation , reward , terminated , truncated , info
437+
438+
439+ # Foundationpose interface wrapper
440+ # on reset:
441+ # reset with custom homing position, gripper open
442+
443+ # A)
444+ # take observation,
445+ # PE:
446+ # pass to SAM3 for segmentation,
447+ # look at average pixel color to determine which block is which => error: set rollout unusable flag -> add to step info
448+ # 2x:
449+ # set param in fp and fpt to correct mesh
450+ # get pose from fp,
451+ # set orientation to prior,
452+ # pass 4x to fpt
453+ # compute relative transform from observed pose to demo pose for grasped block in global frame
454+
455+ # B) compute from stored demo
456+
457+ # go to demo pose with other controller config
458+ # grasp block and move up
459+
460+ # A)
461+ # PE of grasped block
462+ # compute relative transform from grasped block to placed block in global frame
463+
464+ # B)
465+ # use known demo pose of grasped block for delta xy
466+
467+ # switch controller back
468+ # make delta xy 0
469+ # move down until contact (e.g. force threshold)
470+ # while stepping:
471+ # use estimated pose for safety box; step size as parameter
472+ # apply z-force
473+
474+
475+ # crisp gym: no image cropping in ManipulatorEnv
476+
477+ # global safety box: make wrapper instead of modifying env directly
478+
479+ # env = ContainerWatcherWrapper(env, ctx=multiprocessing.get_context("spawn"))
480+ # env = CLIWrapper(env, termination_fn=lambda _obs: False)
481+ # env = StepLimitEnforcerWrapper(env, max_steps=150)
482+
483+ # env = ImageEncoderWrapper(env, n_cameras=1, image_size=(256, 256)) -> add custom cropping
484+ # env = DictObservationToInfoMover(env)
485+ # env = ObservationFormatterWrapper(
486+ # env,
487+ # "cuda",
488+ # keys_ranges_scales=[
489+ # ("observation.previous.action", (0, 2), 10.0),
490+ # ("observation.previous.error.cartesian", (0, 3), 10.0),
491+ # ("observation.velocity.cartesian", (0, 3), 100.0),
492+ # ("observation.error.cartesian", (0, 3), 10.0),
493+ # ("observation.images.wrist_camera", (0, 512), 1.0),
494+ # # ('observation.images.side_camera', (0, 512), 1.0)
495+ # ],
496+ # )
497+
498+ # automatic termination? -> not yet
0 commit comments