diff --git a/torchtitan/experiments/rldriving/config_registry.py b/torchtitan/experiments/rldriving/config_registry.py index dc4a8b6436b..0f6bf1152c5 100644 --- a/torchtitan/experiments/rldriving/config_registry.py +++ b/torchtitan/experiments/rldriving/config_registry.py @@ -96,12 +96,12 @@ def rldriving() -> RLDrivingTrainer.Config: gamma=0.95, fps=fps, smooth_lat_cost=0.15, - smooth_long_cost=0.1, - curv_cost=100.0, + smooth_long_cost=0.05, + curv_rate_cost=20.0, ), warm_start_checkpoint=os.getenv( "RLDRIVING_WARM_START_CHECKPOINT", - "849a624a-8a7d-8946-bf04-86148e5e0ef8/56320", + "b9facbcc-4d47-410e-b3ce-dfcbad12ba92/56320", ), tokenizer=NoOpTokenizer.Config(), dataloader=RLDrivingDataLoader.Config( diff --git a/torchtitan/experiments/rldriving/loss.py b/torchtitan/experiments/rldriving/loss.py index d53ec3ef9f9..29cf1d14f74 100644 --- a/torchtitan/experiments/rldriving/loss.py +++ b/torchtitan/experiments/rldriving/loss.py @@ -38,6 +38,7 @@ def _sample_fixed_noise_policy( def _critic_loss( *, + config: RLDrivingLoss.Config, bootstrap_actor_outputs: ActorOutputs, targets: Targets, online_critic: nn.Module, @@ -45,7 +46,6 @@ def _critic_loss( current_inputs: ModelInputs, bootstrap_inputs: ModelInputs, action_noise_A: torch.Tensor, - gamma: float, ) -> LossResult: action_reward_B = targets["action_reward"] rollout_action_BA = action_reward_B[:, 0:2] @@ -66,11 +66,11 @@ def _critic_loss( action=bootstrap_action_BA, ) bootstrap_B = torch.minimum(q1_target_B, q2_target_B) - discounts_N = gamma ** torch.arange( + discounts_N = config.gamma ** torch.arange( rewards_BN.shape[1], device=rewards_BN.device, dtype=rewards_BN.dtype ) discounted_reward_B = (rewards_BN * discounts_N).sum(dim=1) - bootstrap_discount = gamma ** rewards_BN.shape[1] + bootstrap_discount = config.gamma ** rewards_BN.shape[1] target_B = discounted_reward_B + bootstrap_discount * bootstrap_B q_target_abs_gap_B = torch.abs(q1_target_B - q2_target_B) q_rollout_abs_gap_B = torch.abs(q1_rollout_B - q2_rollout_B) @@ -94,19 +94,15 @@ def _critic_loss( def _actor_loss( *, + config: RLDrivingLoss.Config, actor_outputs: ActorOutputs, next_actor_outputs: ActorOutputs, online_critic: nn.Module, current_inputs: ModelInputs, targets: Targets, - fps: float, - smooth_lat_cost: float, - smooth_long_cost: float, - curv_cost: float, - action_bound: float, - action_bound_loss_weight: float, ) -> LossResult: action_pred_BA = actor_outputs[ACTION_OUTPUT] + next_action_pred_BA = next_actor_outputs[ACTION_OUTPUT] q1_new_B, q2_new_B = online_critic( inputs=current_inputs, action=action_pred_BA[:, :2], @@ -115,18 +111,20 @@ def _actor_loss( actor_q_abs_gap_B = torch.abs(q1_new_B - q2_new_B) curvature_B = action_pred_BA[:, 0] / targets["speed"].squeeze(-1).square() - curvature_loss_B = curv_cost * curvature_B.square() + next_curvature_B = next_action_pred_BA[:, 0] / targets["next_speed"].squeeze(-1).square() + curvature_rate_B = (next_curvature_B - curvature_B) * config.fps + curvature_rate_loss_B = config.curv_rate_cost * curvature_rate_B.square() - command_jerk_BA = (next_actor_outputs[ACTION_OUTPUT][:, :2] - action_pred_BA[:, :2]).abs() * fps - smooth_lat_B = smooth_lat_cost * command_jerk_BA[:, 0].square() - smooth_long_B = smooth_long_cost * command_jerk_BA[:, 1].square() + command_jerk_BA = (next_action_pred_BA[:, :2] - action_pred_BA[:, :2]).abs() * config.fps + smooth_lat_B = config.smooth_lat_cost * command_jerk_BA[:, 0].square() + smooth_long_B = config.smooth_long_cost * command_jerk_BA[:, 1].square() smooth_B = smooth_lat_B + smooth_long_B - actor_loss_B = actor_pi_B + curvature_loss_B + smooth_B + actor_loss_B = actor_pi_B + curvature_rate_loss_B + smooth_B action_abs_BA = torch.abs(action_pred_BA[..., :2]) - action_bound_excess_BA = torch.clamp(action_abs_BA - action_bound, min=0.0) + action_bound_excess_BA = torch.clamp(action_abs_BA - config.action_bound, min=0.0) action_bound_loss_B = action_bound_excess_BA.square().mean(dim=-1) - loss_B = actor_loss_B + action_bound_loss_weight * action_bound_loss_B + loss_B = actor_loss_B + config.action_bound_loss_weight * action_bound_loss_B metrics = { "loss": loss_B.detach(), @@ -134,7 +132,9 @@ def _actor_loss( "actor_pi": actor_pi_B.detach(), "actor_q_abs_gap": actor_q_abs_gap_B.detach(), "actor_curv": curvature_B.detach(), - "actor_curv_loss": curvature_loss_B.detach(), + "actor_curv_rate": curvature_rate_B.detach(), + "actor_curv_rate_abs": curvature_rate_B.abs().detach(), + "actor_curv_rate_loss": curvature_rate_loss_B.detach(), "actor_cmd_lat_jerk": command_jerk_BA[:, 0].detach(), "actor_cmd_long_jerk": command_jerk_BA[:, 1].detach(), "actor_smooth_lat_loss": smooth_lat_B.detach(), @@ -155,7 +155,7 @@ class Config(BaseLoss.Config): fps: float smooth_lat_cost: float = 0.0 smooth_long_cost: float = 0.0 - curv_cost: float = 0.0 + curv_rate_cost: float = 0.0 action_bound: float = 10.0 action_bound_loss_weight: float = 1.0 @@ -165,14 +165,8 @@ def __init__( *, compile_config: CompileConfig | None = None, ) -> None: + self.config = config self.action_noise_A = torch.tensor(config.action_noise) - self.gamma = config.gamma - self.fps = config.fps - self.smooth_lat_cost = config.smooth_lat_cost - self.smooth_long_cost = config.smooth_long_cost - self.curv_cost = config.curv_cost - self.action_bound = config.action_bound - self.action_bound_loss_weight = config.action_bound_loss_weight self.critic_fn = _critic_loss self.actor_fn = _actor_loss @@ -204,6 +198,7 @@ def critic_loss( bootstrap_inputs: ModelInputs, ) -> LossResult: return self.critic_fn( + config=self.config, bootstrap_actor_outputs=bootstrap_actor_outputs, targets=targets, online_critic=online_critic, @@ -211,7 +206,6 @@ def critic_loss( current_inputs=current_inputs, bootstrap_inputs=bootstrap_inputs, action_noise_A=self.action_noise_A, - gamma=self.gamma, ) def actor_loss( @@ -224,15 +218,10 @@ def actor_loss( targets: Targets, ) -> LossResult: return self.actor_fn( + config=self.config, actor_outputs=actor_outputs, next_actor_outputs=next_actor_outputs, online_critic=online_critic, current_inputs=current_inputs, targets=targets, - fps=self.fps, - smooth_lat_cost=self.smooth_lat_cost, - smooth_long_cost=self.smooth_long_cost, - curv_cost=self.curv_cost, - action_bound=self.action_bound, - action_bound_loss_weight=self.action_bound_loss_weight, )