Skip to content
Merged
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
26 changes: 20 additions & 6 deletions conf/offpolicy/task/flashsac/g1_walk_flat/mujoco.yaml
Original file line number Diff line number Diff line change
Expand Up @@ -5,17 +5,31 @@ training:
algo:
num_envs: 2048
learning_starts: 49
max_iterations: 5000
max_iterations: 6000
save_interval: 1000
updates_per_step: 8
#use_symmetry: true
replay_buffer_n: 1024
env:
control_config:
action_scale: 1.0
gait_phase_init_mode: "offset_phase"
reset_base_qvel_limit: 0.5
noise_config:
scale_gyro: 0.0
scale_gravity: 0.0
scale_joint_angle: 0.01
scale_joint_vel: 0.1
scale_linvel: 0.0
reward:
scales:
tracking_lin_vel: 2.0
tracking_ang_vel: 1.5
penalty_ang_vel_xy: -1.0
penalty_ang_vel_xy: -1.5
penalty_orientation: -10.0
penalty_action_rate: -2.0
pose: -0.5
penalty_feet_ori: -25.0
penalty_action_rate: -4.0
pose: -0.8
penalty_feet_ori: -20.0
feet_phase: 5.0
alive: 10.0
tracking_sigma: 0.25
Expand All @@ -24,6 +38,6 @@ reward:
max_tilt_deg: 65.0
gait_frequency: 1.5
feet_phase_swing_height: 0.09
feet_phase_tracking_sigma: 0.008
feet_phase_tracking_sigma: 0.025
close_feet_threshold: 0.15
pose_weights: [0.01, 1.0, 5.0, 0.01, 5.0, 5.0, 0.01, 1.0, 5.0, 0.01, 5.0, 5.0, 50.0, 50.0, 50.0, 50.0, 50.0, 50.0, 50.0, 50.0, 50.0, 50.0, 50.0, 50.0, 50.0, 50.0, 50.0, 50.0, 50.0]
29 changes: 19 additions & 10 deletions src/unilab/algos/torch/flash_sac/learner.py
Original file line number Diff line number Diff line change
Expand Up @@ -298,28 +298,32 @@ def update_critic(self, batch: dict[str, torch.Tensor]) -> dict[str, float]:

gamma = self.gamma**self.n_step

obs_all = torch.cat([critic_obs, critic_next_obs], dim=0)

with torch.no_grad():
with self._autocast():
next_actions, actor_info = self.actor(next_obs, training=False)
entropy_bonus = -self.temperature().detach() * actor_info["log_prob"]
next_q_values, next_q_info = self.target_critic(
critic_next_obs, next_actions, training=False
)
next_q_log_probs = select_min_q_log_probs(next_q_values, next_q_info["log_prob"])
actor_entropy = self.temperature().detach() * actor_info["log_prob"]
act_all = torch.cat([actions, next_actions], dim=0)
qs_all, q_info_all = self.target_critic(obs_all, act_all, training=True)
next_q_values = qs_all.chunk(2, dim=1)[1]
next_q_log_probs_full = q_info_all["log_prob"].chunk(2, dim=1)[1]
next_q_log_probs = select_min_q_log_probs(next_q_values, next_q_log_probs_full)
support = cast(torch.Tensor, self.target_critic.predictor.support)
target_probs = compute_categorical_td_target(
support=support,
target_log_probs=next_q_log_probs,
reward=rewards,
terminated=terminated,
truncated=truncated,
actor_entropy=entropy_bonus,
actor_entropy=actor_entropy,
gamma=gamma,
)

with self._autocast():
_, pred_info = self.critic(critic_obs, actions, training=True)
critic_loss = -(target_probs.unsqueeze(0) * pred_info["log_prob"]).sum(dim=-1).mean()
_, pred_info_all = self.critic(obs_all, act_all, training=True)
pred_log_probs = pred_info_all["log_prob"].chunk(2, dim=1)[0]
critic_loss = -(target_probs.unsqueeze(0) * pred_log_probs).sum(dim=-1).mean()

self.critic_optimizer.zero_grad(set_to_none=True)
if self.scaler is not None:
Expand All @@ -343,16 +347,21 @@ def update_critic(self, batch: dict[str, torch.Tensor]) -> dict[str, float]:

def update_actor(self, batch: dict[str, torch.Tensor]) -> dict[str, float]:
obs = batch["obs"].to(self.device)
next_obs = batch["next_obs"].to(self.device)
expert_actions = batch["actions"].to(self.device)
critic_obs = batch.get("critic")
critic_obs = critic_obs.to(self.device) if critic_obs is not None else None

obs = self._maybe_normalize_obs(obs, update=False)
next_obs = self._maybe_normalize_obs(next_obs, update=False)
critic_obs = critic_obs if critic_obs is not None else obs

obs_all = torch.cat([obs, next_obs], dim=0)

with self._autocast():
actions, actor_info = self.actor(obs, training=True)
log_probs = actor_info["log_prob"]
actions_all, actor_info_all = self.actor(obs_all, training=True)
actions = actions_all.chunk(2, dim=0)[0]
log_probs = actor_info_all["log_prob"].chunk(2, dim=0)[0]

self._set_requires_grad(self.critic, False)
q_values, _ = self.critic(critic_obs, actions, training=False)
Expand Down
Loading