Skip to content
6 changes: 3 additions & 3 deletions torchtitan/experiments/rldriving/config_registry.py
Original file line number Diff line number Diff line change
Expand Up @@ -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(
Expand Down
53 changes: 21 additions & 32 deletions torchtitan/experiments/rldriving/loss.py
Original file line number Diff line number Diff line change
Expand Up @@ -38,14 +38,14 @@ def _sample_fixed_noise_policy(

def _critic_loss(
*,
config: RLDrivingLoss.Config,
bootstrap_actor_outputs: ActorOutputs,
targets: Targets,
online_critic: nn.Module,
target_critic: nn.Module,
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]
Expand All @@ -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)
Expand All @@ -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],
Expand All @@ -115,26 +111,30 @@ 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(),
"actor_loss": actor_loss_B.detach(),
"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(),
Expand All @@ -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

Expand All @@ -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
Expand Down Expand Up @@ -204,14 +198,14 @@ 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,
target_critic=target_critic,
current_inputs=current_inputs,
bootstrap_inputs=bootstrap_inputs,
action_noise_A=self.action_noise_A,
gamma=self.gamma,
)

def actor_loss(
Expand All @@ -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,
)