@@ -47,8 +47,6 @@ def gui_loop(gui_uuid: str, close_event):
4747
4848
4949class Sim (_Sim ):
50- STATE_SPEC = mj .mjtState .mjSTATE_INTEGRATION
51-
5250 def __init__ (self , mjmdl : str | PathLike | ModelComposer , cfg : SimConfig | None = None ):
5351 if isinstance (mjmdl , ModelComposer ):
5452 self .model = mjmdl .get_model ()
@@ -73,31 +71,38 @@ def __init__(self, mjmdl: str | PathLike | ModelComposer, cfg: SimConfig | None
7371 if cfg is not None :
7472 self .set_config (cfg )
7573
76- def get_state_spec (self ) -> int :
77- return int ( self .STATE_SPEC )
74+ def get_state_spec (self ) -> dict [ str , list [ str ] | list [ int ]] :
75+ return self .get_dynamic_joint_schema ( )
7876
79- def get_state_size (self , spec : int | None = None ) -> int :
80- state_spec = self .STATE_SPEC if spec is None else mj .mjtState (spec )
81- return mj .mj_stateSize (self .model , state_spec )
77+ def get_state_size (self , spec : dict [str , list [str ] | list [int ]] | None = None ) -> int :
78+ state_spec = self .get_state_spec () if spec is None else spec
79+ qpos_size = sum (int (value ) for value in state_spec ["qpos_sizes" ])
80+ qvel_size = sum (int (value ) for value in state_spec ["qvel_sizes" ])
81+ return qpos_size + qvel_size
8282
83- def get_state (self , spec : int | None = None ) -> np .ndarray :
84- state_spec = self .STATE_SPEC if spec is None else mj .mjtState (spec )
85- state = np .empty (self .get_state_size (int (state_spec )), dtype = np .float64 )
86- mj .mj_getState (self .model , self .data , state , state_spec )
87- return state
83+ def get_state (self , spec : dict [str , list [str ] | list [int ]] | None = None ) -> np .ndarray :
84+ del spec
85+ dynamic_joint_state = self .get_dynamic_joint_state ()
86+ return np .concatenate ((dynamic_joint_state ["qpos" ], dynamic_joint_state ["qvel" ]))
8887
89- def set_state (self , state : np .ndarray , spec : int | None = None ):
90- state_spec = self .STATE_SPEC if spec is None else mj .mjtState (spec )
88+ def set_state (
89+ self ,
90+ state : np .ndarray ,
91+ spec : dict [str , list [str ] | list [int ]] | None = None ,
92+ ):
93+ state_spec = self .get_state_spec () if spec is None else spec
9194 state_array = np .asarray (state , dtype = np .float64 )
92- expected_size = self .get_state_size (int ( state_spec ) )
95+ expected_size = self .get_state_size (state_spec )
9396 if state_array .shape != (expected_size ,):
94- msg = (
95- f"Expected MuJoCo state with shape ({ expected_size } ,), "
96- f"got { state_array .shape } for spec { int (state_spec )} ."
97- )
97+ msg = f"Expected state with shape ({ expected_size } ,), got { state_array .shape } ."
9898 raise ValueError (msg )
99- mj .mj_setState (self .model , self .data , state_array , state_spec )
100- mj .mj_forward (self .model , self .data )
99+
100+ qpos_size = sum (int (value ) for value in state_spec ["qpos_sizes" ])
101+ dynamic_joint_state = {
102+ "qpos" : state_array [:qpos_size ],
103+ "qvel" : state_array [qpos_size :],
104+ }
105+ self .set_dynamic_joint_state (state_spec , dynamic_joint_state )
101106
102107 def get_dynamic_joint_schema (self ) -> dict [str , list [str ] | list [int ]]:
103108 schema = super ().get_dynamic_joint_schema ()
0 commit comments