@@ -1042,10 +1042,11 @@ def close(self):
10421042class GripperWrapper (ActObsInfoWrapper ):
10431043 # TODO: sticky gripper, like in aloha
10441044
1045+ GRIPPER_THRESHOLD = 0.5
10451046 BINARY_GRIPPER_CLOSED : ClassVar [list [float ]] = [0 ]
10461047 BINARY_GRIPPER_OPEN : ClassVar [list [float ]] = [1 ]
10471048
1048- def __init__ (self , env , gripper : common .Gripper , binary : bool = True ):
1049+ def __init__ (self , env , gripper : common .Gripper , binary : bool = True , prev_action_obs : bool = False ):
10491050 super ().__init__ (env )
10501051 self .binary = binary
10511052 self .observation_space : gym .spaces .Dict
@@ -1055,6 +1056,7 @@ def __init__(self, env, gripper: common.Gripper, binary: bool = True):
10551056 self .gripper_key = get_space_keys (GripperDictType )[0 ]
10561057 self .gripper = gripper
10571058 self ._last_gripper_cmd = None
1059+ self .prev_action_obs = prev_action_obs
10581060
10591061 def _command_changed (self , gripper_action : np .ndarray ) -> bool :
10601062 if self ._last_gripper_cmd is None :
@@ -1098,8 +1100,8 @@ def action(self, action: dict[str, Any]) -> dict[str, Any]:
10981100 gripper_action = np .clip (np .asarray (gripper_action , dtype = np .float32 ), 0.0 , 1.0 )
10991101
11001102 if self ._command_changed (gripper_action ):
1101- if self .binary :
1102- self .gripper .grasp () if gripper_action [0 ] < 0.5 else self .gripper .open ()
1103+ if self .prev_action_obs :
1104+ self .gripper .grasp () if gripper_action [0 ] < self . GRIPPER_THRESHOLD else self .gripper .open ()
11031105 else :
11041106 self .gripper .set_normalized_width (float (gripper_action [0 ]))
11051107 self ._last_gripper_cmd = gripper_action .tolist ()
0 commit comments