Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
34 changes: 27 additions & 7 deletions mujoco_playground/_src/manipulation/tetheria_hand/rotate_z.py
Original file line number Diff line number Diff line change
Expand Up @@ -75,6 +75,18 @@ def __init__(

def _post_init(self) -> None:
self._hand_qids = mjx_env.get_qpos_ids(self.mj_model, consts.JOINT_NAMES)

# Change: this is the control joints that are used for the policy
self._control_qids = mjx_env.get_qpos_ids(
self.mj_model, consts.CONTROL_JOINT_NAMES
)

control_set = set(self._control_qids)
self._control_qids_bool = jp.array(
[qid in control_set for qid in self._hand_qids],
dtype=bool,
) # Change:boolean mask to check if the joint is a control joint

self._hand_dqids = mjx_env.get_qvel_ids(self.mj_model, consts.JOINT_NAMES)
self._cube_qids = mjx_env.get_qpos_ids(self.mj_model, ["cube_freejoint"])
self._floor_geom_id = self._mj_model.geom("floor").id
Expand All @@ -83,7 +95,7 @@ def _post_init(self) -> None:
home_key = self._mj_model.keyframe("home")
self._init_q = jp.array(home_key.qpos)
self._default_pose = self._init_q[self._hand_qids]
self._lowers, self._uppers = self.mj_model.actuator_ctrlrange.T
self._lowers, self._uppers = self.mj_model.jnt_range[self._hand_qids].T

def reset(self, rng: jax.Array) -> mjx_env.State:
# Randomize hand qpos and qvel.
Expand All @@ -110,7 +122,7 @@ def reset(self, rng: jax.Array) -> mjx_env.State:
self.mjx_model,
qpos=qpos,
qvel=qvel,
ctrl=q_hand,
ctrl=q_hand[self._control_qids_bool], # Change: only use the control joints
mocap_pos=jp.array([-100, -100, -100]), # Hide goal for this task.
)

Expand All @@ -126,13 +138,19 @@ def reset(self, rng: jax.Array) -> mjx_env.State:
for k in self._config.reward_config.scales.keys():
metrics[f"reward/{k}"] = jp.zeros(())

obs_history = jp.zeros(self._config.history_len * 40)
# Change: 35 is the sum of the number of the joints (20) and the number of the control actions (15)
obs_history = jp.zeros(self._config.history_len * 35)
obs = self._get_obs(data, info, obs_history)
reward, done = jp.zeros(2) # pylint: disable=redefined-outer-name
return mjx_env.State(data, obs, reward, done, metrics, info)

def step(self, state: mjx_env.State, action: jax.Array) -> mjx_env.State:
motor_targets = self._default_pose + action * self._config.action_scale
motor_targets = (
self._default_pose[
self._control_qids_bool
] # Change: use the control joints
+ +action * self._config.action_scale
)
# NOTE: no clipping.
data = mjx_env.step(self.mjx_model, state.data, motor_targets, self.n_substeps)
state.info["motor_targets"] = motor_targets
Expand Down Expand Up @@ -174,8 +192,8 @@ def _get_obs(

state = jp.concatenate(
[
noisy_joint_angles, # 16
info["last_act"], # 16
noisy_joint_angles, # Change: 16 (leap hand) to 20 (tetheria hand)
info["last_act"], # Change: 16 (leap hand) to 15 (tetheria hand)
]
) # 48
obs_history = jp.roll(obs_history, state.size)
Expand Down Expand Up @@ -241,7 +259,9 @@ def _cost_torques(self, torques: jax.Array) -> jax.Array:
return jp.sum(jp.square(torques))

def _cost_energy(self, qvel: jax.Array, qfrc_actuator: jax.Array) -> jax.Array:
return jp.sum(jp.abs(qvel) * jp.abs(qfrc_actuator))
return jp.sum(
jp.abs(qvel[self._control_qids_bool]) * jp.abs(qfrc_actuator)
) # Change: only use the control joints

def _cost_linvel(self, cube_linvel: jax.Array) -> jax.Array:
return jp.linalg.norm(cube_linvel, ord=1, axis=-1)
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -56,6 +56,39 @@
# "th_ipl",
]

CONTROL_JOINT_NAMES = [
# index
"right_index_mcp_abd",
"right_index_mcp_flex",
"right_index_pip",
# "right_index_dip",
# "if_dip",
# middle
"right_middle_mcp_abd",
"right_middle_mcp_flex",
"right_middle_pip",
# "right_middle_dip",
# "mf_dip",
# ring
"right_ring_mcp_abd",
"right_ring_mcp_flex",
"right_ring_pip",
# "right_ring_dip",
# "rf_dip",
# pinky
"right_pinky_mcp_abd",
"right_pinky_mcp_flex",
"right_pinky_pip",
# "right_pinky_dip",
# "th_dip",
# thumb
"right_thumb_cmc_abd",
"right_thumb_cmc_flex",
"right_thumb_mcp",
# "right_thumb_ip",
# "th_ipl",
]

# CONTROLJOINT_NAMES = [
# # index
# "right_index_mcp_flex",
Expand Down Expand Up @@ -102,31 +135,31 @@
"right_index_A_mcp_abd",
"right_index_A_mcp_flex",
"right_index_A_pip",
"right_index_A_dip",
# "right_index_A_dip",
# "if_dip_act",
# middle
"right_middle_A_mcp_abd",
"right_middle_A_mcp_flex",
"right_middle_A_pip",
"right_middle_A_dip",
# "right_middle_A_dip",
# "mf_dip_act",
# ring
"right_ring_A_mcp_abd",
"right_ring_A_mcp_flex",
"right_ring_A_pip",
"right_ring_A_dip",
# "right_ring_A_dip",
# "rf_dip_act",
# pinky
"right_pinky_A_mcp_abd",
"right_pinky_A_mcp_flex",
"right_pinky_A_pip",
"right_pinky_A_dip",
# "right_pinky_A_dip",
# "th_dip_act",
# thumb
"right_thumb_A_mcp_abd",
"right_thumb_A_cmc_flex",
"right_thumb_A_pip",
"right_thumb_A_dip",
# "right_thumb_A_dip",
# "th_ipl_act",
# "th_mcp_act",
# "th_ipl_act",
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -361,37 +361,36 @@
<position name="right_index_A_mcp_abd" joint="right_index_mcp_abd" class="mcp" />
<position name="right_index_A_mcp_flex" joint="right_index_mcp_flex" class="rot" />
<position name="right_index_A_pip" joint="right_index_pip" class="dip" />
<position name="right_index_A_dip" joint="right_index_dip" class="pip" />
<!-- <position name="right_index_A_dip" joint="right_index_dip" class="pip" /> -->

<position name="right_middle_A_mcp_abd" joint="right_middle_mcp_abd" class="mcp" />
<position name="right_middle_A_mcp_flex" joint="right_middle_mcp_flex" class="rot" />
<position name="right_middle_A_pip" joint="right_middle_pip" class="dip" />
<position name="right_middle_A_dip" joint="right_middle_dip" class="pip" />
<!-- <position name="right_middle_A_dip" joint="right_middle_dip" class="pip" /> -->

<position name="right_ring_A_mcp_abd" joint="right_ring_mcp_abd" class="mcp" />
<position name="right_ring_A_mcp_flex" joint="right_ring_mcp_flex" class="rot" />
<position name="right_ring_A_pip" joint="right_ring_pip" class="dip" />
<position name="right_ring_A_dip" joint="right_ring_dip" class="pip" />
<!-- <position name="right_ring_A_dip" joint="right_ring_dip" class="pip" /> -->

<position name="right_pinky_A_mcp_abd" joint="right_pinky_mcp_abd" class="mcp" />
<position name="right_pinky_A_mcp_flex" joint="right_pinky_mcp_flex" class="rot" />
<position name="right_pinky_A_pip" joint="right_pinky_pip" class="dip" />
<position name="right_pinky_A_dip" joint="right_pinky_dip" class="pip" />
<!-- <position name="right_pinky_A_dip" joint="right_pinky_dip" class="pip" /> -->

<!-- <position name="rh_A_TBJ0" tendon="rh_TBJ0" class="thumb_ipl" /> -->
<position name="right_thumb_A_mcp_abd" joint="right_thumb_cmc_abd" class="thumb_cmc" />
<position name="right_thumb_A_cmc_flex" joint="right_thumb_cmc_flex" class="thumb_axl" />
<position name="right_thumb_A_pip" joint="right_thumb_mcp" class="thumb_mcp" />
<position name="right_thumb_A_dip" joint="right_thumb_ip" class="thumb_ipl" />
<!-- <position name="right_thumb_A_dip" joint="right_thumb_ip" class="thumb_ipl" /> -->
</actuator>

<!-- <equality> -->
<!-- <weld body1="mocap" body2="tetheria_mount" solref="0.01 1" solimp=".9 .9 0.01" /> -->
<!-- <joint joint1="right_thumb_mcp" joint2="right_thumb_ip"/>
<equality>
<joint joint1="right_thumb_mcp" joint2="right_thumb_ip"/>
<joint joint1="right_index_pip" joint2="right_index_dip"/>
<joint joint1="right_middle_pip" joint2="right_middle_dip"/>
<joint joint1="right_ring_pip" joint2="right_ring_dip"/>
<joint joint1="right_pinky_pip" joint2="right_pinky_dip"/> -->
<!-- </equality> -->
<joint joint1="right_pinky_pip" joint2="right_pinky_dip"/>
</equality>

</mujoco>
Original file line number Diff line number Diff line change
Expand Up @@ -57,11 +57,11 @@
0 0.75 0.75 0.75
0.75 0.2 0.6 0.6
0.12 0.0 0.05 0.810967 -0.00262895 -0.585086 -0.000254303" ctrl="
0 0.75 0.75 0.75
0 0.75 0.75 0.75
0 0.75 0.75 0.75
0 0.75 0.75 0.75
0.75 0.2 0.6 0.6 " mpos="0.25 0.16 0" mquat="1 0 0 0"/>
0 0.75 0.75
0 0.75 0.75
0 0.75 0.75
0 0.75 0.75
0.75 0.2 0.6 " mpos="0.25 0.16 0" mquat="1 0 0 0"/>
</keyframe>
<!-- <keyframe>
<key name="home" qpos="
Expand Down