@@ -255,12 +255,9 @@ def get_obs(self) -> ArmObsType:
255255 def step (self , action : CartOrJointContType ) -> tuple [ArmObsType , float , bool , bool , dict ]:
256256 action_dict = cast (dict , action )
257257 if (
258- self .get_base_control_mode () == ControlMode .CARTESIAN_TQuat
259- and self .tquat_key not in action_dict
260- or self .get_base_control_mode () == ControlMode .CARTESIAN_TRPY
261- and self .trpy_key not in action_dict
262- or self .get_base_control_mode () == ControlMode .JOINTS
263- and self .joints_key not in action_dict
258+ (self .get_base_control_mode () == ControlMode .CARTESIAN_TQuat and self .tquat_key not in action_dict )
259+ or (self .get_base_control_mode () == ControlMode .CARTESIAN_TRPY and self .trpy_key not in action_dict )
260+ or (self .get_base_control_mode () == ControlMode .JOINTS and self .joints_key not in action_dict )
264261 ):
265262 msg = "Given type is not matching control mode!"
266263 raise RuntimeError (msg )
@@ -336,8 +333,8 @@ def reset(
336333 obs = {}
337334 info = {}
338335
339- seed_ = seed if seed is not None else { key : None for key in self .envs } # type: ignore
340- options_ = options if options is not None else { key : None for key in self .envs } # type: ignore
336+ seed_ = seed if seed is not None else dict . fromkeys ( self .envs ) # type: ignore
337+ options_ = options if options is not None else dict . fromkeys ( self .envs ) # type: ignore
341338 for key , env in self .envs .items ():
342339 obs [key ], info [key ] = env .reset (seed = seed_ [key ], options = options_ [key ])
343340 return obs , info
@@ -459,8 +456,10 @@ def set_origin_to_current(self):
459456 else :
460457 self ._origin = self .unwrapped .robot .get_cartesian_position ()
461458
462- def reset (self , ** kwargs ) -> tuple [dict , dict [str , Any ]]:
463- obs , info = super ().reset (** kwargs )
459+ def reset (
460+ self , * , seed : int | None = None , options : dict [str , Any ] | None = None
461+ ) -> tuple [dict [str , Any ], dict [str , Any ]]:
462+ obs , info = super ().reset (seed = seed , options = options )
464463 self .initial_obs = obs
465464 self .set_origin_to_current ()
466465 self ._last_action = None
@@ -629,7 +628,9 @@ def __init__(self, env, camera_set: BaseCameraSet, include_depth: bool = False):
629628 )
630629 self .camera_key = get_space_keys (CameraDictType )[0 ]
631630
632- def reset (self , seed : int | None = None , options : dict [str , Any ] | None = None ) -> tuple [dict , dict [str , Any ]]:
631+ def reset (
632+ self , * , seed : int | None = None , options : dict [str , Any ] | None = None
633+ ) -> tuple [dict [str , Any ], dict [str , Any ]]:
633634 self .camera_set .clear_buffer ()
634635 return super ().reset (seed = seed , options = options )
635636
0 commit comments