Skip to content

Commit e32beee

Browse files
fix a few issues introduced with refactoring
1 parent 8b54175 commit e32beee

3 files changed

Lines changed: 65 additions & 29 deletions

File tree

crisp_gym/envs/env_wrapper.py

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -17,7 +17,7 @@
1717
import rclpy
1818
from numpy.typing import NDArray
1919

20-
from crisp_gym.envs.envs.manipulator_env import ManipulatorBaseEnv
20+
from crisp_gym.envs.manipulator_env import ManipulatorBaseEnv
2121

2222

2323
def stack_gym_space(space: gym.Space, repeat: int) -> gym.Space:

crisp_gym/envs/manipulator_env.py

Lines changed: 13 additions & 9 deletions
Original file line numberDiff line numberDiff line change
@@ -28,13 +28,14 @@
2828
from pathlib import Path
2929
from typing import Any, List, Tuple
3030

31+
from crisp_py.utils.geometry import OrientationRepresentation
3132
import gymnasium as gym
3233
import numpy as np
3334
import rclpy
3435
from crisp_py.camera import Camera
3536
from crisp_py.gripper import Gripper
3637
from crisp_py.robot import Pose, Robot
37-
from crisp_py.sensors.sensor import make_sensor
38+
from crisp_py.sensors.sensor import Sensor, make_sensor
3839
from numpy.typing import NDArray
3940
from scipy.spatial.transform import Rotation
4041
from typing_extensions import override
@@ -90,7 +91,7 @@ def __init__(self, config: ManipulatorEnvConfig, namespace: str = ""):
9091
for camera_config in self.config.camera_configs
9192
]
9293
self.sensors: List = [
93-
make_sensor(
94+
Sensor(
9495
namespace=namespace,
9596
sensor_config=sensor_config,
9697
)
@@ -432,7 +433,7 @@ def action_to_rotation(self, rot_action: np.ndarray) -> Rotation:
432433
433434
Returns:
434435
Rotation: A scipy Rotation object representing the rotation.
435-
""" # noqa: D205
436+
"""
436437
if self.config.orientation_representation == OrientationRepresentation.EULER:
437438
return Rotation.from_euler("xyz", rot_action)
438439
elif self.config.orientation_representation == OrientationRepresentation.QUATERNION:
@@ -446,8 +447,10 @@ def action_to_rotation(self, rot_action: np.ndarray) -> Rotation:
446447

447448
def clip_position_for_safety(self, position: np.ndarray) -> np.ndarray:
448449
"""Clip the position to ensure safety.
450+
449451
Args:
450452
position (np.ndarray): The position to be clipped.
453+
451454
Returns:
452455
np.ndarray: The clipped position.
453456
"""
@@ -534,7 +537,8 @@ def step(self, action: np.ndarray, block: bool = True) -> Tuple[dict, float, boo
534537
"""Step the environment with a Cartesian action.
535538
536539
Args:
537-
action (np.ndarray): Cartesian delta action [dx, dy, dz, roll, pitch, yaw, gripper_action].
540+
action (np.ndarray): Cartesian delta action [dx, dy, dz, *d_rot_action, gripper_action],
541+
where d_rot_action dimension depends on the chosen orientation representation.
538542
block (bool): If True, block to maintain the control rate.
539543
540544
Returns:
@@ -543,12 +547,12 @@ def step(self, action: np.ndarray, block: bool = True) -> Tuple[dict, float, boo
543547
assert action.shape == self.action_space.shape, (
544548
f"Action shape {action.shape} does not match expected shape {self.action_space.shape}"
545549
)
546-
# assert self.action_space.contains(action), f"Action {action} is not in the action space {self.action_space}"
550+
translation = action[:3]
551+
rotation = self.action_to_rotation(action[3:-1])
547552

548-
translation, rotation = action[:3], Rotation.from_euler("xyz", action[3:6])
549-
550-
target_position = self.robot.target_pose.position + translation
551-
target_position[2] = max(target_position[2], self._min_z_height)
553+
target_position = self.clip_position_for_safety(
554+
self.robot.target_pose.position + translation
555+
)
552556
target_orientation = rotation * self.robot.target_pose.orientation
553557

554558
target_pose = Pose(position=target_position, orientation=target_orientation)

crisp_gym/envs/manipulator_env_config.py

Lines changed: 51 additions & 19 deletions
Original file line numberDiff line numberDiff line change
@@ -27,7 +27,7 @@ class ObservationKeys:
2727
JOINT_OBS = STATE_OBS + ".joints"
2828
CARTESIAN_OBS = STATE_OBS + ".cartesian"
2929
TARGET_OBS = STATE_OBS + ".target"
30-
SENSOR_OBS = STATE_OBS + ".sensors"
30+
SENSOR_OBS = STATE_OBS + ".sensor"
3131

3232
IMAGE_OBS = "observation.images"
3333

@@ -172,22 +172,32 @@ def from_yaml(cls, yaml_path: Path, **overrides) -> "ManipulatorEnvConfig": # n
172172
Returns:
173173
ManipulatorEnvConfig: Configured environment instance
174174
"""
175-
# TODO: @danielsanjosepro Better validation of YAML contents
176175
with open(yaml_path, "r") as f:
177-
data = yaml.safe_load(f) or {}
176+
original_data = yaml.safe_load(f) or {}
178177

179-
# Apply overrides
180-
data.update(overrides)
178+
original_data.update(overrides)
181179

182-
# Handle nested configs that need special treatment
183-
if "robot_config" in data and isinstance(data["robot_config"], dict):
184-
# Use make_robot_config to handle different robot types
185-
data["robot_config"] = make_robot_config(**data["robot_config"])
180+
data = dict(original_data) # Make a shallow copy to modify
186181

187-
if "gripper_config" in data and isinstance(data["gripper_config"], dict):
182+
if "robot_config" in data:
183+
if not isinstance(data["robot_config"], dict):
184+
raise ValueError("robot_config must be a dictionary in the YAML file.")
185+
186+
if "from_yaml" in data["robot_config"]:
187+
robot_yaml_path = find_config(data["robot_config"]["from_yaml"])
188+
if robot_yaml_path is None:
189+
raise FileNotFoundError(
190+
f"Robot config file '{data['robot_config']['from_yaml']}' not found in any CRISP config paths"
191+
)
192+
data["robot_config"] = RobotConfig.from_yaml(yaml_path=robot_yaml_path.resolve())
193+
else:
194+
data["robot_config"] = make_robot_config(**data["robot_config"])
195+
196+
if "gripper_config" in data:
188197
gripper_cfg = data["gripper_config"]
198+
if not isinstance(gripper_cfg, dict):
199+
raise ValueError("gripper_config must be a dictionary in the YAML file.")
189200
if "from_yaml" in gripper_cfg:
190-
# Load from external YAML file
191201
gripper_yaml_path = find_config(gripper_cfg["from_yaml"])
192202
if gripper_yaml_path is None:
193203
raise FileNotFoundError(
@@ -198,16 +208,38 @@ def from_yaml(cls, yaml_path: Path, **overrides) -> "ManipulatorEnvConfig": # n
198208
data["gripper_config"] = GripperConfig(**gripper_cfg)
199209

200210
if "camera_configs" in data and isinstance(data["camera_configs"], list):
201-
data["camera_configs"] = [
202-
CameraConfig(**cam_cfg) if isinstance(cam_cfg, dict) else cam_cfg
203-
for cam_cfg in data["camera_configs"]
204-
]
211+
data["camera_configs"] = [] # Reset to fill in properly
212+
for camera_cfg in original_data["camera_configs"]:
213+
if "from_yaml" in camera_cfg:
214+
camera_yaml_path = find_config(camera_cfg["from_yaml"])
215+
if camera_yaml_path is None:
216+
raise FileNotFoundError(
217+
f"Camera config file '{camera_cfg['from_yaml']}' not found in any CRISP config paths"
218+
)
219+
cam_config = CameraConfig.from_yaml(yaml_path=camera_yaml_path.resolve())
220+
data["camera_configs"].append(cam_config)
221+
else:
222+
data["camera_configs"].append(
223+
CameraConfig(**camera_cfg) if isinstance(camera_cfg, dict) else camera_cfg
224+
)
205225

206226
if "sensor_configs" in data and isinstance(data["sensor_configs"], list):
207-
data["sensor_configs"] = [
208-
SensorConfig(**sensor_cfg) if isinstance(sensor_cfg, dict) else sensor_cfg
209-
for sensor_cfg in data["sensor_configs"]
210-
]
227+
data["sensor_configs"] = []
228+
for sensor_config in original_data["sensor_configs"]:
229+
if "from_yaml" in sensor_config:
230+
sensor_yaml_path = find_config(sensor_config["from_yaml"])
231+
if sensor_yaml_path is None:
232+
raise FileNotFoundError(
233+
f"Sensor config file '{sensor_config['from_yaml']}' not found in any CRISP config paths"
234+
)
235+
sensor_cfg = SensorConfig.from_yaml(yaml_path=sensor_yaml_path.resolve())
236+
data["sensor_configs"].append(sensor_cfg)
237+
else:
238+
data["sensor_configs"].append(
239+
SensorConfig(**sensor_config)
240+
if isinstance(sensor_config, dict)
241+
else sensor_config
242+
)
211243

212244
return cls(**data)
213245

0 commit comments

Comments
 (0)