Skip to content

Commit a7e1013

Browse files
committed
gabor handover
1 parent 25debe4 commit a7e1013

6 files changed

Lines changed: 372 additions & 15 deletions

File tree

crisp_gym/envs/env_wrapper.py

Lines changed: 276 additions & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -10,6 +10,7 @@
1010
- stack_gym_space: Helper function to repeat/stack Gym spaces
1111
"""
1212

13+
import time
1314
from typing import Any, Dict, Optional, Tuple
1415

1516
import 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

crisp_gym/envs/manipulator_env.py

Lines changed: 9 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -333,7 +333,7 @@ def _set_gripper_action(self, action: float):
333333
raise ValueError(f"Unsupported gripper mode: {self.config.gripper_mode}")
334334

335335
@override
336-
def step(self, action: np.ndarray, block: bool = False) -> Tuple[dict, float, bool, bool, dict]:
336+
def step(self, action: np.ndarray, block: bool = None) -> Tuple[dict, float, bool, bool, dict]:
337337
"""Step the environment.
338338
339339
Args:
@@ -372,8 +372,8 @@ def reset(
372372
if self.gripper is not None:
373373
self.gripper.enable_torque()
374374

375-
for sensor in self.sensors:
376-
sensor.reset()
375+
# for sensor in self.sensors:
376+
# sensor.reset()
377377

378378
return self._get_obs(), {}
379379

@@ -411,7 +411,7 @@ def home(self, home_config: list[float] | None = None, blocking: bool = True):
411411
blocking (bool): If True, wait until the robot reaches the home position.
412412
"""
413413
if self.config.gripper_mode != GripperMode.NONE:
414-
self.gripper.open()
414+
self.gripper.home()
415415
self.robot.home(home_config=home_config, blocking=blocking)
416416

417417
if not blocking:
@@ -620,7 +620,7 @@ def _get_obs(self) -> dict:
620620
return obs
621621

622622
@override
623-
def step(self, action: np.ndarray, block: bool = True) -> Tuple[dict, float, bool, bool, dict]:
623+
def step(self, action: np.ndarray, block: bool = None) -> Tuple[dict, float, bool, bool, dict]:
624624
"""Step the environment with a Cartesian action.
625625
626626
Args:
@@ -654,9 +654,13 @@ def step(self, action: np.ndarray, block: bool = True) -> Tuple[dict, float, boo
654654

655655
t0 = time.perf_counter()
656656

657+
# print(f"[ManipulatorCartesianEnv] Step with block={block}")
658+
if block is None:
659+
block = self.config.is_blocking
657660
if block:
658661
# FIXME: This control rate sleep is never used and if used by the user
659662
# unexpected behavior occurs.
663+
# print(f"Blocking step for {1.0 / self.config.control_frequency:.3f} seconds")
660664
time.sleep(1.0 / self.config.control_frequency)
661665
# if self.start_time is None:
662666
# self.start_time = time.time()

0 commit comments

Comments
 (0)